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

1"""GLM-MoE-DSA architecture adapter.""" 

2 

3from typing import Any 

4 

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) 

21 

22 

23class GlmMoeDsaArchitectureAdapter(ArchitectureAdapter): 

24 """Architecture adapter for Z.ai GLM-5 / GLM-5.1 DSA models. 

25 

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

29 

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

31 super().__init__(cfg) 

32 

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 

37 

38 self.weight_processing_conversions = {} 

39 

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 }