[Model][1/N] Support multiple poolers at model level (#21227)
Signed-off-by: DarkLight1337 <tlleungac@connect.ust.hk>
This commit is contained in:
@@ -1,15 +1,16 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Mapping, Set
|
||||
from dataclasses import dataclass
|
||||
from enum import IntEnum
|
||||
from itertools import groupby
|
||||
from typing import Callable, 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
|
||||
@@ -21,6 +22,10 @@ from vllm.utils import resolve_obj_by_qualname
|
||||
from vllm.v1.pool.metadata import PoolingMetadata as V1PoolingMetadata
|
||||
|
||||
PoolingMetadata = Union[V0PoolingMetadata, V1PoolingMetadata]
|
||||
PoolingFn = Callable[
|
||||
[Union[torch.Tensor, list[torch.Tensor]], PoolingMetadata],
|
||||
Union[torch.Tensor, list[torch.Tensor]]]
|
||||
ClassifierFn = Callable[[torch.Tensor], torch.Tensor]
|
||||
|
||||
|
||||
class PoolingType(IntEnum):
|
||||
@@ -79,37 +84,81 @@ class Pooler(nn.Module, ABC):
|
||||
"""The interface required for all poolers used in pooling models in vLLM."""
|
||||
|
||||
@staticmethod
|
||||
def from_config_with_defaults(
|
||||
def for_encode(
|
||||
pooler_config: PoolerConfig,
|
||||
pooling_type: PoolingType,
|
||||
normalize: bool,
|
||||
softmax: bool,
|
||||
step_tag_id: Optional[int] = None,
|
||||
returned_token_ids: Optional[list[int]] = None,
|
||||
) -> "Pooler":
|
||||
*,
|
||||
default_pooling_type: PoolingType = PoolingType.ALL,
|
||||
default_normalize: bool = False,
|
||||
default_softmax: bool = False,
|
||||
default_step_tag_id: Optional[int] = None,
|
||||
default_returned_token_ids: Optional[list[int]] = None,
|
||||
):
|
||||
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,
|
||||
pooling_type=default_pooling_type,
|
||||
normalize=default_normalize,
|
||||
softmax=default_softmax,
|
||||
step_tag_id=default_step_tag_id,
|
||||
returned_token_ids=default_returned_token_ids,
|
||||
)
|
||||
|
||||
if pooling_type == PoolingType.STEP:
|
||||
if resolved_config.pooling_type == PoolingType.STEP:
|
||||
return StepPooler.from_config(resolved_config)
|
||||
|
||||
return SimplePooler.from_config(resolved_config)
|
||||
|
||||
def get_pooling_updates(
|
||||
self,
|
||||
task: PoolingTask,
|
||||
) -> Optional[PoolingParamsUpdate]:
|
||||
@staticmethod
|
||||
def for_embed(
|
||||
pooler_config: PoolerConfig,
|
||||
*,
|
||||
default_pooling_type: PoolingType = PoolingType.LAST,
|
||||
default_normalize: bool = True,
|
||||
default_softmax: bool = False,
|
||||
):
|
||||
resolved_config = ResolvedPoolingConfig.from_config_with_defaults(
|
||||
pooler_config=pooler_config,
|
||||
pooling_type=default_pooling_type,
|
||||
normalize=default_normalize,
|
||||
softmax=default_softmax,
|
||||
)
|
||||
|
||||
return SimplePooler.from_config(resolved_config)
|
||||
|
||||
@staticmethod
|
||||
def for_classify(
|
||||
pooler_config: PoolerConfig,
|
||||
classifier: Optional[ClassifierFn],
|
||||
*,
|
||||
default_pooling_type: PoolingType = PoolingType.LAST,
|
||||
default_normalize: bool = False,
|
||||
default_softmax: bool = True,
|
||||
):
|
||||
resolved_config = ResolvedPoolingConfig.from_config_with_defaults(
|
||||
pooler_config=pooler_config,
|
||||
pooling_type=default_pooling_type,
|
||||
normalize=default_normalize,
|
||||
softmax=default_softmax,
|
||||
)
|
||||
base_pooler = SimplePooler.from_config(resolved_config)
|
||||
if classifier is None:
|
||||
return base_pooler
|
||||
|
||||
return ClassifierPooler(
|
||||
pooling=base_pooler.pooling,
|
||||
classifier=classifier,
|
||||
act_fn=base_pooler.head.activation,
|
||||
)
|
||||
|
||||
@abstractmethod
|
||||
def get_supported_tasks(self) -> Set[PoolingTask]:
|
||||
"""Determine which pooling tasks are supported."""
|
||||
raise NotImplementedError
|
||||
|
||||
def get_pooling_updates(self, task: PoolingTask) -> PoolingParamsUpdate:
|
||||
"""
|
||||
Construct the pooling parameters to use for a task,
|
||||
or `None` if the task is not supported.
|
||||
Construct the updated pooling parameters to use for a supported task.
|
||||
"""
|
||||
return None
|
||||
return PoolingParamsUpdate()
|
||||
|
||||
@abstractmethod
|
||||
def forward(
|
||||
@@ -127,9 +176,8 @@ def get_prompt_lens(
|
||||
if isinstance(pooling_metadata, V1PoolingMetadata):
|
||||
return pooling_metadata.prompt_lens
|
||||
|
||||
assert isinstance(hidden_states, torch.Tensor)
|
||||
return PoolingTensors.from_pooling_metadata(
|
||||
pooling_metadata, hidden_states.device).prompt_lens
|
||||
pooling_metadata, hidden_states[0].device).prompt_lens
|
||||
|
||||
|
||||
def get_prompt_token_ids(
|
||||
@@ -149,6 +197,21 @@ def get_prompt_token_ids(
|
||||
]
|
||||
|
||||
|
||||
def get_tasks(pooling_metadata: PoolingMetadata) -> list[PoolingTask]:
|
||||
if isinstance(pooling_metadata, V0PoolingMetadata):
|
||||
pooling_params = [p for _, p in pooling_metadata.seq_groups]
|
||||
else:
|
||||
pooling_params = pooling_metadata.pooling_params
|
||||
|
||||
tasks: list[PoolingTask] = [
|
||||
task for pooling_param in pooling_params
|
||||
if (task := pooling_param.task) is not None
|
||||
]
|
||||
assert len(pooling_params) == len(tasks)
|
||||
|
||||
return tasks
|
||||
|
||||
|
||||
def get_classification_activation_function(config: PretrainedConfig):
|
||||
return PoolerClassify()
|
||||
|
||||
@@ -172,7 +235,8 @@ def get_cross_encoder_activation_function(config: PretrainedConfig):
|
||||
return PoolerScore()
|
||||
|
||||
|
||||
def build_output(all_data: torch.Tensor) -> PoolerOutput:
|
||||
def build_output(
|
||||
all_data: Union[torch.Tensor, list[torch.Tensor]], ) -> PoolerOutput:
|
||||
all_outputs = [PoolingSequenceGroupOutput(data) for data in all_data]
|
||||
return PoolerOutput(outputs=all_outputs)
|
||||
|
||||
@@ -193,12 +257,12 @@ class PoolingMethod(nn.Module, ABC):
|
||||
raise NotImplementedError(f"Unsupported method: {pooling_type}")
|
||||
|
||||
@abstractmethod
|
||||
def get_pooling_updates(
|
||||
self,
|
||||
task: PoolingTask,
|
||||
) -> Optional[PoolingParamsUpdate]:
|
||||
def get_supported_tasks(self) -> Set[PoolingTask]:
|
||||
raise NotImplementedError
|
||||
|
||||
def get_pooling_updates(self, task: PoolingTask) -> PoolingParamsUpdate:
|
||||
return PoolingParamsUpdate()
|
||||
|
||||
@abstractmethod
|
||||
def forward_one(
|
||||
self,
|
||||
@@ -237,16 +301,8 @@ class PoolingMethod(nn.Module, ABC):
|
||||
|
||||
class CLSPool(PoolingMethod):
|
||||
|
||||
def get_pooling_updates(
|
||||
self,
|
||||
task: PoolingTask,
|
||||
) -> Optional[PoolingParamsUpdate]:
|
||||
# The equalities are split up to keep mypy happy
|
||||
if (task == "encode" or task == "embed" or task == "classify"
|
||||
or task == "score"):
|
||||
return PoolingParamsUpdate()
|
||||
|
||||
assert_never(task)
|
||||
def get_supported_tasks(self) -> Set[PoolingTask]:
|
||||
return {"encode", "embed", "classify", "score"}
|
||||
|
||||
def forward_one(
|
||||
self,
|
||||
@@ -270,16 +326,8 @@ class CLSPool(PoolingMethod):
|
||||
|
||||
class LastPool(PoolingMethod):
|
||||
|
||||
def get_pooling_updates(
|
||||
self,
|
||||
task: PoolingTask,
|
||||
) -> Optional[PoolingParamsUpdate]:
|
||||
# The equalities are split up to keep mypy happy
|
||||
if (task == "encode" or task == "embed" or task == "classify"
|
||||
or task == "score"):
|
||||
return PoolingParamsUpdate()
|
||||
|
||||
assert_never(task)
|
||||
def get_supported_tasks(self) -> Set[PoolingTask]:
|
||||
return {"encode", "embed", "classify", "score"}
|
||||
|
||||
def forward_one(
|
||||
self,
|
||||
@@ -299,18 +347,8 @@ class LastPool(PoolingMethod):
|
||||
|
||||
class AllPool(PoolingMethod):
|
||||
|
||||
def get_pooling_updates(
|
||||
self,
|
||||
task: PoolingTask,
|
||||
) -> Optional[PoolingParamsUpdate]:
|
||||
if task == "encode":
|
||||
return PoolingParamsUpdate()
|
||||
|
||||
# The equalities are split up to keep mypy happy
|
||||
if task == "embed" or task == "classify" or task == "score":
|
||||
return None
|
||||
|
||||
assert_never(task)
|
||||
def get_supported_tasks(self) -> Set[PoolingTask]:
|
||||
return {"encode"}
|
||||
|
||||
def forward_one(
|
||||
self,
|
||||
@@ -327,28 +365,13 @@ class AllPool(PoolingMethod):
|
||||
hidden_states: torch.Tensor,
|
||||
prompt_lens: torch.Tensor,
|
||||
) -> Union[list[torch.Tensor], torch.Tensor]:
|
||||
offset = 0
|
||||
pooled_data = list[torch.Tensor]()
|
||||
|
||||
for prompt_len in prompt_lens:
|
||||
pooled_data.append(hidden_states[offset:offset + prompt_len])
|
||||
offset += prompt_len
|
||||
|
||||
return pooled_data
|
||||
return list(hidden_states.split_with_sizes(prompt_lens.tolist()))
|
||||
|
||||
|
||||
class MeanPool(PoolingMethod):
|
||||
|
||||
def get_pooling_updates(
|
||||
self,
|
||||
task: PoolingTask,
|
||||
) -> Optional[PoolingParamsUpdate]:
|
||||
# The equalities are split up to keep mypy happy
|
||||
if (task == "encode" or task == "embed" or task == "classify"
|
||||
or task == "score"):
|
||||
return PoolingParamsUpdate()
|
||||
|
||||
assert_never(task)
|
||||
def get_supported_tasks(self) -> Set[PoolingTask]:
|
||||
return {"encode", "embed", "classify", "score"}
|
||||
|
||||
def forward_one(
|
||||
self,
|
||||
@@ -529,24 +552,6 @@ class SimplePooler(Pooler):
|
||||
3. Returns structured results as `PoolerOutput`.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def from_config_with_defaults( # type: ignore[override]
|
||||
cls,
|
||||
pooler_config: PoolerConfig,
|
||||
pooling_type: PoolingType,
|
||||
normalize: bool,
|
||||
softmax: bool,
|
||||
) -> "SimplePooler":
|
||||
resolved_config = ResolvedPoolingConfig.from_config_with_defaults(
|
||||
pooler_config=pooler_config,
|
||||
pooling_type=pooling_type,
|
||||
normalize=normalize,
|
||||
softmax=softmax,
|
||||
)
|
||||
assert resolved_config.pooling_type != PoolingType.STEP
|
||||
|
||||
return cls.from_config(resolved_config)
|
||||
|
||||
@classmethod
|
||||
def from_config(
|
||||
cls,
|
||||
@@ -563,10 +568,10 @@ class SimplePooler(Pooler):
|
||||
self.pooling = pooling
|
||||
self.head = head
|
||||
|
||||
def get_pooling_updates(
|
||||
self,
|
||||
task: PoolingTask,
|
||||
) -> Optional[PoolingParamsUpdate]:
|
||||
def get_supported_tasks(self) -> Set[PoolingTask]:
|
||||
return self.pooling.get_supported_tasks()
|
||||
|
||||
def get_pooling_updates(self, task: PoolingTask) -> PoolingParamsUpdate:
|
||||
return self.pooling.get_pooling_updates(task)
|
||||
|
||||
def forward(
|
||||
@@ -627,18 +632,11 @@ class StepPooler(Pooler):
|
||||
|
||||
return pooled_data
|
||||
|
||||
def get_pooling_updates(
|
||||
self,
|
||||
task: PoolingTask,
|
||||
) -> Optional[PoolingParamsUpdate]:
|
||||
if task == "encode":
|
||||
return PoolingParamsUpdate(requires_token_ids=True)
|
||||
def get_supported_tasks(self) -> Set[PoolingTask]:
|
||||
return {"encode"}
|
||||
|
||||
# The equalities are split up to keep mypy happy
|
||||
if task == "embed" or task == "classify" or task == "score":
|
||||
return None
|
||||
|
||||
assert_never(task)
|
||||
def get_pooling_updates(self, task: PoolingTask) -> PoolingParamsUpdate:
|
||||
return PoolingParamsUpdate(requires_token_ids=True)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -650,68 +648,43 @@ class StepPooler(Pooler):
|
||||
return build_output(pooled_data)
|
||||
|
||||
|
||||
PoolingFn = Callable[
|
||||
[Union[torch.Tensor, list[torch.Tensor]], PoolingMetadata],
|
||||
Union[torch.Tensor, list[torch.Tensor]]]
|
||||
ClassifierFn = Callable[[torch.Tensor], torch.Tensor]
|
||||
|
||||
|
||||
class ClassifierPooler(nn.Module):
|
||||
class ClassifierPooler(Pooler):
|
||||
"""A pooling layer for classification tasks.
|
||||
|
||||
This layer does the following:
|
||||
1. Applies a classification layer to the hidden states.
|
||||
2. Optionally applies a pooler layer.
|
||||
3. Applies an activation function to the output. In the case of
|
||||
classification models it is either sigmoid or softmax. In the
|
||||
case of scoring models, the same behavior is configuration
|
||||
dependent, as in the sentence-transformers library.
|
||||
3. Applies an activation function to the output.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def act_fn_for_seq_cls(config: ModelConfig):
|
||||
return get_classification_activation_function(config.hf_config)
|
||||
|
||||
@staticmethod
|
||||
def act_fn_for_cross_encoder(config: ModelConfig):
|
||||
return get_cross_encoder_activation_function(config.hf_config)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: ModelConfig,
|
||||
pooling: PoolingFn,
|
||||
classifier: ClassifierFn,
|
||||
act_fn: Optional[PoolerActivation] = None,
|
||||
act_fn: PoolerActivation,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.pooling = pooling
|
||||
self.classifier = classifier
|
||||
self.act_fn = act_fn
|
||||
|
||||
self.classification_act_fn = get_classification_activation_function(
|
||||
config.hf_config) if act_fn is None else act_fn
|
||||
self.cross_encoder_act_fn = get_cross_encoder_activation_function(
|
||||
config.hf_config) if act_fn is None else act_fn
|
||||
|
||||
def _get_act_fn(self, task: PoolingTask):
|
||||
if task == "encode" or task == "classify":
|
||||
return self.classification_act_fn
|
||||
if task == "score":
|
||||
return self.cross_encoder_act_fn
|
||||
|
||||
raise ValueError(f"Unsupported task: {task!r}")
|
||||
|
||||
def get_pooling_updates(
|
||||
self,
|
||||
task: PoolingTask,
|
||||
) -> Optional[PoolingParamsUpdate]:
|
||||
# The equalities are split up to keep mypy happy
|
||||
if task == "encode" or task == "classify" or task == "score":
|
||||
return PoolingParamsUpdate()
|
||||
|
||||
if task == "embed":
|
||||
return None
|
||||
|
||||
assert_never(task)
|
||||
def get_supported_tasks(self) -> Set[PoolingTask]:
|
||||
return {"classify", "score"}
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: Union[torch.Tensor, list[torch.Tensor]],
|
||||
pooling_metadata: PoolingMetadata,
|
||||
) -> PoolerOutput:
|
||||
"""Pools sentence pair scores from the hidden_states."""
|
||||
pooled_data = self.pooling(hidden_states, pooling_metadata)
|
||||
|
||||
# apply classifier once on the full batch if possible
|
||||
@@ -722,28 +695,59 @@ class ClassifierPooler(nn.Module):
|
||||
else:
|
||||
pooled_output = [self.classifier(data) for data in pooled_data]
|
||||
|
||||
task_list: list[PoolingTask]
|
||||
if isinstance(pooling_metadata, V0PoolingMetadata):
|
||||
task_list = [
|
||||
task for _, pooling_param in pooling_metadata.seq_groups
|
||||
if (task := pooling_param.task) is not None
|
||||
]
|
||||
else:
|
||||
task_list = [
|
||||
task for pooling_param in pooling_metadata.pooling_params
|
||||
if (task := pooling_param.task) is not None
|
||||
]
|
||||
|
||||
assert len(task_list) == len(pooled_output)
|
||||
|
||||
# shape of scores: (batch_size, num_labels)
|
||||
if len(set(task_list)) <= 1:
|
||||
act_fn = self._get_act_fn(task_list[0])
|
||||
scores = act_fn(pooled_output)
|
||||
else:
|
||||
scores = torch.stack([
|
||||
self._get_act_fn(task)(vecs)
|
||||
for task, vecs in zip(task_list, pooled_output)
|
||||
])
|
||||
scores = self.act_fn(pooled_output)
|
||||
|
||||
return build_output(scores)
|
||||
|
||||
|
||||
class DispatchPooler(Pooler):
|
||||
"""Dispatches calls to a sub-pooler based on the pooling task."""
|
||||
|
||||
def __init__(self, poolers_by_task: Mapping[PoolingTask, Pooler]) -> None:
|
||||
super().__init__()
|
||||
|
||||
for task, pooler in poolers_by_task.items():
|
||||
if task not in pooler.get_supported_tasks():
|
||||
raise ValueError(
|
||||
f"{pooler=} does not support {task=}. "
|
||||
f"Supported tasks: {pooler.get_supported_tasks()}")
|
||||
|
||||
self.poolers_by_task = poolers_by_task
|
||||
|
||||
def get_supported_tasks(self) -> Set[PoolingTask]:
|
||||
return set(self.poolers_by_task)
|
||||
|
||||
def get_pooling_updates(self, task: PoolingTask) -> PoolingParamsUpdate:
|
||||
return self.poolers_by_task[task].get_pooling_updates(task)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: Union[torch.Tensor, list[torch.Tensor]],
|
||||
pooling_metadata: PoolingMetadata,
|
||||
) -> PoolerOutput:
|
||||
poolers_by_task = self.poolers_by_task
|
||||
|
||||
if isinstance(hidden_states, list):
|
||||
hidden_states_lst = hidden_states
|
||||
else:
|
||||
prompt_lens = get_prompt_lens(hidden_states, pooling_metadata)
|
||||
hidden_states_lst = list(hidden_states.split(prompt_lens.tolist()))
|
||||
|
||||
outputs = list[PoolingSequenceGroupOutput]()
|
||||
offset = 0
|
||||
for task, group in groupby(get_tasks(pooling_metadata)):
|
||||
if not (pooler := poolers_by_task.get(task)):
|
||||
raise ValueError(
|
||||
f"Unsupported task: {task} "
|
||||
f"Supported tasks: {self.get_supported_tasks()}")
|
||||
|
||||
num_items = len(list(group))
|
||||
group_output: PoolerOutput = pooler(
|
||||
hidden_states_lst[offset:offset + num_items],
|
||||
pooling_metadata[offset:offset + num_items],
|
||||
)
|
||||
|
||||
outputs.extend(group_output.outputs)
|
||||
offset += num_items
|
||||
|
||||
return PoolerOutput(outputs)
|
||||
|
||||
Reference in New Issue
Block a user