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
« 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)."""
5from typing import Any
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)
19class NanoChatArchitectureAdapter(ArchitectureAdapter):
20 """Architecture adapter for NanoChatForCausalLM models."""
22 # Weightless RMSNorms: no scale to fold into downstream projections.
23 supports_fold_ln = False
25 def __init__(self, cfg: Any) -> None:
26 """Initialize the NanoChat architecture adapter."""
27 super().__init__(cfg)
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)
41 self.weight_processing_conversions = {
42 **self._qkvo_weight_conversions(),
43 }
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 }