From 78739e3bda0466ecbd63de1bda07d5ac09d88dca Mon Sep 17 00:00:00 2001 From: Maxwill Lin <0312fs3@gmail.com> Date: Mon, 22 Jun 2026 03:16:35 -0700 Subject: [PATCH] [Bugfix] Reject matryoshka embedding dimensions above hidden size (#46313) Signed-off-by: EazyReal <8047065+EazyReal@users.noreply.github.com> Co-authored-by: EazyReal <8047065+EazyReal@users.noreply.github.com> --- tests/test_pooling_params.py | 23 +++++++++++++++++++++++ vllm/pooling_params.py | 5 +++++ 2 files changed, 28 insertions(+) diff --git a/tests/test_pooling_params.py b/tests/test_pooling_params.py index 6cf2a82d2ff..6bd97db03dc 100644 --- a/tests/test_pooling_params.py +++ b/tests/test_pooling_params.py @@ -74,6 +74,29 @@ def test_embed_dimensions(model_info: EmbedModelInfo): pooling_params.verify(model_config) +@dataclass() +class MockMatryoshkaModelConfig: + pooler_config: PoolerConfig + is_matryoshka: bool = True + matryoshka_dimensions: list[int] | None = None + served_model_name: str = "mock-matryoshka-model" + embedding_size: int = 32 + + +def test_embed_dimensions_matryoshka_without_list_upper_bound(): + task = "embed" + model_config = MockMatryoshkaModelConfig( + pooler_config=PoolerConfig(seq_pooling_type="CLS"), + matryoshka_dimensions=None, + embedding_size=32, + ) + + PoolingParams(task=task, dimensions=16).verify(model_config) + + with pytest.raises(ValueError): + PoolingParams(task=task, dimensions=64).verify(model_config) + + @pytest.mark.parametrize("task", ["classify"]) def test_classify(task): model_config = MockModelConfig(pooler_config=PoolerConfig(seq_pooling_type="CLS")) diff --git a/vllm/pooling_params.py b/vllm/pooling_params.py index 3cfe9b427bd..240c999ab4b 100644 --- a/vllm/pooling_params.py +++ b/vllm/pooling_params.py @@ -182,6 +182,11 @@ class PoolingParams( ) elif self.dimensions < 1: raise ValueError("Dimensions must be greater than 0") + elif self.dimensions > model_config.embedding_size: + raise ValueError( + "Dimensions must be less than or equal to the model's " + f"embedding size ({model_config.embedding_size})" + ) elif self.task in ["classify", "token_classify"]: if self.use_activation is None: