[Model] Update pooling model interface (#21058)

Signed-off-by: DarkLight1337 <tlleungac@connect.ust.hk>
This commit is contained in:
Cyrus Leung
2025-07-18 00:05:40 +08:00
committed by GitHub
parent 9fb2d22032
commit 90bd2ab6e3
17 changed files with 247 additions and 345 deletions

View File

@@ -3,22 +3,25 @@
from abc import ABC, abstractmethod
from dataclasses import dataclass
from enum import IntEnum
from typing import Callable, Optional, TypeVar, Union
from typing import Callable, Literal, Optional, TypeVar, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import PretrainedConfig
from typing_extensions import assert_never
from vllm.config import ModelConfig, PoolerConfig
from vllm.model_executor.pooling_metadata import ( # noqa: E501
PoolingMetadata as V0PoolingMetadata)
from vllm.model_executor.pooling_metadata import PoolingTensors
from vllm.pooling_params import PoolingParams
from vllm.sequence import PoolerOutput, PoolingSequenceGroupOutput
from vllm.utils import resolve_obj_by_qualname
from vllm.v1.pool.metadata import PoolingMetadata as V1PoolingMetadata
PoolingMetadata = Union[V0PoolingMetadata, V1PoolingMetadata]
PoolingTask = Literal["encode", "embed", "classify", "score"]
class PoolingType(IntEnum):
@@ -64,6 +67,48 @@ class ResolvedPoolingConfig:
)
class Pooler(nn.Module, ABC):
"""The interface required for all poolers used in pooling models in vLLM."""
@staticmethod
def from_config_with_defaults(
pooler_config: PoolerConfig,
pooling_type: PoolingType,
normalize: bool,
softmax: bool,
step_tag_id: Optional[int] = None,
returned_token_ids: Optional[list[int]] = None,
) -> "Pooler":
resolved_config = ResolvedPoolingConfig.from_config_with_defaults(
pooler_config=pooler_config,
pooling_type=pooling_type,
normalize=normalize,
softmax=softmax,
step_tag_id=step_tag_id,
returned_token_ids=returned_token_ids,
)
if pooling_type == PoolingType.STEP:
return StepPooler.from_config(resolved_config)
return SimplePooler.from_config(resolved_config)
def get_pooling_params(self, task: PoolingTask) -> Optional[PoolingParams]:
"""
Construct the pooling parameters to use for a task,
or `None` if the task is not supported.
"""
return None
@abstractmethod
def forward(
self,
hidden_states: Union[list[torch.Tensor], torch.Tensor],
pooling_metadata: PoolingMetadata,
) -> PoolerOutput:
raise NotImplementedError
def get_prompt_lens(
hidden_states: Union[torch.Tensor, list[torch.Tensor]],
pooling_metadata: PoolingMetadata,
@@ -104,17 +149,6 @@ def build_output(all_data: torch.Tensor) -> PoolerOutput:
return PoolerOutput(outputs=all_outputs)
class BasePooler(nn.Module):
@abstractmethod
def forward(
self,
hidden_states: Union[torch.Tensor, list[torch.Tensor]],
pooling_metadata: PoolingMetadata,
) -> PoolerOutput:
raise NotImplementedError
class PoolingMethod(nn.Module, ABC):
@staticmethod
@@ -130,6 +164,10 @@ class PoolingMethod(nn.Module, ABC):
raise NotImplementedError(f"Unsupported method: {pooling_type}")
@abstractmethod
def get_pooling_params(self, task: PoolingTask) -> Optional[PoolingParams]:
raise NotImplementedError
@abstractmethod
def forward_one(
self,
@@ -168,6 +206,14 @@ class PoolingMethod(nn.Module, ABC):
class CLSPool(PoolingMethod):
def get_pooling_params(self, task: PoolingTask) -> Optional[PoolingParams]:
# The equalities are split up to keep mypy happy
if (task == "encode" or task == "embed" or task == "classify"
or task == "score"):
return PoolingParams()
assert_never(task)
def forward_one(
self,
hidden_states: torch.Tensor,
@@ -190,6 +236,14 @@ class CLSPool(PoolingMethod):
class LastPool(PoolingMethod):
def get_pooling_params(self, task: PoolingTask) -> Optional[PoolingParams]:
# The equalities are split up to keep mypy happy
if (task == "encode" or task == "embed" or task == "classify"
or task == "score"):
return PoolingParams()
assert_never(task)
def forward_one(
self,
hidden_states: torch.Tensor,
@@ -208,6 +262,16 @@ class LastPool(PoolingMethod):
class AllPool(PoolingMethod):
def get_pooling_params(self, task: PoolingTask) -> Optional[PoolingParams]:
if task == "encode":
return PoolingParams()
# The equalities are split up to keep mypy happy
if task == "embed" or task == "classify" or task == "score":
return None
assert_never(task)
def forward_one(
self,
hidden_states: torch.Tensor,
@@ -235,6 +299,14 @@ class AllPool(PoolingMethod):
class MeanPool(PoolingMethod):
def get_pooling_params(self, task: PoolingTask) -> Optional[PoolingParams]:
# The equalities are split up to keep mypy happy
if (task == "encode" or task == "embed" or task == "classify"
or task == "score"):
return PoolingParams()
assert_never(task)
def forward_one(
self,
hidden_states: torch.Tensor,
@@ -345,25 +417,6 @@ class LambdaPoolerActivation(PoolerActivation):
class PoolerHead(nn.Module):
@classmethod
def from_config_with_defaults(
cls,
pooler_config: PoolerConfig,
pooling_type: PoolingType,
normalize: bool,
softmax: bool,
) -> "PoolerHead":
resolved_config = ResolvedPoolingConfig.from_config_with_defaults(
pooler_config=pooler_config,
pooling_type=pooling_type,
normalize=normalize,
softmax=softmax,
step_tag_id=None,
returned_token_ids=None,
)
return cls.from_config(resolved_config)
@classmethod
def from_config(cls, pooler_config: ResolvedPoolingConfig) -> "PoolerHead":
if pooler_config.normalize and pooler_config.softmax:
@@ -424,21 +477,17 @@ class PoolerHead(nn.Module):
return self.activation(pooled_data)
class SimplePooler(BasePooler):
class SimplePooler(Pooler):
"""A layer that pools specific information from hidden states.
This layer does the following:
1. Extracts specific tokens or aggregates data based on pooling method.
2. Normalizes output if specified.
3. Returns structured results as `PoolerOutput`.
Attributes:
pooling_type: The type of pooling to use.
normalize: Whether to normalize the pooled data.
"""
@classmethod
def from_config_with_defaults(
def from_config_with_defaults( # type: ignore[override]
cls,
pooler_config: PoolerConfig,
pooling_type: PoolingType,
@@ -471,6 +520,9 @@ class SimplePooler(BasePooler):
self.pooling = pooling
self.head = head
def get_pooling_params(self, task: PoolingTask) -> Optional[PoolingParams]:
return self.pooling.get_pooling_params(task)
def forward(
self,
hidden_states: Union[torch.Tensor, list[torch.Tensor]],
@@ -481,7 +533,7 @@ class SimplePooler(BasePooler):
return build_output(pooled_data)
class StepPooler(BasePooler):
class StepPooler(Pooler):
@classmethod
def from_config(cls, pooler_config: ResolvedPoolingConfig) -> "StepPooler":
@@ -543,6 +595,16 @@ class StepPooler(BasePooler):
return pooled_data
def get_pooling_params(self, task: PoolingTask) -> Optional[PoolingParams]:
if task == "encode":
return PoolingParams(logits_processing_needs_token_ids=True)
# The equalities are split up to keep mypy happy
if task == "embed" or task == "classify" or task == "score":
return None
assert_never(task)
def forward(
self,
hidden_states: Union[torch.Tensor, list[torch.Tensor]],
@@ -553,32 +615,6 @@ class StepPooler(BasePooler):
return build_output(pooled_data)
class Pooler(nn.Module):
@staticmethod
def from_config_with_defaults(
pooler_config: PoolerConfig,
pooling_type: PoolingType,
normalize: bool,
softmax: bool,
step_tag_id: Optional[int] = None,
returned_token_ids: Optional[list[int]] = None,
) -> BasePooler:
resolved_config = ResolvedPoolingConfig.from_config_with_defaults(
pooler_config=pooler_config,
pooling_type=pooling_type,
normalize=normalize,
softmax=softmax,
step_tag_id=step_tag_id,
returned_token_ids=returned_token_ids,
)
if pooling_type == PoolingType.STEP:
return StepPooler.from_config(resolved_config)
return SimplePooler.from_config(resolved_config)
PoolingFn = Callable[
[Union[torch.Tensor, list[torch.Tensor]], PoolingMetadata],
Union[torch.Tensor, list[torch.Tensor]]]
@@ -618,6 +654,18 @@ class ClassifierPooler(nn.Module):
return (self.cross_encoder_act_fn
if use_cross_encoder else self.classification_act_fn)
def get_pooling_params(self, task: PoolingTask) -> Optional[PoolingParams]:
if task == "encode":
return PoolingParams()
if task == "embed":
return None
if task == "classify":
return PoolingParams()
if task == "score":
return PoolingParams(use_cross_encoder=True)
assert_never(task)
def forward(
self,
hidden_states: Union[torch.Tensor, list[torch.Tensor]],