Coverage for transformer_lens/model_bridge/supported_architectures/llava.py: 100%

20 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-01 16:23 +0000

1"""LLava architecture adapter. 

2 

3This adapter supports LlavaForConditionalGeneration, the vision-language 

4model combining a CLIP vision encoder with a LLaMA language model. 

5""" 

6 

7from typing import Any 

8 

9from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter 

10from transformer_lens.model_bridge.generalized_components import ( 

11 BlockBridge, 

12 CLIPVisionEncoderBridge, 

13 EmbeddingBridge, 

14 LinearBridge, 

15 RMSNormalizationBridge, 

16 RotaryEmbeddingBridge, 

17 SiglipVisionEncoderBridge, 

18 UnembeddingBridge, 

19 VisionProjectionBridge, 

20) 

21from transformer_lens.model_bridge.generalized_components.base import ( 

22 GeneralizedComponent, 

23) 

24from transformer_lens.model_bridge.generalized_components.position_embeddings_attention import ( 

25 PositionEmbeddingsAttentionBridge, 

26) 

27 

28 

29class LlavaArchitectureAdapter(ArchitectureAdapter): 

30 """Architecture adapter for LLava multimodal models (LlavaForConditionalGeneration). 

31 

32 This adapter handles vision-language models like LLava 1.5. 

33 The model structure is: 

34 - model.vision_tower: CLIP vision encoder 

35 - model.multi_modal_projector: 2-layer MLP (Linear -> GELU -> Linear) 

36 - model.language_model: LlamaForCausalLM 

37 - model.language_model.model.embed_tokens 

38 - model.language_model.model.layers[]: LLaMA transformer blocks 

39 - model.language_model.model.norm 

40 - model.language_model.lm_head 

41 

42 The language model component follows the same patterns as LlamaArchitectureAdapter. 

43 """ 

44 

45 _testing_lm_attr = "model.language_model" 

46 

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

48 """Initialize the LLava architecture adapter.""" 

49 super().__init__(cfg) 

50 

51 self.cfg.is_multimodal = True 

52 

53 # Language model configuration (same as LLaMA) 

54 self._set_rms_rotary_defaults() 

55 self.cfg.attn_implementation = "eager" 

56 

57 # Store vision-related config 

58 self._extract_vision_dims(cfg) 

59 

60 # Weight processing conversions (same as LLaMA - Q/K/V/O rearrangements) 

61 self.weight_processing_conversions = { 

62 **self._qkvo_weight_conversions(), 

63 } 

64 

65 # Select vision encoder bridge based on vision model type 

66 vision_cfg = getattr(cfg, "vision_config", None) 

67 vision_type = getattr(vision_cfg, "model_type", "clip_vision_model") 

68 vision_bridge: GeneralizedComponent 

69 if vision_type in ("siglip_vision_model", "siglip"): 

70 vision_bridge = SiglipVisionEncoderBridge(name="model.vision_tower", config=self.cfg) 

71 else: 

72 vision_bridge = CLIPVisionEncoderBridge(name="model.vision_tower", config=self.cfg) 

73 

74 # Component mapping for the full multimodal model 

75 # LlavaForConditionalGeneration wraps: 

76 # model.vision_tower, model.multi_modal_projector, model.language_model 

77 # The language_model is a *Model (LlamaModel, Qwen2Model, MistralModel) 

78 # with embed_tokens, layers, norm, rotary_emb directly (no nested .model). 

79 # lm_head sits at the top level of LlavaForConditionalGeneration. 

80 self.component_mapping = { 

81 # Vision components 

82 "vision_encoder": vision_bridge, 

83 "vision_projector": VisionProjectionBridge(name="model.multi_modal_projector"), 

84 # Language model components 

85 "embed": EmbeddingBridge(name="model.language_model.embed_tokens"), 

86 "rotary_emb": RotaryEmbeddingBridge(name="model.language_model.rotary_emb"), 

87 "blocks": BlockBridge( 

88 name="model.language_model.layers", 

89 submodules={ 

90 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg), 

91 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg), 

92 "attn": PositionEmbeddingsAttentionBridge( 

93 name="self_attn", 

94 config=self.cfg, 

95 submodules={ 

96 "q": LinearBridge(name="q_proj"), 

97 "k": LinearBridge(name="k_proj"), 

98 "v": LinearBridge(name="v_proj"), 

99 "o": LinearBridge(name="o_proj"), 

100 # The text tower decides: Llama and Qwen2 towers have no 

101 # QK-norm, a Qwen3 one does (NCSOFT/VARCO-VISION-2.0). 

102 "q_norm": RMSNormalizationBridge( 

103 name="q_norm", config=self.cfg, optional=True 

104 ), 

105 "k_norm": RMSNormalizationBridge( 

106 name="k_norm", config=self.cfg, optional=True 

107 ), 

108 }, 

109 requires_attention_mask=True, 

110 requires_position_embeddings=True, 

111 ), 

112 "mlp": self._gated_mlp(), 

113 }, 

114 ), 

115 "ln_final": RMSNormalizationBridge(name="model.language_model.norm", config=self.cfg), 

116 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg), 

117 }