Coverage for transformer_lens/model_bridge/supported_architectures/ernie4_5_moe.py: 100%
13 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"""ERNIE 4.5 MoE architecture adapter.
3Baidu's ERNIE 4.5 MoE (``Ernie4_5_MoeForCausalLM``): the dense ERNIE
4attention (GLM-style interleaved RoPE, config-gated biases) with a sparse
5MoE MLP — sigmoid-corrected top-k router, batched fused gate_up experts,
6optional shared experts, and a dense-MLP prefix before
7``moe_layer_start_index``.
8"""
10from typing import Any
12from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
13from transformer_lens.model_bridge.generalized_components import (
14 BlockBridge,
15 EmbeddingBridge,
16 LinearBridge,
17 MoEBridge,
18 PositionEmbeddingsAttentionBridge,
19 RMSNormalizationBridge,
20 RotaryEmbeddingBridge,
21 UnembeddingBridge,
22)
23from transformer_lens.model_bridge.generalized_components.base import (
24 GeneralizedComponent,
25)
28class Ernie4_5_MoeArchitectureAdapter(ArchitectureAdapter):
29 """Architecture adapter for Ernie4_5_MoeForCausalLM models."""
31 _testing_eager = "config"
33 def __init__(self, cfg: Any) -> None:
34 """Initialize the ERNIE 4.5 MoE architecture adapter."""
35 super().__init__(cfg)
37 self._set_rms_rotary_defaults()
38 # Same conventions as dense ERNIE 4.5.
39 self.cfg.rotary_adjacent_pairs = True
40 self.cfg.default_prepend_bos = False
42 # Biases are config-gated (use_bias); reshape them so a use_bias=True
43 # GQA checkpoint gets the (n_kv, d_head) K/V bias layout. No-op when
44 # the checkpoint carries no attention biases.
45 self.weight_processing_conversions = {
46 **self._qkvo_weight_conversions(include_biases=True),
47 }
49 self.component_mapping = {
50 "embed": EmbeddingBridge(name="model.embed_tokens"),
51 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb", config=self.cfg),
52 "blocks": BlockBridge(
53 name="model.layers",
54 submodules={
55 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg),
56 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg),
57 "attn": PositionEmbeddingsAttentionBridge(
58 name="self_attn",
59 config=self.cfg,
60 submodules={
61 "q": LinearBridge(name="q_proj"),
62 "k": LinearBridge(name="k_proj"),
63 "v": LinearBridge(name="v_proj"),
64 "o": LinearBridge(name="o_proj"),
65 },
66 requires_attention_mask=True,
67 requires_position_embeddings=True,
68 ),
69 # Layers before moe_layer_start_index hold a plain gated
70 # MLP; router and shared experts are absent there.
71 "mlp": MoEBridge(
72 name="mlp",
73 config=self.cfg,
74 submodules={
75 # Raw-Parameter router; tuple-safe hook via base.
76 "gate": GeneralizedComponent(name="gate", optional=True),
77 "shared_experts": self._gated_mlp(name="shared_experts", optional=True),
78 },
79 ),
80 },
81 ),
82 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg),
83 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg),
84 }