From c60ad48e4f7174089d2561e56dfb6542b48832b7 Mon Sep 17 00:00:00 2001 From: kiroxu <148877251+BabyDrangoner@users.noreply.github.com> Date: Wed, 12 Aug 2026 18:54:40 +0000 Subject: [PATCH 1/2] [Model] Skip unused Jina V5 output layers Co-authored-by: Codex Signed-off-by: kiroxu <148877251+BabyDrangoner@users.noreply.github.com> --- .../pooling/test_jina_embeddings_v5.py | 45 +++++++++++++++++++ vllm/model_executor/models/jina.py | 24 ++++++++-- 2 files changed, 66 insertions(+), 3 deletions(-) diff --git a/tests/models/language/pooling/test_jina_embeddings_v5.py b/tests/models/language/pooling/test_jina_embeddings_v5.py index fbfebcc5cedb..9e26dc4a2fd4 100644 --- a/tests/models/language/pooling/test_jina_embeddings_v5.py +++ b/tests/models/language/pooling/test_jina_embeddings_v5.py @@ -14,6 +14,7 @@ from typing import cast import pytest +import torch from transformers import PretrainedConfig from vllm.config import ModelConfig @@ -21,6 +22,13 @@ MODELS_CONFIG_MAP, JinaEmbeddingsV5ModelConfig, ) +from vllm.model_executor.models.jina import ( + JinaEmbeddingsV5DecoderModel, + JinaEmbeddingsV5EncoderModel, +) +from vllm.model_executor.models.llama import LlamaForCausalLM +from vllm.model_executor.models.qwen3 import Qwen3ForCausalLM +from vllm.model_executor.models.utils import StageMissingLayer def _model_config(hf_config: PretrainedConfig) -> ModelConfig: @@ -65,3 +73,40 @@ def test_supported_decoder_backbone_is_accepted(): JinaEmbeddingsV5ModelConfig.verify_and_update_model_config( _model_config(PretrainedConfig(is_decoder=True)) ) + + +@pytest.mark.cpu_test +@pytest.mark.parametrize( + ("model_cls", "base_cls"), + [ + (JinaEmbeddingsV5DecoderModel, Qwen3ForCausalLM), + (JinaEmbeddingsV5EncoderModel, LlamaForCausalLM), + ], +) +def test_pooling_model_skips_output_layer(monkeypatch, model_cls, base_cls): + class FakeLMHead(torch.nn.Linear): + pass + + class FakeLogitsProcessor(torch.nn.Module): + pass + + def fake_base_init(self, *, vllm_config, prefix=""): + torch.nn.Module.__init__(self) + self.model = torch.nn.Linear(2, 2, bias=False) + self.lm_head = FakeLMHead(2, 1024, bias=False) + self.logits_processor = FakeLogitsProcessor() + + import vllm.model_executor.models.jina as jina_module + + monkeypatch.setattr(jina_module, "ParallelLMHead", FakeLMHead) + monkeypatch.setattr(jina_module, "LogitsProcessor", FakeLogitsProcessor) + monkeypatch.setattr(jina_module, "_setup_jina_v5_task_and_pooler", lambda *_: None) + monkeypatch.setattr(base_cls, "__init__", fake_base_init) + + model = model_cls(vllm_config=object()) + + assert isinstance(model.lm_head, StageMissingLayer) + assert isinstance(model.logits_processor, StageMissingLayer) + params = dict(model.named_parameters()) + assert list(params) == ["model.weight"] + assert params["model.weight"] is model.model.weight diff --git a/vllm/model_executor/models/jina.py b/vllm/model_executor/models/jina.py index 63d690bb3d99..5fae3031542a 100644 --- a/vllm/model_executor/models/jina.py +++ b/vllm/model_executor/models/jina.py @@ -11,6 +11,8 @@ from torch import nn from vllm.config import VllmConfig +from vllm.model_executor.layers.logits_processor import LogitsProcessor +from vllm.model_executor.layers.vocab_parallel_embedding import ParallelLMHead from vllm.sequence import IntermediateTensors from vllm.tasks import PoolingTask from vllm.transformers_utils.repo_utils import get_hf_file_bytes @@ -26,7 +28,13 @@ from .interfaces_base import VllmModelForPooling from .llama import LlamaForCausalLM from .qwen3 import Qwen3ForCausalLM, Qwen3Model -from .utils import AutoWeightsLoader, WeightsMapper, maybe_prefix +from .utils import ( + AutoWeightsLoader, + StageMissingLayer, + WeightsMapper, + maybe_prefix, + no_init_weights, +) logger = logging.getLogger(__name__) @@ -267,7 +275,12 @@ class JinaEmbeddingsV5DecoderModel(Qwen3ForCausalLM, VllmModelForPooling): ) def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): - super().__init__(vllm_config=vllm_config, prefix=prefix) + with no_init_weights( + self, + lambda mod: StageMissingLayer("output", mod), + targets=(LogitsProcessor, ParallelLMHead), + ): + super().__init__(vllm_config=vllm_config, prefix=prefix) _setup_jina_v5_task_and_pooler(self, vllm_config) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: @@ -289,7 +302,12 @@ class JinaEmbeddingsV5EncoderModel(LlamaForCausalLM, VllmModelForPooling): ) def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): - super().__init__(vllm_config=vllm_config, prefix=prefix) + with no_init_weights( + self, + lambda mod: StageMissingLayer("output", mod), + targets=(LogitsProcessor, ParallelLMHead), + ): + super().__init__(vllm_config=vllm_config, prefix=prefix) _setup_jina_v5_task_and_pooler(self, vllm_config) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: From 2f5766a43365bba5e997ae36af4f25560c41787a Mon Sep 17 00:00:00 2001 From: kiroxu <148877251+BabyDrangoner@users.noreply.github.com> Date: Thu, 13 Aug 2026 02:05:59 +0000 Subject: [PATCH 2/2] [Model] Remove unnecessary Jina V5 regression test Co-authored-by: Codex Signed-off-by: kiroxu <148877251+BabyDrangoner@users.noreply.github.com> --- .../pooling/test_jina_embeddings_v5.py | 45 ------------------- 1 file changed, 45 deletions(-) diff --git a/tests/models/language/pooling/test_jina_embeddings_v5.py b/tests/models/language/pooling/test_jina_embeddings_v5.py index 9e26dc4a2fd4..fbfebcc5cedb 100644 --- a/tests/models/language/pooling/test_jina_embeddings_v5.py +++ b/tests/models/language/pooling/test_jina_embeddings_v5.py @@ -14,7 +14,6 @@ from typing import cast import pytest -import torch from transformers import PretrainedConfig from vllm.config import ModelConfig @@ -22,13 +21,6 @@ MODELS_CONFIG_MAP, JinaEmbeddingsV5ModelConfig, ) -from vllm.model_executor.models.jina import ( - JinaEmbeddingsV5DecoderModel, - JinaEmbeddingsV5EncoderModel, -) -from vllm.model_executor.models.llama import LlamaForCausalLM -from vllm.model_executor.models.qwen3 import Qwen3ForCausalLM -from vllm.model_executor.models.utils import StageMissingLayer def _model_config(hf_config: PretrainedConfig) -> ModelConfig: @@ -73,40 +65,3 @@ def test_supported_decoder_backbone_is_accepted(): JinaEmbeddingsV5ModelConfig.verify_and_update_model_config( _model_config(PretrainedConfig(is_decoder=True)) ) - - -@pytest.mark.cpu_test -@pytest.mark.parametrize( - ("model_cls", "base_cls"), - [ - (JinaEmbeddingsV5DecoderModel, Qwen3ForCausalLM), - (JinaEmbeddingsV5EncoderModel, LlamaForCausalLM), - ], -) -def test_pooling_model_skips_output_layer(monkeypatch, model_cls, base_cls): - class FakeLMHead(torch.nn.Linear): - pass - - class FakeLogitsProcessor(torch.nn.Module): - pass - - def fake_base_init(self, *, vllm_config, prefix=""): - torch.nn.Module.__init__(self) - self.model = torch.nn.Linear(2, 2, bias=False) - self.lm_head = FakeLMHead(2, 1024, bias=False) - self.logits_processor = FakeLogitsProcessor() - - import vllm.model_executor.models.jina as jina_module - - monkeypatch.setattr(jina_module, "ParallelLMHead", FakeLMHead) - monkeypatch.setattr(jina_module, "LogitsProcessor", FakeLogitsProcessor) - monkeypatch.setattr(jina_module, "_setup_jina_v5_task_and_pooler", lambda *_: None) - monkeypatch.setattr(base_cls, "__init__", fake_base_init) - - model = model_cls(vllm_config=object()) - - assert isinstance(model.lm_head, StageMissingLayer) - assert isinstance(model.logits_processor, StageMissingLayer) - params = dict(model.named_parameters()) - assert list(params) == ["model.weight"] - assert params["model.weight"] is model.model.weight