Coverage for transformer_lens/model_bridge/supported_architectures/afmoe.py: 100%
11 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"""AFMoE (Arcee Trinity, ``AfmoeForCausalLM``) adapter: sandwich norms, QK-norm
2attention with sigmoid gating and NoPE/sliding RoPE (so attention delegates to HF),
3dense + sparse-MoE MLP layers split at ``num_dense_layers``."""
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 MoEBridge,
14 RMSNormalizationBridge,
15 UnembeddingBridge,
16)
19class AfmoeArchitectureAdapter(ArchitectureAdapter):
20 """Architecture adapter for AfmoeForCausalLM models."""
22 # Sandwich norms scale sublayer outputs before the residual add; folding
23 # ln1/ln2 into the projections changes the function (Trinity-Nano compat
24 # mode diverged to loss 10.9 vs 2.3 before this was disabled).
25 supports_fold_ln = False
27 def __init__(self, cfg: Any) -> None:
28 """Initialize the AFMoE architecture adapter."""
29 super().__init__(cfg)
31 self._set_rms_rotary_defaults()
32 self.cfg.attn_implementation = "eager"
34 self.weight_processing_conversions = {
35 **self._qkvo_weight_conversions(),
36 }
38 self.component_mapping = {
39 "embed": EmbeddingBridge(name="model.embed_tokens"),
40 "blocks": BlockBridge(
41 name="model.layers",
42 submodules={
43 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg),
44 "ln1_post": RMSNormalizationBridge(
45 name="post_attention_layernorm", config=self.cfg
46 ),
47 "ln2": RMSNormalizationBridge(name="pre_mlp_layernorm", config=self.cfg),
48 "ln2_post": RMSNormalizationBridge(name="post_mlp_layernorm", config=self.cfg),
49 # Per-head QK-norm before RoPE, RoPE only on sliding
50 # layers, and sigmoid output gating live in HF's forward.
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 "gate": LinearBridge(name="gate_proj"),
60 "q_norm": RMSNormalizationBridge(name="q_norm", config=self.cfg),
61 "k_norm": RMSNormalizationBridge(name="k_norm", config=self.cfg),
62 },
63 maintain_native_attention=True,
64 requires_attention_mask=True,
65 ),
66 # Dense layers (< num_dense_layers) hold a plain gated MLP
67 # under the same name; router and shared experts are
68 # optional. The tuple-returning router stays unwrapped —
69 # only its inner gate Linear is hookable.
70 "mlp": MoEBridge(
71 name="mlp",
72 config=self.cfg,
73 submodules={
74 "router_gate": LinearBridge(name="router.gate", optional=True),
75 "shared_experts": self._gated_mlp(name="shared_experts", optional=True),
76 },
77 ),
78 },
79 ),
80 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg),
81 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg),
82 }