Coverage for transformer_lens/model_bridge/supported_architectures/qwen3_5_multimodal.py: 100%
19 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
1"""Qwen3.5 multimodal (vision-language) adapter for ``Qwen3_5ForConditionalGeneration``.
3Reuses the text-only Qwen3.5 hybrid backbone nested under ``model.language_model`` and adds
4the vision tower (``model.visual``) + merger. The HF model runs the vision computation during
5forward; this adapter only supplies the component mapping (hooks + weights).
6"""
8from typing import Any
10import torch
12from transformer_lens.model_bridge.generalized_components import VisionProjectionBridge
13from transformer_lens.model_bridge.generalized_components.qwen3_5_vision_encoder import (
14 Qwen3_5VisionEncoderBridge,
15)
16from transformer_lens.model_bridge.supported_architectures.qwen3 import (
17 Qwen3ArchitectureAdapter,
18)
21class Qwen3_5MultimodalArchitectureAdapter(Qwen3ArchitectureAdapter):
22 """Full vision-language adapter for Qwen3_5ForConditionalGeneration."""
24 _testing_lm_attr = "model.language_model"
25 _testing_hybrid = True
27 # Qwen3.5's image/video processor (Qwen3VLProcessor) requires torchvision.
28 required_libraries: list[str] = ["torchvision"]
29 required_libraries_group: str = "multimodal"
31 def __init__(self, cfg: Any) -> None:
32 setattr(cfg, "gated_q_proj", True)
33 super().__init__(cfg, hybrid=True, lm_prefix="model.language_model")
35 self.cfg.is_multimodal = True
37 self._extract_vision_dims(cfg)
38 self.components["vision_encoder"] = Qwen3_5VisionEncoderBridge(
39 name="model.visual", config=self.cfg
40 )
41 self.components["vision_projector"] = VisionProjectionBridge(name="model.visual.merger")
43 def preprocess_weights(self, state_dict: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
44 """Slice query half from gated q_proj.weight (matcher is path-prefix-agnostic)."""
45 return self._preprocess_gated_q_proj(state_dict, self.cfg.n_heads, self.cfg.d_head)