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

1"""Qwen3-VL architecture adapter. 

2 

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

17 

18from typing import Any 

19 

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) 

36 

37 

38class _DeepStackMergerBridge(GeneralizedComponent): 

39 """Per-level DeepStack patch merger (multi-scale visual features).""" 

40 

41 is_list_item: bool = True 

42 

43 

44class Qwen3VLArchitectureAdapter(ArchitectureAdapter): 

45 """Architecture adapter for Qwen3VLForConditionalGeneration models.""" 

46 

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

48 required_libraries_group: str = "multimodal" 

49 

50 def __init__(self, cfg: Any) -> None: 

51 """Initialize the Qwen3-VL architecture adapter.""" 

52 super().__init__(cfg) 

53 

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 

60 

61 self._extract_vision_dims(cfg) 

62 

63 self.weight_processing_conversions = { 

64 **self._qkvo_weight_conversions(), 

65 } 

66 

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 } 

106 

107 def _build_mlp_bridge(self) -> Any: 

108 """Dense gated MLP; the MoE variant overrides this.""" 

109 return self._gated_mlp()