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
« prev ^ index » next coverage.py v7.10.1, created at 2026-08-11 18:50 +0000
1"""JetMoE architecture adapter.
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"""
12from typing import Any
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)
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."""
35 hook_aliases = {
36 "hook_kv": "kv.hook_out",
37 }
40class JetMoeArchitectureAdapter(ArchitectureAdapter):
41 """Architecture adapter for JetMoeForCausalLM models."""
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")
49 def __init__(self, cfg: Any) -> None:
50 """Initialize the JetMoE architecture adapter."""
51 super().__init__(cfg)
53 self._set_rms_rotary_defaults()
55 self.weight_processing_conversions = {}
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 }
96 def setup_component_testing(self, hf_model: Any, bridge_model: Any = None) -> None:
97 """Delegated attention computes rotary inside HF; nothing to wire."""