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

15 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-08-11 18:50 +0000

1"""JetMoE architecture adapter. 

2 

3MIT-IBM's JetMoE (``JetMoeForCausalLM``, native in transformers): the only 

4open at-scale Mixture-of-Attention-heads model — attention Q and output 

5projections are per-expert parallel 3D tensors behind a top-k router 

6(``experts``: JetMoeMoA), with a shared fused KV projection, alongside a 

7conventional parallel-experts MoE MLP. Both routers are hookable; the 

8mixers delegate to HF (per-expert 3D projections have no uniform 

9reconstruction, so no fold target exists either). 

10""" 

11 

12from typing import Any 

13 

14from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter 

15from transformer_lens.model_bridge.generalized_components import ( 

16 AttentionBridge, 

17 BlockBridge, 

18 EmbeddingBridge, 

19 LinearBridge, 

20 MoEBridge, 

21 MoERouterBridge, 

22 RMSNormalizationBridge, 

23 RotaryEmbeddingBridge, 

24 UnembeddingBridge, 

25) 

26from transformer_lens.model_bridge.generalized_components.base import ( 

27 GeneralizedComponent, 

28) 

29 

30 

31class _JetMoeAttentionBridge(AttentionBridge): 

32 """Mixture-of-Attention: no separate q/k/v/o Linears to alias — Q and O 

33 live inside the per-expert MoA; only the shared fused KV is a Linear.""" 

34 

35 hook_aliases = { 

36 "hook_kv": "kv.hook_out", 

37 } 

38 

39 

40class JetMoeArchitectureAdapter(ArchitectureAdapter): 

41 """Architecture adapter for JetMoeForCausalLM models.""" 

42 

43 # Per-expert 3D Q/O projections: nothing to fold a norm into. 

44 supports_fold_ln = False 

45 # TopKGating's forward sorts/scatters expert assignments and crashes on 

46 # the harness's isolated probes; routers stay hookable at runtime. 

47 component_test_skip_suffixes = ("mlp.gate", "attn.experts.router") 

48 

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

50 """Initialize the JetMoE architecture adapter.""" 

51 super().__init__(cfg) 

52 

53 self._set_rms_rotary_defaults() 

54 

55 self.weight_processing_conversions = {} 

56 

57 self.component_mapping = { 

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

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

60 "blocks": BlockBridge( 

61 name="model.layers", 

62 config=self.cfg, 

63 submodules={ 

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

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

66 # MoA: delegate; the attention router and shared KV are 

67 # hookable, per-expert Q/O stay inside the delegated MoA. 

68 "attn": _JetMoeAttentionBridge( 

69 name="self_attention", 

70 config=self.cfg, 

71 submodules={ 

72 "kv": LinearBridge(name="kv_proj"), 

73 "experts": GeneralizedComponent( 

74 name="experts", 

75 submodules={ 

76 # JetMoeTopKGating puts logits last in its 5-tuple. 

77 "router": MoERouterBridge(name="router", logits_index=-1), 

78 }, 

79 ), 

80 }, 

81 maintain_native_attention=True, 

82 ), 

83 "mlp": MoEBridge( 

84 name="mlp", 

85 config=self.cfg, 

86 submodules={ 

87 "gate": MoERouterBridge(name="router", logits_index=-1), 

88 }, 

89 ), 

90 }, 

91 ), 

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

93 "unembed": UnembeddingBridge(name="lm_head"), 

94 } 

95 

96 def setup_component_testing(self, hf_model: Any, bridge_model: Any = None) -> None: 

97 """Delegated attention computes rotary inside HF; nothing to wire."""