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

1"""ERNIE 4.5 MoE architecture adapter. 

2 

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

9 

10from typing import Any 

11 

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) 

26 

27 

28class Ernie4_5_MoeArchitectureAdapter(ArchitectureAdapter): 

29 """Architecture adapter for Ernie4_5_MoeForCausalLM models.""" 

30 

31 _testing_eager = "config" 

32 

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

34 """Initialize the ERNIE 4.5 MoE architecture adapter.""" 

35 super().__init__(cfg) 

36 

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 

41 

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 } 

48 

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 }