Coverage for transformer_lens/model_bridge/supported_architectures/mamba.py: 92%
24 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"""Architecture adapter for HF's MambaForCausalLM (Mamba-1)."""
2from typing import Any
4import torch
6from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
7from transformer_lens.model_bridge.generalized_components import (
8 DepthwiseConv1DBridge,
9 EmbeddingBridge,
10 LinearBridge,
11 RMSNormalizationBridge,
12 SSMBlockBridge,
13 SSMMixerBridge,
14 UnembeddingBridge,
15)
18class MambaArchitectureAdapter(ArchitectureAdapter):
19 """Wraps HF's MambaForCausalLM. No attention, no positional embeddings.
21 SSM config fields (state_size, conv_kernel, expand, time_step_rank,
22 intermediate_size) are propagated from the HF config via
23 ``_HF_PASSTHROUGH_ATTRS`` in sources/transformers.py.
24 """
26 # White-box forward: P1 is exact vs raw HF (mixer delegates to HF); P2/P3 skip
27 # without a HookedTransformer; P4 is generation.
28 applicable_phases: list[int] = [1, 2, 3, 4]
30 def __init__(self, cfg: Any) -> None:
31 super().__init__(cfg)
33 self.cfg.normalization_type = "RMS"
34 self.cfg.uses_rms_norm = True
35 self.cfg.positional_embedding_type = "none"
36 self.cfg.gated_mlp = False
37 self.cfg.attn_only = False
38 self.cfg.final_rms = True
40 # Routes bridge.generate() through the dedicated SSM cache loop.
41 self.cfg.is_stateful = True
43 # No Q/K/V/O weights to rearrange.
44 self.weight_processing_conversions = {}
46 self.component_mapping = {
47 "embed": EmbeddingBridge(name="backbone.embeddings"),
48 "blocks": SSMBlockBridge(
49 name="backbone.layers",
50 submodules={
51 "norm": RMSNormalizationBridge(name="norm", config=self.cfg),
52 "mixer": SSMMixerBridge(
53 name="mixer",
54 config=self.cfg,
55 submodules={
56 "in_proj": LinearBridge(name="in_proj"),
57 "conv1d": DepthwiseConv1DBridge(name="conv1d"),
58 "x_proj": LinearBridge(name="x_proj"),
59 "dt_proj": LinearBridge(name="dt_proj"),
60 "out_proj": LinearBridge(name="out_proj"),
61 },
62 ),
63 },
64 ),
65 "ln_final": RMSNormalizationBridge(name="backbone.norm_f", config=self.cfg),
66 "unembed": UnembeddingBridge(name="lm_head"),
67 }
69 def create_stateful_cache(
70 self,
71 hf_model: Any,
72 batch_size: int,
73 device: Any,
74 dtype: torch.dtype,
75 ) -> Any:
76 """Build a cache for the stateful generation loop."""
77 from transformers.cache_utils import DynamicCache
78 from transformers.models.mamba import modeling_mamba
80 cache_cls = getattr(modeling_mamba, "MambaCache", None)
81 if cache_cls is not None: 81 ↛ 82line 81 didn't jump to line 82 because the condition on line 81 was never true
82 return cache_cls(hf_model.config, batch_size, device=device, dtype=dtype)
84 return DynamicCache(config=hf_model.config)