[Misc] Aligning tokwise pooler heads for consistency (#43041)

Signed-off-by: Taneem Ibrahim <taneem.ibrahim@gmail.com>
This commit is contained in:
Taneem Ibrahim
2026-05-19 06:16:42 +00:00
committed by GitHub
parent f1e3f0e6d6
commit 4a4fdabe28
4 changed files with 9 additions and 4 deletions
@@ -89,7 +89,7 @@ class SequencePooler(Pooler):
return pooled_data
def pooler_for_embed(pooler_config: PoolerConfig):
def pooler_for_embed(pooler_config: PoolerConfig) -> SequencePooler:
pooling = get_seq_pooling_method(pooler_config.get_seq_pooling_type())
vllm_config = get_current_vllm_config()
@@ -109,7 +109,7 @@ def pooler_for_classify(
pooling: SequencePoolingMethod | SequencePoolingFn | None = None,
classifier: ClassifierFn | None = None,
act_fn: PoolerActivation | None = None,
):
) -> SequencePooler:
if pooling is None:
pooling = get_seq_pooling_method(pooler_config.get_seq_pooling_type())
@@ -18,6 +18,8 @@ from .methods import (
from .poolers import (
TokenPooler,
TokenPoolerOutput,
TokenPoolingFn,
TokenPoolingHeadFn,
pooler_for_token_classify,
pooler_for_token_embed,
)
@@ -34,6 +36,8 @@ __all__ = [
"get_tok_pooling_method",
"TokenPooler",
"TokenPoolerOutput",
"TokenPoolingFn",
"TokenPoolingHeadFn",
"pooler_for_token_classify",
"pooler_for_token_embed",
]
@@ -78,7 +78,8 @@ class TokenEmbeddingPoolerHead(TokenPoolerHead):
# embeddings shape: [n_tokens, embedding_size]
# for matryoshka representation
embeddings = embeddings[..., : pooling_param.dimensions]
if pooling_param.dimensions is not None:
embeddings = embeddings[..., : pooling_param.dimensions]
# for normalize
if self.activation is not None and pooling_param.use_activation:
@@ -118,7 +118,7 @@ def pooler_for_token_classify(
pooling: TokenPoolingMethod | TokenPoolingFn | None = None,
classifier: ClassifierFn | None = None,
act_fn: PoolerActivation | None = None,
):
) -> TokenPooler:
if pooling is None:
pooling = get_tok_pooling_method(pooler_config.get_tok_pooling_type())