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

1"""GLM-4 MoE Lite architecture adapter. 

2 

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

9 

10from typing import Any 

11 

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) 

29 

30 

31class Glm4MoeLiteArchitectureAdapter(ArchitectureAdapter): 

32 """GLM-4.7-Flash (Glm4MoeLiteForCausalLM) adapter: DeepSeek-V2 MLA + GLM-4-MoE 

33 routing (dense/sparse per mlp_layer_types).""" 

34 

35 _testing_eager = None 

36 

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

38 super().__init__(cfg) 

39 

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 

47 

48 # MLA has no per-head q/k/v to fold into; skip LN folding. 

49 self.supports_fold_ln = False 

50 

51 # MLA weights keep their HF layout; no QKVO rearrangements apply. 

52 self.weight_processing_conversions = {} 

53 

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 }