Coverage for transformer_lens/model_bridge/supported_architectures/afmoe.py: 100%
11 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +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; the dense_* projections below make
68 # MoEBridge bind gated-MLP neuron hooks there. The
69 # tuple-returning router stays unwrapped — only its inner
70 # gate Linear is hookable.
71 "mlp": MoEBridge(
72 name="mlp",
73 config=self.cfg,
74 sparse_required=("router_gate",),
75 submodules={
76 "router_gate": LinearBridge(name="router.gate", optional=True),
77 "shared_experts": self._gated_mlp(name="shared_experts", optional=True),
78 # Dense-layer projections (present only on the
79 # dense layers of this interleaved stack); their
80 # presence is what makes MoEBridge bind gated-MLP
81 # neuron hooks there (#1645).
82 "dense_gate": LinearBridge(name="gate_proj", optional=True),
83 "dense_in": LinearBridge(name="up_proj", optional=True),
84 "dense_out": LinearBridge(name="down_proj", optional=True),
85 },
86 ),
87 },
88 ),
89 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg),
90 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg),
91 }