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

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``.""" 

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 MoEBridge, 

14 RMSNormalizationBridge, 

15 UnembeddingBridge, 

16) 

17 

18 

19class AfmoeArchitectureAdapter(ArchitectureAdapter): 

20 """Architecture adapter for AfmoeForCausalLM models.""" 

21 

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 

26 

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

28 """Initialize the AFMoE architecture adapter.""" 

29 super().__init__(cfg) 

30 

31 self._set_rms_rotary_defaults() 

32 self.cfg.attn_implementation = "eager" 

33 

34 self.weight_processing_conversions = { 

35 **self._qkvo_weight_conversions(), 

36 } 

37 

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 }