Skip to content

vllm.model_executor.models.transformers.pooling

Transformers modeling backend mixins for pooling models.

Classes:

ClassifierWithReshape

Bases: Module

Token extraction has already been applied in pooler.pooling.

Add dim to match expected input shape of classifier.forward.

Source code in vllm/model_executor/models/transformers/pooling.py
class ClassifierWithReshape(nn.Module):
    """Token extraction has already been applied in `pooler.pooling`.

    Add dim to match expected input shape of `classifier.forward`.
    """

    def forward(self, *args, **kwargs):
        if len(args) > 0:
            args = (args[0].unsqueeze(1), *args[1:])
        return super().forward(*args, **kwargs)