已合并
Feat: HyperParallel Trainer新增GLM5系列模型 #2099 #885
Feat: HyperParallel Trainer新增GLM5系列模型 #2099 #885
已合并
Moy创建于 6月22日
共 16 个文件变更+3422-8
@@ -0,0 +1,145 @@
1+# GLM5 training example
2+ 
3+This directory provides the GLM5 Trainer entry point for `scripts/train_lm.py`.
4+It follows the same example layout as the existing Qwen examples: one user-facing
5+`train.yaml` for GLM5 training.
6+ 
7+## Scope
8+ 
9+The GLM5 Trainer path supports:
10+ 
11+- dense GLM5 causal LM forward/loss/backward training;
12+- checkpoint save, load, and resume with model, optimizer, scheduler, RNG, and
13+ dataloader state;
14+- MLA attention and 2D/4D attention mask handling;
15+- MoE construction and EP=2 expert dispatch;
16+- DSA sparse-attention construction and CP=2 context-parallel training.
17+ 
18+Tensor parallelism is not implemented for GLM5. Setting
19+`train.accelerator.tp > 1` raises `NotImplementedError`.
20+ 
21+Cached decode supports append-only autoregressive decoding with contiguous
22+history positions. Packed sequences, non-contiguous cached positions, cache
23+reuse, and prefix stitching are not part of this training example.
24+ 
25+## Configuration
26+ 
27+`train.yaml` is the only training configuration committed under
28+`examples/glm5`. It is intended as the stable user-facing entry point.
29+Validation scenarios such as MoE, EP, CP, DSA, scaled architecture alignment,
30+and cross-framework comparison are reproduced with command-line overrides and
31+external validation materials, rather than additional YAML files under
32+`examples`.
33+ 
34+## Unit Tests
35+ 
36+GLM5 Trainer tests live under `tests/torch/trainer/test_glm5_trainer.py`.
37+Trainer callback configuration regression tests live under `tests/ut/trainer`.
38+Together they cover model discovery, GLM5 batch preparation through
39+`BaseTrainer`, shifted CausalLM loss semantics, CP batch sharding, checkpoint
40+callback round-trip, parallelization guards, and nested train logging/checkpoint
41+configuration.
42+ 
43+```bash
44+python -m pytest \
45+ tests/torch/trainer/test_glm5_trainer.py \
46+ tests/ut/trainer/test_checkpoint_callback_config.py \
47+ tests/ut/trainer/test_logging_callback_config.py -q
48+```
49+ 
50+Validated result:
51+ 
52+```text
53+13 passed
54+```
55+ 
56+## Cross-Framework Validation
57+ 
58+The full official GLM5 checkpoint is too large for the minimum validation
59+environment, so cross-framework checks use scaled GLM5 dense and MoE+MLA+DSA
60+variants. The variants keep GLM5 module signatures and the GLM-5 tokenizer
61+vocabulary size while reducing layer width, layer count, and expert count.
62+ 
63+The validation compares the same fixed batch and same exported weights. The
64+Transformers side uses the official `GlmMoeDsaConfig` and
65+`GlmMoeDsaForCausalLM` classes; it does not import
66+`hyper_parallel.models.glm5`.
67+ 
68+- Transformers official `GlmMoeDsaForCausalLM` vs HyperParallel GLM5;
69+- LLaMAFactory `CustomSeq2SeqTrainer.compute_loss` vs HyperParallel GLM5.
70+ 
71+These checks validate single-card loss/logits semantics for the GLM5 model
72+structure paths covered by Dense, MoE, MLA, and the official DSA indexer. CP and
73+EP are HyperParallel parallel-training strategies and are validated by
74+in-repository tests and separate ST validation materials instead of
75+external-framework comparison.
76+ 
77+## Verified Result
78+ 
79+The following results were collected on Ascend NPU with GLM5 validation commands
80+based on `train.yaml` and fixed validation materials.
81+ 
82+```text
83+Scaled dense GLM5 fp32 1c vs DP2, 100 steps:
84+common_steps: 100
85+avg_diff: 6.55615000e-05
86+max_diff: 2.77690000e-04
87+pass_avg_5e-3: True
88+ 
89+Scaled MoE+MLA+DSA GLM5 fp32 1c vs DP2, 100 steps:
90+common_steps: 100
91+avg_diff: 6.64436000e-05
92+max_diff: 2.55470000e-04
93+pass_avg_5e-3: True
94+ 
95+MoE preset fp32 1c vs EP2, 100 steps:
96+avg_diff: 4.83024000e-05
97+max_diff: 2.59350000e-04
98+pass_avg_5e-3: True
99+ 
100+DSA preset fp32 1c vs CP2, 100 steps:
101+avg_diff: 3.80690000e-05
102+max_diff: 2.98080000e-04
103+pass_avg_5e-3: True
104+ 
105+Transformers scaled dense:
106+hf_loss: 12.262624740600586
107+hyper_loss: 12.262624740600586
108+logits_max_abs_diff: 0.0
109+loss_abs_diff: 0.0
110+pass_logits: True
111+pass_loss: True
112+ 
113+Transformers scaled MoE+MLA+DSA:
114+hf_loss: 12.128127098083496
115+hyper_loss: 12.128125190734863
116+logits_max_abs_diff: 0.0
117+loss_abs_diff: 1.9073486328125e-06
118+pass_logits: True
119+pass_loss: True
120+ 
121+LLaMAFactory scaled dense:
122+llamafactory_loss: 12.262624740600586
123+hyper_loss: 12.262624740600586
124+logits_max_abs_diff: 0.0
125+loss_abs_diff: 0.0
126+pass_logits: True
127+pass_loss: True
128+ 
129+LLaMAFactory scaled MoE+MLA+DSA:
130+llamafactory_loss: 12.128129005432129
131+hyper_loss: 12.128125190734863
132+logits_max_abs_diff: 0.0
133+loss_abs_diff: 3.814697265625e-06
134+pass_logits: True
135+pass_loss: True
136+ 
137+Cross-framework state_dict diagnostics:
138+missing_hp_keys: []
139+unexpected_hf_keys: []
140+shape_mismatch_keys: {}
141+```
142+ 
143+The fp32 settings are validation-only overrides. They remove low-precision
144+rounding from strict alignment checks and do not change the normal mixed
145+precision training path exposed by `train.yaml`.
@@ -0,0 +1,69 @@
1+model:
2+ name: glm5
3+ weights_path: null
4+ tokenizer_path: null
5+ config_overrides:
6+ vocab_size: 154856
7+ hidden_size: 1024
8+ intermediate_size: 3072
9+ num_hidden_layers: 4
10+ num_dense_layers: 4
11+ num_attention_heads: 16
12+ num_key_value_heads: 4
13+ head_dim: 64
14+ max_position_embeddings: 131072
15+ 
16+data:
17+ type: preset_pt
18+ train_path: /path/to/preset_batches.pt
19+ max_seq_len: 64
20+ shuffle: false
21+ num_workers: 0
22+ pin_memory: true
23+ 
24+train:
25+ max_steps: 100
26+ num_train_epochs: 1
27+ global_batch_size: 4
28+ micro_batch_size: 1
29+ seed: 1234
30+ backend: torch
31+ init_device: meta
32+ 
33+ accelerator:
34+ dp_shard: 1
35+ tp: 1
36+ cp: 1
37+ ep: 1
38+ comm_fusion: true
39+ 
40+ optimizer:
41+ type: adamw
42+ lr: 1.0e-4
43+ lr_min: 0.0
44+ lr_decay_style: cosine
45+ lr_warmup_ratio: 0.1
46+ max_grad_norm: 1.0e9
47+ weight_decay: 0.0
48+ loss_aggregation: rank_average
49+ 
50+ mixed_precision:
51+ enabled: true
52+ param_dtype: bfloat16
53+ reduce_dtype: float32
54+ output_dtype: bfloat16
55+ 
56+ gradient_checkpointing:
57+ activation_checkpoint: full
58+ 
59+ checkpoint:
60+ output_dir: outputs/glm5
61+ save_steps: 50
62+ save_hf_weights: false
63+ 
64+ logging:
65+ log_steps: 1
66+ report_throughput: false
67+ 
68+ debug:
69+ deterministic: true
@@ -0,0 +1,74 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
atomgit-bot
atomgit-botatomgit-bot6月22日

🔵 Low Priority

对以下所有变更文件进行了全面审查:

  • README.md:无危险命令、无硬编码密钥、无不安全的配置指导。
  • model_spec.py:新增的 prepare_batch_fn 字段具有默认值 None,完全向后兼容。
  • base.py:token 计数守卫的重构修复了一个小问题(旧代码在分布式情况下将 local_tokens 填充为 1,这会稍微抬高 all-reduce 的 token 计数)——更改是正确的。
  • callbacks/base.py:LoggingCallback、CheckpointCallback 和 SafetensorsExportCallback 的嵌套配置读取(train.checkpoint、train.logging)会优雅地回退到旧版顶层键——向后兼容。
  • 3 个测试文件(test_glm5_train.py、test_checkpoint_callback_config.py、test_logging_callback_config.py):集成测试覆盖了前向/后向、缓存解码、左填充、MoE/MLA/DSA/MTP 模块、CP 分片、EP 应用、推测解码、checkpoint 路径验证和确定性初始化;单元测试验证了嵌套和回退配置读取。

建议:无需操作。此审查结果确认未发现正确性、安全性、可靠性或破坏性变更问题。

likedislike
不准确?
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+"""GLM5 model registration."""
16+from hyper_parallel.models.glm5.model import (
17+ GLM5Config,
18+ GLM5Decoder,
19+ GLM5ForCausalLM,
20+ prepare_glm5_batch,
21+)
22+from hyper_parallel.models.glm5.parallelize import parallelize_glm5
23+from hyper_parallel.models.glm5.state_dict import GLM5StateDictAdapter
24+from hyper_parallel.models.spec import ModelSpec, register_spec
25+ 
26+_UNIVERSAL_FIELDS = (
27+ "vocab_size",
28+ "hidden_size",
29+ "intermediate_size",
30+ "num_hidden_layers",
31+ "num_attention_heads",
32+ "num_key_value_heads",
33+ "max_position_embeddings",
34+)
35+ 
36+ 
37+def _resolve_overrides(model_cfg) -> dict:
38+ """Collect GLM5 config overrides from model args."""
39+ overrides = {}
40+ if model_cfg is None:
41+ return overrides
42+ for field in _UNIVERSAL_FIELDS:
43+ value = getattr(model_cfg, field, None)
44+ if value is not None:
45+ overrides[field] = value
46+ extra = getattr(model_cfg, "config_overrides", None)
47+ if isinstance(extra, dict):
48+ overrides.update(extra)
49+ return overrides
50+ 
51+ 
52+def _build(cfg) -> GLM5ForCausalLM:
53+ overrides = _resolve_overrides(getattr(cfg, "model", None))
54+ config = GLM5Config(**overrides) if overrides else GLM5Config()
55+ return GLM5ForCausalLM(config)
56+ 
57+ 
58+register_spec(
59+ "glm5",
60+ ModelSpec(
61+ name="glm5",
62+ build_model_fn=_build,
63+ parallelize_fn=parallelize_glm5,
64+ state_dict_adapter=GLM5StateDictAdapter,
65+ prepare_batch_fn=prepare_glm5_batch,
66+ ),
67+)
68+ 
69+__all__ = [
70+ "GLM5Config",
71+ "GLM5Decoder",
72+ "GLM5ForCausalLM",
73+ "prepare_glm5_batch",
74+]
@@ -0,0 +1,280 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+"""HuggingFace GLM5/GLM4 safetensors to hyper state-dict conversion."""
16+from __future__ import annotations
17+ 
18+import json
19+import logging
20+import os
21+import re
22+from typing import Dict, Optional, Tuple
23+ 
24+import torch
25+from safetensors import safe_open
26+ 
27+logger = logging.getLogger(__name__)
28+ 
29+_PER_EXPERT_RE = re.compile(
30+ r"^model\.layers\.(?P<layer>\d+)\.mlp\.experts\.(?P<expert>\d+)\."
31+ r"(?P<kind>gate_proj|up_proj|down_proj)\.weight$"
32+)
33+ 
34+_SUPPORTED_LAYER_SUFFIXES = (
35+ "input_layernorm.weight",
36+ "post_attention_layernorm.weight",
37+ "self_attn.q_proj.weight",
38+ "self_attn.q_proj.bias",
39+ "self_attn.k_proj.weight",
40+ "self_attn.k_proj.bias",
41+ "self_attn.v_proj.weight",
42+ "self_attn.v_proj.bias",
43+ "self_attn.o_proj.weight",
44+ "self_attn.o_proj.bias",
45+ "self_attn.kv_lora_a_proj.weight",
46+ "self_attn.kv_lora_a_proj.bias",
47+ "self_attn.kv_lora_norm.weight",
48+ "self_attn.kv_lora_b_proj.weight",
49+ "self_attn.kv_lora_b_proj.bias",
50+ "dsa_indexer.query_proj.weight",
51+ "dsa_indexer.key_proj.weight",
52+ "mlp.gate_proj.weight",
53+ "mlp.up_proj.weight",
54+ "mlp.down_proj.weight",
55+ "mlp.gate.weight",
56+ "mlp.experts.gate_up_proj",
57+ "mlp.experts.down_proj",
58+)
59+ 
60+ 
61+def _resolve_shard_path(weights_path: str, shard: str) -> str:
62+ """Return a shard path constrained to the checkpoint directory."""
63+ if os.path.isabs(shard):
64+ raise ValueError(f"GLM5 checkpoint shard must be relative: {shard}")
65+ base_path = os.path.realpath(weights_path)
66+ shard_path = os.path.realpath(os.path.join(base_path, shard))
67+ if os.path.commonpath([base_path, shard_path]) != base_path:
68+ raise ValueError(f"GLM5 checkpoint shard escapes weights_path: {shard}")
69+ return shard_path
70+ 
71+ 
72+def _remap_key(hf_key: str, max_layer: int) -> Optional[str]:
73+ """Map a GLM-family HF key to the hyper GLM5 dense module layout."""
74+ if hf_key == "lm_head.weight":
75+ return "lm_head.weight"
76+ standard_prefix = "model."
77+ if hf_key.startswith(standard_prefix):
78+ tail = hf_key[len(standard_prefix):]
79+ if tail.startswith("layers."):
80+ try:
81+ layer_i = int(tail.split(".")[1])
82+ except (IndexError, ValueError):
83+ return None
84+ if layer_i > max_layer:
85+ return None
86+ if tail.startswith(("embed_tokens.", "layers.", "norm.")):
X
Xxuxinglei6月25日

_remap_key 只透传 / 映射 dense 命名(lm_head / embed_tokens / layers.* / norm.*),对 GLM5 真实存在的以下权重完全不处理:

  • MoE packed experts gate_up_proj / down_proj(HF 通常是 per-expert,需按 num_experts 拆 / 拼)
  • MLA kv_lora_a/b_proj(与 HF q_a / kv_a / kv_b 命名不同)
  • DSA indexer query_proj / key_proj、MTP head

这些 HF key 经此处原样透传后与 hyper 参数名对不上,落入 base.py 的 load_state_dict(strict=False) → 仅 logger.warning("randomly initialised") 不报错。净效果:用真实 GLM5(MoE / MLA / MTP)权重训练时,这些子模块静默停留在随机初始化,训练"看起来正常跑通"但精度静默错误。对照 qwen3_5_moe/checkpoint.py 是有完整 packed-expert 拆分的。当前所有 yaml weights_path=null 掩盖了这个缺口。

建议:补齐各变体 HF→hyper 映射 + expert 打包;至少对 allowlist 之外的结构性 missing 改为 raise(而非 warning);并补一个真正调用 load_hf_glm5_state_dict 的 mini HF safetensors 测试,断言 strict=True 无 missing / unexpected。

likedislike
Moy
Moy
6月26日 评论:
87+ return hf_key
88+ 
89+ legacy_map = {
90+ "transformer.embedding.word_embeddings.": "model.embed_tokens.",
91+ "transformer.encoder.layers.": "model.layers.",
92+ "transformer.encoder.final_layernorm.": "model.norm.",
93+ "transformer.output_layer.": "lm_head.",
94+ }
95+ for old_prefix, new_prefix in legacy_map.items():
96+ if hf_key.startswith(old_prefix):
97+ tail = hf_key[len(old_prefix):]
98+ mapped = f"{new_prefix}{tail}"
99+ if mapped.startswith("model.layers."):
100+ try:
101+ layer_i = int(mapped.split(".")[2])
102+ except (IndexError, ValueError):
103+ return None
104+ if layer_i > max_layer:
105+ return None
106+ if mapped == "lm_head.weight.weight":
107+ return "lm_head.weight"
108+ return mapped
109+ 
110+ logger.debug("Unmapped GLM-family key dropped: %s", hf_key)
111+ return None
112+ 
113+ 
114+def _is_structural_glm5_key(hf_key: str, max_layer: int) -> bool:
115+ """Return True for in-range GLM5 model weights that must not be dropped."""
116+ if not hf_key.startswith("model."):
117+ return hf_key == "lm_head.weight"
118+ tail = hf_key[len("model."):]
119+ if tail.startswith("layers."):
120+ try:
121+ layer_i = int(tail.split(".")[1])
122+ except (IndexError, ValueError):
123+ return True
124+ return layer_i <= max_layer
125+ return tail.startswith(("embed_tokens.", "norm."))
126+ 
127+ 
128+def _is_supported_mapped_key(mapped_key: str) -> bool:
129+ """Return True if mapped key belongs to the implemented GLM5 layout."""
130+ if mapped_key in (
131+ "model.embed_tokens.weight",
132+ "model.norm.weight",
133+ "lm_head.weight",
134+ ):
135+ return True
136+ layer_prefix = "model.layers."
137+ if not mapped_key.startswith(layer_prefix):
138+ return False
139+ tail = mapped_key[len(layer_prefix):]
140+ parts = tail.split(".", 1)
141+ if len(parts) != 2:
142+ return False
143+ return parts[1] in _SUPPORTED_LAYER_SUFFIXES
144+ 
145+ 
146+def _pack_per_expert_tensors(
147+ expert_tensors: Dict[Tuple[int, int, str], torch.Tensor],
148+ num_experts: int,
149+) -> Dict[str, torch.Tensor]:
150+ """Pack per-expert HF MoE tensors into GLM5 expert-major parameters."""
151+ packed: Dict[str, torch.Tensor] = {}
152+ layers = sorted({layer_i for layer_i, _, _ in expert_tensors})
153+ for layer_i in layers:
154+ gate = []
155+ up = []
156+ down = []
157+ for expert_i in range(num_experts):
158+ try:
159+ gate_w = expert_tensors[(layer_i, expert_i, "gate_proj")]
160+ up_w = expert_tensors[(layer_i, expert_i, "up_proj")]
161+ down_w = expert_tensors[(layer_i, expert_i, "down_proj")]
162+ except KeyError as exc:
163+ raise ValueError(
164+ f"GLM5 HF MoE layer {layer_i} missing expert {expert_i} "
165+ f"tensor for packed expert conversion"
166+ ) from exc
167+ if gate_w.shape != up_w.shape:
168+ raise ValueError(
169+ f"GLM5 HF MoE layer {layer_i} expert {expert_i} gate/up "
170+ f"shapes differ: {tuple(gate_w.shape)} vs {tuple(up_w.shape)}"
171+ )
172+ gate.append(gate_w)
173+ up.append(up_w)
174+ target_down_shape = (gate_w.shape[1], gate_w.shape[0])
175+ if down_w.shape == target_down_shape:
176+ down.append(down_w)
177+ continue
178+ if down_w.shape == gate_w.shape:
179+ down.append(down_w.transpose(0, 1).contiguous())
180+ continue
181+ raise ValueError(
182+ f"GLM5 HF MoE layer {layer_i} expert {expert_i} down_proj "
183+ f"shape {tuple(down_w.shape)} does not match gate/up shape "
184+ f"{tuple(gate_w.shape)}"
185+ )
186+ 
187+ prefix = f"model.layers.{layer_i}.mlp.experts"
188+ packed[f"{prefix}.gate_up_proj"] = torch.stack(
189+ [torch.cat([g, u], dim=0) for g, u in zip(gate, up)],
190+ dim=0,
191+ )
192+ packed[f"{prefix}.down_proj"] = torch.stack(down, dim=0)
193+ return packed
194+ 
195+ 
196+def load_hf_glm5_state_dict(
197+ weights_path: str,
198+ num_hidden_layers: int,
199+ num_experts: Optional[int] = None,
200+ dtype: Optional[torch.dtype] = None,
201+) -> Dict[str, torch.Tensor]:
202+ """Load a GLM5/GLM4 safetensors checkpoint into hyper key names."""
203+ idx_path = os.path.join(weights_path, "model.safetensors.index.json")
204+ if not os.path.isfile(idx_path):
205+ raise FileNotFoundError(
206+ f"GLM5 loader needs {idx_path}; pass a directory containing "
207+ "model.safetensors.index.json."
208+ )
209+ with open(idx_path, "r", encoding="utf-8") as f:
210+ idx = json.load(f)
211+ weight_map: Dict[str, str] = idx["weight_map"]
212+ shard_to_keys: Dict[str, list] = {}
213+ for hf_key, shard in weight_map.items():
214+ shard_to_keys.setdefault(shard, []).append(hf_key)
215+ 
216+ max_layer = num_hidden_layers - 1
217+ 
218+ def _cast(tensor: torch.Tensor) -> torch.Tensor:
219+ return tensor.to(dtype) if dtype is not None and tensor.dtype != dtype else tensor
220+ 
221+ hyper_sd: Dict[str, torch.Tensor] = {}
222+ expert_tensors: Dict[Tuple[int, int, str], torch.Tensor] = {}
223+ unsupported = []
224+ skipped = 0
225+ for shard in sorted(shard_to_keys.keys()):
226+ shard_path = _resolve_shard_path(weights_path, shard)
227+ with safe_open(shard_path, framework="pt", device="cpu") as f:
228+ for hf_key in shard_to_keys[shard]:
229+ expert_match = _PER_EXPERT_RE.match(hf_key)
230+ if expert_match:
231+ layer_i = int(expert_match.group("layer"))
232+ if layer_i > max_layer:
233+ skipped += 1
234+ continue
235+ expert_tensors[
236+ (
237+ layer_i,
238+ int(expert_match.group("expert")),
239+ expert_match.group("kind"),
240+ )
241+ ] = _cast(f.get_tensor(hf_key))
242+ continue
243+ mapped = _remap_key(hf_key, max_layer)
244+ if mapped is None:
245+ if _is_structural_glm5_key(hf_key, max_layer):
246+ unsupported.append(hf_key)
247+ else:
248+ skipped += 1
249+ continue
250+ if not _is_supported_mapped_key(mapped):
251+ unsupported.append(hf_key)
252+ continue
253+ hyper_sd[mapped] = _cast(f.get_tensor(hf_key))
254+ 
255+ if expert_tensors:
256+ if num_experts is None:
257+ raise ValueError(
258+ "GLM5 HF checkpoint has per-expert MoE weights; pass "
259+ "num_experts for packed expert conversion."
260+ )
261+ hyper_sd.update(_pack_per_expert_tensors(expert_tensors, num_experts))
262+ if unsupported:
263+ examples = ", ".join(unsupported[:5])
264+ raise ValueError(
265+ "Unsupported GLM5 HF structural keys encountered; refusing to "
266+ f"silently random-initialize model weights: {examples}"
267+ )
268+ 
269+ if "lm_head.weight" not in hyper_sd and "model.embed_tokens.weight" in hyper_sd:
270+ hyper_sd["lm_head.weight"] = hyper_sd["model.embed_tokens.weight"].clone()
271+ 
272+ logger.info(
273+ "GLM-family HF -> hyper state_dict ready: %d keys (%d skipped)",
274+ len(hyper_sd),
275+ skipped,
276+ )
277+ return hyper_sd
278+ 
279+ 
280+__all__ = ["load_hf_glm5_state_dict"]
@@ -0,0 +1,340 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+"""GLM5 DSA indexer and sparse-attention boundaries."""
16+from typing import Optional
17+ 
18+import torch
19+from torch import nn
20+from torch.nn import functional as F
21+ 
22+ 
23+def _local_tensor(value):
24+ return value.to_local() if hasattr(value, "to_local") else value
25+ 
26+ 
27+def _prepare_query_positions(
28+ query_positions: Optional[torch.Tensor],
29+ batch_size: int,
30+ query_len: int,
31+ key_len: int,
32+ device: torch.device,
33+) -> torch.Tensor:
34+ """Return batch-major query positions."""
35+ if query_positions is None:
36+ query_positions = torch.arange(
37+ key_len - query_len,
38+ key_len,
39+ device=device,
40+ dtype=torch.long,
41+ )
42+ if query_positions.ndim == 1:
43+ query_positions = query_positions.view(1, -1).expand(batch_size, -1)
44+ if query_positions.ndim != 2:
45+ raise ValueError("GLM5 DSA positions must have shape (seq,) or (batch, seq)")
46+ if query_positions.shape[0] != batch_size or query_positions.shape[1] != query_len:
47+ raise ValueError("GLM5 DSA positions must match query shape")
48+ return query_positions.to(device=device, dtype=torch.long)
49+ 
50+ 
51+def _infer_key_positions(query_positions: torch.Tensor, key_len: int) -> torch.Tensor:
52+ """Infer key positions for append-only cached decoding."""
53+ query_len = query_positions.shape[1]
54+ past_len = key_len - query_len
55+ if past_len <= 0:
56+ return query_positions.unsqueeze(1)
57+ if query_len > 1:
58+ expected = query_positions[:, :1] + torch.arange(
59+ query_len,
60+ device=query_positions.device,
61+ dtype=query_positions.dtype,
62+ ).view(1, -1)
63+ if not torch.equal(query_positions, expected):
64+ raise ValueError(
65+ "GLM5 DSA cached decode requires contiguous query positions; "
66+ "packed or non-contiguous cached position metadata is not supported."
67+ )
68+ offsets = torch.arange(
69+ key_len,
70+ device=query_positions.device,
71+ dtype=query_positions.dtype,
72+ ).view(1, 1, key_len)
73+ return (query_positions[:, :1].unsqueeze(1) - past_len + offsets).clamp_min(0)
74+ 
75+ 
76+class GLM5DSAIndexerBoundary(nn.Module):
77+ """Select global causal key positions for each local query token."""
78+ 
79+ def __init__(self, topk: int, query_chunk_size: int = 64) -> None:
80+ super().__init__()
81+ self.topk = topk
82+ self.query_chunk_size = query_chunk_size
83+ 
84+ def forward(
85+ self,
86+ query: torch.Tensor,
87+ key: torch.Tensor,
88+ query_positions: Optional[torch.Tensor] = None,
89+ ) -> torch.Tensor:
90+ """Return sparse key indices for each query token."""
91+ query = _local_tensor(query)
92+ key = _local_tensor(key)
93+ if query_positions is not None:
94+ query_positions = _local_tensor(query_positions)
95+ batch_size, query_len, _, indexer_dim = query.shape
96+ key_len = key.shape[1]
97+ topk = min(self.topk, key_len)
98+ if query_len == 0:
99+ return query.new_empty(batch_size, query_len, topk)
100+ query_positions = _prepare_query_positions(
101+ query_positions,
102+ batch_size,
103+ query_len,
104+ key_len,
105+ query.device,
106+ )
107+ key_positions = _infer_key_positions(query_positions, key_len)
108+ topk_chunks = []
109+ for start in range(0, query_len, self.query_chunk_size):
110+ end = min(start + self.query_chunk_size, query_len)
111+ query_chunk = query[:, start:end]
112+ scores = torch.einsum(
113+ "bsnd,btnd->bst", query_chunk, key,
114+ ) * (indexer_dim ** -0.5)
115+ query_pos = query_positions[:, start:end].reshape(
116+ batch_size, end - start, 1,
117+ )
118+ scores = scores.masked_fill(
119+ ~(key_positions <= query_pos), float("-inf"),
120+ )
121+ topk_scores, topk_indices = torch.topk(scores, k=topk, dim=-1)
122+ topk_chunks.append(
123+ topk_indices.masked_fill(~torch.isfinite(topk_scores), -1)
124+ )
125+ if not topk_chunks:
126+ return query.new_empty(batch_size, query_len, topk)
127+ return torch.cat(topk_chunks, dim=1).reshape(
128+ batch_size, query_len, topk,
129+ )
130+ 
131+ 
132+class GLM5DSAIndexer(nn.Module):
133+ """Project hidden states and invoke the CP-compatible indexer boundary."""
134+ 
135+ def __init__(self, hidden_size: int, indexer_dim: int, topk: int) -> None:
136+ super().__init__()
137+ self.query_proj = nn.Linear(hidden_size, indexer_dim, bias=False)
138+ self.key_proj = nn.Linear(hidden_size, indexer_dim, bias=False)
139+ self.boundary = GLM5DSAIndexerBoundary(topk)
140+ 
141+ def forward(
142+ self,
143+ hidden_states: torch.Tensor,
144+ position_ids: Optional[torch.Tensor] = None,
145+ past_key: Optional[torch.Tensor] = None,
146+ ) -> tuple[torch.Tensor, torch.Tensor]:
147+ """Return selected positions and indexer key cache."""
148+ query = self.query_proj(hidden_states).unsqueeze(2)
149+ current_key = self.key_proj(hidden_states).unsqueeze(2)
150+ key = (
151+ torch.cat([past_key, current_key], dim=1)
152+ if past_key is not None
153+ else current_key
154+ )
155+ topk_indices = self.boundary(query, key, position_ids)
156+ return topk_indices, key
157+ 
158+ 
159+def _rotate_half(x: torch.Tensor) -> torch.Tensor:
160+ """Rotate the split-half RoPE dimensions."""
161+ x1 = x[..., : x.shape[-1] // 2]
162+ x2 = x[..., x.shape[-1] // 2:]
163+ return torch.cat((-x2, x1), dim=-1)
164+ 
165+ 
166+def _apply_single_rotary(
167+ x: torch.Tensor,
168+ cos: torch.Tensor,
169+ sin: torch.Tensor,
170+ unsqueeze_dim: int) -> torch.Tensor:
171+ """Apply official split-half GLM-MoE-DSA RoPE to one tensor."""
172+ return (x * cos.unsqueeze(unsqueeze_dim)) + (
173+ _rotate_half(x) * sin.unsqueeze(unsqueeze_dim)
174+ )
175+ 
176+ 
177+class GLM5OfficialDSAIndexer(nn.Module):
178+ """Transformers GLM-MoE-DSA compatible sparse-attention indexer."""
179+ 
180+ def __init__(
181+ self,
182+ hidden_size: int,
183+ q_lora_rank: int,
184+ qk_rope_head_dim: int,
185+ index_topk: int,
186+ index_head_dim: int,
187+ index_n_heads: int) -> None:
188+ super().__init__()
189+ if index_head_dim < qk_rope_head_dim:
190+ raise ValueError("index_head_dim must be >= qk_rope_head_dim")
191+ self.n_heads = index_n_heads
192+ self.head_dim = index_head_dim
193+ self.qk_rope_head_dim = qk_rope_head_dim
194+ self.index_topk = index_topk
195+ self.wq_b = nn.Linear(q_lora_rank, index_n_heads * index_head_dim, bias=False)
196+ self.wk = nn.Linear(hidden_size, index_head_dim, bias=False)
197+ self.k_norm = nn.LayerNorm(index_head_dim, eps=1e-6)
198+ self.weights_proj = nn.Linear(hidden_size, index_n_heads, bias=False)
199+ self.softmax_scale = index_head_dim ** -0.5
200+ self.register_buffer("_cached_keys", None, persistent=False)
201+ 
202+ @torch.no_grad()
203+ def forward(
204+ self,
205+ hidden_states: torch.Tensor,
206+ q_resid: torch.Tensor,
207+ position_embeddings: tuple[torch.Tensor, torch.Tensor],
208+ attention_mask: Optional[torch.Tensor] = None,
209+ use_cache: bool = False) -> torch.Tensor:
210+ """Return official GLM-MoE-DSA top-k key indices."""
211+ batch_size, seq_len, _ = hidden_states.shape
212+ cos, sin = position_embeddings
213+ 
214+ q = self.wq_b(q_resid)
215+ q = q.view(batch_size, seq_len, self.n_heads, self.head_dim)
216+ q_pe, q_nope = torch.split(
217+ q,
218+ [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim],
219+ dim=-1,
220+ )
221+ q_pe = _apply_single_rotary(q_pe, cos, sin, unsqueeze_dim=2)
222+ q = torch.cat([q_pe, q_nope], dim=-1)
223+ 
224+ k = self.k_norm(self.wk(hidden_states))
225+ k_pe, k_nope = torch.split(
226+ k,
227+ [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim],
228+ dim=-1,
229+ )
230+ k_pe = _apply_single_rotary(
231+ k_pe.unsqueeze(2), cos, sin, unsqueeze_dim=2,
232+ ).squeeze(2)
233+ k = torch.cat([k_pe, k_nope], dim=-1)
234+ 
235+ if seq_len > 1:
236+ self._cached_keys = None
237+ if use_cache:
238+ k = (
239+ torch.cat([self._cached_keys, k], dim=1)
240+ if self._cached_keys is not None
241+ else k
242+ )
243+ self._cached_keys = k
244+ 
245+ weights = self.weights_proj(hidden_states).float() * (self.n_heads ** -0.5)
246+ scores = torch.einsum("bshd,btd->bsht", q.float(), k.float())
247+ scores = F.relu(scores * self.softmax_scale)
248+ index_scores = torch.einsum("bsht,bsh->bst", scores, weights)
249+ if attention_mask is not None:
250+ index_scores = index_scores + attention_mask
251+ 
252+ topk = min(self.index_topk, index_scores.shape[-1])
253+ return index_scores.topk(topk, dim=-1).indices
254+ 
255+ 
256+class GLM5SparseAttentionCore(nn.Module):
257+ """Sparse attention over selected global key/value positions."""
258+ 
259+ def __init__(self, scale: float, query_chunk_size: int = 64) -> None:
260+ super().__init__()
261+ self.scale = scale
262+ self.query_chunk_size = query_chunk_size
263+ 
264+ def forward(
265+ self,
266+ query: torch.Tensor,
267+ key: torch.Tensor,
268+ value: torch.Tensor,
269+ topk_indices: torch.Tensor,
270+ query_positions: Optional[torch.Tensor] = None,
271+ attention_mask: Optional[torch.Tensor] = None,
272+ ) -> torch.Tensor:
273+ """Run sparse attention over BSHD tensors."""
274+ query = _local_tensor(query)
275+ key = _local_tensor(key)
276+ value = _local_tensor(value)
277+ topk_indices = _local_tensor(topk_indices)
278+ if query_positions is not None:
279+ query_positions = _local_tensor(query_positions)
280+ if attention_mask is not None:
281+ attention_mask = _local_tensor(attention_mask)
282+ 
283+ batch_size, query_len, num_heads, head_dim = query.shape
284+ batch_indices = torch.arange(
285+ batch_size, device=query.device,
286+ ).reshape(batch_size, 1, 1)
287+ mask = attention_mask
288+ if attention_mask is not None and attention_mask.ndim == 4:
289+ if mask.shape[1] == 1 and num_heads > 1:
290+ mask = mask.expand(batch_size, num_heads, -1, -1)
291+ mask = mask.transpose(1, 2)
292+ outputs = []
293+ for start in range(0, query_len, self.query_chunk_size):
294+ end = min(start + self.query_chunk_size, query_len)
295+ query_chunk = query[:, start:end]
296+ topk_chunk = topk_indices[:, start:end]
297+ safe_indices = topk_chunk.clamp_min(0)
298+ selected_key = key[batch_indices, safe_indices].permute(
299+ 0, 1, 3, 2, 4,
300+ )
301+ selected_value = value[batch_indices, safe_indices].permute(
302+ 0, 1, 3, 2, 4,
303+ )
304+ scores = (
305+ query_chunk.unsqueeze(-2) * selected_key
306+ ).sum(dim=-1) * self.scale
307+ invalid = topk_chunk.lt(0).unsqueeze(2)
308+ scores = scores.masked_fill(invalid, float("-inf"))
309+ 
310+ if mask is not None:
311+ mask_chunk = (
312+ mask[:, start:end]
313+ if mask.shape[1] == query_len
314+ else mask
315+ )
316+ gather_indices = safe_indices.unsqueeze(2).expand(
317+ batch_size, end - start, num_heads, safe_indices.shape[-1],
318+ )
319+ scores = scores + torch.gather(mask_chunk, -1, gather_indices)
320+ 
321+ probabilities = torch.softmax(scores, dim=-1, dtype=torch.float32).to(
322+ query.dtype
323+ )
324+ probabilities = probabilities.masked_fill(invalid, 0.0)
325+ outputs.append(
326+ (probabilities.unsqueeze(-1) * selected_value)
327+ .sum(dim=-2)
328+ .reshape(batch_size, end - start, num_heads, head_dim)
329+ )
330+ if not outputs:
331+ return query.new_empty(batch_size, query_len, num_heads, head_dim)
332+ return torch.cat(outputs, dim=1)
atomgit-bot
atomgit-botatomgit-bot6月22日

🔵 Low Priority

与 GLM5DSAIndexerBoundary.forward() 相同模式:GLM5SparseAttentionCore.forward() 第 233 行的 torch.cat(outputs, dim=1) 在 query_len == 0 时也会因为 outputs 为空列表而抛出 RuntimeError。

变更行:dsa.py 第 233 行 return torch.cat(outputs, dim=1)

触发条件:当 query.shape[1] == 0 时,循环 for start in range(0, query_len, self.query_chunk_size) 不执行,outputs 保持为空列表。虽然正常训练/推理中 query_len 不会为 0,但缺少防御性检查使得该模块在异常调用(如空序列测试)下存在崩溃风险。

建议:在 torch.cat(outputs, dim=1) 前添加空列表保护,返回形状正确的空张量。

likedislike
不准确?
333+ 
334+ 
335+__all__ = [
336+ "GLM5DSAIndexer",
337+ "GLM5DSAIndexerBoundary",
338+ "GLM5OfficialDSAIndexer",
339+ "GLM5SparseAttentionCore",
340+]
@@ -0,0 +1,580 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+"""GLM5 causal language model."""
16+from dataclasses import dataclass, field
17+from typing import Any, List, Optional
18+ 
19+import torch
20+from torch import nn
21+from torch.nn import functional as F
22+ 
23+from hyper_parallel.models.glm5.attention import (
24+ GLM5GQAAttention,
25+ GLM5MLAAttention,
26+ GLM5OfficialMLAAttention,
27+)
28+from hyper_parallel.models.glm5.dsa import GLM5DSAIndexer
29+from hyper_parallel.models.glm5.moe import GLM5MoE
30+from hyper_parallel.models.modules.feed_forward import SwiGLUMLP
31+from hyper_parallel.models.modules.rmsnorm import RMSNorm
32+from hyper_parallel.models.modules.rope import RotaryEmbedding
33+ 
34+ 
35+def _init_glm5_moe_experts(module: GLM5MoE) -> None:
36+ nn.init.kaiming_uniform_(module.experts.gate_up_proj, a=5 ** 0.5)
37+ nn.init.kaiming_uniform_(module.experts.down_proj, a=5 ** 0.5)
38+ 
39+ 
40+def _masked_cross_entropy(logits: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:
41+ """Return a finite mean when a sequence shard has no valid labels."""
42+ flat_logits = logits.view(-1, logits.size(-1))
43+ flat_labels = labels.view(-1)
44+ loss_sum = F.cross_entropy(
45+ flat_logits,
46+ flat_labels,
47+ ignore_index=-100,
48+ reduction="sum",
49+ )
50+ valid_tokens = (flat_labels != -100).sum().clamp_min(1)
51+ return loss_sum / valid_tokens
52+ 
53+ 
54+def prepare_glm5_batch(batch: dict, model: Any) -> dict:
55+ """Pad and shard sequence inputs for GLM5 context parallel training."""
56+ cp_size = getattr(model, "_cp_size", 1)
57+ if cp_size <= 1:
58+ return batch
59+ 
60+ cp_rank = getattr(model, "_cp_rank", 0)
61+ input_ids = batch["input_ids"]
62+ labels = batch.get("labels")
63+ seq_len = input_ids.shape[1]
64+ pad_len = (-seq_len) % cp_size
65+ 
66+ if pad_len:
67+ input_ids = F.pad(input_ids, (0, pad_len), value=0)
68+ if labels is not None:
69+ labels = F.pad(labels, (0, pad_len), value=-100)
70+ 
71+ padded_len = input_ids.shape[1]
72+ local_len = padded_len // cp_size
73+ start = cp_rank * local_len
74+ end = start + local_len
75+ 
76+ prepared = dict(batch)
77+ prepared["input_ids"] = input_ids[:, start:end].contiguous()
78+ prepared["position_ids"] = torch.arange(
79+ start, end, dtype=torch.long, device=input_ids.device,
80+ )
81+ 
82+ if labels is not None:
83+ shifted_labels = F.pad(labels[:, 1:], (0, 1), value=-100)
84+ prepared["labels"] = shifted_labels[:, start:end].contiguous()
85+ 
86+ attention_mask = batch.get("attention_mask")
87+ if attention_mask is not None:
88+ if attention_mask.ndim == 2:
89+ if pad_len:
90+ attention_mask = F.pad(attention_mask, (0, pad_len), value=0)
91+ prepared["attention_mask"] = attention_mask[:, start:end].contiguous()
92+ elif attention_mask.ndim == 4:
93+ if pad_len:
94+ attention_mask = F.pad(attention_mask, (0, pad_len), value=float("-inf"))
95+ if attention_mask.shape[-2] != 1:
96+ attention_mask = F.pad(
97+ attention_mask, (0, 0, 0, pad_len), value=float("-inf")
98+ )
99+ query_slice = slice(None) if attention_mask.shape[-2] == 1 else slice(start, end)
100+ prepared["attention_mask"] = attention_mask[
101+ :, :, query_slice, start:end
102+ ].contiguous()
103+ 
104+ return prepared
105+ 
106+ 
107+@dataclass
108+class GLM5Config:
109+ """GLM5 model configuration."""
110+ 
111+ vocab_size: int = 154856
112+ hidden_size: int = 1024
113+ intermediate_size: int = 3072
114+ num_hidden_layers: int = 24
115+ num_attention_heads: int = 16
116+ num_key_value_heads: int = 4
117+ head_dim: int = 64
118+ q_lora_rank: Optional[int] = None
119+ max_position_embeddings: int = 131072
120+ rms_norm_eps: float = 1e-6
121+ rope_theta: float = 500000.0
122+ tie_word_embeddings: bool = True
123+ attention_bias: bool = False
124+ 
125+ num_experts: int = 256
126+ num_experts_per_tok: int = 8
127+ num_dense_layers: int = 3
128+ moe_intermediate_size: int = 1024
129+ kv_lora_rank: int = 576
130+ qk_nope_head_dim: int = 0
131+ qk_rope_head_dim: int = 64
132+ v_head_dim: int = 128
133+ index_topk: int = 2048
134+ index_head_dim: int = 64
135+ index_n_heads: int = 16
136+ dsa_topk: int = 2048
137+ dsa_indexer_dim: int = 64
138+ attention_type: str = "gqa"
139+ use_dsa: bool = False
140+ moe_router_type: str = "softmax"
141+ n_shared_experts: int = 0
142+ routed_scaling_factor: float = 1.0
143+ n_group: int = 1
144+ topk_group: int = 1
145+ norm_topk_prob: bool = True
146+ layer_types: Optional[List[str]] = field(default=None)
147+ 
148+ def __post_init__(self):
149+ self._validate_positive_fields()
150+ self._validate_attention_shape()
151+ self._validate_expert_shape()
152+ self._validate_layer_types()
153+ 
154+ def _validate_positive_fields(self) -> None:
155+ """Validate scalar fields that must stay positive."""
156+ for field_name in (
157+ "vocab_size", "hidden_size", "intermediate_size",
158+ "num_hidden_layers", "num_attention_heads",
159+ "num_key_value_heads", "head_dim", "moe_intermediate_size",
160+ "kv_lora_rank", "qk_rope_head_dim", "v_head_dim",
161+ "index_topk", "index_head_dim", "index_n_heads",
162+ "dsa_topk", "dsa_indexer_dim", "max_position_embeddings"):
163+ if getattr(self, field_name) <= 0:
164+ raise ValueError(f"{field_name} must be positive")
165+ if self.num_dense_layers < 0:
166+ raise ValueError("num_dense_layers must be non-negative")
167+ if self.q_lora_rank is not None and self.q_lora_rank <= 0:
168+ raise ValueError("q_lora_rank must be positive when set")
169+ if self.qk_nope_head_dim < 0:
170+ raise ValueError("qk_nope_head_dim must be non-negative")
171+ 
172+ def _validate_attention_shape(self) -> None:
173+ """Validate attention dimensions and attention type."""
174+ expected_hidden = self.num_attention_heads * self.head_dim
175+ if self.hidden_size != expected_hidden:
176+ raise ValueError(
177+ f"hidden_size ({self.hidden_size}) must equal "
178+ f"num_attention_heads * head_dim ({expected_hidden})"
179+ )
180+ if self.num_attention_heads % self.num_key_value_heads != 0:
181+ raise ValueError(
182+ "num_attention_heads must be divisible by num_key_value_heads"
183+ )
184+ if self.attention_type == "mla" and self.qk_rope_head_dim > self.head_dim:
185+ raise ValueError("qk_rope_head_dim must be <= head_dim")
186+ if self.attention_type not in ("gqa", "mla", "glm_moe_dsa_mla"):
187+ raise ValueError(f"Unsupported GLM5 attention_type: {self.attention_type}")
188+ if self.attention_type == "glm_moe_dsa_mla":
189+ if self.q_lora_rank is None:
190+ raise ValueError("q_lora_rank is required for glm_moe_dsa_mla")
191+ if self.qk_nope_head_dim + self.qk_rope_head_dim != self.head_dim:
192+ raise ValueError(
193+ "head_dim must equal qk_nope_head_dim + qk_rope_head_dim"
194+ )
195+ if self.index_head_dim < self.qk_rope_head_dim:
196+ raise ValueError("index_head_dim must be >= qk_rope_head_dim")
197+ 
198+ def _validate_expert_shape(self) -> None:
199+ if self.num_experts <= 0:
200+ raise ValueError("num_experts must be positive")
201+ if self.num_experts_per_tok <= 0:
202+ raise ValueError("num_experts_per_tok must be positive")
203+ if self.num_experts_per_tok > self.num_experts:
204+ raise ValueError("num_experts_per_tok must be <= num_experts")
205+ 
206+ def _validate_layer_types(self) -> None:
207+ """Validate dense and MoE layer layout."""
208+ if not 0 <= self.num_dense_layers <= self.num_hidden_layers:
209+ raise ValueError("num_dense_layers must be in [0, num_hidden_layers]")
210+ if self.layer_types is None:
211+ self.layer_types = [
212+ "dense" if layer_idx < self.num_dense_layers else "moe"
213+ for layer_idx in range(self.num_hidden_layers)
214+ ]
215+ if len(self.layer_types) != self.num_hidden_layers:
216+ raise ValueError(
217+ f"layer_types length {len(self.layer_types)} != "
218+ f"num_hidden_layers {self.num_hidden_layers}"
219+ )
220+ invalid_layer_types = [
221+ layer_type for layer_type in self.layer_types
222+ if layer_type not in ("dense", "moe")
223+ ]
224+ if invalid_layer_types:
225+ raise ValueError(
226+ f"Unsupported GLM5 layer_types: {invalid_layer_types}"
227+ )
228+ 
229+ 
230+class GLM5Decoder(nn.Module):
231+ """One GLM5 decoder layer: RMSNorm -> Attention -> RMSNorm -> MLP/MoE."""
232+ 
233+ def __init__(
234+ self,
235+ config: GLM5Config,
236+ layer_idx: int,
237+ ):
238+ super().__init__()
239+ self.layer_idx = layer_idx
240+ self.layer_type = config.layer_types[layer_idx]
241+ self.input_layernorm = RMSNorm(
242+ config.hidden_size, eps=config.rms_norm_eps,
243+ )
244+ if config.attention_type == "gqa":
245+ self.rotary_emb = RotaryEmbedding(
246+ dim=config.head_dim,
247+ max_seq_len=config.max_position_embeddings,
248+ theta=config.rope_theta,
249+ )
250+ self.self_attn = GLM5GQAAttention(
251+ hidden_size=config.hidden_size,
252+ num_heads=config.num_attention_heads,
253+ num_kv_heads=config.num_key_value_heads,
254+ head_dim=config.head_dim,
255+ qkv_bias=config.attention_bias,
256+ out_bias=config.attention_bias,
257+ rope=self.rotary_emb,
258+ rms_norm_eps=config.rms_norm_eps,
259+ use_dsa=config.use_dsa,
260+ )
261+ elif config.attention_type == "mla":
262+ self.self_attn = GLM5MLAAttention(
263+ hidden_size=config.hidden_size,
264+ num_heads=config.num_attention_heads,
265+ num_kv_heads=config.num_key_value_heads,
266+ head_dim=config.head_dim,
267+ kv_lora_rank=config.kv_lora_rank,
268+ qk_rope_head_dim=config.qk_rope_head_dim,
269+ v_head_dim=config.v_head_dim,
270+ max_position_embeddings=config.max_position_embeddings,
271+ rope_theta=config.rope_theta,
272+ bias=config.attention_bias,
273+ rms_norm_eps=config.rms_norm_eps,
274+ use_dsa=config.use_dsa,
275+ )
276+ self.rotary_emb = self.self_attn.rotary_emb
277+ else:
278+ self.self_attn = GLM5OfficialMLAAttention(
279+ hidden_size=config.hidden_size,
280+ num_heads=config.num_attention_heads,
281+ head_dim=config.head_dim,
282+ q_lora_rank=config.q_lora_rank,
283+ kv_lora_rank=config.kv_lora_rank,
284+ qk_nope_head_dim=config.qk_nope_head_dim,
285+ qk_rope_head_dim=config.qk_rope_head_dim,
286+ v_head_dim=config.v_head_dim,
287+ index_topk=config.index_topk,
288+ index_head_dim=config.index_head_dim,
289+ index_n_heads=config.index_n_heads,
290+ max_position_embeddings=config.max_position_embeddings,
291+ rope_theta=config.rope_theta,
292+ bias=config.attention_bias,
293+ rms_norm_eps=config.rms_norm_eps,
294+ )
295+ self.rotary_emb = self.self_attn.rotary_emb
296+ self.dsa_indexer = (
297+ GLM5DSAIndexer(
298+ hidden_size=config.hidden_size,
299+ indexer_dim=config.dsa_indexer_dim,
300+ topk=config.dsa_topk,
301+ )
302+ if config.use_dsa
303+ else None
304+ )
305+ self.post_attention_layernorm = RMSNorm(
306+ config.hidden_size, eps=config.rms_norm_eps,
307+ )
308+ if self.layer_type == "dense":
309+ self.mlp = SwiGLUMLP(
310+ hidden_size=config.hidden_size,
311+ intermediate_size=config.intermediate_size,
312+ bias=False,
313+ )
314+ elif self.layer_type == "moe":
315+ self.mlp = GLM5MoE(
316+ hidden_size=config.hidden_size,
317+ moe_intermediate_size=config.moe_intermediate_size,
318+ num_experts=config.num_experts,
319+ top_k=config.num_experts_per_tok,
320+ router_type=config.moe_router_type,
321+ n_shared_experts=config.n_shared_experts,
322+ routed_scaling_factor=config.routed_scaling_factor,
323+ n_group=config.n_group,
324+ topk_group=config.topk_group,
325+ norm_topk_prob=config.norm_topk_prob,
326+ )
327+ _init_glm5_moe_experts(self.mlp)
328+ else:
329+ raise ValueError(
330+ f"Unknown GLM5 layer_type '{self.layer_type}' at layer {layer_idx}"
331+ )
332+ 
333+ def forward(
334+ self,
335+ hidden_states: torch.Tensor,
336+ position_ids: Optional[torch.Tensor] = None,
337+ attention_mask: Optional[torch.Tensor] = None,
338+ past_key_value: Optional[torch.Tensor] = None,
339+ use_cache: bool = False,
340+ **kwargs,
341+ ):
342+ """Run one decoder layer."""
343+ del kwargs
344+ attention_past = past_key_value
345+ indexer_past_key = None
346+ if self.dsa_indexer is not None and past_key_value is not None:
347+ attention_past, indexer_past_key = past_key_value
348+ residual = hidden_states
349+ hidden_states = self.input_layernorm(hidden_states)
350+ topk_indices = None
351+ indexer_key_cache = None
352+ if self.dsa_indexer is not None:
353+ topk_indices, indexer_key_cache = self.dsa_indexer(
354+ hidden_states,
355+ position_ids,
356+ past_key=indexer_past_key,
357+ )
358+ attn_output = self.self_attn(
359+ hidden_states,
360+ position_ids=position_ids,
361+ attention_mask=attention_mask,
362+ past_key_value=attention_past,
363+ use_cache=use_cache,
364+ topk_indices=topk_indices,
365+ )
366+ present_key_value = None
367+ if use_cache and isinstance(attn_output, tuple):
368+ hidden_states, present_key_value = attn_output
369+ else:
370+ hidden_states = attn_output
371+ hidden_states = residual + hidden_states
372+ 
373+ residual = hidden_states
374+ hidden_states = self.post_attention_layernorm(hidden_states)
375+ hidden_states = self.mlp(hidden_states)
376+ hidden_states = residual + hidden_states
377+ if use_cache:
378+ if self.dsa_indexer is not None:
379+ present_key_value = (present_key_value, indexer_key_cache)
380+ return hidden_states, present_key_value
381+ return hidden_states
382+ 
383+ 
384+class GLM5TextModel(nn.Module):
385+ """Inner GLM5 decoder stack."""
386+ 
387+ def __init__(self, config: GLM5Config):
388+ super().__init__()
389+ self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size)
390+ self.layers = nn.ModuleList([
391+ GLM5Decoder(config, layer_idx)
392+ for layer_idx in range(config.num_hidden_layers)
393+ ])
394+ self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
395+ 
396+ @property
397+ def rotary_emb(self):
398+ return self.layers[0].rotary_emb
399+ 
400+ 
401+class GLM5ForCausalLM(nn.Module):
402+ """GLM5 causal language model for Trainer integration."""
403+ 
404+ _tp_plan = {
405+ "*.self_attn.q_proj": "colwise",
406+ "*.self_attn.k_proj": "colwise",
407+ "*.self_attn.v_proj": "colwise",
408+ "*.self_attn.o_proj": "rowwise",
409+ "*.mlp.gate_proj": "colwise",
410+ "*.mlp.up_proj": "colwise",
411+ "*.mlp.down_proj": "rowwise",
412+ }
413+ _cp_modules = ["*.self_attn.attention_core"]
414+ _ep_modules = ["*.mlp.experts"]
415+ 
416+ def __init__(self, config: GLM5Config):
417+ super().__init__()
418+ self.config = config
419+ self._cp_size = 1
420+ self._cp_rank = 0
421+ self.model = GLM5TextModel(config)
422+ self.lm_head = nn.Linear(
423+ config.hidden_size,
424+ config.vocab_size,
425+ bias=False,
426+ )
427+ self.tie_weights()
428+ 
429+ @property
430+ def layers(self):
431+ return self.model.layers
432+ 
433+ @property
434+ def embed_tokens(self):
435+ return self.model.embed_tokens
436+ 
437+ @property
438+ def norm(self):
439+ return self.model.norm
440+ 
441+ @property
442+ def rotary_emb(self):
443+ return self.model.rotary_emb
444+ 
445+ def tie_weights(self) -> None:
446+ if getattr(self.config, "tie_word_embeddings", False):
447+ self.lm_head.weight = self.model.embed_tokens.weight
448+ 
449+ def _past_length(self, past_key_values: list) -> int:
450+ """Return cached sequence length from the first layer cache."""
451+ if not past_key_values or past_key_values[0] is None:
452+ return 0
453+ first_past = past_key_values[0]
454+ if (
455+ isinstance(first_past, tuple)
456+ and first_past
457+ and isinstance(first_past[0], tuple)
458+ ):
459+ return first_past[0][0].shape[1]
460+ if isinstance(first_past, tuple):
461+ return first_past[0].shape[1]
462+ return first_past.shape[1]
463+ 
464+ def _default_position_ids(
465+ self,
466+ input_ids: torch.Tensor,
467+ seq_len: int,
468+ past_len: int,
469+ ) -> torch.Tensor:
470+ return torch.arange(
471+ past_len,
472+ past_len + seq_len,
473+ device=input_ids.device,
474+ dtype=torch.long,
475+ )
476+ 
477+ def _forward_layers(
478+ self,
479+ hidden_states: torch.Tensor,
480+ position_ids: torch.Tensor,
481+ attention_mask: Optional[torch.Tensor],
482+ past_key_values: list,
483+ use_cache: bool,
484+ ) -> tuple:
485+ """Run all decoder layers and collect cache outputs when requested."""
486+ next_past_key_values = [] if use_cache else None
487+ for layer, layer_past in zip(self.model.layers, past_key_values):
488+ layer_output = layer(
489+ hidden_states,
490+ position_ids=position_ids,
491+ attention_mask=attention_mask,
492+ past_key_value=layer_past,
493+ use_cache=use_cache,
494+ )
495+ if use_cache:
496+ hidden_states, present = layer_output
497+ next_past_key_values.append(present)
498+ else:
499+ hidden_states = layer_output
500+ return hidden_states, next_past_key_values
501+ 
502+ def _loss(
503+ self,
504+ logits: torch.Tensor,
505+ labels: Optional[torch.Tensor],
506+ ) -> Optional[torch.Tensor]:
507+ """Compute shifted causal LM loss, accounting for CP-prepared labels."""
508+ if labels is None:
509+ return None
510+ if self._cp_size > 1:
511+ shift_logits = logits.contiguous().float()
512+ shift_labels = labels.contiguous()
513+ else:
514+ shift_logits = logits[..., :-1, :].contiguous().float()
515+ shift_labels = labels[..., 1:].contiguous()
516+ return _masked_cross_entropy(shift_logits, shift_labels)
517+ 
518+ def _build_output(
519+ self,
520+ loss: Optional[torch.Tensor],
521+ logits: torch.Tensor,
522+ next_past_key_values: Optional[list],
523+ hidden_states: torch.Tensor,
524+ return_hidden_states: bool,
525+ ) -> dict:
526+ """Pack model outputs in the Trainer-compatible dictionary format."""
527+ output: dict[str, Any] = {"loss": loss, "logits": logits}
528+ if next_past_key_values is not None:
529+ output["past_key_values"] = next_past_key_values
530+ if return_hidden_states:
531+ output["hidden_states"] = hidden_states
532+ return output
533+ 
534+ def forward(
535+ self,
536+ input_ids: torch.Tensor,
537+ labels: Optional[torch.Tensor] = None,
538+ position_ids: Optional[torch.Tensor] = None,
539+ attention_mask: Optional[torch.Tensor] = None,
540+ past_key_values: Optional[list] = None,
541+ use_cache: bool = False,
542+ **kwargs,
543+ ):
544+ """Run GLM5 causal LM forward."""
545+ return_hidden_states = kwargs.pop("return_hidden_states", False)
546+ del kwargs
547+ _, seq_len = input_ids.shape
548+ if past_key_values is None:
549+ past_key_values = [None] * len(self.model.layers)
550+ if position_ids is None:
551+ position_ids = self._default_position_ids(
552+ input_ids, seq_len, self._past_length(past_key_values)
553+ )
554+ 
555+ input_embeds = self.model.embed_tokens(input_ids)
556+ hidden_states, next_past_key_values = self._forward_layers(
557+ input_embeds,
558+ position_ids,
559+ attention_mask,
560+ past_key_values,
561+ use_cache,
562+ )
563+ hidden_states = self.model.norm(hidden_states)
564+ logits = self.lm_head(hidden_states)
565+ 
566+ return self._build_output(
567+ self._loss(logits, labels),
568+ logits,
569+ next_past_key_values,
570+ hidden_states,
571+ return_hidden_states,
572+ )
573+ 
574+ 
575+__all__ = [
576+ "GLM5Config",
577+ "GLM5Decoder",
578+ "GLM5ForCausalLM",
579+ "GLM5TextModel",
580+]
@@ -0,0 +1,329 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+"""GLM5 mixture-of-experts modules."""
16+import importlib
17+ 
18+import torch
19+from torch import nn
20+from torch.nn import functional as F
21+ 
22+ 
23+def _run_experts_for_loop(
24+ gate_up_proj: torch.Tensor,
25+ down_proj: torch.Tensor,
26+ hidden_states: torch.Tensor,
27+ num_tokens_per_expert: list[int],
28+) -> torch.Tensor:
29+ """Run packed experts with one batched matmul per non-empty expert."""
30+ outputs = []
31+ offset = 0
32+ for expert_idx, count in enumerate(num_tokens_per_expert):
33+ if count == 0:
34+ continue
35+ current_state = hidden_states[offset:offset + count]
36+ gate, up = F.linear(
37+ current_state, gate_up_proj[expert_idx],
38+ ).chunk(2, dim=-1)
39+ current_state = F.silu(gate) * up
40+ outputs.append(F.linear(current_state, down_proj[expert_idx]))
41+ offset += count
42+ if not outputs:
43+ return hidden_states * 0.0
44+ return torch.cat(outputs, dim=0)
45+ 
46+ 
47+def _run_experts_grouped_mm_npu(
48+ gate_up_proj: torch.Tensor,
49+ down_proj: torch.Tensor,
50+ hidden_states: torch.Tensor,
51+ num_tokens_per_expert: list[int],
52+) -> torch.Tensor:
53+ """Run packed experts with Ascend grouped matmul kernels."""
54+ torch_npu = importlib.import_module("torch_npu")
55+ 
56+ expert_inputs = list(torch.split(hidden_states, num_tokens_per_expert, dim=0))
57+ gate_proj, up_proj = gate_up_proj.chunk(2, dim=1)
58+ gate_weights = [
59+ gate_proj[idx].T.contiguous() for idx in range(gate_proj.shape[0])
60+ ]
61+ up_weights = [
62+ up_proj[idx].T.contiguous() for idx in range(up_proj.shape[0])
63+ ]
64+ down_weights = [
65+ down_proj[idx].T.contiguous() for idx in range(down_proj.shape[0])
66+ ]
67+ 
68+ gate_outputs = torch_npu.npu_grouped_matmul(
69+ expert_inputs, gate_weights, group_type=-1,
70+ )
71+ up_outputs = torch_npu.npu_grouped_matmul(
72+ expert_inputs, up_weights, group_type=-1,
73+ )
74+ hidden_outputs = [
75+ F.silu(gate) * up for gate, up in zip(gate_outputs, up_outputs)
76+ ]
77+ output_list = torch_npu.npu_grouped_matmul(
78+ hidden_outputs, down_weights, group_type=-1,
79+ )
80+ return torch.cat(output_list, dim=0)
81+ 
82+ 
83+class GLM5MoEExperts(nn.Module):
84+ """Packed GLM5 experts with an expert-major forward interface."""
85+ 
86+ def __init__(
87+ self,
88+ num_experts: int,
89+ hidden_size: int,
90+ intermediate_size: int,
91+ ) -> None:
92+ """Initialize packed expert parameters."""
93+ super().__init__()
94+ self.num_experts = num_experts
95+ self.hidden_size = hidden_size
96+ self.intermediate_size = intermediate_size
97+ self.gate_up_proj = nn.Parameter(
98+ torch.empty(num_experts, 2 * intermediate_size, hidden_size)
99+ )
100+ self.down_proj = nn.Parameter(
101+ torch.empty(num_experts, hidden_size, intermediate_size)
102+ )
103+ 
104+ def forward(
105+ self,
106+ hidden_states: torch.Tensor,
107+ num_tokens_per_expert: torch.Tensor,
108+ ) -> torch.Tensor:
109+ """Run local experts on expert-major routed tokens."""
110+ gate_up_proj = (
111+ self.gate_up_proj.to_local()
112+ if hasattr(self.gate_up_proj, "to_local")
113+ else self.gate_up_proj
114+ )
115+ down_proj = (
116+ self.down_proj.to_local()
117+ if hasattr(self.down_proj, "to_local")
118+ else self.down_proj
119+ )
120+ if num_tokens_per_expert.ndim != 1:
121+ raise ValueError("num_tokens_per_expert must be a 1-D tensor")
122+ if num_tokens_per_expert.shape[0] != gate_up_proj.shape[0]:
123+ raise ValueError(
124+ "num_tokens_per_expert length must match local experts"
125+ )
126+ if hidden_states.shape[0] == 0:
127+ return hidden_states * 0.0
128+ tokens_per_expert = num_tokens_per_expert.tolist()
129+ if sum(tokens_per_expert) != hidden_states.shape[0]:
130+ raise ValueError(
131+ "sum(num_tokens_per_expert) must match routed token count"
132+ )
133+ 
134+ if hidden_states.device.type == "npu":
135+ return _run_experts_grouped_mm_npu(
136+ gate_up_proj, down_proj, hidden_states, tokens_per_expert,
137+ )
138+ return _run_experts_for_loop(
139+ gate_up_proj, down_proj, hidden_states, tokens_per_expert,
140+ )
141+ 
142+ 
143+class GLM5MoE(nn.Module):
144+ """GLM5 top-k router with expert-major token dispatch boundaries."""
145+ 
146+ def __init__(
147+ self,
148+ hidden_size: int,
149+ moe_intermediate_size: int,
150+ num_experts: int,
151+ top_k: int,
152+ router_type: str = "softmax",
153+ n_shared_experts: int = 0,
154+ routed_scaling_factor: float = 1.0,
155+ n_group: int = 1,
156+ topk_group: int = 1,
157+ norm_topk_prob: bool = True,
158+ ) -> None:
159+ """Initialize the router and packed experts."""
160+ super().__init__()
161+ self.num_experts = num_experts
162+ self.top_k = top_k
163+ self.router_type = router_type
164+ self.routed_scaling_factor = routed_scaling_factor
165+ self.n_group = n_group
166+ self.topk_group = topk_group
167+ self.norm_topk_prob = norm_topk_prob
168+ self.gate = (
169+ GLM5MoERouter(hidden_size, num_experts)
170+ if router_type == "glm_moe_dsa"
171+ else nn.Linear(hidden_size, num_experts, bias=False)
172+ )
173+ self.experts = GLM5MoEExperts(
174+ num_experts, hidden_size, moe_intermediate_size,
175+ )
176+ self.shared_experts = (
177+ GLM5SharedExperts(
178+ hidden_size,
179+ moe_intermediate_size * n_shared_experts,
180+ )
181+ if n_shared_experts > 0
182+ else None
183+ )
184+ 
185+ def _forward_softmax(self, hidden_states: torch.Tensor) -> torch.Tensor:
186+ """Run the original softmax top-k expert router."""
187+ batch_size, seq_len, hidden_size = hidden_states.shape
188+ flat_states = hidden_states.reshape(-1, hidden_size)
189+ router_logits = self.gate(flat_states)
190+ topk_logits, selected_experts = torch.topk(
191+ router_logits, self.top_k, dim=-1,
192+ )
193+ topk_weights = F.softmax(
194+ topk_logits, dim=-1, dtype=torch.float32,
195+ ).to(hidden_states.dtype)
196+ 
197+ flat_experts = selected_experts.flatten()
198+ permutation = flat_experts.argsort(stable=True)
199+ token_indices = permutation // self.top_k
200+ routed_states = flat_states[token_indices]
201+ tokens_per_expert = torch.bincount(
202+ flat_experts, minlength=self.num_experts,
203+ )
204+ 
205+ expert_output = self.experts(routed_states, tokens_per_expert)
206+ sorted_weights = topk_weights.flatten()[permutation]
207+ expert_output = expert_output * sorted_weights.unsqueeze(-1)
208+ combined = torch.zeros(
209+ flat_states.shape,
210+ dtype=expert_output.dtype,
211+ device=expert_output.device,
212+ ).scatter_add(
213+ 0,
214+ token_indices.unsqueeze(-1).expand(-1, hidden_size),
215+ expert_output,
216+ )
217+ return combined.view(batch_size, seq_len, hidden_size)
218+ 
219+ def _select_glm_moe_dsa_experts(
220+ self,
221+ router_scores: torch.Tensor,
222+ ) -> tuple[torch.Tensor, torch.Tensor]:
223+ """Select experts using the Transformers GLM-MoE-DSA router rule."""
224+ scores_for_choice = router_scores
225+ correction = getattr(self.gate, "e_score_correction_bias", None)
226+ if correction is not None:
227+ scores_for_choice = scores_for_choice + correction
228+ if self.n_group > 1:
229+ group_scores = scores_for_choice.view(
230+ -1, self.n_group, self.num_experts // self.n_group,
231+ )
232+ group_scores = group_scores.topk(2, dim=-1)[0].sum(dim=-1)
233+ group_indices = torch.topk(
234+ group_scores, k=self.topk_group, dim=-1, sorted=False,
235+ )[1]
236+ group_mask = torch.zeros_like(group_scores)
237+ group_mask.scatter_(1, group_indices, 1)
238+ score_mask = group_mask.unsqueeze(-1).expand(
239+ -1, self.n_group, self.num_experts // self.n_group,
240+ ).reshape(-1, self.num_experts)
241+ scores_for_choice = scores_for_choice.masked_fill(
242+ ~score_mask.bool(), float("-inf"),
243+ )
244+ topk_indices = torch.topk(
245+ scores_for_choice, k=self.top_k, dim=-1, sorted=False,
246+ )[1]
247+ topk_weights = router_scores.gather(1, topk_indices)
248+ if self.norm_topk_prob:
249+ denominator = topk_weights.sum(dim=-1, keepdim=True) + 1e-20
250+ topk_weights = topk_weights / denominator
251+ topk_weights = topk_weights * self.routed_scaling_factor
252+ return topk_indices, topk_weights
253+ 
254+ def _forward_glm_moe_dsa(self, hidden_states: torch.Tensor) -> torch.Tensor:
255+ """Run the official GLM-MoE-DSA sigmoid router."""
256+ batch_size, seq_len, hidden_size = hidden_states.shape
257+ flat_states = hidden_states.reshape(-1, hidden_size)
258+ router_logits = F.linear(flat_states.float(), self.gate.weight.float())
259+ router_scores = router_logits.sigmoid()
260+ topk_indices, topk_weights = self._select_glm_moe_dsa_experts(router_scores)
261+ final_hidden_states = torch.zeros_like(flat_states)
262+ 
263+ gate_up_proj = (
264+ self.experts.gate_up_proj.to_local()
265+ if hasattr(self.experts.gate_up_proj, "to_local")
266+ else self.experts.gate_up_proj
267+ )
268+ down_proj = (
269+ self.experts.down_proj.to_local()
270+ if hasattr(self.experts.down_proj, "to_local")
271+ else self.experts.down_proj
272+ )
273+ expert_mask = torch.nn.functional.one_hot(
274+ topk_indices, num_classes=self.num_experts,
275+ ).permute(2, 1, 0)
276+ for expert_idx in range(self.num_experts):
277+ expert_slots, token_idx = torch.where(expert_mask[expert_idx])
278+ if token_idx.numel() == 0:
279+ continue
280+ current_state = flat_states[token_idx]
281+ gate, up = F.linear(
282+ current_state, gate_up_proj[expert_idx],
283+ ).chunk(2, dim=-1)
284+ current_state = F.silu(gate) * up
285+ current_state = F.linear(current_state, down_proj[expert_idx])
286+ current_state = current_state * topk_weights[token_idx, expert_slots, None]
287+ final_hidden_states.index_add_(0, token_idx, current_state)
288+ output = final_hidden_states.view(batch_size, seq_len, hidden_size)
289+ if self.shared_experts is not None:
290+ output = output + self.shared_experts(hidden_states)
291+ return output
292+ 
293+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
294+ """Route tokens, run experts, and combine weighted outputs."""
295+ if self.router_type == "glm_moe_dsa":
296+ return self._forward_glm_moe_dsa(hidden_states)
297+ return self._forward_softmax(hidden_states)
298+ 
299+ 
300+class GLM5SharedExperts(nn.Module):
301+ """Shared GLM-MoE-DSA experts."""
302+ 
303+ def __init__(self, hidden_size: int, intermediate_size: int) -> None:
304+ """Initialize shared expert projections."""
305+ super().__init__()
306+ self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
307+ self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
308+ self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False)
309+ 
310+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
311+ """Run the shared expert MLP."""
312+ hidden_states = F.silu(self.gate_proj(hidden_states)) * self.up_proj(
313+ hidden_states,
314+ )
315+ return self.down_proj(hidden_states)
316+ 
317+ 
318+class GLM5MoERouter(nn.Module):
319+ """Official GLM-MoE-DSA router parameters."""
320+ 
321+ def __init__(self, hidden_size: int, num_experts: int) -> None:
322+ """Initialize router score parameters."""
323+ super().__init__()
324+ self.weight = nn.Parameter(torch.empty(num_experts, hidden_size))
325+ self.e_score_correction_bias = nn.Parameter(torch.zeros(num_experts))
326+ nn.init.kaiming_uniform_(self.weight, a=5 ** 0.5)
327+ 
328+ 
329+__all__ = ["GLM5MoE", "GLM5MoEExperts"]
@@ -0,0 +1,182 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+"""GLM5 parallelization."""
16+import logging
17+ 
18+import torch
19+ 
20+from hyper_parallel import (
21+ ContextParallel,
22+ DSAIndexerContextParallel,
23+ DSASparseAttentionContextParallel,
24+ fully_shard,
25+)
26+from hyper_parallel.core.activation_checkpoint import checkpoint_wrapper
27+from hyper_parallel.core.expert_parallel.expert_parallel import ExpertParallel
28+from hyper_parallel.core.fully_shard.utils import MixedPrecisionPolicy
29+from hyper_parallel.models.glm5.model import GLM5ForCausalLM
30+ 
31+logger = logging.getLogger(__name__)
32+ 
33+ 
34+def _apply_ac(model, cfg) -> None:
35+ """Apply activation checkpointing to GLM5 decoder layers."""
36+ ac_mode = getattr(cfg.train.gradient_checkpointing, "activation_checkpoint", "off")
37+ if ac_mode in ("off", "none", None, False, ""):
38+ return
39+ if not hasattr(model, "layers"):
40+ logger.warning("AC enabled but GLM5 model has no .layers; skipping.")
41+ return
42+ layers = list(model.layers)
43+ for i, layer in enumerate(layers):
44+ model.layers[i] = checkpoint_wrapper(layer)
45+ logger.info_rank0("AC applied to %d GLM5 layers (mode=%s)", len(layers), ac_mode)
46+ 
47+ 
48+def _get_mesh(mesh, name):
49+ try:
50+ return mesh[name]
51+ except (KeyError, TypeError):
52+ return None
53+ 
54+ 
55+def _validate_tp_support(cfg) -> None:
56+ tp = getattr(cfg.train.accelerator, "tp", 1)
57+ if tp > 1:
58+ raise NotImplementedError(
59+ "GLM5 TP is not supported yet. Set train.accelerator.tp=1 "
60+ "or add GLM5 TP parallelization before enabling tensor parallel."
61+ )
62+ 
63+ 
64+def _apply_cp(model, mesh, cfg) -> None:
65+ """Apply context parallel styles to GLM5 attention modules."""
66+ cp = getattr(cfg.train.accelerator, "cp", 1)
67+ if cp <= 1:
68+ return
69+ cp_mesh = _get_mesh(mesh, "cp")
70+ if cp_mesh is None:
71+ raise ValueError("GLM5 CP requested but mesh has no 'cp' dimension.")
72+ cp_size = cp_mesh.size()
73+ cp_rank = mesh.get_local_rank("cp")
74+ setattr(model, "_cp_size", cp_size)
75+ setattr(model, "_cp_rank", cp_rank)
76+ dsa_layers = 0
77+ dense_layers = 0
78+ for layer in list(model.layers):
79+ if layer.dsa_indexer is not None:
80+ DSAIndexerContextParallel(
81+ layout="BSND",
82+ weights_index=None,
83+ use_local_output=True,
84+ ).apply(layer.dsa_indexer.boundary, cp_mesh)
85+ DSASparseAttentionContextParallel(
86+ layout="BSND",
87+ query_rope_index=None,
88+ key_rope_index=None,
89+ query_rope_kwarg_name=None,
90+ key_rope_kwarg_name=None,
91+ use_local_output=True,
92+ ).apply(layer.self_attn.sparse_attention_core, cp_mesh)
93+ dsa_layers += 1
94+ else:
95+ setattr(layer.self_attn.attention_core, "_cp_size", cp_size)
96+ setattr(layer.self_attn.attention_core, "_cp_rank", cp_rank)
97+ ContextParallel(
98+ seq_dim=1, head_dim=2, ulysses_degree=1,
99+ ).apply(layer.self_attn.attention_core, cp_mesh)
100+ dense_layers += 1
101+ logger.info_rank0(
102+ "CP applied to GLM5 attention cores: dense=%d dsa=%d",
103+ dense_layers,
104+ dsa_layers,
105+ )
106+ 
107+ 
108+def _apply_ep(model, mesh, cfg) -> None:
109+ """Apply expert parallelism to GLM5 MoE experts."""
110+ ep = getattr(cfg.train.accelerator, "ep", 1)
111+ if ep <= 1:
112+ return
113+ ep_mesh = _get_mesh(mesh, "ep")
114+ if ep_mesh is None:
115+ raise ValueError("GLM5 EP requested but mesh has no 'ep' dimension.")
116+ moe_layers = [layer for layer in model.layers if layer.layer_type == "moe"]
117+ if not moe_layers:
118+ raise ValueError("GLM5 EP requires at least one MoE decoder layer.")
119+ for layer in moe_layers:
120+ ExpertParallel().apply(layer.mlp.experts, ep_mesh)
121+ logger.info_rank0("EP applied to %d GLM5 MoE expert modules", len(moe_layers))
122+ 
123+ 
124+def _apply_fsdp(model, mesh, cfg) -> None:
125+ """Apply FSDP to GLM5 when a data-shard mesh is active."""
126+ dp_mesh = _get_mesh(mesh, "fsdp")
127+ if dp_mesh is None:
128+ dp_mesh = _get_mesh(mesh, "dp_shard")
129+ accelerator = cfg.train.accelerator
130+ dp_size = dp_mesh.size() if dp_mesh is not None else 1
131+ if dp_size <= 1:
132+ logger.info_rank0(
133+ "FSDP skipped for GLM5 because dp_shard/fsdp size is one"
134+ )
135+ return
136+ 
137+ fsdp_kwargs = {
138+ "mesh": dp_mesh,
139+ "reshard_after_forward": getattr(
140+ accelerator, "reshard_after_forward", True,
141+ ),
142+ "comm_fusion": getattr(accelerator, "comm_fusion", True),
143+ }
144+ 
145+ mp_cfg = getattr(cfg.train, "mixed_precision", None)
146+ if mp_cfg is not None and getattr(mp_cfg, "enabled", False):
147+ dtype_map = {
148+ "bfloat16": torch.bfloat16,
149+ "bf16": torch.bfloat16,
150+ "float16": torch.float16,
151+ "fp16": torch.float16,
152+ "float32": torch.float32,
153+ "fp32": torch.float32,
154+ }
155+ param_dtype = dtype_map.get(getattr(mp_cfg, "param_dtype", "bfloat16"))
156+ reduce_dtype = dtype_map.get(getattr(mp_cfg, "reduce_dtype", "float32"))
157+ output_dtype_str = getattr(mp_cfg, "output_dtype", None)
158+ output_dtype = dtype_map.get(output_dtype_str) if output_dtype_str else None
159+ fsdp_kwargs["mp_policy"] = MixedPrecisionPolicy(
160+ param_dtype=param_dtype,
161+ reduce_dtype=reduce_dtype,
162+ output_dtype=output_dtype,
163+ )
164+ 
165+ if hasattr(model, "layers"):
166+ for layer in list(model.layers):
167+ fully_shard(layer, **fsdp_kwargs)
168+ fully_shard(model, **fsdp_kwargs)
169+ logger.info_rank0("FSDP applied to GLM5")
170+ 
171+ 
172+def parallelize_glm5(model: GLM5ForCausalLM, mesh, cfg) -> GLM5ForCausalLM:
173+ """Apply EP, CP, AC and FSDP to the GLM5 model."""
174+ _validate_tp_support(cfg)
175+ _apply_ep(model, mesh, cfg)
176+ _apply_cp(model, mesh, cfg)
177+ _apply_ac(model, cfg)
178+ _apply_fsdp(model, mesh, cfg)
179+ return model
180+ 
181+ 
182+__all__ = ["parallelize_glm5"]
@@ -0,0 +1,48 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+"""State-dict adapter for GLM5 dense Phase-1 checkpoints."""
16+from typing import Dict, Optional
17+ 
18+import torch
19+ 
20+from hyper_parallel.models.glm5.checkpoint import load_hf_glm5_state_dict
21+ 
22+ 
23+class GLM5StateDictAdapter:
24+ """HF ↔ hyper state-dict adapter for the GLM5 dense training skeleton."""
25+ 
26+ def load_hf_state_dict(
27+ self,
28+ weights_path: str,
29+ model_config,
30+ dtype: Optional[torch.dtype] = None,
31+ ) -> Dict[str, torch.Tensor]:
32+ return load_hf_glm5_state_dict(
33+ weights_path,
34+ num_hidden_layers=model_config.num_hidden_layers,
35+ num_experts=getattr(model_config, "num_experts", None),
36+ dtype=dtype,
37+ )
38+ 
39+ def save_hf_state_dict(
40+ self,
41+ state_dict: Dict[str, torch.Tensor],
42+ model_config,
43+ ) -> Dict[str, torch.Tensor]:
44+ del model_config
45+ return dict(state_dict)
46+ 
47+ 
48+__all__ = ["GLM5StateDictAdapter"]
@@ -37,6 +37,8 @@ class ModelSpec:
37 When ``None``, hyper's DTensor-aware ``clip_grad_norm_`` is used.37 When ``None``, hyper's DTensor-aware ``clip_grad_norm_`` is used.
38 pipelining_fn: Optional PP setup ``(model, mesh, cfg) -> (schedule, stages)``.38 pipelining_fn: Optional PP setup ``(model, mesh, cfg) -> (schedule, stages)``.
39 state_dict_adapter: Optional class for HF ↔ hyper weight key translation.39 state_dict_adapter: Optional class for HF ↔ hyper weight key translation.
40+ prepare_batch_fn: Optional model-specific transform ``(batch, model) -> batch``
41+ applied before token counting and forward.
40 42 
41 Example:43 Example:
42 Standard transformer (uses defaults, ~15 lines to register)::44 Standard transformer (uses defaults, ~15 lines to register)::
@@ -66,3 +68,4 @@ class ModelSpec:
66 clip_grad_fn: Optional[Callable] = None68 clip_grad_fn: Optional[Callable] = None
67 pipelining_fn: Optional[Callable] = None69 pipelining_fn: Optional[Callable] = None
68 state_dict_adapter: Optional[Type] = None70 state_dict_adapter: Optional[Type] = None
71+ prepare_batch_fn: Optional[Callable] = None
@@ -795,10 +795,10 @@ class BaseTrainer:
795 self.profiler_callback.on_train_begin(self.state)795 self.profiler_callback.on_train_begin(self.state)
796 self.wandb_callback.on_train_begin(self.state)796 self.wandb_callback.on_train_begin(self.state)
797 self.tensorboard_callback.on_train_begin(self.state)797 self.tensorboard_callback.on_train_begin(self.state)
798- self.progress_callback.on_train_begin(self.state)798+ # Checkpoint runs after log writers are armed and before progress so
799- # Checkpoint runs LAST so resume sees an already-armed TB writer799+ # resumed ``global_step`` is reflected in the tqdm initial position.
800- # (it'll record the load event via dispatch_load_event).
801 self.checkpoint_callback.on_train_begin(self.state)800 self.checkpoint_callback.on_train_begin(self.state)
801+ self.progress_callback.on_train_begin(self.state)
X
Xxuxinglei7月2日

调换顺序的考虑是什么

likedislike
Moy
Moy
7月2日 评论:
802 for cb in self.user_callbacks:802 for cb in self.user_callbacks:
803 cb.on_train_begin(self.state)803 cb.on_train_begin(self.state)
804 804 
@@ -970,6 +970,12 @@ class BaseTrainer:
970 # step counter so checkpoint dirs / log indices match the steps that970 # step counter so checkpoint dirs / log indices match the steps that
971 # actually trained.971 # actually trained.
972 micro_batches = next(data_iterator)972 micro_batches = next(data_iterator)
973+ prepare_batch_fn = getattr(self.spec, "prepare_batch_fn", None)
X
Xxuxinglei7月2日

这块修改是一定必要的吗?

likedislike
Moy
Moy
7月2日 评论:
974+ if prepare_batch_fn is not None:
975+ micro_batches = [
976+ prepare_batch_fn(batch, self.model)
977+ for batch in micro_batches
978+ ]
973 self.state.global_step += 1979 self.state.global_step += 1
974 num_micro = len(micro_batches)980 num_micro = len(micro_batches)
975 981 
@@ -978,13 +984,13 @@ class BaseTrainer:
978 984 
979 token_counts = [count_loss_token(mb) for mb in micro_batches]985 token_counts = [count_loss_token(mb) for mb in micro_batches]
980 local_tokens = sum(token_counts)986 local_tokens = sum(token_counts)
981- if local_tokens == 0:
982- local_tokens = 1
983 global_tokens = local_tokens987 global_tokens = local_tokens
984 if platform.get_world_size() > 1:988 if platform.get_world_size() > 1:
X
Xxuxinglei6月25日

去掉 all_reduce 前的 local_tokens == 0 → 1 clamp 是对的:后续归一化分母一律是被 max(..., 1) 守护的 global_tokens,不存在除零;多卡时把真实的 0 送进 all_reduce,让全局 token 分母变为真实值(原来会把"整段被 -100 mask 的 rank"虚增 1),这正是 GLM5 CP 空 label shard 需要的。对 Qwen 常用的无 -100 dummy 数据是 no-op。

建议在 PR 描述里说明这是对多卡全局 token 分母的修正、且对现有 Qwen 对齐结果无影响,方便导师确认不是回归。

likedislike
Moy
Moy
6月26日 评论:
985 gt = platform.full((1,), local_tokens).to(self.device)989 gt = platform.full((1,), local_tokens).to(self.device)
986 platform.all_reduce(gt, self._dp_group_info)990 platform.all_reduce(gt, self._dp_group_info)
987 global_tokens = max(int(gt.item()), 1)991 global_tokens = max(int(gt.item()), 1)
992+ else:
993+ global_tokens = max(global_tokens, 1)
988 # Expose for callbacks (e.g. LoggingCallback throughput).994 # Expose for callbacks (e.g. LoggingCallback throughput).
989 self._last_global_tokens = global_tokens995 self._last_global_tokens = global_tokens
990 996 
@@ -160,7 +160,10 @@ class LoggingCallback(Callback):
160 160 
161 def __init__(self, trainer: "BaseTrainer") -> None:161 def __init__(self, trainer: "BaseTrainer") -> None:
162 super().__init__(trainer)162 super().__init__(trainer)
163- log_cfg = getattr(trainer.args, 'logging', None)163+ train_cfg = getattr(trainer.args, 'train', None)
X
Xxuxinglei6月25日

把 config 查找从 args.logging 改成 args.train.logging 是对的:本仓库 config 是严格三层 schema(顶层只允许 model / data / train,顶层 logging 是 parser 拒绝的 forbidden key),所以旧代码 getattr(trainer.args, 'logging') 永远拿到 None,现有 Qwen 配置里的 train.logging / train.checkpoint 其实一直被静默忽略走默认值。新代码让它们真正生效。

两点提醒:(1) 这意味着会改变现有 Qwen 的运行行为(例如 save_steps / log_steps 从被忽略变成按 yaml 生效),回归 Qwen 时请确认 checkpoint / 日志频率符合预期,建议在 PR 描述里点明这是跨模型的行为变更;(2) if log_cfg is None: log_cfg = getattr(trainer.args, 'logging', None) 这个 fallback 分支因顶层 logging 被禁是 dead code,留着无害,但容易让人误以为顶层 logging 仍被支持。

likedislike
Moy
Moy
6月26日 评论:
164+ log_cfg = getattr(train_cfg, 'logging', None)
165+ if log_cfg is None:
166+ log_cfg = getattr(trainer.args, 'logging', None)
164 self.log_steps = getattr(log_cfg, 'log_steps', 10) if log_cfg else 10167 self.log_steps = getattr(log_cfg, 'log_steps', 10) if log_cfg else 10
165 self.report_global_loss = (168 self.report_global_loss = (
166 getattr(log_cfg, 'report_global_loss', False) if log_cfg else False169 getattr(log_cfg, 'report_global_loss', False) if log_cfg else False
@@ -261,7 +264,10 @@ class CheckpointCallback(Callback):
261 264 
262 def __init__(self, trainer: "BaseTrainer") -> None:265 def __init__(self, trainer: "BaseTrainer") -> None:
263 super().__init__(trainer)266 super().__init__(trainer)
264- ckpt_cfg = getattr(trainer.args, 'checkpoint', None)267+ train_cfg = getattr(trainer.args, 'train', None)
268+ ckpt_cfg = getattr(train_cfg, 'checkpoint', None)
269+ if ckpt_cfg is None:
270+ ckpt_cfg = getattr(trainer.args, 'checkpoint', None)
265 self.save_steps = getattr(ckpt_cfg, 'save_steps', 0) if ckpt_cfg else 0271 self.save_steps = getattr(ckpt_cfg, 'save_steps', 0) if ckpt_cfg else 0
266 self.output_dir = (272 self.output_dir = (
267 getattr(ckpt_cfg, 'output_dir', 'outputs') if ckpt_cfg else 'outputs'273 getattr(ckpt_cfg, 'output_dir', 'outputs') if ckpt_cfg else 'outputs'
@@ -477,7 +483,10 @@ class SafetensorsExportCallback(Callback):
477 483 
478 def __init__(self, trainer: "BaseTrainer") -> None:484 def __init__(self, trainer: "BaseTrainer") -> None:
479 super().__init__(trainer)485 super().__init__(trainer)
480- ckpt_cfg = getattr(trainer.args, 'checkpoint', None)486+ train_cfg = getattr(trainer.args, 'train', None)
487+ ckpt_cfg = getattr(train_cfg, 'checkpoint', None)
488+ if ckpt_cfg is None:
489+ ckpt_cfg = getattr(trainer.args, 'checkpoint', None)
481 self.enabled = getattr(ckpt_cfg, 'save_hf_weights', False) if ckpt_cfg else False490 self.enabled = getattr(ckpt_cfg, 'save_hf_weights', False) if ckpt_cfg else False
482 self.save_steps = getattr(ckpt_cfg, 'save_steps', 0) if ckpt_cfg else 0491 self.save_steps = getattr(ckpt_cfg, 'save_steps', 0) if ckpt_cfg else 0
483 self.output_dir = getattr(ckpt_cfg, 'output_dir', 'outputs') if ckpt_cfg else 'outputs'492 self.output_dir = getattr(ckpt_cfg, 'output_dir', 'outputs') if ckpt_cfg else 'outputs'
@@ -0,0 +1,373 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+"""GLM5 Trainer integration tests."""
16+import os
17+from contextlib import nullcontext
18+from types import SimpleNamespace
19+ 
20+import pytest
21+import torch
22+ 
23+from hyper_parallel.models.glm5 import GLM5Config, GLM5ForCausalLM, prepare_glm5_batch
24+from hyper_parallel.models.glm5.attention import GLM5AttentionCore
25+from hyper_parallel.models.glm5.parallelize import parallelize_glm5
26+from hyper_parallel.models.spec import get_spec
27+from hyper_parallel.trainer import base as trainer_base
28+from hyper_parallel.trainer.base import BaseTrainer, TrainerState
29+from hyper_parallel.trainer.callbacks import base as callback_base
30+from hyper_parallel.trainer.callbacks.base import CheckpointCallback
31+from hyper_parallel.trainer.utils.discovery import discover_model_spec
32+ 
33+ 
34+class _RecordingDataloader:
35+ """Minimal stateful dataloader stand-in."""
36+ 
37+ def __init__(self, position: int = 0) -> None:
38+ """Initialize the in-memory dataloader position."""
39+ self.position = position
40+ 
41+ def state_dict(self) -> dict:
42+ """Return a resumable dataloader state."""
43+ return {"position": self.position}
44+ 
45+ def load_state_dict(self, state: dict) -> None:
46+ """Restore the dataloader position."""
47+ self.position = state["position"]
48+ 
49+ 
50+def _tiny_config_kwargs() -> dict:
51+ """Return the small GLM5 config used by Trainer tests."""
52+ return {
53+ "vocab_size": 32,
54+ "hidden_size": 16,
55+ "intermediate_size": 32,
56+ "num_hidden_layers": 2,
57+ "num_attention_heads": 4,
58+ "num_key_value_heads": 2,
59+ "head_dim": 4,
60+ "num_dense_layers": 2,
61+ "max_position_embeddings": 64,
62+ }
63+ 
64+ 
65+def _build_tiny_model() -> GLM5ForCausalLM:
66+ """Build a tiny dense GLM5 model for trainer-path checks."""
67+ return GLM5ForCausalLM(GLM5Config(**_tiny_config_kwargs()))
68+ 
69+ 
70+def _causal_lm_loss(logits: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:
71+ """Compute the shifted CausalLM loss convention used by Trainer baselines."""
72+ shift_logits = logits[..., :-1, :].contiguous().float()
73+ shift_labels = labels[..., 1:].contiguous()
74+ return torch.nn.functional.cross_entropy(
75+ shift_logits.view(-1, shift_logits.size(-1)),
76+ shift_labels.view(-1),
77+ ignore_index=-100,
78+ )
79+ 
80+ 
81+def test_glm5_discovery_and_build():
82+ """
83+ Feature: GLM5 Trainer model discovery
84+ Description: Discover GLM5 ModelSpec and build a tiny model from config overrides.
85+ Expectation: The registered builder creates a GLM5 model with requested values.
86+ """
87+ discover_model_spec("glm5")
88+ spec = get_spec("glm5")
89+ cfg = SimpleNamespace(
90+ model=SimpleNamespace(
91+ vocab_size=None,
92+ hidden_size=None,
93+ intermediate_size=None,
94+ num_hidden_layers=None,
95+ num_attention_heads=None,
96+ num_key_value_heads=None,
97+ max_position_embeddings=None,
98+ config_overrides=_tiny_config_kwargs(),
99+ )
100+ )
101+ 
102+ model = spec.build_model_fn(cfg)
103+ 
104+ assert isinstance(model, GLM5ForCausalLM)
105+ assert model.config.vocab_size == 32
106+ assert model.config.num_hidden_layers == 2
107+ 
108+ 
109+def test_glm5_trainer_step_applies_prepare_batch_fn(monkeypatch):
110+ """
111+ Feature: GLM5 Trainer train_step integration
112+ Description: Run one BaseTrainer.train_step through the GLM5 prepare-batch hook.
113+ Expectation: CP-prepared labels/positions are used and one optimizer step runs.
114+ """
115+ torch.manual_seed(0)
116+ prepared = {}
117+ 
118+ def _record_prepare_batch(batch: dict, model: GLM5ForCausalLM) -> dict:
119+ prepared_batch = prepare_glm5_batch(batch, model)
120+ prepared["position_ids"] = prepared_batch["position_ids"]
121+ prepared["labels"] = prepared_batch["labels"]
122+ return prepared_batch
123+ 
124+ spec = SimpleNamespace(prepare_batch_fn=_record_prepare_batch, clip_grad_fn=None)
125+ monkeypatch.setattr(trainer_base, "get_spec", lambda _: spec)
126+ monkeypatch.setattr(trainer_base.platform, "get_world_size", lambda: 1)
127+ monkeypatch.setattr(trainer_base, "hsdp_sync_stream", lambda: None)
128+ args = SimpleNamespace(
129+ model=SimpleNamespace(name="glm5"),
130+ train=SimpleNamespace(
131+ max_steps=2,
132+ optimizer=SimpleNamespace(
133+ loss_aggregation="token_weighted",
134+ max_grad_norm=1.0,
135+ ),
136+ ),
137+ )
138+ trainer = BaseTrainer(args)
139+ model = _build_tiny_model()
140+ setattr(model, "_cp_size", 2)
141+ setattr(model, "_cp_rank", 0)
142+ trainer.model = model
143+ trainer.device = torch.device("cpu")
144+ trainer.model_fwd_context = nullcontext()
145+ trainer.model_bwd_context = nullcontext()
146+ trainer.optimizer = torch.optim.SGD(model.parameters(), lr=1e-3)
147+ trainer.lr_scheduler = None
148+ trainer.parallel_dims = SimpleNamespace(dp_size=1)
149+ setattr(trainer, "_dp_group_info", SimpleNamespace(rank_size=1))
150+ trainer.on_substep_end = lambda: None
151+ trainer.on_pre_optimizer_step = lambda grad_norm=None: None
152+ with torch.no_grad():
153+ first_weight = model.model.embed_tokens.weight.clone()
154+ input_ids = torch.tensor([[1, 2, 3, 4, 5]])
155+ batch = {"input_ids": input_ids, "labels": input_ids.clone()}
156+ 
157+ metrics = trainer.train_step(iter([[batch]]))
158+ 
159+ assert trainer.state.global_step == 1
160+ assert torch.isfinite(torch.tensor(metrics["loss"]))
161+ assert prepared["position_ids"].tolist() == [0, 1, 2]
162+ assert prepared["labels"].tolist() == [[2, 3, 4]]
163+ assert getattr(trainer, "_last_global_tokens") == 3
164+ assert not torch.equal(model.model.embed_tokens.weight, first_weight)
165+ 
166+ 
167+def test_glm5_loss_matches_causal_lm_shifted_ce():
168+ """
169+ Feature: GLM5 Trainer loss semantics
170+ Description: Compare GLM5 loss with shifted CausalLM loss.
171+ Expectation: GLM5 loss matches the Transformers/LLaMAFactory label shift.
172+ """
173+ torch.manual_seed(0)
174+ model = _build_tiny_model()
175+ input_ids = torch.tensor([
176+ [0, 0, 3, 4, 5, 6],
177+ [0, 7, 8, 9, 10, 11],
178+ ])
179+ attention_mask = torch.tensor([
180+ [0, 0, 1, 1, 1, 1],
181+ [0, 1, 1, 1, 1, 1],
182+ ])
183+ position_ids = attention_mask.cumsum(dim=-1).sub(1).clamp_min(0)
184+ labels = input_ids.masked_fill(attention_mask == 0, -100)
185+ 
186+ output = model(
187+ input_ids=input_ids,
188+ attention_mask=attention_mask,
189+ position_ids=position_ids,
190+ labels=labels,
191+ )
192+ 
193+ assert torch.allclose(
194+ output["loss"],
195+ _causal_lm_loss(output["logits"], labels),
196+ atol=1e-6,
197+ rtol=0,
198+ )
199+ 
200+ 
201+def test_glm5_cp_batch_shards_inputs_and_shifted_labels():
202+ """
203+ Feature: GLM5 CP batch preparation
204+ Description: Shard input tokens and labels across two CP ranks.
205+ Expectation: Shifted labels and position ids remain globally aligned.
206+ """
207+ batch = {
208+ "input_ids": torch.tensor([[10, 11, 12, 13, 14]]),
209+ "labels": torch.tensor([[10, 11, 12, 13, 14]]),
210+ }
211+ rank0_model = SimpleNamespace(**{"_cp_size": 2, "_cp_rank": 0})
212+ rank1_model = SimpleNamespace(**{"_cp_size": 2, "_cp_rank": 1})
213+ 
214+ rank0 = prepare_glm5_batch(batch, rank0_model)
215+ rank1 = prepare_glm5_batch(batch, rank1_model)
216+ 
217+ assert rank0["input_ids"].tolist() == [[10, 11, 12]]
218+ assert rank0["labels"].tolist() == [[11, 12, 13]]
219+ assert rank0["position_ids"].tolist() == [0, 1, 2]
220+ assert rank1["input_ids"].tolist() == [[13, 14, 0]]
221+ assert rank1["labels"].tolist() == [[14, -100, -100]]
222+ assert rank1["position_ids"].tolist() == [3, 4, 5]
223+ 
224+ 
225+def test_glm5_cp_batch_shards_4d_attention_mask():
226+ """
227+ Feature: GLM5 CP batch preparation
228+ Description: Shard a 4D additive attention mask across two CP ranks.
229+ Expectation: Local mask query/key dimensions match the local token shard.
230+ """
231+ batch = {
232+ "input_ids": torch.tensor([[10, 11, 12, 13]]),
233+ "attention_mask": torch.zeros(1, 1, 4, 4),
234+ }
235+ rank1_model = SimpleNamespace(**{"_cp_size": 2, "_cp_rank": 1})
236+ 
237+ rank1 = prepare_glm5_batch(batch, rank1_model)
238+ 
239+ assert rank1["input_ids"].tolist() == [[12, 13]]
240+ assert rank1["attention_mask"].shape == (1, 1, 2, 2)
241+ 
242+ 
243+def test_glm5_attention_core_slices_cp_4d_attention_mask():
244+ """
245+ Feature: GLM5 CP attention core
246+ Description: Run a local core shard with a full 4D additive mask.
247+ Expectation: The mask is sliced to the core-local query/key dimensions.
248+ """
249+ core = GLM5AttentionCore(scale=1.0)
250+ setattr(core, "_cp_size", 2)
251+ setattr(core, "_cp_rank", 1)
252+ query = torch.randn(1, 2, 1, 4)
253+ key = torch.randn(1, 2, 1, 4)
254+ value = torch.randn(1, 2, 1, 4)
255+ attention_mask = torch.zeros(1, 1, 4, 4)
256+ 
257+ output = core(query, key, value, attention_mask=attention_mask)
258+ 
259+ assert output.shape == query.shape
260+ 
261+ 
262+def test_glm5_parallelize_rejects_tp_until_supported():
263+ """
264+ Feature: GLM5 TP guard
265+ Description: Request tensor parallel before GLM5 TP apply is implemented.
266+ Expectation: Parallelization raises NotImplementedError.
267+ """
268+ model = _build_tiny_model()
269+ cfg = SimpleNamespace(
270+ train=SimpleNamespace(
271+ accelerator=SimpleNamespace(tp=2, cp=1, ep=1),
272+ gradient_checkpointing=SimpleNamespace(activation_checkpoint="off"),
273+ ),
274+ )
275+ 
276+ with pytest.raises(NotImplementedError, match="GLM5 TP is not supported"):
277+ parallelize_glm5(model, mesh={}, cfg=cfg)
278+ 
279+ 
280+def test_glm5_checkpoint_callback_round_trip(tmp_path, monkeypatch):
281+ """
282+ Feature: GLM5 checkpoint save and resume
283+ Description: Save a tiny GLM5 training state and restore it.
284+ Expectation: Model, optimizer, scheduler, RNG, dataloader, and step restore.
285+ """
286+ 
287+ def _save_state_dict(state_dict, checkpoint_id, use_collectives=False):
288+ del use_collectives
289+ torch.save(
290+ {
291+ key: value.detach().cpu().clone()
292+ for key, value in state_dict.items()
293+ },
294+ os.path.join(checkpoint_id, "model_state.pt"),
295+ )
296+ 
297+ def _load_state_dict(state_dict, checkpoint_id, use_collectives=False):
298+ del use_collectives
299+ payload = torch.load(
300+ os.path.join(checkpoint_id, "model_state.pt"),
301+ map_location="cpu",
302+ weights_only=True,
303+ )
304+ for key, value in state_dict.items():
305+ value.copy_(payload[key])
306+ 
307+ set_rng_calls = []
308+ monkeypatch.setattr(callback_base, "dcp_save", _save_state_dict)
309+ monkeypatch.setattr(callback_base, "dcp_load", _load_state_dict)
310+ monkeypatch.setattr(callback_base.platform, "get_rank", lambda: 0)
311+ monkeypatch.setattr(
312+ callback_base.platform,
313+ "get_rng_state",
314+ lambda: torch.tensor([1, 2, 3]),
315+ )
316+ monkeypatch.setattr(callback_base.platform, "set_rng_state", set_rng_calls.append)
317+ 
318+ torch.manual_seed(1)
319+ model = _build_tiny_model()
320+ optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
321+ scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lambda _: 1.0)
322+ dataloader = _RecordingDataloader(position=12)
323+ checkpoint_cfg = SimpleNamespace(
324+ output_dir=str(tmp_path),
325+ save_steps=1,
326+ save_async=False,
327+ load_path=None,
328+ save_hf_weights=False,
329+ )
330+ trainer = SimpleNamespace(
331+ args=SimpleNamespace(train=SimpleNamespace(checkpoint=checkpoint_cfg)),
332+ model=model,
333+ optimizer=optimizer,
334+ lr_scheduler=scheduler,
335+ train_dataloader=dataloader,
336+ dispatch_save_event=lambda *_args, **_kwargs: None,
337+ dispatch_load_event=lambda *_args, **_kwargs: None,
338+ )
339+ callback = CheckpointCallback(trainer)
340+ 
341+ input_ids = torch.randint(0, model.config.vocab_size, (2, 8))
342+ loss = model(input_ids=input_ids, labels=input_ids)["loss"]
343+ loss.backward()
344+ optimizer.step()
345+ scheduler.step()
346+ optimizer.zero_grad()
347+ expected_state = {
348+ key: value.detach().clone() for key, value in model.state_dict().items()
349+ }
350+ 
351+ state = TrainerState(max_steps=10)
352+ state.global_step = 3
353+ state.epoch = 1
354+ callback.on_step_end(state, loss=0.0, grad_norm=0.0)
355+ save_dir = tmp_path / "step_3"
356+ 
357+ with torch.no_grad():
358+ for param in model.parameters():
359+ param.zero_()
360+ dataloader.position = 0
361+ state.global_step = 0
362+ state.epoch = 0
363+ checkpoint_cfg.load_path = str(save_dir)
364+ callback.load_path = str(save_dir)
365+ 
366+ callback.on_train_begin(state)
367+ 
368+ for key, value in model.state_dict().items():
369+ assert torch.allclose(value, expected_state[key])
370+ assert state.global_step == 3
371+ assert state.epoch == 1
372+ assert dataloader.position == 12
373+ assert set_rng_calls and torch.equal(set_rng_calls[-1], torch.tensor([1, 2, 3]))
@@ -0,0 +1,89 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+"""Checkpoint callback configuration tests."""
16+from types import SimpleNamespace
17+ 
18+from hyper_parallel.trainer.callbacks.base import (
19+ CheckpointCallback,
20+ SafetensorsExportCallback,
21+)
22+ 
23+ 
24+def _trainer(args):
25+ return SimpleNamespace(args=args)
26+ 
27+ 
28+def test_checkpoint_callback_reads_nested_train_checkpoint():
29+ """
30+ Feature: nested checkpoint config
31+ Description: CheckpointCallback reads checkpoint settings from train.checkpoint.
32+ Expectation: save/load fields match the nested config values.
33+ """
34+ ckpt = SimpleNamespace(
35+ save_steps=7,
36+ output_dir="/tmp/glm5_ckpt",
37+ load_path="/tmp/glm5_ckpt/step_7",
38+ save_async=True,
39+ )
40+ args = SimpleNamespace(train=SimpleNamespace(checkpoint=ckpt))
41+ 
42+ callback = CheckpointCallback(_trainer(args))
43+ 
44+ assert callback.save_steps == 7
45+ assert callback.output_dir == "/tmp/glm5_ckpt"
46+ assert callback.load_path == "/tmp/glm5_ckpt/step_7"
47+ assert callback.save_async is True
48+ 
49+ 
50+def test_checkpoint_callback_keeps_top_level_fallback():
51+ """
52+ Feature: checkpoint config fallback
53+ Description: Legacy top-level checkpoint config remains supported.
54+ Expectation: callback uses args.checkpoint when train.checkpoint is absent.
55+ """
56+ ckpt = SimpleNamespace(
57+ save_steps=3,
58+ output_dir="/tmp/legacy_ckpt",
59+ load_path=None,
60+ save_async=False,
61+ )
62+ args = SimpleNamespace(checkpoint=ckpt)
63+ 
64+ callback = CheckpointCallback(_trainer(args))
65+ 
66+ assert callback.save_steps == 3
67+ assert callback.output_dir == "/tmp/legacy_ckpt"
68+ assert callback.load_path is None
69+ assert callback.save_async is False
70+ 
71+ 
72+def test_hf_export_callback_reads_nested_train_checkpoint():
73+ """
74+ Feature: nested HF export config
75+ Description: SafetensorsExportCallback reads settings from train.checkpoint.
76+ Expectation: export controls match the nested checkpoint config.
77+ """
78+ ckpt = SimpleNamespace(
79+ save_hf_weights=True,
80+ save_steps=11,
81+ output_dir="/tmp/glm5_hf",
82+ )
83+ args = SimpleNamespace(train=SimpleNamespace(checkpoint=ckpt))
84+ 
85+ callback = SafetensorsExportCallback(_trainer(args))
86+ 
87+ assert callback.enabled is True
88+ assert callback.save_steps == 11
89+ assert callback.output_dir == "/tmp/glm5_hf"
@@ -0,0 +1,76 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+"""Logging callback configuration tests."""
16+from types import SimpleNamespace
17+ 
18+from hyper_parallel.trainer.callbacks.base import LoggingCallback
19+ 
20+ 
21+def _trainer(args):
22+ return SimpleNamespace(args=args, lr_scheduler=None)
23+ 
24+ 
25+def _args_with_train_logging(logging_cfg):
26+ train = SimpleNamespace(
27+ logging=logging_cfg,
28+ global_batch_size=4,
29+ )
30+ data = SimpleNamespace(max_seq_len=128)
31+ return SimpleNamespace(train=train, data=data)
32+ 
33+ 
34+def test_logging_callback_reads_nested_train_logging():
35+ """
36+ Feature: nested logging config
37+ Description: LoggingCallback reads logging settings from train.logging.
38+ Expectation: per-step logging config is applied from the nested config.
39+ """
40+ logging_cfg = SimpleNamespace(
41+ log_steps=1,
42+ report_global_loss=True,
43+ report_throughput=False,
44+ model_flops_per_token=None,
45+ peak_tflops=None,
46+ )
47+ 
48+ callback = LoggingCallback(_trainer(_args_with_train_logging(logging_cfg)))
49+ 
50+ assert callback.log_steps == 1
51+ assert callback.report_global_loss is True
52+ assert callback.report_throughput is False
53+ 
54+ 
55+def test_logging_callback_keeps_top_level_fallback():
56+ """
57+ Feature: logging config fallback
58+ Description: Legacy top-level logging config remains supported.
59+ Expectation: callback uses args.logging when train.logging is absent.
60+ """
61+ logging_cfg = SimpleNamespace(
62+ log_steps=5,
63+ report_global_loss=False,
64+ report_throughput=True,
65+ model_flops_per_token=None,
66+ peak_tflops=None,
67+ )
68+ train = SimpleNamespace(global_batch_size=4)
69+ data = SimpleNamespace(max_seq_len=128)
70+ args = SimpleNamespace(train=train, data=data, logging=logging_cfg)
71+ 
72+ callback = LoggingCallback(_trainer(args))
73+ 
74+ assert callback.log_steps == 5
75+ assert callback.report_global_loss is False
76+ assert callback.report_throughput is True