Coverage for transformer_lens/model_bridge/supported_architectures/glm_moe_dsa.py: 100%
14 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
1"""GLM-MoE-DSA architecture adapter."""
3from typing import Any
5from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
6from transformer_lens.model_bridge.generalized_components import (
7 EmbeddingBridge,
8 LinearBridge,
9 MLABlockBridge,
10 MoEBridge,
11 RMSNormalizationBridge,
12 RotaryEmbeddingBridge,
13 UnembeddingBridge,
14)
15from transformer_lens.model_bridge.generalized_components.base import (
16 GeneralizedComponent,
17)
18from transformer_lens.model_bridge.generalized_components.glm_moe_dsa_attention import (
19 GlmMoeDsaAttentionBridge,
20)
23class GlmMoeDsaArchitectureAdapter(ArchitectureAdapter):
24 """Architecture adapter for Z.ai GLM-5 / GLM-5.1 DSA models.
26 GLM-MoE-DSA combines MLA-style latent attention, a learned sparse-attention
27 indexer, dense early MLP layers, and sparse MoE later layers.
28 """
30 def __init__(self, cfg: Any) -> None:
31 super().__init__(cfg)
33 self.supports_fold_ln = False
34 self._set_rms_rotary_defaults()
35 self.cfg.attn_implementation = "eager"
36 self.cfg.default_prepend_bos = False
38 self.weight_processing_conversions = {}
40 self.component_mapping = {
41 "embed": EmbeddingBridge(name="model.embed_tokens"),
42 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb", config=self.cfg),
43 "blocks": MLABlockBridge(
44 name="model.layers",
45 submodules={
46 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg),
47 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg),
48 "attn": GlmMoeDsaAttentionBridge(
49 name="self_attn",
50 config=self.cfg,
51 submodules={
52 "q_a_proj": LinearBridge(name="q_a_proj"),
53 "q_a_layernorm": RMSNormalizationBridge(
54 name="q_a_layernorm", config=self.cfg
55 ),
56 "q_b_proj": LinearBridge(name="q_b_proj"),
57 "kv_a_proj_with_mqa": LinearBridge(name="kv_a_proj_with_mqa"),
58 "kv_a_layernorm": RMSNormalizationBridge(
59 name="kv_a_layernorm", config=self.cfg
60 ),
61 "kv_b_proj": LinearBridge(name="kv_b_proj"),
62 "o": LinearBridge(name="o_proj"),
63 },
64 ),
65 "mlp": MoEBridge(
66 name="mlp",
67 config=self.cfg,
68 sparse_required=("gate",),
69 submodules={
70 "gate": GeneralizedComponent(name="gate", optional=True),
71 "shared_experts": self._gated_mlp(name="shared_experts", optional=True),
72 # Dense-layer projections (present only on the
73 # dense layers of this interleaved stack); their
74 # presence is what makes MoEBridge bind gated-MLP
75 # neuron hooks there (#1645).
76 "dense_gate": LinearBridge(name="gate_proj", optional=True),
77 "dense_in": LinearBridge(name="up_proj", optional=True),
78 "dense_out": LinearBridge(name="down_proj", optional=True),
79 },
80 ),
81 },
82 ),
83 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg),
84 "unembed": UnembeddingBridge(name="lm_head"),
85 }