Coverage for transformer_lens/model_bridge/supported_architectures/glm4_moe_lite.py: 100%
18 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"""GLM-4 MoE Lite architecture adapter.
3Supports the GLM-4.7-Flash family (`Glm4MoeLiteForCausalLM`): DeepSeek-style
4Multi-head Latent Attention (LoRA-compressed Q and KV, nope/rope split heads,
5interleaved partial RoPE) combined with GLM's sparse MoE — sigmoid router with
6e_score_correction_bias, batched routed experts, one shared expert — and a
7per-layer dense/sparse MLP mix declared in ``config.mlp_layer_types``.
8"""
10from typing import Any
12from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
13from transformer_lens.model_bridge.generalized_components import (
14 EmbeddingBridge,
15 LinearBridge,
16 MLAAttentionBridge,
17 MLABlockBridge,
18 MoEBridge,
19 RMSNormalizationBridge,
20 RotaryEmbeddingBridge,
21 UnembeddingBridge,
22)
23from transformer_lens.model_bridge.generalized_components.base import (
24 GeneralizedComponent,
25)
26from transformer_lens.model_bridge.supported_architectures.glm4_moe import (
27 Glm4MoeRouterBridge,
28)
31class Glm4MoeLiteArchitectureAdapter(ArchitectureAdapter):
32 """GLM-4.7-Flash (Glm4MoeLiteForCausalLM) adapter: DeepSeek-V2 MLA + GLM-4-MoE
33 routing (dense/sparse per mlp_layer_types)."""
35 _testing_eager = None
37 def __init__(self, cfg: Any) -> None:
38 super().__init__(cfg)
40 self.cfg.normalization_type = "RMS"
41 self.cfg.positional_embedding_type = "rotary"
42 self.cfg.gated_mlp = True
43 self.cfg.final_rms = True
44 self.cfg.uses_rms_norm = True
45 # Verified against zai-org/GLM-4.7-Flash: tokenizer has no BOS token.
46 self.cfg.default_prepend_bos = False
48 # MLA has no per-head q/k/v to fold into; skip LN folding.
49 self.supports_fold_ln = False
51 # MLA weights keep their HF layout; no QKVO rearrangements apply.
52 self.weight_processing_conversions = {}
54 self.component_mapping = {
55 "embed": EmbeddingBridge(name="model.embed_tokens"),
56 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb", config=self.cfg),
57 "blocks": MLABlockBridge(
58 name="model.layers",
59 submodules={
60 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg),
61 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg),
62 "attn": MLAAttentionBridge(
63 name="self_attn",
64 config=self.cfg,
65 submodules={
66 # Public GLM-4.7 checkpoints set q_lora_rank — two-stage
67 # LoRA Q compression; direct q_proj kept optional for
68 # hypothetical uncompressed variants.
69 "q_a_proj": LinearBridge(name="q_a_proj", optional=True),
70 "q_a_layernorm": GeneralizedComponent(
71 name="q_a_layernorm", optional=True
72 ),
73 "q_b_proj": LinearBridge(name="q_b_proj", optional=True),
74 "q_proj": LinearBridge(name="q_proj", optional=True),
75 "kv_a_proj_with_mqa": LinearBridge(name="kv_a_proj_with_mqa"),
76 "kv_a_layernorm": RMSNormalizationBridge(
77 name="kv_a_layernorm", config=self.cfg
78 ),
79 "kv_b_proj": LinearBridge(name="kv_b_proj"),
80 "o": LinearBridge(name="o_proj"),
81 },
82 ),
83 # Layers marked "dense" in mlp_layer_types hold a plain gated MLP:
84 # router and shared expert absent, so both are optional.
85 "mlp": MoEBridge(
86 name="mlp",
87 config=self.cfg,
88 submodules={
89 "gate": Glm4MoeRouterBridge(name="gate", optional=True),
90 "shared_experts": self._gated_mlp(name="shared_experts", optional=True),
91 },
92 ),
93 },
94 ),
95 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg),
96 "unembed": UnembeddingBridge(name="lm_head"),
97 }