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

18 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-08-11 18:50 +0000

1"""NanoChat (``NanoChatForCausalLM``) adapter: Llama-style decoder with weightless 

2RMSNorm (so nothing for fold_ln to fold), ungated relu^2 MLP, and soft-capped logits; 

3attention delegated (full-width q/k norm after RoPE).""" 

4 

5from typing import Any 

6 

7from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter 

8from transformer_lens.model_bridge.generalized_components import ( 

9 AttentionBridge, 

10 BlockBridge, 

11 EmbeddingBridge, 

12 LinearBridge, 

13 RMSNormalizationBridge, 

14 RotaryEmbeddingBridge, 

15 UnembeddingBridge, 

16) 

17 

18 

19class NanoChatArchitectureAdapter(ArchitectureAdapter): 

20 """Architecture adapter for NanoChatForCausalLM models.""" 

21 

22 # Weightless RMSNorms: no scale to fold into downstream projections. 

23 supports_fold_ln = False 

24 

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

26 """Initialize the NanoChat architecture adapter.""" 

27 super().__init__(cfg) 

28 

29 self.cfg.normalization_type = "RMS" 

30 self.cfg.positional_embedding_type = "rotary" 

31 self.cfg.final_rms = True 

32 self.cfg.gated_mlp = False # ungated relu^2 MLP (fc1 -> act -> fc2) 

33 self.cfg.attn_only = False 

34 self.cfg.uses_rms_norm = True 

35 soft_cap = getattr(cfg, "final_logit_softcapping", None) or getattr( 

36 cfg, "logits_soft_cap", None 

37 ) 

38 if soft_cap: 

39 self.cfg.output_logits_soft_cap = float(soft_cap) 

40 

41 self.weight_processing_conversions = { 

42 **self._qkvo_weight_conversions(), 

43 } 

44 

45 self.component_mapping = { 

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

47 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb"), 

48 "blocks": BlockBridge( 

49 name="model.layers", 

50 config=self.cfg, 

51 submodules={ 

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

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

54 # q/k norms run AFTER rope here; the bridge reimplementation 

55 # applies them before, so delegate attention to HF. 

56 "attn": AttentionBridge( 

57 name="self_attn", 

58 config=self.cfg, 

59 submodules={ 

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

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

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

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

64 "q_norm": RMSNormalizationBridge(name="q_norm", config=self.cfg), 

65 "k_norm": RMSNormalizationBridge(name="k_norm", config=self.cfg), 

66 }, 

67 maintain_native_attention=True, 

68 ), 

69 "mlp": self._ungated_mlp(up="fc1", down="fc2"), 

70 }, 

71 ), 

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

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

74 }