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
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
1"""LLava architecture adapter.
3This adapter supports LlavaForConditionalGeneration, the vision-language
4model combining a CLIP vision encoder with a LLaMA language model.
5"""
7from typing import Any
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)
29class LlavaArchitectureAdapter(ArchitectureAdapter):
30 """Architecture adapter for LLava multimodal models (LlavaForConditionalGeneration).
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
42 The language model component follows the same patterns as LlamaArchitectureAdapter.
43 """
45 _testing_lm_attr = "model.language_model"
47 def __init__(self, cfg: Any) -> None:
48 """Initialize the LLava architecture adapter."""
49 super().__init__(cfg)
51 self.cfg.is_multimodal = True
53 # Language model configuration (same as LLaMA)
54 self._set_rms_rotary_defaults()
55 self.cfg.attn_implementation = "eager"
57 # Store vision-related config
58 self._extract_vision_dims(cfg)
60 # Weight processing conversions (same as LLaMA - Q/K/V/O rearrangements)
61 self.weight_processing_conversions = {
62 **self._qkvo_weight_conversions(),
63 }
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)
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 }