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

13 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-21 19:27 +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._set_rms_rotary_defaults(gated=False) 

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

31 cfg, "logits_soft_cap", None 

32 ) 

33 if soft_cap: 

34 self.cfg.output_logits_soft_cap = float(soft_cap) 

35 

36 self.weight_processing_conversions = { 

37 **self._qkvo_weight_conversions(), 

38 } 

39 

40 self.component_mapping = { 

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

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

43 "blocks": BlockBridge( 

44 name="model.layers", 

45 config=self.cfg, 

46 submodules={ 

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

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

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

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

51 "attn": AttentionBridge( 

52 name="self_attn", 

53 config=self.cfg, 

54 submodules={ 

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

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

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

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

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

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

61 }, 

62 maintain_native_attention=True, 

63 ), 

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

65 }, 

66 ), 

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

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

69 }