[Model] Update pooling model interface (#21058)
Signed-off-by: DarkLight1337 <tlleungac@connect.ust.hk>
This commit is contained in:
@@ -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]],
|
||||
|
||||
Reference in New Issue
Block a user