Coverage for transformer_lens/model_bridge/supported_architectures/qwen3_vl.py: 100%
21 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-08-11 18:50 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-08-11 18:50 +0000
1"""Qwen3-VL architecture adapter.
3Alibaba's Qwen3-VL (``Qwen3VLForConditionalGeneration``): a ViT tower at
4``model.visual`` matching the Qwen3.5 vision layout (learned pos_embed +
52D rotary, qkv/proj attention, fc1/fc2 MLP) plus DeepStack — extra patch
6mergers on early vision blocks whose features the text model injects
7into the residual stream at visual token positions during the first
8decoder layers. The injection itself is a tensor add inside the HF text
9loop (not a module call); the per-level DeepStack mergers are wrapped so
10their features are hookable at the source. Because the injection happens
11BETWEEN block calls, blocks.{i}.hook_out omits it while blocks.{i+1}'s
12input contains it on image runs — resid_post[i] != resid_pre[i+1] at
13visual positions for the first DeepStack layers, and patching resid_post
14there drops the injection (prefer blocks.{i+1}.hook_in). Text attention
15is Qwen3-style (per-head QK RMS-norm) with interleaved mRoPE, HF-native.
16"""
18from typing import Any
20from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
21from transformer_lens.model_bridge.generalized_components import (
22 AttentionBridge,
23 BlockBridge,
24 EmbeddingBridge,
25 LinearBridge,
26 RMSNormalizationBridge,
27 UnembeddingBridge,
28 VisionProjectionBridge,
29)
30from transformer_lens.model_bridge.generalized_components.base import (
31 GeneralizedComponent,
32)
33from transformer_lens.model_bridge.generalized_components.qwen3_5_vision_encoder import (
34 Qwen3_5VisionEncoderBridge,
35)
38class _DeepStackMergerBridge(GeneralizedComponent):
39 """Per-level DeepStack patch merger (multi-scale visual features)."""
41 is_list_item: bool = True
44class Qwen3VLArchitectureAdapter(ArchitectureAdapter):
45 """Architecture adapter for Qwen3VLForConditionalGeneration models."""
47 required_libraries: list[str] = ["torchvision"]
48 required_libraries_group: str = "multimodal"
50 def __init__(self, cfg: Any) -> None:
51 """Initialize the Qwen3-VL architecture adapter."""
52 super().__init__(cfg)
54 self.cfg.is_multimodal = True
55 self._set_rms_rotary_defaults()
56 self.cfg.attn_implementation = "eager"
57 # Qwen tokenizers have no BOS; the prepend fallback would inject
58 # <|im_end|>, which reads as an ended turn.
59 self.cfg.default_prepend_bos = False
61 self._extract_vision_dims(cfg)
63 self.weight_processing_conversions = {
64 **self._qkvo_weight_conversions(),
65 }
67 self.component_mapping = {
68 # The Qwen3.5 vision bridge defaults (patch_embed, learned
69 # pos_embed, qkv/proj + fc1/fc2 blocks) match this tower exactly;
70 # Qwen3-VL adds the DeepStack merger stack on top.
71 "vision_encoder": Qwen3_5VisionEncoderBridge(
72 name="model.visual",
73 config=self.cfg,
74 submodules={
75 "deepstack_mergers": _DeepStackMergerBridge(name="deepstack_merger_list"),
76 },
77 ),
78 "vision_projector": VisionProjectionBridge(name="model.visual.merger"),
79 "embed": EmbeddingBridge(name="model.language_model.embed_tokens"),
80 "blocks": BlockBridge(
81 name="model.language_model.layers",
82 submodules={
83 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg),
84 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg),
85 # Interleaved mRoPE + per-head QK-norm live in HF's forward.
86 "attn": AttentionBridge(
87 name="self_attn",
88 config=self.cfg,
89 submodules={
90 "q": LinearBridge(name="q_proj"),
91 "k": LinearBridge(name="k_proj"),
92 "v": LinearBridge(name="v_proj"),
93 "o": LinearBridge(name="o_proj"),
94 "q_norm": RMSNormalizationBridge(name="q_norm", config=self.cfg),
95 "k_norm": RMSNormalizationBridge(name="k_norm", config=self.cfg),
96 },
97 maintain_native_attention=True,
98 requires_attention_mask=True,
99 ),
100 "mlp": self._build_mlp_bridge(),
101 },
102 ),
103 "ln_final": RMSNormalizationBridge(name="model.language_model.norm", config=self.cfg),
104 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg),
105 }
107 def _build_mlp_bridge(self) -> Any:
108 """Dense gated MLP; the MoE variant overrides this."""
109 return self._gated_mlp()