Coverage for transformer_lens/model_bridge/supported_architectures/llama4_multimodal.py: 100%
13 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"""Llama 4 multimodal architecture adapter.
3Meta's Llama 4 composite (``Llama4ForConditionalGeneration``): a Llama4
4vision transformer at ``vision_model`` and a projector feeding the full
5Llama4ForCausalLM at ``language_model`` (text stack at
6``language_model.model.*`` with ``language_model.lm_head``). The vision
7side is delegated opaquely; the text mapping is the Llama4 text adapter's,
8re-prefixed.
9"""
11from typing import Any
13from transformer_lens.model_bridge.generalized_components.base import (
14 GeneralizedComponent,
15)
16from transformer_lens.model_bridge.supported_architectures.llama4 import (
17 Llama4ArchitectureAdapter,
18)
21class Llama4MultimodalArchitectureAdapter(Llama4ArchitectureAdapter):
22 """Architecture adapter for Llama4ForConditionalGeneration models."""
24 def __init__(self, cfg: Any) -> None:
25 """Initialize the Llama 4 multimodal architecture adapter."""
26 super().__init__(cfg)
28 self.cfg.is_multimodal = True
29 if hasattr(cfg, "vision_config"):
30 self.cfg.vision_hidden_size = getattr(cfg.vision_config, "hidden_size", None)
32 self._reprefix_components("model.", "language_model.model.")
33 self._reprefix_components("lm_head", "language_model.lm_head")
35 self.components["vision_encoder"] = GeneralizedComponent(name="vision_model")
36 self.components["vision_projector"] = GeneralizedComponent(name="multi_modal_projector")