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

1"""Qwen3.5 multimodal (vision-language) adapter for ``Qwen3_5ForConditionalGeneration``. 

2 

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""" 

7 

8from typing import Any 

9 

10import torch 

11 

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) 

19 

20 

21class Qwen3_5MultimodalArchitectureAdapter(Qwen3ArchitectureAdapter): 

22 """Full vision-language adapter for Qwen3_5ForConditionalGeneration.""" 

23 

24 _testing_lm_attr = "model.language_model" 

25 _testing_hybrid = True 

26 

27 # Qwen3.5's image/video processor (Qwen3VLProcessor) requires torchvision. 

28 required_libraries: list[str] = ["torchvision"] 

29 required_libraries_group: str = "multimodal" 

30 

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") 

34 

35 self.cfg.is_multimodal = True 

36 

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") 

42 

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)