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

1"""Architecture adapter for HF's MambaForCausalLM (Mamba-1).""" 

2from typing import Any 

3 

4import torch 

5 

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) 

16 

17 

18class MambaArchitectureAdapter(ArchitectureAdapter): 

19 """Wraps HF's MambaForCausalLM. No attention, no positional embeddings. 

20 

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

25 

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] 

29 

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

31 super().__init__(cfg) 

32 

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 

39 

40 # Routes bridge.generate() through the dedicated SSM cache loop. 

41 self.cfg.is_stateful = True 

42 

43 # No Q/K/V/O weights to rearrange. 

44 self.weight_processing_conversions = {} 

45 

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 } 

68 

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 

79 

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) 

83 

84 return DynamicCache(config=hf_model.config)