| @@ -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 | ||
| 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
这些 HF key 经此处原样透传后与 hyper 参数名对不上,落入 建议:补齐各变体 HF→hyper 映射 + expert 打包;至少对 allowlist 之外的结构性 missing 改为 raise(而非 warning);并补一个真正调用 ![]() ![]() | |||
| 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 | + | ||
| 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) | ||
🔵 Low Priority 与 变更行: 触发条件:当 建议:在 ![]() ![]() 不准确? | |||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 430 | + def layers(self): | ||
| 431 | + return self.model.layers | ||
| 432 | + | ||
| 433 | + | ||
| 434 | + def embed_tokens(self): | ||
| 435 | + return self.model.embed_tokens | ||
| 436 | + | ||
| 437 | + | ||
| 438 | + def norm(self): | ||
| 439 | + return self.model.norm | ||
| 440 | + | ||
| 441 | + | ||
| 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] = None | 68 | clip_grad_fn: Optional[Callable] = None |
| 67 | pipelining_fn: Optional[Callable] = None | 69 | pipelining_fn: Optional[Callable] = None |
| 68 | state_dict_adapter: Optional[Type] = None | 70 | 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 writer | 799 | + # 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) | ||
| 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 that | 970 | # 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) | ||
| 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 += 1 | 979 | 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_tokens | 987 | global_tokens = local_tokens |
| 984 | if platform.get_world_size() > 1: | 988 | if platform.get_world_size() > 1: |
X 去掉 all_reduce 前的 建议在 PR 描述里说明这是对多卡全局 token 分母的修正、且对现有 Qwen 对齐结果无影响,方便导师确认不是回归。 ![]() ![]() | |||
| 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_tokens | 995 | 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 把 config 查找从 两点提醒:(1) 这意味着会改变现有 Qwen 的运行行为(例如 ![]() ![]() | |||
| 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 10 | 167 | 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 False | 169 | 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 0 | 271 | 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 False | 490 | 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 0 | 491 | 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 | ||


🔵 Low Priority
对以下所有变更文件进行了全面审查:
建议:无需操作。此审查结果确认未发现正确性、安全性、可靠性或破坏性变更问题。