forked from Karylab-cklius/vllm
[Misc] Aligning tokwise pooler heads for consistency (#43041)
Signed-off-by: Taneem Ibrahim <taneem.ibrahim@gmail.com>
This commit is contained in:
@@ -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())
|
||||
|
||||
|
||||
Reference in New Issue
Block a user