diff --git a/.gitignore b/.gitignore index 0f4bfe2800..e5c51fda61 100644 --- a/.gitignore +++ b/.gitignore @@ -118,6 +118,11 @@ data *.pkl.json *.log.json work_dirs/ + +# Local training-analysis artifacts generated by plot_xtuner_losses.py. +/loss_comparison*/ +/examples/v1/scripts/plot_xtuner_losses.py +/tests/scripts/test_plot_xtuner_losses.py work_dir/ # Pytorch diff --git a/.rjobignore b/.rjobignore new file mode 100644 index 0000000000..8c6f1f41e2 --- /dev/null +++ b/.rjobignore @@ -0,0 +1,6 @@ +.git/ +.mypy_cache/ +.third_party/ +work_dirs/ +**/__pycache__/ +*.pyc diff --git a/docs/design/model/glm47_flash_dsa.md b/docs/design/model/glm47_flash_dsa.md new file mode 100644 index 0000000000..e449c4dda9 --- /dev/null +++ b/docs/design/model/glm47_flash_dsa.md @@ -0,0 +1,237 @@ +# GLM-4.7-Flash DSA and IndexShare integration design + +## 1. Objective and invariants + +Convert the pretrained `zai-org/GLM-4.7-Flash` MLA model into a deployable +DeepSeek Sparse Attention model without changing the pretrained dense function +before training. Training is split into two separate jobs: + +1. `dense_indexer_warmup`: keep the original dense MLA path, freeze every + pretrained parameter, and train only newly added indexers against full + causal-attention teachers. +2. `sparse_train`: reload the warm-up checkpoint, run SparseMLA with shared + Top-K indices, and train the complete model. Language-model losses update + the backbone; the indexer KL path updates indexers only. + +The implementation must preserve these invariants: + +- Loading a GLM-4.7-Flash checkpoint may miss only explicitly enumerated new + `indexer.*` parameters and may not ignore any unexpected key. +- `dense_baseline` and `dense_indexer_warmup` produce the same attention output + and logits as the original Hugging Face model before an optimizer update. +- Teacher tensors and indexer inputs are detached from the KL autograd path. +- Top-K selection is always non-differentiable. +- A source indexer is supervised by every attention layer that consumes its + indices, not only by the source layer. +- Dense warm-up never retains an `O(S^2)` tensor for backward. +- The main-attention and indexer RoPE layouts are independently configurable. + +## 2. Model topology + +GLM-4.7-Flash has 47 main decoder layers and one physical MTP layer. The main +stack uses uniform one-in-four IndexShare: + +```text +0F 1S 2S 3S | 4F 5S 6S 7S | ... | 44F 45S 46S +``` + +`F` owns an indexer and `S` reuses the nearest preceding `F` result. This gives +12 main-stack indexers. The physical MTP layer owns and trains a separate +indexer; it must not freeze an uninitialised random indexer. + +The initial indexer shape follows GLM-5.2: + +```text +index_n_heads = 32 +index_head_dim = 128 +index_topk = 2048 +index_topk_freq = 4 +index_skip_topk_offset = 1 +``` + +For source layer `r`, indexer projections are + +```math +q^I_t = W^I_q q^{resid}_t, +\qquad +k^I_s = \operatorname{Norm}(W^I_k h_s), +\qquad +w_t = W^I_w h_t, +``` + +and the score is + +```math +I_{t,s} = \frac{1}{\sqrt{H_I D_I}} +\sum_{j=1}^{H_I} w_{t,j} +\operatorname{ReLU}\left((q^I_{t,j})^\top k^I_s\right). +``` + +Selection and training must implement the same effective scale even if their +kernel APIs divide the head and head-dimension factors at different points. + +## 3. Training modes + +`DSAMLAConfig` exposes an architecture/runtime mode independent of whether the +module is in `train()` or `eval()`: + +```python +attention_mode: Literal[ + "dense_baseline", + "dense_indexer_warmup", + "sparse_train", +] +``` + +The two trainable modes are intentionally run as separate Trainer jobs. The +warm-up job builds an optimizer from indexer parameters only. The sparse job +reloads the hand-off checkpoint and builds a new optimizer from all trainable +parameters. No optimizer or FSDP membership changes in the middle of a job. + +### 3.1 Dense frozen-teacher warm-up + +For attention layer `l`, the full causal teacher is the head-aggregated dense +attention distribution + +```math +p^{(l)}_{t,s} = \frac{1}{H_A}\sum_h +\operatorname{softmax}(A^{(l,h)}_{t,:})_s. +``` + +For source indexer `r`, the full causal student is + +```math +q^{(r)}_{t,s} = \operatorname{softmax}(I^{(r)}_{t,:})_s. +``` + +Let `G_r` be the main layers served by source `r`. Its loss is + +```math +L^{(r)}_{dense} = \frac{1}{|G_r|N} +\sum_{l\in G_r}\sum_t +D_{KL}\left(p^{(l)}_{t,:}\|q^{(r)}_{t,:}\right). +``` + +All pretrained parameters have `requires_grad=False`. Hidden states and +Q-LoRA residuals entering the indexer are detached, so only indexer projection +parameters receive gradients. + +### 3.2 Sparse full-model training + +The source selects `T_t = TopK(I_{t,:}, K)`. Every layer in `G_r` consumes the +same `T_t` through SparseMLA and contributes a selected-support teacher: + +```math +L^{(r)}_{sparse} = \frac{1}{|G_r|N} +\sum_{l\in G_r}\sum_t +D_{KL}\left(p^{(l)}_{t,T_t}\|q^{(r)}_{t,T_t}\right). +``` + +The total training loss is + +```math +L = L_{LM} + \lambda_{indexer}L_{indexer} + L_{MTP}. +``` + +The sparse teacher is stop-gradient but moves with the fully trainable model. +A second frozen dense model is not retained during sparse training. + +## 4. Dense-loss memory design + +A single FP32 score matrix at `S=200000` occupies about 160 GB, so dense score +recompute is query-blocked. `dense_query_block_size` starts at 256 and remains +configurable. + +For each local query block: + +1. Recompute dense attention raw scores and denominator. +2. Recompute dense indexer raw scores and denominator. +3. Accumulate the KL value. +4. Run cuDNN DenseIndexerBackward and accumulate feature gradients. +5. Release both block score tensors immediately. + +The full score tensors are never saved in an autograd context. Instead, a +first-order custom autograd node attaches the accumulated `dQ`, `dK`, and `dW` +buffers to the original indexer features. Its backward multiplies them by the +upstream scalar and returns ordinary PyTorch gradients, allowing FSDP and the +optimizer to handle indexer parameters normally. Higher-order derivatives are +unsupported by design. + +Under sequence parallelism, queries remain sharded, K is gathered as required, +and every query block passes its global `q_causal_offsets`. Loss normalisation +uses the global number of valid query rows. Tests compare one-rank and +multi-rank loss and gradients. + +## 5. Layering and ownership + +- `xtuner/v1/model/moe/glm47_flash.py` + - HF config conversion, model defaults, strict base-to-DSA conversion, MTP + policy, HF key mapping and export. +- `xtuner/v1/module/attention/dsa_mla.py` + - dense/sparse runtime modes, feature production, teacher capture and + per-layer calls into the loss interface. +- `xtuner/v1/module/attention/dsa_topk_sharing.py` + - source resolution, Top-K sharing, group lifetime and checkpoint replay. +- `xtuner/v1/data_proto/sequence_context.py` + - per-call source features, Top-K entries and group-loss accumulator state. +- `xtuner/v1/loss/dsa_indexer_loss.py` + - backend-neutral dense/sparse KL definitions and reduction semantics. +- `xtuner/v1/ops/sparse_mla/` + - cuDNN dense/sparse score recompute and manual first-order gradient ops. + +The math/reduction API belongs under `v1/loss`; cuDNN-specific implementation +details remain under `v1/ops`. + +## 6. Checkpoint and export contract + +The original `glm4_moe_lite` model is first supported as an exact dense XTuner +model. DSA conversion reuses every pretrained tensor and initialises only the +new indexers. Conversion asserts the exact missing-key set and rejects every +unexpected key. + +The warm-up hand-off is an XTuner distributed checkpoint to avoid an extra full +HF save. Final export writes a Hugging Face DSA checkpoint only after its RoPE +layout and weight names are proven compatible with the target Transformers +runtime. If upstream `GlmMoeDsaForCausalLM` cannot express the preserved +GLM-4.7 main-attention RoPE layout, export uses a dedicated +`glm4_moe_lite_dsa` config/runtime rather than silently changing numerics. + +## 7. Verification gates + +1. Dense baseline: per-layer forward, loss and input-gradient parity with HF; + full-model forward parity where hardware permits; byte-equal save/load. +2. Dense loss: PyTorch reference parity for loss and `dQ/dK/dW`; invariant + across query block sizes; causal-offset tests. +3. Frozen warm-up: only indexers are in the optimizer, every teacher gradient is + `None`, teacher parameter checksums do not change, and logits equal the dense + baseline before/after indexer-only updates. +4. IndexShare: each source receives every served layer exactly once; group loss + and gradients match an explicit reference. +5. Sparse path: one Top-K computation per group, shared indices are identical, + sparse loss updates only indexers, and LM loss updates the full model. +6. Parallel path: SP 1/2/8 parity, checkpoint original/replay equivalence, + compiled/eager comparison, activation-offload smoke test. +7. Integration: 4K fixed-sample overfit, 16K/32K smoke training, then 128K and + 200K forward/backward runs on a GPU node. + +For a single teacher, the fixed-sample KL must approach zero. For multi-layer +supervision, the optimum is the mean teacher distribution and the irreducible +minimum is the corresponding Jensen-Shannon divergence, so the acceptance gate +compares against that oracle floor rather than requiring zero. + +## 8. Training defaults + +The initial reproduction schedule is: + +```text +dense warm-up: 1000 steps, indexers only, LR 1e-3 starting point +sparse train: 4000 steps, full model, LR 7.3e-6 starting point +``` + +Training ramps through 4K correctness, 16K/32K smoke tests, and only then +128K/200K. Exact batch size and gradient accumulation are selected from the +available node memory without changing the loss normalisation. + +The existing sparse-backward `-1` padding change is outside this integration's +scope unless it blocks a required correctness test; any such failure is reported +separately rather than silently broadening the patch. diff --git a/examples/v1/config/rl_grpo_gsm8k_async_glm47.py b/examples/v1/config/rl_grpo_gsm8k_async_glm47.py new file mode 100644 index 0000000000..09caec4355 --- /dev/null +++ b/examples/v1/config/rl_grpo_gsm8k_async_glm47.py @@ -0,0 +1,322 @@ +"""Generic GSM8K async GRPO config used by the GLM-4.7 rollout job.""" + +import json +import os +from pathlib import Path + +from xtuner.v1.config import AdamWConfig, FSDPConfig, LRConfig +from xtuner.v1.data_proto.rl_data import RolloutState, SampleParams +from xtuner.v1.datasets.config import DataloaderConfig, DatasetConfig +from xtuner.v1.datasets.rl_tokenize_fn import RLTextTokenizeFnConfig +from xtuner.v1.float8 import Float8Config, ScalingGranularity +from xtuner.v1.model import get_model_config_from_hf +from xtuner.v1.module.mtp import MTPConfig +from xtuner.v1.rl.advantage import GRPOAdvantageConfig +from xtuner.v1.rl.agent_loop import SingleTurnAgentLoopConfig +from xtuner.v1.rl.agent_loop_manager import ( + AgentLoopManagerConfig, + AsyncProduceStrategyConfig, + SamplerConfig, + TaskSpecConfig, +) +from xtuner.v1.rl.evaluator import EvaluatorConfig +from xtuner.v1.rl.judger import GSM8KJudgerConfig +from xtuner.v1.rl.loss import GRPOLossConfig +from xtuner.v1.rl.replay_buffer import AsyncReplayBufferConfig +from xtuner.v1.rl.rollout.worker import RolloutConfig +from xtuner.v1.rl.rollout_is import RolloutImportanceSampling +from xtuner.v1.rl.trainer import WorkerConfig +from xtuner.v1.rl.utils import AcceleratorResourcesConfig, CPUResourcesConfig +from xtuner.v1.train.rl_trainer import RLColocateTrainerConfig + + +work_dir = os.environ["WORK_DIR"] +model_path = os.environ["MODEL_PATH"] +data_path = os.environ["DATA_PATH"] +eval_data_path = os.environ["EVAL_DATA_PATH"] +num_nodes = int(os.environ.get("WORLD_SIZE", "1")) +enable_return_routed_experts = ( + os.environ.get("ENABLE_RETURN_ROUTED_EXPERTS", "1") == "1" +) + +experimental_name = "grpo_gsm8k" +total_train_steps = int(os.environ.get("TOTAL_TRAIN_STEPS", "45")) +evaluate_step = int(os.environ.get("EVALUATE_STEP", str(total_train_steps))) +train_optimizer_steps = int(os.environ.get("TRAIN_OPTIMIZER_STEPS", "1")) +train_batch_size = int(os.environ.get("TRAIN_BATCH_SIZE", "64")) +prompt_repeat_k = int(os.environ.get("PROMPT_REPEAT_K", "5")) +rollout_tp_size = int(os.environ.get("ROLLOUT_TP_SIZE", "1")) +rollout_ep_size = int(os.environ.get("ROLLOUT_EP_SIZE", "4")) +train_ep_size = int(os.environ.get("TRAIN_EP_SIZE", "4")) +max_prompt_length = int(os.environ.get("MAX_PROMPT_LENGTH", "512")) +max_response_length = int(os.environ.get("MAX_RESPONSE_LENGTH", "1024")) +pack_max_length = int(os.environ.get("PACK_MAX_LENGTH", str(32 * 1024))) +enable_evaluate = os.environ.get("ENABLE_EVALUATE", "1") == "1" +enable_fp8 = os.environ.get("FP8", "0") == "1" +enable_mtp = os.environ.get("ENABLE_MTP", "1") == "1" +mtp_num_layers = int(os.environ.get("MTP_NUM_LAYERS", "2")) +over_sample_threshold = float(os.environ.get("OVER_SAMPLE_THRESHOLD", "0.8")) +partial_rollout = os.environ.get("PARTIAL_ROLLOUT", "1") == "1" +max_staleness = int(os.environ.get("MAX_STALENESS", "0")) +tail_batch_trigger_size = int(os.environ.get("TAIL_BATCH_TRIGGER_SIZE", "64")) +enable_group_filter = os.environ.get("ENABLE_GROUP_FILTER", "1") == "1" + +if enable_fp8: + os.environ.setdefault("XTUNER_RL_FP8_QUANTIZE_IN_BF16", "1") + +model_cfg = get_model_config_from_hf(Path(model_path)) +language_model_cfg = getattr(model_cfg, "text_config", model_cfg) +language_model_cfg.float8_cfg = ( + Float8Config( + scaling_granularity_gemm=None, + scaling_granularity_grouped_gemm=ScalingGranularity.TILEWISE, + ) + if enable_fp8 + else None +) +language_model_cfg.ep_size = train_ep_size +language_model_cfg.z_loss_cfg = None +language_model_cfg.balancing_loss_cfg = None +language_model_cfg.freeze_routers = True +if hasattr(language_model_cfg.attention, "sparse_mla_backend"): + language_model_cfg.attention.sparse_mla_backend = os.environ.get( + "SPARSE_MLA_BACKEND", "tilelang" + ) +language_model_cfg.mtp_config = ( + MTPConfig( + num_layers=mtp_num_layers, + loss_scaling_factor=1.0, + detach_mtp_lm_head_weight=True, + detach_mtp_inputs=True, + share_weights=True, + ) + if enable_mtp + else None +) +model_cfg.compile_cfg = None + +with (Path(model_path) / "config.json").open() as config_file: + hf_model_type = json.load(config_file)["model_type"] +default_speculative_algorithm = ( + "qwen3_5_mtp" if hf_model_type == "qwen3_5_moe" else "deepseek_mtp" +) + +resources = AcceleratorResourcesConfig( + accelerator="GPU", + num_workers=8 * num_nodes, + num_cpus_per_worker=12, + cpu_memory_per_worker=int( + os.environ.get("CPU_MEMORY_PER_WORKER_GIB", "32") + ) + * 1024**3, +) + +extra_rollout_config = { + "lmdeploy_backend": "pytorch", + "lmdeploy_log_level": os.environ.get("LMDEPLOY_LOG_LEVEL", "ERROR"), + "lmdeploy_uvicorn_log_level": os.environ.get( + "LMDEPLOY_UVICORN_LOG_LEVEL", "ERROR" + ), +} +if enable_mtp: + extra_rollout_config.update( + lmdeploy_speculative_algorithm=os.environ.get( + "ROLLOUT_SPECULATIVE_ALGORITHM", default_speculative_algorithm + ), + lmdeploy_speculative_num_draft_tokens=int( + os.environ.get("ROLLOUT_SPECULATIVE_NUM_DRAFT_TOKENS", "3") + ), + ) + +rollout_config = RolloutConfig( + env=experimental_name, + device=resources.accelerator, + model_path=model_path, + dtype="bfloat16", + tensor_parallel_size=rollout_tp_size, + expert_parallel_size=rollout_ep_size, + gpu_memory_utilization=float( + os.environ.get("ROLLOUT_GPU_MEMORY_UTILIZATION", "0.8") + ), + context_length=max_response_length + max_prompt_length, + enable_float8=enable_fp8, + skip_load_weights=os.environ.get("SKIP_LOAD_WEIGHTS", "0") == "1", + enable_return_routed_experts=enable_return_routed_experts, + fp32_lm_head=True, + rollout_timeout=36000, + rollout_max_batch_size_per_instance=32 * rollout_ep_size, + extra_rollout_config=extra_rollout_config, +) + +judger_resources = CPUResourcesConfig(num_workers=1, num_cpus_per_worker=1) +train_judger_config = GSM8KJudgerConfig( + judger_name="openai/gsm8k", + cpu_resources=judger_resources, +) +eval_judger_config = GSM8KJudgerConfig( + judger_name="openai/gsm8k", + cpu_resources=judger_resources, +) + +lr_cfg = LRConfig(lr_type="constant", warmup_ratio=0, lr_min=1e-6) +fsdp_cfg = FSDPConfig( + torch_compile=False, + cpu_offload=False, + ep_size=train_ep_size, + fp32_lm_head=True, +) +optim_cfg = AdamWConfig( + lr=1e-6, + betas=(0.9, 0.95), + max_grad_norm=1.0, + weight_decay=0.1, + foreach=False, + skip_grad_norm_threshold=None, + eps=1e-15, +) +loss_cfg = GRPOLossConfig( + policy_loss_cfg={ + "cliprange_high": 0.28, + "cliprange_low": 0.2, + "loss_type": "vanilla", + "clip_ratio_c": 3.0, + "log_prob_diff_min": -20.0, + "log_prob_diff_max": 20.0, + }, + ignore_idx=-100, + use_kl_loss=False, + kl_loss_coef=0.0, + kl_loss_type="low_var_kl", + mode="chunk", + chunk_size=512, + rollout_is=RolloutImportanceSampling( + rollout_is_level="token", + rollout_is_mode="mask", + rollout_is_threshold=(5.0, 0.5), + ), +) +train_worker_cfg = WorkerConfig( + model_cfg=model_cfg, + load_from=model_path, + optim_cfg=optim_cfg, + loss_cfg=loss_cfg, + lr_cfg=lr_cfg, + fsdp_cfg=fsdp_cfg, + sp_size=int(os.environ.get("SP_SIZE", "1")), + optimizer_steps=train_optimizer_steps, + pack_max_length=pack_max_length, +) + +tokenizer_config = RLTextTokenizeFnConfig(max_length=max_prompt_length) +train_dataset = DatasetConfig(name=experimental_name, anno_path=data_path) +train_dataloader_cfg = DataloaderConfig( + dataset_config_list=[ + {"dataset": train_dataset, "tokenize_fn": tokenizer_config} + ], + pack_max_length=pack_max_length, + collator="fake_collator", + pack_level="none", +) +training_sample_params = SampleParams( + max_tokens=max_response_length, + top_k=0, + top_p=1.0, + temperature=1.0, + min_tokens=0, + return_routed_experts=enable_return_routed_experts, +) + + +def group_samples_filter_func(rollout_states: list[RolloutState]) -> bool: + rewards = [ + state.reward["score"] + for state in rollout_states + if state.response_ids is not None + ] + return len(set(rewards)) != 1 + + +produce_strategy_kwargs = { + "over_sample_threshold": over_sample_threshold, + "enable_partial_rollout": partial_rollout, + "max_staleness": max_staleness, + "tail_batch_trigger_size": tail_batch_trigger_size, +} +if enable_group_filter: + produce_strategy_kwargs["is_valid_sample_fn"] = group_samples_filter_func + +agent_loop_manager_cfg = AgentLoopManagerConfig( + tasks=TaskSpecConfig( + task_name="train_task", + agent_loop_config=SingleTurnAgentLoopConfig( + hf_checkpoint=model_path, + sample_params=training_sample_params, + ), + judger_config=train_judger_config, + produce_strategy_config=AsyncProduceStrategyConfig( + **produce_strategy_kwargs + ), + sampler_config=SamplerConfig( + dataloader_cfg=train_dataloader_cfg, + prompt_repeat_k=prompt_repeat_k, + ), + ), +) + +eval_dataset = DatasetConfig( + name=experimental_name, + anno_path=eval_data_path, + sample_ratio=1.0, +) +eval_dataloader_cfg = DataloaderConfig( + dataset_config_list=[ + {"dataset": eval_dataset, "tokenize_fn": tokenizer_config} + ], + pack_max_length=pack_max_length, + collator="fake_collator", + pack_level="none", +) +evaluation_sample_params = SampleParams( + max_tokens=max_response_length, + top_k=1, + top_p=1.0, + temperature=0.0, + min_tokens=0, + return_routed_experts=False, +) +eval_agent_loop_manager_cfg = AgentLoopManagerConfig( + tasks=TaskSpecConfig( + task_name="eval_task", + agent_loop_config=SingleTurnAgentLoopConfig( + hf_checkpoint=model_path, + sample_params=evaluation_sample_params, + ), + judger_config=eval_judger_config, + sampler_config=SamplerConfig( + dataloader_cfg=eval_dataloader_cfg, + prompt_repeat_k=1, + ), + ), +) + +trainer = RLColocateTrainerConfig( + resources=resources, + train_worker_cfg=train_worker_cfg, + rollout_config=rollout_config, + tokenizer_path=model_path, + replay_buffer_config=AsyncReplayBufferConfig(), + agent_loop_manager_cfg=agent_loop_manager_cfg, + eval_agent_loop_manager_cfg=eval_agent_loop_manager_cfg, + evaluator_config=EvaluatorConfig(compute_metric_func=None), + load_from=model_path, + total_train_steps=total_train_steps, + train_batch_size=train_batch_size, + advantage_estimator_config=GRPOAdvantageConfig(eps=1e-8), + enable_evaluate=enable_evaluate, + enable_initial_evaluate=False, + evaluate_step=evaluate_step, + hf_interval=10, + work_dir=work_dir, + seed=int(os.environ.get("SEED", "123")), + debug_rollout=False, +) diff --git a/examples/v1/config/sft_glm47_flash_dsa.py b/examples/v1/config/sft_glm47_flash_dsa.py new file mode 100644 index 0000000000..e30b13ad40 --- /dev/null +++ b/examples/v1/config/sft_glm47_flash_dsa.py @@ -0,0 +1,177 @@ +"""Two-stage GLM-4.7 Flash main/MTP indexer training.""" + +import os + +from xtuner.v1.config import AdamWConfig, FSDPConfig, LRConfig +from xtuner.v1.datasets import LongTextPretrainTokenizeFunctionConfig, OpenaiTokenizeFunctionConfig +from xtuner.v1.datasets.config import DataloaderConfig, DatasetConfig +from xtuner.v1.loss import CELossConfig +from xtuner.v1.model import Glm47FlashDSAConfig +from xtuner.v1.module.attention import DSAIndexerTrainingConfig +from xtuner.v1.train import TrainerConfig +from xtuner.v1.train.trainer import LoadCheckpointConfig + + +def _bool_env(name: str, default: bool = False) -> bool: + return os.environ.get(name, "1" if default else "0").lower() in ("1", "true", "yes", "on") + + +stage = os.environ.get("STAGE", "dense_warmup").lower() +if stage not in ("dense_warmup", "sparse"): + raise ValueError(f"Unsupported STAGE={stage!r}.") + +model_path = os.environ["GLM47_FLASH_MODEL_PATH"] +work_dir = os.environ.get("WORK_DIR", f"work_dirs/glm47_flash_dsa_{stage}") +world_size = int(os.environ.get("WORLD_SIZE", "8")) +ep_size = int(os.environ.get("EP_SIZE", str(world_size))) +sample_max_length = int(os.environ.get("SAMPLE_MAX_LENGTH", "4096")) +pack_max_length = int(os.environ.get("PACK_MAX_LENGTH", str(sample_max_length))) +total_step = int(os.environ.get("TOTAL_STEP", "350" if stage == "dense_warmup" else "3000")) + +model_cfg = Glm47FlashDSAConfig.from_hf(model_path) +model_cfg.ep_size = ep_size +model_cfg.dispatcher = os.environ.get("DISPATCHER", "all2all") +model_cfg.compile_cfg = False +model_cfg.float8_cfg = None +model_cfg.attention.index_topk = int(os.environ.get("INDEX_TOPK", "512")) +model_cfg.attention.sparse_mla_backend = ( + "torch" if stage == "dense_warmup" else os.environ.get("SPARSE_MLA_BACKEND", "cudnn_dsa") +) +model_cfg.attention.indexer_training = DSAIndexerTrainingConfig( + stage=stage, + loss_coeff=float(os.environ.get("INDEXER_LOSS_COEFF", "1.0")), + train_mtp_indexer=True, + indexer_only=True, + dense_query_block_size=int(os.environ.get("DENSE_QUERY_BLOCK_SIZE", "128")), + debug_interval=int(os.environ.get("INDEXER_DEBUG_INTERVAL", "0")), +) +model_cfg.freeze_routers = True + +if stage == "sparse" and model_cfg.attention.sparse_mla_backend != "cudnn_dsa": + raise ValueError("Sparse indexer training requires SPARSE_MLA_BACKEND=cudnn_dsa.") +if stage == "sparse" and model_cfg.attention.index_topk % 128 != 0: + raise ValueError("Sparse cuDNN DSA requires INDEX_TOPK divisible by 128.") +if int(os.environ.get("SP_SIZE", "1")) != 1: + raise ValueError("Indexer training requires SP_SIZE=1.") +if float(os.environ.get("RECOMPUTE_RATIO", "0")) != 0: + raise ValueError("Indexer training requires RECOMPUTE_RATIO=0.") + +loss_cfg = CELossConfig(mode="chunk", chunk_size=int(os.environ.get("LOSS_CHUNK_SIZE", "1024"))) +model_cfg.lm_loss_cfg = loss_cfg + +pretrain_tokenize_cfg = LongTextPretrainTokenizeFunctionConfig( + chunk_size=sample_max_length, + tokenizer_chunk_chars=int(os.environ.get("TOKENIZER_CHUNK_CHARS", "32768")), + overlap_chars=int(os.environ.get("TOKENIZER_OVERLAP_CHARS", "512")), + min_chunk_tokens=0, + max_length=sample_max_length, + add_bos_token=False, + add_eos_token=True, +) +cache_tag = f"glm47_flash_{sample_max_length}" +if stage == "dense_warmup": + dataset_config = [ + { + "dataset": DatasetConfig( + name="warmup_pretrain_4k", + anno_path=os.environ["WARMUP_DATASET_PATH"], + sample_ratio=float(os.environ.get("WARMUP_DATASET_SAMPLE_RATIO", "1.0")), + cache_dir=os.path.join(work_dir, "jsonl_cache"), + cache_tag=f"{cache_tag}_warmup", + ), + "tokenize_fn": pretrain_tokenize_cfg, + } + ] +else: + dataset_config = [ + { + "dataset": DatasetConfig( + name="sft_4k", + anno_path=os.environ["SFT_DATASET_PATH"], + sample_ratio=float(os.environ.get("SFT_DATASET_SAMPLE_RATIO", "1.0")), + cache_dir=os.path.join(work_dir, "jsonl_cache"), + cache_tag=f"{cache_tag}_sft", + ), + "tokenize_fn": OpenaiTokenizeFunctionConfig( + chat_template=os.environ.get("CHAT_TEMPLATE", "glm5.2"), + max_length=sample_max_length, + ), + }, + { + "dataset": DatasetConfig( + name="pretrain_4k", + anno_path=os.environ["PRETRAIN_DATASET_PATH"], + sample_ratio=float(os.environ.get("PRETRAIN_DATASET_SAMPLE_RATIO", "1.0")), + cache_dir=os.path.join(work_dir, "jsonl_cache"), + cache_tag=f"{cache_tag}_pretrain", + ), + "tokenize_fn": pretrain_tokenize_cfg, + }, + ] + +dataloader_config = DataloaderConfig( + dataset_config_list=dataset_config, + pack_level="soft", + pack_max_length=pack_max_length, + pack_chunk_size=int(os.environ.get("PACK_CHUNK_SIZE", "10000")), + pack_workers=int(os.environ.get("PACK_WORKERS", "4")), + global_pack=True, + group_by_length=True, + num_workers=int(os.environ.get("DATALOADER_NUM_WORKERS", "4")), +) + +optim_cfg = AdamWConfig( + lr=float(os.environ.get("LR", "1e-3" if stage == "dense_warmup" else "7.3e-6")), + weight_decay=float(os.environ.get("WEIGHT_DECAY", "0.0" if stage == "dense_warmup" else "0.01")), + foreach=False, +) +lr_cfg = LRConfig( + lr_type="constant" if stage == "dense_warmup" else "cosine", + warmup_ratio=0, +) +fsdp_cfg = FSDPConfig( + cpu_offload=False, + ep_size=ep_size, + torch_compile=False, + recompute_ratio=0, +) + +indexer_checkpoint_path = os.environ.get("INDEXER_CHECKPOINT_PATH") +# A full warm-up HF export already contains base, MTP, and indexer weights, so +# it must go through the normal ``load_from`` path rather than the DCP loader. +if indexer_checkpoint_path and os.path.isfile(os.path.join(indexer_checkpoint_path, "model.safetensors.index.json")): + indexer_checkpoint_path = None +load_checkpoint_cfg = LoadCheckpointConfig( + checkpoint_path=indexer_checkpoint_path, + load_optimizer_states=False, + load_optimizer_args=False, + load_dataset=False, + load_scheduler=False, +) + +trainer = TrainerConfig( + model_cfg=model_cfg, + load_from=model_path, + tokenizer_path=model_path, + strict_load=True, + optim_cfg=optim_cfg, + dataloader_cfg=dataloader_config, + lr_cfg=lr_cfg, + loss_cfg=loss_cfg, + fsdp_cfg=fsdp_cfg, + global_batch_size=int(os.environ.get("GLOBAL_BATCH_SIZE", str(world_size))), + total_step=total_step, + intra_layer_micro_batch=1, + sp_size=1, + auto_resume=_bool_env("AUTO_RESUME", True), + load_checkpoint_cfg=load_checkpoint_cfg, + checkpoint_interval=int(os.environ.get("CHECKPOINT_INTERVAL", str(total_step))), + checkpoint_maxkeep=int(os.environ.get("CHECKPOINT_MAX_KEEP", "2")), + hf_interval=int(os.environ.get("HF_INTERVAL", "0")) or None, + work_dir=work_dir, + profile_memory=_bool_env("PROFILE_MEMORY", False), + profile_time=_bool_env("PROFILE_TIME", False), + profile_step=[2, 3], + debug_skip_save=False, + do_clip=True, +) diff --git a/examples/v1/config/sft_glm5p2.py b/examples/v1/config/sft_glm5p2.py index cab158ec2f..b5a74ce159 100644 --- a/examples/v1/config/sft_glm5p2.py +++ b/examples/v1/config/sft_glm5p2.py @@ -6,6 +6,7 @@ from xtuner.v1.float8.config import Float8Config, ScalingGranularity from xtuner.v1.loss import CELossConfig from xtuner.v1.model import get_model_config_from_hf +from xtuner.v1.module.attention import DSAIndexerTrainingConfig from xtuner.v1.train import TrainerConfig from xtuner.v1.train.trainer import LoadCheckpointConfig @@ -36,6 +37,7 @@ def _get_float8_config() -> Float8Config | None: # On single-node 8-GPU SFT, EP=8 leaves FSDP size at 1 and replicates non-expert params. ep_size = int(os.environ.get("EP_SIZE", "1")) intra_layer_micro_batch = int(os.environ.get("INTRA_LAYER_MICRO_BATCH", "1")) +sp_size = int(os.environ.get("SP_SIZE", "1")) global_batch_size = int(os.environ.get("GLOBAL_BATCH_SIZE", os.environ.get("WORLD_SIZE", "8"))) sample_max_length = int(os.environ.get("SAMPLE_MAX_LENGTH", "4096")) pack_max_length = int(os.environ.get("PACK_MAX_LENGTH", "16384")) @@ -54,6 +56,16 @@ def _get_float8_config() -> Float8Config | None: model_cfg.lm_loss_cfg = loss_cfg if hasattr(model_cfg.attention, "sparse_mla_backend"): model_cfg.attention.sparse_mla_backend = os.environ.get("SPARSE_MLA_BACKEND", "tilelang") +train_dsa_indexer = _get_bool_env("TRAIN_DSA_INDEXER", False) +if train_dsa_indexer: + if model_cfg.attention.sparse_mla_backend != "cudnn_dsa": + raise ValueError("DSA indexer training requires SPARSE_MLA_BACKEND=cudnn_dsa.") + model_cfg.attention.indexer_training = DSAIndexerTrainingConfig( + loss_coeff=float(os.environ.get("INDEXER_LOSS_COEFF", "1.0")), + train_mtp_indexer=_get_bool_env("TRAIN_MTP_INDEXER", False), + indexer_only=_get_bool_env("INDEXER_ONLY", False), + debug_interval=int(os.environ.get("INDEXER_DEBUG_INTERVAL", "0")), + ) cache_dir = os.path.join(work_dir, "jsonl_cache") cache_tag = os.environ.get("CACHE_TAG", f"glm52_{sample_max_length}") @@ -98,16 +110,30 @@ def _get_float8_config() -> Float8Config | None: elif optimizer == "adamw": optim_cfg = AdamWConfig( lr=lr, + weight_decay=float(os.environ.get("WEIGHT_DECAY", "0.01")), foreach=_get_bool_env("ADAMW_FOREACH", False), swap_optimizer=_get_bool_env("SWAP_OPTIMIZER", False), ) else: raise ValueError(f"Unsupported OPTIMIZER={optimizer!r}. Use adamw or muon.") lr_cfg = LRConfig(lr_type=os.environ.get("LR_TYPE", "cosine"), warmup_ratio=float(os.environ.get("WARMUP_RATIO", "0"))) +recompute_ratio = float(os.environ.get("RECOMPUTE_RATIO", "1.0")) +torch_compile = _get_bool_env("TORCH_COMPILE", False) +if train_dsa_indexer: + if sp_size != 1: + raise ValueError("DSA indexer training requires SP_SIZE=1.") + if intra_layer_micro_batch != 1: + raise ValueError("DSA indexer training requires INTRA_LAYER_MICRO_BATCH=1.") + if model_cfg.compile_cfg or torch_compile: + raise ValueError("DSA indexer training requires MODEL_COMPILE=0 and TORCH_COMPILE=0.") + if recompute_ratio != 0: + raise ValueError("DSA indexer training requires RECOMPUTE_RATIO=0 (no activation checkpointing).") + fsdp_cfg = FSDPConfig( cpu_offload=_get_bool_env("CPU_OFFLOAD", False), ep_size=ep_size, - torch_compile=_get_bool_env("TORCH_COMPILE", False), + torch_compile=torch_compile, + recompute_ratio=recompute_ratio, ) trainer = TrainerConfig( @@ -123,7 +149,7 @@ def _get_float8_config() -> Float8Config | None: global_batch_size=global_batch_size, total_step=total_step, intra_layer_micro_batch=intra_layer_micro_batch, - sp_size=int(os.environ.get("SP_SIZE", "1")), + sp_size=sp_size, load_checkpoint_cfg=LoadCheckpointConfig(checkpoint_path=os.environ.get("LOAD_CHECKPOINT_PATH")), checkpoint_interval=int(os.environ.get("CHECKPOINT_INTERVAL", "200")), checkpoint_maxkeep=int(os.environ.get("CHECKPOINT_MAX_KEEP", "3")), @@ -134,4 +160,5 @@ def _get_float8_config() -> Float8Config | None: profile_time=_get_bool_env("PROFILE_TIME", False), profile_step=[int(x) for x in os.environ.get("PROFILE_STEP", "2,3").split(",") if x], debug_skip_save=_get_bool_env("DEBUG_SKIP_SAVE", False), + do_clip=_get_bool_env("DO_CLIP", True), ) diff --git a/examples/v1/scripts/train_glm52_indexer.sh b/examples/v1/scripts/train_glm52_indexer.sh new file mode 100755 index 0000000000..e8be37da39 --- /dev/null +++ b/examples/v1/scripts/train_glm52_indexer.sh @@ -0,0 +1,91 @@ +#!/usr/bin/env bash +set -euo pipefail + +# Train GLM-5.2 and its main-stack source indexers jointly. Activate the +# intended Python environment before invoking this script. +: "${GLM5_2_MODEL_PATH:?GLM5_2_MODEL_PATH is required}" + +export DATASET_TYPE="${DATASET_TYPE:-alpaca}" +case "${DATASET_TYPE}" in + alpaca) + : "${ALPACA_PATH:?ALPACA_PATH is required when DATASET_TYPE=alpaca}" + ;; + alpaca_long) + : "${ALPACA_LONG_PATH:?ALPACA_LONG_PATH is required when DATASET_TYPE=alpaca_long}" + ;; + *) + echo "Unsupported DATASET_TYPE=${DATASET_TYPE}; use alpaca or alpaca_long." >&2 + exit 2 + ;; +esac + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +REPO_ROOT="$(cd "${SCRIPT_DIR}/../../.." && pwd)" +CONFIG_PATH="${1:-${REPO_ROOT}/examples/v1/config/sft_glm5p2.py}" +export WORK_DIR="${2:-${WORK_DIR:-work_dirs/glm52_indexer_sft}}" +export PYTHONPATH="${REPO_ROOT}${PYTHONPATH:+:${PYTHONPATH}}" + +export TRAIN_DSA_INDEXER="${TRAIN_DSA_INDEXER:-1}" +export INDEXER_LOSS_COEFF="${INDEXER_LOSS_COEFF:-1.0}" +export INDEXER_ONLY="${INDEXER_ONLY:-0}" +export INDEXER_DEBUG_INTERVAL="${INDEXER_DEBUG_INTERVAL:-0}" +export SPARSE_MLA_BACKEND="${SPARSE_MLA_BACKEND:-cudnn_dsa}" + +# These values reflect the constraints enforced by sft_glm5p2.py while +# source-indexer training is enabled. +export SP_SIZE="${SP_SIZE:-1}" +export INTRA_LAYER_MICRO_BATCH="${INTRA_LAYER_MICRO_BATCH:-1}" +export RECOMPUTE_RATIO="${RECOMPUTE_RATIO:-0}" +export MODEL_COMPILE="${MODEL_COMPILE:-0}" +export TORCH_COMPILE="${TORCH_COMPILE:-0}" + +export EP_SIZE="${EP_SIZE:-4}" +export TOTAL_STEP="${TOTAL_STEP:-300}" +export GLOBAL_BATCH_SIZE="${GLOBAL_BATCH_SIZE:-8}" +export LR="${LR:-1e-6}" +export WEIGHT_DECAY="${WEIGHT_DECAY:-0.01}" +export DO_CLIP="${DO_CLIP:-1}" + +export DATASET_SAMPLE_RATIO="${DATASET_SAMPLE_RATIO:-1.0}" +export SAMPLE_MAX_LENGTH="${SAMPLE_MAX_LENGTH:-4096}" +export PACK_MAX_LENGTH="${PACK_MAX_LENGTH:-4096}" +export CACHE_TAG="${CACHE_TAG:-glm52_indexer_4096}" + +export FP8="${FP8:-1}" +export DEBUG_SKIP_SAVE="${DEBUG_SKIP_SAVE:-0}" +export CHECKPOINT_INTERVAL="${CHECKPOINT_INTERVAL:-200}" +export HF_INTERVAL="${HF_INTERVAL:-${TOTAL_STEP}}" +export HF_MAX_KEEP="${HF_MAX_KEEP:-1}" +export PROFILE_TIME="${PROFILE_TIME:-0}" +export PROFILE_MEMORY="${PROFILE_MEMORY:-0}" + +NNODES="${NNODES:-${NODE_COUNT:-1}}" +NODE_RANK="${NODE_RANK:-0}" +MASTER_ADDR="${MASTER_ADDR:-127.0.0.1}" +MASTER_PORT="${MASTER_PORT:-6000}" +NPROC_PER_NODE="${NPROC_PER_NODE:-8}" + +cd "${REPO_ROOT}" +test -f "${CONFIG_PATH}" +mkdir -p "${WORK_DIR}" +ulimit -n 65536 + +command=( + torchrun + "--nproc-per-node=${NPROC_PER_NODE}" + "--master-addr=${MASTER_ADDR}" + "--master-port=${MASTER_PORT}" + "--nnodes=${NNODES}" + "--node-rank=${NODE_RANK}" + --tee 3 + -m xtuner.v1.train.cli.sft + --config "${CONFIG_PATH}" +) + +if [[ "${DRY_RUN:-0}" != "0" ]]; then + printf '%q ' "${command[@]}" + printf '\n' + exit 0 +fi + +"${command[@]}" 2>&1 | tee -a "${WORK_DIR}/node_${NODE_RANK}.txt" diff --git a/tests/loss/test_dsa_indexer_loss.py b/tests/loss/test_dsa_indexer_loss.py new file mode 100644 index 0000000000..f1fac4c372 --- /dev/null +++ b/tests/loss/test_dsa_indexer_loss.py @@ -0,0 +1,34 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import torch + +from xtuner.v1.data_proto import SequenceContext +from xtuner.v1.loss import dense_dsa_indexer_kl_loss + + +def test_dense_dsa_indexer_kl_is_finite_and_only_updates_indexer(): + torch.manual_seed(7) + seq_ctx = SequenceContext.from_input_ids((torch.tensor([[1, 2, 3, 4]]),), device="cpu") + index_q = torch.randn(1, 4, 2, 4, requires_grad=True) + index_k = torch.randn(1, 4, 4, requires_grad=True) + index_weights = torch.randn(1, 4, 2, requires_grad=True) + teacher_q = torch.randn(1, 4, 3, 6, requires_grad=True) + teacher_k = torch.randn(1, 4, 3, 6, requires_grad=True) + + loss = dense_dsa_indexer_kl_loss( + index_q, + index_k, + index_weights, + teacher_q, + teacher_k, + seq_ctx, + softmax_scale=0.25, + query_block_size=2, + ) + loss.backward() + + assert torch.isfinite(loss) + assert all( + tensor.grad is not None and torch.isfinite(tensor.grad).all() for tensor in (index_q, index_k, index_weights) + ) + assert teacher_q.grad is None + assert teacher_k.grad is None diff --git a/tests/model/test_glm47_flash.py b/tests/model/test_glm47_flash.py new file mode 100644 index 0000000000..f8acd70a38 --- /dev/null +++ b/tests/model/test_glm47_flash.py @@ -0,0 +1,168 @@ +from unittest import mock + +import torch + +from xtuner.v1.model import Glm47FlashConfig, Glm47FlashDSAConfig, get_model_config +from xtuner.v1.module.attention import DSAIndexerTrainingConfig, DSAMLAConfig, MLAConfig +from xtuner.v1.module.attention.mla import ( + mla_apply_rotary_pos_emb, + mla_apply_rotary_pos_emb_non_interleaved, +) +from xtuner.v1.module.mtp import MTPConfig +from xtuner.v1.module.router.noaux_router import NoAuxRouterConfig + + +def _tiny_glm47_flash_config() -> Glm47FlashConfig: + return Glm47FlashConfig( + vocab_size=32, + max_position_embeddings=64, + pad_token_id=0, + eos_token_id=1, + hf_eos_token_id=[1, 2], + num_hidden_layers=3, + first_k_dense_replace=1, + hidden_size=16, + intermediate_size=32, + moe_intermediate_size=4, + attention=MLAConfig( + num_attention_heads=2, + head_dim=4, + kv_lora_rank=4, + q_lora_rank=8, + qk_nope_head_dim=4, + qk_rope_head_dim=4, + v_head_dim=4, + rope_interleave=True, + ), + hf_head_dim=4, + qk_head_dim=8, + n_routed_experts=4, + n_shared_experts=1, + num_experts_per_tok=2, + router=NoAuxRouterConfig( + n_group=1, + topk_group=1, + scoring_func="sigmoid", + norm_topk_prob=True, + router_scaling_factor=1.8, + ), + mlp_layer_types=["dense", "sparse", "sparse"], + num_nextn_predict_layers=0, + mtp_config=None, + compile_cfg=False, + ) + + +def _tiny_glm47_flash_dsa_config() -> Glm47FlashDSAConfig: + dense = _tiny_glm47_flash_config() + dense.num_nextn_predict_layers = 1 + dense.mtp_config = MTPConfig(num_layers=1, share_weights=False) + values = { + name: getattr(dense, name) for name in Glm47FlashConfig.model_fields if name not in ("attention", "model_type") + } + return Glm47FlashDSAConfig( + **values, + attention=DSAMLAConfig( + **dense.attention.model_dump(), + index_topk=4, + index_head_dim=8, + index_n_heads=2, + index_topk_freq=2, + index_skip_topk_offset=1, + indexer_rope_interleave=True, + sparse_mla_backend="torch", + indexer_training=DSAIndexerTrainingConfig( + stage="dense_warmup", + train_mtp_indexer=True, + indexer_only=True, + dense_query_block_size=2, + ), + ), + ) + + +class TestGlm47FlashConfig: + def test_alias_and_default_architecture(self): + config = get_model_config("glm-4.7-flash") + + assert isinstance(config, Glm47FlashConfig) + assert config.model_type == "glm4_moe_lite" + assert config.num_hidden_layers == 47 + assert config.first_k_dense_replace == 1 + assert config.attention.rope_interleave + assert config.n_routed_experts == 64 + assert config.num_experts_per_tok == 4 + + def test_hf_key_mapping_uses_official_experts_and_mtp_layer(self): + config = _tiny_glm47_flash_config() + with mock.patch("torch.cuda.Stream"): + model = config.build() + + assert model.to_hf_key_list("layers.1.experts.fused_w1w3.weight") == [ + f"model.layers.1.mlp.experts.{expert_idx}.{projection}_proj.weight" + for expert_idx in range(4) + for projection in ("gate", "up") + ] + assert model.to_hf_key_list("layers.1.experts.fused_w2.weight") == [ + f"model.layers.1.mlp.experts.{expert_idx}.down_proj.weight" for expert_idx in range(4) + ] + assert model.to_hf_key_list("layers.1.gate.router.e_score_correction_bias") == [ + "model.layers.1.mlp.gate.e_score_correction_bias" + ] + assert model.to_hf_key_list("mtp_block.layers.0.final_layernorm.weight") == [ + "model.layers.3.shared_head.norm.weight" + ] + + def test_dense_warmup_trains_main_and_mtp_source_indexers_only(self): + config = _tiny_glm47_flash_dsa_config() + with mock.patch("torch.cuda.Stream"): + model = config.build() + + trainable = {name for name, parameter in model.named_parameters() if parameter.requires_grad} + assert trainable + assert all(".self_attn.indexer." in name for name in trainable) + assert hasattr(model.layers["0"].self_attn, "indexer") + assert not hasattr(model.layers["1"].self_attn, "indexer") + assert hasattr(model.layers["2"].self_attn, "indexer") + assert model.mtp_block is not None + assert hasattr(model.mtp_block.layers[0].decoder_layer.self_attn, "indexer") + + def test_official_expert_layout_round_trip(self): + config = _tiny_glm47_flash_config() + with mock.patch("torch.cuda.Stream"): + model = config.build() + + gate_up_hf = [ + torch.arange(offset, offset + 4 * 16, dtype=torch.float32).reshape(4, 16) + for offset in range(0, 8 * 4 * 16, 4 * 16) + ] + gate_up_local = torch.empty(32, 16) + model.safetensors_to_params( + gate_up_hf, gate_up_local, "layers.1.experts.fused_w1w3.weight", None, None, 0 + ) + torch.testing.assert_close(gate_up_local, torch.cat(gate_up_hf)) + + down_hf = [ + torch.arange(offset, offset + 16 * 4, dtype=torch.float32).reshape(16, 4) + for offset in range(0, 4 * 16 * 4, 16 * 4) + ] + down_local = torch.empty(64, 4) + model.safetensors_to_params(down_hf, down_local, "layers.1.experts.fused_w2.weight", None, None, 0) + torch.testing.assert_close(down_local, torch.cat(down_hf)) + + +class TestGlm47FlashRope: + def test_interleaved_and_half_split_layouts(self): + q = torch.tensor([[[[1.0, 2.0, 3.0, 4.0]]]]) + k = q + 1 + cos = torch.tensor([[[0.8, 0.6, 0.8, 0.6]]]) + sin = torch.tensor([[[0.6, 0.8, 0.6, 0.8]]]) + + q_interleaved, k_interleaved = mla_apply_rotary_pos_emb(q, k, cos, sin) + q_half, k_half = mla_apply_rotary_pos_emb_non_interleaved(q, k, cos, sin) + + expected_q_interleaved = torch.tensor([[[[-0.4, -1.4, 2.2, 4.8]]]]) + expected_q_half = torch.tensor([[[[-1.0, -2.0, 3.0, 4.0]]]]) + torch.testing.assert_close(q_interleaved, expected_q_interleaved) + torch.testing.assert_close(q_half, expected_q_half) + assert not torch.equal(k_interleaved, k_half) diff --git a/tests/model/test_glm52_moe.py b/tests/model/test_glm52_moe.py index a09aeca5e0..8bce1c06da 100644 --- a/tests/model/test_glm52_moe.py +++ b/tests/model/test_glm52_moe.py @@ -4,6 +4,8 @@ test_save_hf_matches_transformers_and_engine_contracts: HF round-trip 与推理引擎字段契约一致。 test_from_hf_preserves_glm_specific_behavior: HF 配置转换保留 DSA、router 与 MTP 语义。 test_rejects_shared_physical_mtp_indexer: 非法的 physical MTP indexer 计划会被拒绝。 + test_indexer_training_can_include_physical_mtp_indexer: MTP indexer 可显式加入训练。 + test_indexer_only_can_train_main_and_mtp_source_indexers: overfit 模式可同时训练两类 indexer。 TestGlm52CheckpointConversion test_tiny_model_round_trips_through_hf: tiny 主干与 MTP 参数可经公共 HF API 无损往返。 TestGlm52RouterBias @@ -29,7 +31,7 @@ from xtuner.v1.data_proto import SequenceContext from xtuner.v1.loss.ce_loss import CELossConfig from xtuner.v1.model import Glm52MoEConfig, get_model_config, get_model_config_from_hf -from xtuner.v1.module.attention import DSAMLAConfig +from xtuner.v1.module.attention import DSAIndexerTrainingConfig, DSAMLAConfig from xtuner.v1.module.mtp import MTPConfig from xtuner.v1.module.router.noaux_router import NoAuxRouterConfig from xtuner.v1.utils.test_utils import init_data_mesh @@ -182,6 +184,72 @@ def test_rejects_shared_physical_mtp_indexer(self): with pytest.raises(ValueError, match="physical MTP indexer_types"): config.build() + def test_indexer_training_keeps_physical_mtp_indexer_frozen_by_default(self): + # 主干 source indexer 解冻时,physical MTP indexer 仍保持 frozen。 + config = _tiny_glm52_config() + config.attention.indexer_training = DSAIndexerTrainingConfig(loss_coeff=1.0) + config.mtp_config = MTPConfig(num_layers=1, share_weights=True) + + with mock.patch("torch.cuda.Stream"): + model = config.build() + + main_attention = model.layers["0"].self_attn + mtp_attention = model.mtp_block.layers[0].decoder_layer.self_attn # type: ignore[union-attr] + assert all(parameter.requires_grad for parameter in main_attention.indexer.parameters()) + assert all(not parameter.requires_grad for parameter in mtp_attention.indexer.parameters()) + assert mtp_attention.indexer_training is None + + def test_indexer_training_can_include_physical_mtp_indexer(self): + # 显式开启后,checkpoint 中的 physical MTP Full/source indexer 参与训练。 + config = _tiny_glm52_config() + config.attention.indexer_training = DSAIndexerTrainingConfig( + loss_coeff=1.0, + train_mtp_indexer=True, + ) + config.mtp_config = MTPConfig(num_layers=1, share_weights=True) + + with mock.patch("torch.cuda.Stream"): + model = config.build() + + mtp_attention = model.mtp_block.layers[0].decoder_layer.self_attn # type: ignore[union-attr] + assert mtp_attention.indexer_training is not None + assert mtp_attention.indexer_training.train_mtp_indexer + assert all(parameter.requires_grad for parameter in mtp_attention.indexer.parameters()) + + def test_indexer_only_trains_main_source_indexers_exclusively(self): + # 严格过拟合模式固定 attention teacher,只训练主干 source indexer。 + config = _tiny_glm52_config() + config.attention.indexer_training = DSAIndexerTrainingConfig(loss_coeff=1.0, indexer_only=True) + config.mtp_config = MTPConfig(num_layers=1, share_weights=True) + + with mock.patch("torch.cuda.Stream"): + model = config.build() + + trainable_names = [name for name, parameter in model.named_parameters() if parameter.requires_grad] + assert trainable_names + assert all(name.startswith("layers.0.self_attn.indexer.") for name in trainable_names) + assert all(not parameter.requires_grad for parameter in model.layers["1"].parameters()) + assert all(not parameter.requires_grad for parameter in model.mtp_block.parameters()) # type: ignore[union-attr] + + def test_indexer_only_can_train_main_and_mtp_source_indexers(self): + # MTP opt-in 也适用于严格 overfit 模式,其余 MTP 参数仍全部冻结。 + config = _tiny_glm52_config() + config.attention.indexer_training = DSAIndexerTrainingConfig( + loss_coeff=1.0, + train_mtp_indexer=True, + indexer_only=True, + ) + config.mtp_config = MTPConfig(num_layers=1, share_weights=True) + + with mock.patch("torch.cuda.Stream"): + model = config.build() + + trainable_names = [name for name, parameter in model.named_parameters() if parameter.requires_grad] + assert trainable_names + assert all(".self_attn.indexer." in name for name in trainable_names) + assert any(name.startswith("layers.0.") for name in trainable_names) + assert any(name.startswith("mtp_block.layers.0.") for name in trainable_names) + @unittest.skipUnless(torch.cuda.is_available(), "requires CUDA") class TestGlm52CheckpointConversion(DeterministicDDPTestCase): diff --git a/tests/module/attention/test_dsa_mla.py b/tests/module/attention/test_dsa_mla.py index 289a6ec8c4..2074e94724 100644 --- a/tests/module/attention/test_dsa_mla.py +++ b/tests/module/attention/test_dsa_mla.py @@ -3,6 +3,8 @@ TestTorchSparseMLA test_padded_indices_support_int32_and_backward: PyTorch 后端处理 padding、int32 和反向传播。 TestDSAAttention + test_source_layer_passes_real_query_mask_to_indexer_loss: padding query 不进入 indexer KL。 + test_mtp_iteration_trains_indexer_only_at_compute_depth: MTP 复用深度不重复投影或累计 KL。 test_packed_inputs_respect_causal_boundaries_and_backward: packed attention 遵守分段因果边界并可反传。 test_shared_layers_reuse_topk_without_cross_context_leak: shared layer 复用当前样本 top-k 且不跨样本泄漏。 test_reentrant_checkpoint_reuses_and_releases_topk: checkpoint 重算复用并最终释放 top-k。 @@ -28,9 +30,14 @@ from xtuner._testing import DeterministicDDPTestCase from xtuner.v1.data_proto import SequenceContext +from xtuner.v1.model.moe.moe import MoE from xtuner.v1.model.utils import checkpoint_wrapper -from xtuner.v1.module.attention import DSAMLAConfig -from xtuner.v1.module.attention.dsa_topk_sharing import register_dsa_topk_decoder_lifecycle_hooks +from xtuner.v1.module.attention import DSAIndexerTrainingConfig, DSAMLAConfig +from xtuner.v1.module.attention import dsa_mla as dsa_mla_module +from xtuner.v1.module.attention.dsa_topk_sharing import ( + get_dsa_topk_sharing_runtime, + register_dsa_topk_decoder_lifecycle_hooks, +) from xtuner.v1.ops.sparse_mla import dsa_topk_indices, sparse_mla from xtuner.v1.utils.test_utils import init_data_mesh @@ -102,6 +109,7 @@ def _cudnn_dsa_sparse_mla_inputs(): def _tiny_dsa_attention( indexer_types: list[str] | None = None, layer_idx: int = 0, + indexer_training: DSAIndexerTrainingConfig | None = None, ): return DSAMLAConfig( num_attention_heads=2, @@ -115,6 +123,7 @@ def _tiny_dsa_attention( index_head_dim=4, index_n_heads=2, indexer_types=indexer_types, + indexer_training=indexer_training, sparse_mla_backend="torch", ).build(hidden_size=4, layer_idx=layer_idx) @@ -167,6 +176,250 @@ def test_padded_indices_support_int32_and_backward(self): class TestDSAAttention: + def test_source_layer_losses_are_averaged_and_released_at_model_boundary(self): + first_ctx = SequenceContext.from_input_ids((torch.tensor([[1, 2]]),), device="cpu") + second_ctx = SequenceContext.from_input_ids((torch.tensor([[3, 4]]),), device="cpu") + first_ctx.dsa_topk_cache.indexer_losses.extend([torch.tensor(1.0), torch.tensor(3.0)]) + second_ctx.dsa_topk_cache.indexer_losses.append(torch.tensor(5.0)) + + indexer_loss = MoE._consume_indexer_losses([first_ctx, second_ctx]) + + torch.testing.assert_close(indexer_loss, torch.tensor(3.0)) + assert first_ctx.dsa_topk_cache.indexer_losses == [] + assert second_ctx.dsa_topk_cache.indexer_losses == [] + + def test_indexer_training_is_opt_in_and_only_materializes_on_source_layers(self): + # None 是严格 frozen baseline;启用后也只有 Full/source layer 持有可训练 indexer。 + frozen = _tiny_dsa_attention(indexer_types=["full", "shared"], layer_idx=0) + trainable = _tiny_dsa_attention( + indexer_types=["full", "shared"], + layer_idx=0, + indexer_training=DSAIndexerTrainingConfig(loss_coeff=1.0), + ) + shared = _tiny_dsa_attention( + indexer_types=["full", "shared"], + layer_idx=1, + indexer_training=DSAIndexerTrainingConfig(loss_coeff=1.0), + ) + + assert all(not parameter.requires_grad for parameter in frozen.indexer.parameters()) + assert all(parameter.requires_grad for parameter in trainable.indexer.parameters()) + assert not hasattr(shared, "indexer") + + def test_single_weights_projection_preserves_score_scaling(self): + attention = _tiny_dsa_attention( + indexer_types=["full"], + indexer_training=DSAIndexerTrainingConfig(loss_coeff=1.0), + ) + hidden_states = torch.randn(1, 4, 4) + q_resid = torch.randn(1, 4, 4) + position_embeddings = (torch.ones(1, 4, 2), torch.zeros(1, 4, 2)) + seq_ctx = SequenceContext.from_input_ids((torch.tensor([[1, 2, 3, 4]]),), device="cpu") + + projection_calls = [] + handle = attention.indexer.weights_proj.register_forward_hook(lambda *_: projection_calls.append(None)) + try: + features = attention.indexer.project_features(hidden_states, q_resid, position_embeddings, seq_ctx) + finally: + handle.remove() + dot = torch.einsum("bshd,btd->bsht", features.q.float(), features.k.float()) + selection_logits = torch.einsum( + "bsht,bsh->bst", + torch.relu(dot * (attention.index_head_dim**-0.5)), + features.weights.float() * (attention.index_n_heads**-0.5), + ) + training_logits = torch.einsum( + "bsht,bsh->bst", + torch.relu(dot), + (features.weights * ((attention.index_n_heads * attention.index_head_dim) ** -0.5)).float(), + ) + + assert len(projection_calls) == 1 + torch.testing.assert_close(training_logits, selection_logits) + + def test_source_layer_loss_updates_only_indexer_and_detaches_model_inputs(self, monkeypatch): + # Use a differentiable PyTorch oracle to verify the source-layer autograd boundary. + def fake_indexer_loss( + index_q, + index_k, + index_weights, + attention_q, + attention_k, + softmax_lse, + topk_indices, + *, + softmax_scale, + row_coefficient, + valid_query_mask, + debug_name, + debug_interval, + ): + del attention_q, attention_k, softmax_lse, topk_indices, softmax_scale, debug_name, debug_interval + query_mask = valid_query_mask.unsqueeze(-1) + return row_coefficient * ( + index_q.float().square().masked_fill(~query_mask.unsqueeze(-1), 0.0).sum() + + index_k.float().square().sum() + + index_weights.float().square().masked_fill(~query_mask, 0.0).sum() + ) + + monkeypatch.setattr(dsa_mla_module, "dsa_indexer_kl_loss", fake_indexer_loss) + torch.manual_seed(17) + frozen = _tiny_dsa_attention(indexer_types=["full"], layer_idx=0) + trainable = _tiny_dsa_attention( + indexer_types=["full"], + layer_idx=0, + indexer_training=DSAIndexerTrainingConfig(loss_coeff=0.25), + ) + trainable.load_state_dict(frozen.state_dict()) + position_embeddings = (torch.ones(1, 4, 2), torch.zeros(1, 4, 2)) + frozen_hidden = torch.randn(1, 4, 4, requires_grad=True) + trained_hidden = frozen_hidden.detach().clone().requires_grad_() + + frozen_ctx = SequenceContext.from_input_ids((torch.tensor([[1, 2, 3, 4]]),), device="cpu") + frozen_output = frozen(frozen_hidden, position_embeddings, frozen_ctx)["projected_output"] + frozen_output.square().mean().backward() + + trained_ctx = SequenceContext.from_input_ids((torch.tensor([[1, 2, 3, 4]]),), device="cpu") + trained_output = trainable(trained_hidden, position_embeddings, trained_ctx)["projected_output"] + assert len(trained_ctx.dsa_topk_cache.indexer_losses) == 1 + indexer_loss = trained_ctx.dsa_topk_cache.indexer_losses[0] + (trained_output.square().mean() + indexer_loss).backward() + + torch.testing.assert_close(trained_hidden.grad, frozen_hidden.grad) + frozen_parameters = dict(frozen.named_parameters()) + for name, trained_parameter in trainable.named_parameters(): + if name.startswith("indexer."): + continue + frozen_grad = frozen_parameters[name].grad + assert (trained_parameter.grad is None) == (frozen_grad is None) + if trained_parameter.grad is not None: + torch.testing.assert_close(trained_parameter.grad, frozen_grad) + assert all(parameter.grad is not None for parameter in trainable.indexer.parameters()) + assert all(torch.isfinite(parameter.grad).all() for parameter in trainable.indexer.parameters()) + + def test_source_layer_passes_real_query_mask_to_indexer_loss(self, monkeypatch): + captured = {} + + def capture_indexer_loss( + index_q, + index_k, + index_weights, + attention_q, + attention_k, + softmax_lse, + topk_indices, + *, + softmax_scale, + row_coefficient, + valid_query_mask, + debug_name, + debug_interval, + ): + del index_k, index_weights, attention_q, attention_k, softmax_lse, softmax_scale + captured["row_coefficient"] = row_coefficient + captured["valid_query_mask"] = valid_query_mask.detach().clone() + captured["topk_valid_rows"] = (topk_indices != -1).any(dim=-1).detach().clone() + captured["debug_name"] = debug_name + captured["debug_interval"] = debug_interval + return index_q.float().sum() * 0.0 + + monkeypatch.setattr(dsa_mla_module, "dsa_indexer_kl_loss", capture_indexer_loss) + attention = _tiny_dsa_attention( + indexer_types=["full"], + layer_idx=0, + indexer_training=DSAIndexerTrainingConfig(loss_coeff=0.25), + ) + hidden_states = torch.randn(1, 4, 4) + position_embeddings = (torch.ones(1, 4, 2), torch.zeros(1, 4, 2)) + # The final two physical rows form a causal padding chunk. Its top-k + # indices are valid-looking, but num_padding must still exclude it. + seq_ctx = SequenceContext( + input_ids=torch.tensor([[1, 2, 0, 0]]), + cu_seq_lens_q=torch.tensor([0, 2, 4], dtype=torch.int32), + cu_seq_lens_k=torch.tensor([0, 2, 4], dtype=torch.int32), + max_length_q=2, + max_length_k=2, + num_padding=2, + device="cpu", + ) + + attention(hidden_states, position_embeddings, seq_ctx) + + assert captured["topk_valid_rows"].tolist() == [[True, True, True, True]] + torch.testing.assert_close(captured["valid_query_mask"], torch.tensor([[True, True, False, False]])) + assert captured["row_coefficient"] == pytest.approx(0.25 / 2) + assert captured["debug_name"] == "layer0" + assert captured["debug_interval"] == 0 + + def test_mtp_iteration_trains_indexer_only_at_compute_depth(self, monkeypatch): + # GLM-5.2 的物理 MTP 层只在第一个 logical depth 计算 index;后续 + # depth 复用同一 top-k,不能重复跑 projection 或重复加 KL。 + loss_calls = [] + + def fake_indexer_loss(index_q, *args, **kwargs): + del args, kwargs + loss_calls.append(None) + return index_q.float().sum() * 0.0 + + monkeypatch.setattr(dsa_mla_module, "dsa_indexer_kl_loss", fake_indexer_loss) + attention = _tiny_dsa_attention( + indexer_types=["full"], + layer_idx=0, + indexer_training=DSAIndexerTrainingConfig(loss_coeff=1.0, train_mtp_indexer=True), + ) + projection_calls = [] + handle = attention.indexer.weights_proj.register_forward_hook(lambda *_: projection_calls.append(None)) + hidden_states = torch.randn(1, 4, 4) + position_embeddings = (torch.ones(1, 4, 2), torch.zeros(1, 4, 2)) + seq_ctx = SequenceContext.from_input_ids((torch.tensor([[1, 2, 3, 4]]),), device="cpu") + get_dsa_topk_sharing_runtime().register_mtp_iteration_topk_sharing( + seq_ctx=seq_ctx, + source_layer_idx=0, + num_iterations=2, + ) + + try: + attention(hidden_states, position_embeddings, seq_ctx) + first_topk = seq_ctx.dsa_topk_cache.indices[0] + attention(torch.randn_like(hidden_states), position_embeddings, seq_ctx) + finally: + handle.remove() + + assert seq_ctx.dsa_topk_cache.indices[0] is first_topk + assert len(projection_calls) == 1 + assert len(loss_calls) == 1 + assert len(seq_ctx.dsa_topk_cache.indexer_losses) == 1 + + def test_source_layer_training_rejects_no_grad_checkpoint_forward(self): + attention = _tiny_dsa_attention( + indexer_types=["full"], + layer_idx=0, + indexer_training=DSAIndexerTrainingConfig(loss_coeff=1.0), + ) + hidden_states = torch.randn(1, 4, 4) + position_embeddings = (torch.ones(1, 4, 2), torch.zeros(1, 4, 2)) + seq_ctx = SequenceContext.from_input_ids((torch.tensor([[1, 2, 3, 4]]),), device="cpu") + + with torch.no_grad(), pytest.raises(RuntimeError, match="does not support activation checkpointing"): + attention(hidden_states, position_embeddings, seq_ctx) + + def test_zero_loss_coefficient_keeps_trainable_indexer_grad_none(self): + # coeff=0 必须完全绕过 autograd,避免 AdamW 对 zero grad tensor 执行 weight decay。 + attention = _tiny_dsa_attention( + indexer_types=["full"], + layer_idx=0, + indexer_training=DSAIndexerTrainingConfig(loss_coeff=0.0), + ) + hidden_states = torch.randn(1, 4, 4, requires_grad=True) + position_embeddings = (torch.ones(1, 4, 2), torch.zeros(1, 4, 2)) + seq_ctx = SequenceContext.from_input_ids((torch.tensor([[1, 2, 3, 4]]),), device="cpu") + + output = attention(hidden_states, position_embeddings, seq_ctx)["projected_output"] + output.square().mean().backward() + + assert seq_ctx.dsa_topk_cache.indexer_losses == [] + assert all(parameter.grad is None for parameter in attention.indexer.parameters()) + def test_packed_inputs_respect_causal_boundaries_and_backward(self): # 验证 packed attention 不跨子序列取 key,并能对真实输入完成有限反向传播。 torch.manual_seed(0) diff --git a/tests/ops/test_cudnn_dsa_indexer_loss.py b/tests/ops/test_cudnn_dsa_indexer_loss.py new file mode 100644 index 0000000000..002ed1e8fc --- /dev/null +++ b/tests/ops/test_cudnn_dsa_indexer_loss.py @@ -0,0 +1,535 @@ +"""Correctness tests for the cuDNN DSA sparse indexer training loss.""" + +import importlib +import subprocess +import sys +from functools import cache + +import pytest +import torch + +from xtuner.v1.ops.sparse_mla.cudnn_dsa_indexer_loss import ( + _INDEXER_LOSS_DEBUG_CALLS, + _copy_aligned_grad_loss, + _mask_invalid_query_rows, + _maybe_log_indexer_loss_diagnostics, + _pad_indexer_heads_for_cudnn, + _standard_kl_loss, + _xtuner_indexer_backward, + dsa_indexer_kl_from_distribution, + dsa_indexer_kl_loss, + sparse_attention_target, + sparse_indexer_predict, +) + + +@cache +def _cudnn_indexer_training_available() -> bool: + if not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] < 9: + return False + result = subprocess.run( + [ + sys.executable, + "-c", + "from cudnn.deepseek_sparse_attention.indexer_backward import indexer_backward_wrapper; " + "from cudnn.deepseek_sparse_attention.score_recompute import " + "sparse_attn_score_recompute_wrapper, sparse_indexer_score_recompute_wrapper", + ], + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + check=False, + ) + return result.returncode == 0 + + +def _packed_topk_indices(seq_lens: tuple[int, ...], topk: int, device: str) -> torch.Tensor: + seq_len = sum(seq_lens) + indices = torch.full((1, seq_len, topk), -1, dtype=torch.int32, device=device) + row = 0 + for seq_len_i in seq_lens: + for offset in range(seq_len_i): + valid = min(offset + 1, topk) + indices[0, row, :valid] = torch.arange( + row + 1 - valid, + row + 1, + dtype=torch.int32, + device=device, + ) + row += 1 + return indices + + +def _gather_selected_k(k: torch.Tensor, topk_indices: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + assert k.shape[0] == 1, "The test oracle only needs the packed batch-size-one layout." + valid = topk_indices != -1 + selected = k[:, topk_indices.clamp_min(0)[0].long(), :] + return selected, valid + + +def _indexer_predict_oracle( + index_q: torch.Tensor, + index_k: torch.Tensor, + weights: torch.Tensor, + topk_indices: torch.Tensor, +) -> torch.Tensor: + selected_k, valid = _gather_selected_k(index_k, topk_indices) + scores = torch.einsum("bshd,bskd->bshk", index_q.float(), selected_k.float()).relu() + logits = torch.einsum("bshk,bsh->bsk", scores, weights.float()) + logits = logits.masked_fill(~valid, float("-inf")) + return torch.softmax(logits, dim=-1).masked_fill(~valid, 0.0) + + +def _attention_target_oracle( + attn_q: torch.Tensor, + attn_k: torch.Tensor, + topk_indices: torch.Tensor, + softmax_scale: float, +) -> tuple[torch.Tensor, torch.Tensor]: + selected_k, valid = _gather_selected_k(attn_k, topk_indices) + scores = torch.einsum("bshd,bskd->bshk", attn_q.float(), selected_k.float()) * softmax_scale + scores = scores.masked_fill(~valid[:, :, None, :], float("-inf")) + lse = torch.logsumexp(scores, dim=-1) + probs = torch.softmax(scores, dim=-1).masked_fill(~valid[:, :, None, :], 0.0) + target = probs.sum(dim=2) + target = target / target.sum(dim=-1, keepdim=True) + return target.masked_fill(~valid, 0.0), lse + + +def _explicit_standard_kl( + target: torch.Tensor, + predict: torch.Tensor, + topk_indices: torch.Tensor, + row_coefficient: float, +) -> torch.Tensor: + valid_rows = (topk_indices != -1).any(dim=-1) + target_xlogx = torch.special.xlogy(target, target).sum(dim=-1) + cross_entropy = torch.special.xlogy(target, predict).sum(dim=-1) + return row_coefficient * (target_xlogx - cross_entropy).masked_fill(~valid_rows, 0.0).sum() + + +class TestStandardKLLoss: + @staticmethod + def _indexer_logits(q, k, weights): + scores = torch.einsum("bshd,btd->bsht", q, k).relu() + return torch.einsum("bsht,bsh->bst", scores, weights) + + def test_zero_head_padding_preserves_scores_and_original_gradients(self): + torch.manual_seed(5) + q_data = torch.randn(1, 3, 32, 8) + k_data = torch.randn(1, 4, 8) + w_data = torch.randn(1, 3, 32) + + reference_q = q_data.clone().requires_grad_() + reference_k = k_data.clone().requires_grad_() + reference_w = w_data.clone().requires_grad_() + reference_logits = self._indexer_logits(reference_q, reference_k, reference_w) + reference_logits.square().sum().backward() + + actual_q = q_data.clone().requires_grad_() + actual_k = k_data.clone().requires_grad_() + actual_w = w_data.clone().requires_grad_() + padded_q, padded_w, original_heads = _pad_indexer_heads_for_cudnn(actual_q, actual_w) + actual_logits = self._indexer_logits(padded_q, actual_k, padded_w) + actual_logits.square().sum().backward() + + assert original_heads == 32 + assert padded_q.shape[-2] == 64 + assert padded_w.shape[-1] == 64 + torch.testing.assert_close(actual_logits, reference_logits) + torch.testing.assert_close(actual_q.grad, reference_q.grad) + torch.testing.assert_close(actual_k.grad, reference_k.grad) + torch.testing.assert_close(actual_w.grad, reference_w.grad) + + def test_grad_loss_copy_realigns_contiguous_storage_offset_view(self): + storage = torch.tensor([0.0, 0.37], dtype=torch.float32) + misaligned_grad_loss = storage[1:] + assert misaligned_grad_loss.is_contiguous() + assert misaligned_grad_loss.data_ptr() % 16 != 0 + + aligned_grad_loss = _copy_aligned_grad_loss(misaligned_grad_loss, torch.device("cpu")) + + assert aligned_grad_loss.shape == (1,) + assert aligned_grad_loss.dtype == torch.float32 + assert aligned_grad_loss.data_ptr() % 16 == 0 + torch.testing.assert_close(aligned_grad_loss, misaligned_grad_loss) + + def test_backward_adapter_satisfies_cudnn_shape_alignment_and_mutation_contract(self, monkeypatch): + import cudnn.deepseek_sparse_attention.indexer_backward as cudnn_indexer_backward + + seen = {} + + def fake_indexer_backward_wrapper( + index_q, + weights, + index_k, + target, + predict, + topk_indices, + **kwargs, + ): + tensors = (index_q, weights, index_k, target, predict, topk_indices, kwargs["grad_loss"]) + assert all(tensor.is_contiguous() for tensor in tensors) + assert all(tensor.data_ptr() % 16 == 0 for tensor in tensors) + assert index_q.shape[-2] == 64 + assert weights.shape[-1] == 64 + assert topk_indices.shape[-1] == 128 + assert topk_indices.dtype == torch.int32 + seen["loss_coeff"] = kwargs["loss_coeff"] + target.zero_() + predict.zero_() + return { + "d_index_q": torch.ones_like(index_q), + "d_index_k": torch.ones_like(index_k), + "d_weights": torch.ones_like(weights), + } + + monkeypatch.setattr(cudnn_indexer_backward, "indexer_backward_wrapper", fake_indexer_backward_wrapper) + + index_q = torch.randn(1, 2, 32, 8) + index_k = torch.randn(1, 4, 8) + index_weights = torch.randn(1, 2, 32) + target = torch.rand(1, 2, 128) + predict = torch.rand(1, 2, 128) + target_before = target.clone() + predict_before = predict.clone() + safe_topk_indices = torch.zeros(1, 2, 128, dtype=torch.int32) + grad_storage = torch.tensor([0.0, 0.25], dtype=torch.float32) + + d_q, d_k, d_w = _xtuner_indexer_backward( + index_q, + index_k, + index_weights, + target, + predict, + safe_topk_indices, + row_coefficient=0.125, + grad_loss=grad_storage[1], + ) + + assert seen["loss_coeff"] == 0.25 + assert d_q.shape == index_q.shape + assert d_k.shape == index_k.shape + assert d_w.shape == index_weights.shape + torch.testing.assert_close(target, target_before) + torch.testing.assert_close(predict, predict_before) + + def test_matches_explicit_standard_kl_with_zero_probability_slots(self): + target = torch.tensor([[[0.7, 0.3, 0.0], [1.0, 0.0, 0.0]]], dtype=torch.float32) + predict = torch.tensor([[[0.6, 0.4, 0.0], [0.8, 0.2, 0.0]]], dtype=torch.float32) + topk_indices = torch.tensor([[[0, 1, -1], [0, -1, -1]]], dtype=torch.int32) + + actual = _standard_kl_loss(target, predict, topk_indices, row_coefficient=0.25) + expected = _explicit_standard_kl(target, predict, topk_indices, row_coefficient=0.25) + + torch.testing.assert_close(actual, expected) + assert torch.isfinite(actual) + + def test_identical_distributions_have_zero_loss(self): + distribution = torch.tensor([[[0.7, 0.3, 0.0], [1.0, 0.0, 0.0]]], dtype=torch.float32) + topk_indices = torch.tensor([[[0, 1, -1], [0, -1, -1]]], dtype=torch.int32) + + loss = _standard_kl_loss(distribution, distribution, topk_indices, row_coefficient=0.5) + + torch.testing.assert_close(loss, torch.tensor(0.0), atol=0.0, rtol=0.0) + + def test_positive_target_with_zero_predict_is_rejected(self): + target = torch.tensor([[[1.0, 0.0]]], dtype=torch.float32) + predict = torch.tensor([[[0.0, 1.0]]], dtype=torch.float32) + topk_indices = torch.tensor([[[0, 1]]], dtype=torch.int32) + + with pytest.raises((AssertionError, RuntimeError), match="zero probability"): + _standard_kl_loss(target, predict, topk_indices, row_coefficient=1.0) + + def test_query_mask_excludes_causal_padding_chunks(self): + target = torch.tensor( + [[[0.8, 0.2], [0.3, 0.7], [0.6, 0.4]]], + dtype=torch.float32, + ) + predict = torch.tensor( + [[[0.5, 0.5], [0.9, 0.1], [0.1, 0.9]]], + dtype=torch.float32, + ) + # Padding chunks are causal sequences, so their top-k rows can contain + # valid-looking indices even though they must not contribute to loss. + topk_indices = torch.tensor([[[0, 1], [2, 3], [4, 5]]], dtype=torch.int32) + valid_query_mask = torch.tensor([[True, False, False]]) + + masked_target, masked_predict, masked_topk = _mask_invalid_query_rows( + target, + predict, + topk_indices, + valid_query_mask, + ) + actual = _standard_kl_loss(masked_target, masked_predict, masked_topk, row_coefficient=1.0) + expected = _explicit_standard_kl( + target[:, :1], + predict[:, :1], + topk_indices[:, :1], + row_coefficient=1.0, + ) + + torch.testing.assert_close(actual, expected) + assert torch.count_nonzero(masked_target[:, 1:]) == 0 + assert torch.count_nonzero(masked_predict[:, 1:]) == 0 + assert torch.all(masked_topk[:, 1:] == -1) + + def test_query_mask_shape_is_validated(self): + target = torch.ones(1, 2, 1) + predict = torch.ones_like(target) + topk_indices = torch.zeros(1, 2, 1, dtype=torch.int32) + + with pytest.raises(ValueError, match="valid_query_mask must have shape"): + _mask_invalid_query_rows(target, predict, topk_indices, torch.ones(2, dtype=torch.bool)) + + def test_distribution_diagnostics_are_named_and_interval_gated(self, monkeypatch): + target = torch.tensor([[[0.8, 0.2], [0.0, 0.0]]], dtype=torch.float32) + predict = torch.tensor([[[0.5, 0.5], [0.0, 0.0]]], dtype=torch.float32) + topk_indices = torch.tensor([[[0, 1], [-1, -1]]], dtype=torch.int32) + messages = [] + loss_module = importlib.import_module("xtuner.v1.ops.sparse_mla.cudnn_dsa_indexer_loss") + monkeypatch.setattr(loss_module.log_rank0, "info", messages.append) + _INDEXER_LOSS_DEBUG_CALLS.clear() + + for _ in range(3): + _maybe_log_indexer_loss_diagnostics( + target, + predict, + topk_indices, + debug_name="layer2", + debug_interval=3, + ) + + assert len(messages) == 2 + assert "name=layer2 call=1 valid_rows=1" in messages[0] + assert "call=3" in messages[1] + assert "kl_mean=" in messages[0] + assert "target_entropy=" in messages[0] + assert "top1_match=" in messages[0] + + +@pytest.mark.skipif(not _cudnn_indexer_training_available(), reason="requires CUDA SM90+ and cuDNN DSA") +class TestCudnnDSAIndexerLoss: + def test_score_recompute_matches_pytorch_oracles_for_packed_padding(self): + torch.manual_seed(7) + device = "cuda" + seq_len = 128 + topk_indices = _packed_topk_indices((64, 64), topk=128, device=device) + topk_length = (topk_indices != -1).sum(dim=-1, dtype=torch.int32) + + index_q = torch.randn(1, seq_len, 32, 128, dtype=torch.bfloat16, device=device) + index_k = torch.randn(1, seq_len, 128, dtype=torch.bfloat16, device=device) + projected_weights = torch.randn(1, seq_len, 32, dtype=torch.bfloat16, device=device) + index_weights = projected_weights * ((32 * 128) ** -0.5) + + attn_q = torch.randn(1, seq_len, 64, 576, dtype=torch.bfloat16, device=device) + attn_k = torch.randn(1, seq_len, 576, dtype=torch.bfloat16, device=device) + softmax_scale = 576**-0.5 + expected_target, attn_lse = _attention_target_oracle(attn_q, attn_k, topk_indices, softmax_scale) + expected_predict = _indexer_predict_oracle(index_q, index_k, index_weights, topk_indices) + + actual_target = sparse_attention_target( + attn_q, + attn_k, + attn_lse, + topk_indices, + softmax_scale=softmax_scale, + topk_length=topk_length, + ) + actual_predict = sparse_indexer_predict( + index_q, + index_k, + index_weights, + topk_indices, + topk_length=topk_length, + ) + + torch.testing.assert_close(actual_target, expected_target, atol=2e-5, rtol=2e-5) + torch.testing.assert_close(actual_predict, expected_predict, atol=2e-5, rtol=2e-5) + torch.testing.assert_close(actual_target.sum(dim=-1), torch.ones_like(actual_target[..., 0])) + torch.testing.assert_close(actual_predict.sum(dim=-1), torch.ones_like(actual_predict[..., 0])) + assert torch.count_nonzero(actual_target.masked_select(topk_indices == -1)) == 0 + assert torch.count_nonzero(actual_predict.masked_select(topk_indices == -1)) == 0 + + def test_distribution_loss_and_gradients_match_pytorch_with_upstream_scale(self): + torch.manual_seed(11) + device = "cuda" + seq_len = 128 + topk_indices = _packed_topk_indices((128,), topk=128, device=device) + topk_length = (topk_indices != -1).sum(dim=-1, dtype=torch.int32) + row_coefficient = 0.017 + upstream_scale = 0.37 + + q_data = torch.randn(1, seq_len, 32, 128, dtype=torch.bfloat16, device=device) + k_data = torch.randn(1, seq_len, 128, dtype=torch.bfloat16, device=device) + w_data = torch.randn(1, seq_len, 32, dtype=torch.bfloat16, device=device) * ((32 * 128) ** -0.5) + + reference_q = q_data.detach().clone().requires_grad_() + reference_k = k_data.detach().clone().requires_grad_() + reference_w = w_data.detach().clone().requires_grad_() + reference_predict = _indexer_predict_oracle(reference_q, reference_k, reference_w, topk_indices) + target_logits = torch.randn_like(reference_predict) + target_logits = target_logits.masked_fill(topk_indices == -1, float("-inf")) + target = torch.softmax(target_logits, dim=-1) + expected_loss = _explicit_standard_kl( + target, + reference_predict, + topk_indices, + row_coefficient=row_coefficient, + ) + (expected_loss * upstream_scale).backward() + + actual_q = q_data.detach().clone().requires_grad_() + actual_k = k_data.detach().clone().requires_grad_() + actual_w = w_data.detach().clone().requires_grad_() + predict = sparse_indexer_predict( + actual_q, + actual_k, + actual_w, + topk_indices, + topk_length=topk_length, + ) + target_before = target.clone() + predict_before = predict.clone() + actual_loss = dsa_indexer_kl_from_distribution( + actual_q, + actual_k, + actual_w, + target, + predict, + topk_indices, + row_coefficient=row_coefficient, + ) + # Exercise the exact failure mode seen in distributed training: a + # contiguous scalar view can still have a 4-byte-offset data pointer. + upstream_storage = torch.tensor([0.0, upstream_scale], dtype=torch.float32, device=device) + misaligned_upstream_scale = upstream_storage[1] + assert misaligned_upstream_scale.is_contiguous() + assert misaligned_upstream_scale.data_ptr() % 16 != 0 + actual_loss.backward(gradient=misaligned_upstream_scale) + + torch.testing.assert_close(actual_loss, expected_loss, atol=2e-5, rtol=2e-5) + torch.testing.assert_close(target, target_before) + torch.testing.assert_close(predict, predict_before) + torch.testing.assert_close(actual_q.grad, reference_q.grad, atol=5e-2, rtol=5e-2) + torch.testing.assert_close(actual_k.grad, reference_k.grad, atol=8e-2, rtol=8e-2) + torch.testing.assert_close(actual_w.grad, reference_w.grad, atol=5e-2, rtol=5e-2) + + def test_padding_query_mask_zeros_cudnn_loss_and_gradients(self): + torch.manual_seed(13) + device = "cuda" + seq_len = 128 + valid_rows = 64 + topk_indices = _packed_topk_indices((seq_len,), topk=128, device=device) + topk_length = (topk_indices != -1).sum(dim=-1, dtype=torch.int32) + valid_query_mask = torch.arange(seq_len, device=device).unsqueeze(0) < valid_rows + row_coefficient = 1.0 / valid_rows + + q_data = torch.randn(1, seq_len, 32, 128, dtype=torch.bfloat16, device=device) + k_data = torch.randn(1, seq_len, 128, dtype=torch.bfloat16, device=device) + w_data = torch.randn(1, seq_len, 32, dtype=torch.bfloat16, device=device) * ((32 * 128) ** -0.5) + + reference_q = q_data.detach().clone().requires_grad_() + reference_k = k_data.detach().clone().requires_grad_() + reference_w = w_data.detach().clone().requires_grad_() + reference_predict = _indexer_predict_oracle(reference_q, reference_k, reference_w, topk_indices) + target_logits = torch.randn_like(reference_predict).masked_fill(topk_indices == -1, float("-inf")) + target = torch.softmax(target_logits, dim=-1) + masked_target, masked_predict, masked_topk = _mask_invalid_query_rows( + target, + reference_predict, + topk_indices, + valid_query_mask, + ) + expected_loss = _explicit_standard_kl( + masked_target, + masked_predict, + masked_topk, + row_coefficient=row_coefficient, + ) + expected_loss.backward() + + actual_q = q_data.detach().clone().requires_grad_() + actual_k = k_data.detach().clone().requires_grad_() + actual_w = w_data.detach().clone().requires_grad_() + predict = sparse_indexer_predict( + actual_q, + actual_k, + actual_w, + topk_indices, + topk_length=topk_length, + ) + actual_loss = dsa_indexer_kl_from_distribution( + actual_q, + actual_k, + actual_w, + target, + predict, + topk_indices, + row_coefficient=row_coefficient, + valid_query_mask=valid_query_mask, + ) + actual_loss.backward() + + torch.testing.assert_close(actual_loss, expected_loss, atol=2e-5, rtol=2e-5) + torch.testing.assert_close(actual_q.grad, reference_q.grad, atol=5e-2, rtol=5e-2) + torch.testing.assert_close(actual_k.grad, reference_k.grad, atol=8e-2, rtol=8e-2) + torch.testing.assert_close(actual_w.grad, reference_w.grad, atol=5e-2, rtol=5e-2) + assert torch.count_nonzero(actual_q.grad[:, valid_rows:]) == 0 + assert torch.count_nonzero(actual_w.grad[:, valid_rows:]) == 0 + assert torch.count_nonzero(actual_k.grad[:, valid_rows:]) == 0 + + def test_single_layer_convenience_wrapper_matches_distribution_api(self): + torch.manual_seed(19) + device = "cuda" + seq_len = 128 + topk_indices = _packed_topk_indices((128,), topk=128, device=device) + topk_length = (topk_indices != -1).sum(dim=-1, dtype=torch.int32) + index_q = torch.randn(1, seq_len, 32, 128, dtype=torch.bfloat16, device=device, requires_grad=True) + index_k = torch.randn(1, seq_len, 128, dtype=torch.bfloat16, device=device, requires_grad=True) + index_weights = ( + torch.randn(1, seq_len, 32, dtype=torch.bfloat16, device=device) * ((32 * 128) ** -0.5) + ).requires_grad_() + attn_q = torch.randn(1, seq_len, 64, 576, dtype=torch.bfloat16, device=device) + attn_k = torch.randn(1, seq_len, 576, dtype=torch.bfloat16, device=device) + softmax_scale = 576**-0.5 + _, attn_lse = _attention_target_oracle(attn_q, attn_k, topk_indices, softmax_scale) + + target = sparse_attention_target( + attn_q, + attn_k, + attn_lse, + topk_indices, + softmax_scale=softmax_scale, + topk_length=topk_length, + ) + predict = sparse_indexer_predict( + index_q, + index_k, + index_weights, + topk_indices, + topk_length=topk_length, + ) + expected = dsa_indexer_kl_from_distribution( + index_q, + index_k, + index_weights, + target, + predict, + topk_indices, + row_coefficient=0.02, + ) + actual = dsa_indexer_kl_loss( + index_q, + index_k, + index_weights, + attn_q, + attn_k, + attn_lse, + topk_indices, + topk_length=topk_length, + softmax_scale=softmax_scale, + row_coefficient=0.02, + ) + + torch.testing.assert_close(actual, expected, atol=2e-5, rtol=2e-5) diff --git a/xtuner/v1/data_proto/sequence_context.py b/xtuner/v1/data_proto/sequence_context.py index fa1829a91c..78dbfe0e0f 100644 --- a/xtuner/v1/data_proto/sequence_context.py +++ b/xtuner/v1/data_proto/sequence_context.py @@ -28,6 +28,7 @@ class DSATopKCacheState: offload_slot: int # Stable offload slot among concurrently active microbatches. mtp_forward_uses_remaining: dict[int, int] # Original-forward MTP uses left per shared source. mtp_replays_remaining: dict[int, int] # Backward MTP replays left per shared source. + indexer_losses: list[torch.Tensor] # Source-layer sparse KL terms owned by this microbatch. def __init__( self, @@ -39,6 +40,7 @@ def __init__( offload_slot: int = 0, mtp_forward_uses_remaining: dict[int, int] | None = None, mtp_replays_remaining: dict[int, int] | None = None, + indexer_losses: list[torch.Tensor] | None = None, ) -> None: # topk_indices format: {source_layer_idx: [seq_len, kv_group, topk]}. # Invalid/padded sparse slots are represented by -1. @@ -49,6 +51,7 @@ def __init__( self.offload_slot = offload_slot self.mtp_forward_uses_remaining = {} if mtp_forward_uses_remaining is None else mtp_forward_uses_remaining self.mtp_replays_remaining = {} if mtp_replays_remaining is None else mtp_replays_remaining + self.indexer_losses = [] if indexer_losses is None else indexer_losses # Avoid using dataclass decorator here to get rid of extra ops called in pytorch 2.8 and above diff --git a/xtuner/v1/loss/__init__.py b/xtuner/v1/loss/__init__.py index d2f20b3a16..a83ef75a07 100644 --- a/xtuner/v1/loss/__init__.py +++ b/xtuner/v1/loss/__init__.py @@ -2,6 +2,7 @@ from .base_loss_ctx import BaseLossConfig, BaseLossContext, BaseLossKwargs from .ce_loss import CELossConfig, CELossContext, LMHeadLossContext from .chunk_loss import ChunkLoss +from .dsa_indexer_loss import dense_dsa_indexer_kl_loss from .moe_loss import ( BalancingLossConfig, BalancingLossContext, @@ -26,6 +27,7 @@ "CELossContext", "CELossConfig", "ChunkLoss", + "dense_dsa_indexer_kl_loss", "BaseLossConfig", "BaseLossContext", "BaseLossKwargs", diff --git a/xtuner/v1/loss/dsa_indexer_loss.py b/xtuner/v1/loss/dsa_indexer_loss.py new file mode 100644 index 0000000000..46c137a7d5 --- /dev/null +++ b/xtuner/v1/loss/dsa_indexer_loss.py @@ -0,0 +1,121 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""Dense-teacher distillation loss for training a DSA indexer from scratch.""" + +from __future__ import annotations + +from functools import partial + +import torch +from torch import Tensor +from torch.utils.checkpoint import checkpoint + +from xtuner.v1.data_proto import SequenceContext + + +def _dense_indexer_kl_block( + index_q: Tensor, + index_k: Tensor, + index_weights: Tensor, + attn_q: Tensor, + attn_k: Tensor, + query_starts: Tensor, + query_ends: Tensor, + *, + softmax_scale: float, + row_coefficient: float, +) -> Tensor: + """Compute one causal query block without materializing ``S x S``.""" + + batch_size, query_len, _, _ = index_q.shape + key_len = index_k.shape[1] + if batch_size != 1: + raise ValueError(f"Dense DSA indexer warmup expects packed batch size 1, got {batch_size}.") + + key_positions = torch.arange(key_len, device=index_q.device) + causal_mask = (key_positions[None, :] >= query_starts[:, None]) & (key_positions[None, :] < query_ends[:, None]) + causal_mask = causal_mask.unsqueeze(0) + + index_scores = torch.einsum("bqjd,bkd->bqjk", index_q.float(), index_k.float()) + index_logits = torch.einsum("bqjk,bqj->bqk", torch.relu(index_scores), index_weights.float()) + index_logits = index_logits.masked_fill(~causal_mask, float("-inf")) + student_log_probs = torch.log_softmax(index_logits, dim=-1).masked_fill(~causal_mask, 0.0) + + teacher_logits = torch.einsum("bqhd,bkhd->bhqk", attn_q.float(), attn_k.float()) + teacher_logits = teacher_logits.mul(float(softmax_scale)) + teacher_logits = teacher_logits.masked_fill(~causal_mask.unsqueeze(1), float("-inf")) + teacher_probs = torch.softmax(teacher_logits, dim=-1).mean(dim=1) + + row_kl = torch.xlogy(teacher_probs, teacher_probs).sum(dim=-1) + row_kl = row_kl - (teacher_probs * student_log_probs).sum(dim=-1) + if row_kl.shape != (batch_size, query_len): + raise RuntimeError(f"Unexpected dense DSA KL shape: {tuple(row_kl.shape)}") + return row_kl.sum() * float(row_coefficient) + + +def dense_dsa_indexer_kl_loss( + index_q: Tensor, + index_k: Tensor, + index_weights: Tensor, + attn_q: Tensor, + attn_k: Tensor, + seq_ctx: SequenceContext, + *, + softmax_scale: float, + loss_coefficient: float = 1.0, + query_block_size: int = 256, + use_checkpoint: bool = True, +) -> Tensor: + """Distill packed causal dense attention into one DSA source indexer.""" + + if query_block_size <= 0: + raise ValueError(f"query_block_size must be positive, got {query_block_size}.") + if index_q.ndim != 4 or index_k.ndim != 3 or index_weights.ndim != 3: + raise ValueError("Indexer tensors must have shapes [B,S,H,D], [B,S,D], and [B,S,H].") + if attn_q.ndim != 4 or attn_k.ndim != 4: + raise ValueError("Teacher tensors must have shapes [B,S,H,D].") + if index_q.shape[:2] != index_k.shape[:2] or index_q.shape[:2] != index_weights.shape[:2]: + raise ValueError("Indexer Q, K, and weights must have matching batch/sequence dimensions.") + if attn_q.shape != attn_k.shape or attn_q.shape[:2] != index_q.shape[:2]: + raise ValueError("Dense teacher Q/K must match the indexer batch/sequence dimensions.") + if index_q.shape[-2] != index_weights.shape[-1] or index_q.shape[-1] != index_k.shape[-1]: + raise ValueError("Indexer head and feature dimensions are inconsistent.") + + valid_query_rows = index_q.shape[1] - seq_ctx.num_padding + if valid_query_rows <= 0: + return ( + index_q.sum(dtype=torch.float32) + + index_k.sum(dtype=torch.float32) + + index_weights.sum(dtype=torch.float32) + ) * 0.0 + if float(loss_coefficient) == 0.0: + return index_q.new_zeros((), dtype=torch.float32) + + starts, ends = seq_ctx.packed_causal_query_ranges(index_q.shape[1], index_q.device) + block_fn = partial( + _dense_indexer_kl_block, + softmax_scale=float(softmax_scale), + row_coefficient=float(loss_coefficient) / valid_query_rows, + ) + + loss = index_q.new_zeros((), dtype=torch.float32) + for block_start in range(0, valid_query_rows, query_block_size): + block_end = min(block_start + query_block_size, valid_query_rows) + block_args = ( + index_q[:, block_start:block_end], + index_k, + index_weights[:, block_start:block_end], + attn_q[:, block_start:block_end].detach(), + attn_k.detach(), + starts[block_start:block_end], + ends[block_start:block_end], + ) + block_loss = ( + checkpoint(block_fn, *block_args, use_reentrant=True) + if use_checkpoint and torch.is_grad_enabled() + else block_fn(*block_args) + ) + loss = loss + block_loss + return loss + + +__all__ = ["dense_dsa_indexer_kl_loss"] diff --git a/xtuner/v1/model/__init__.py b/xtuner/v1/model/__init__.py index 0c3dc28145..b4cf313ecc 100644 --- a/xtuner/v1/model/__init__.py +++ b/xtuner/v1/model/__init__.py @@ -22,6 +22,7 @@ from .dense.qwen2 import Qwen2Dense7BConfig, Qwen2DenseConfig from .dense.qwen3 import Qwen3Dense0P6BConfig, Qwen3Dense4BConfig, Qwen3Dense8BConfig, Qwen3DenseConfig from .moe.deepseek_v3 import DeepSeekV3Config +from .moe.glm47_flash import Glm47FlashConfig, Glm47FlashDSAConfig from .moe.glm52 import Glm52MoEConfig from .moe.gpt_oss import GptOss21BA3P6Config, GptOss117BA5P8Config, GptOssConfig from .moe.moe import BalancingLossConfig, MoE, MoEConfig, MoEModelOutputs, ZLossConfig @@ -40,6 +41,7 @@ "internvl-3.5-1b-hf": InternVL3P5Dense1BConfig(), "internvl-3.5-30b-a3b-hf": InternVL3P5MoE30BA3Config(), "qwen3.5-vl-4b": Qwen3_5_VLDense4BConfig(), + "glm-4.7-flash": Glm47FlashConfig(), "glm-5.2": Glm52MoEConfig(), } @@ -65,6 +67,8 @@ def get_model_config_from_hf(model_path: Path): return GptOssConfig.from_hf(model_path) elif cfg.model_type == "deepseek_v3": return DeepSeekV3Config.from_hf(model_path) + elif cfg.model_type == "glm4_moe_lite": + return Glm47FlashConfig.from_hf(model_path) elif cfg.model_type == "glm_moe_dsa": return Glm52MoEConfig.from_hf(model_path) else: @@ -79,6 +83,8 @@ def get_model_config_from_hf(model_path: Path): "Qwen3Dense8BConfig", "Qwen3MoEConfig", "Qwen3MoE30BA3Config", + "Glm47FlashConfig", + "Glm47FlashDSAConfig", "Glm52MoEConfig", "InternS1Config", "InternS1MiniConfig", diff --git a/xtuner/v1/model/moe/glm47_flash.py b/xtuner/v1/model/moe/glm47_flash.py new file mode 100644 index 0000000000..2c5ebf4e31 --- /dev/null +++ b/xtuner/v1/model/moe/glm47_flash.py @@ -0,0 +1,414 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import re +from pathlib import Path +from typing import Literal, cast + +import torch +from pydantic import Field, computed_field +from torch.distributed.fsdp import CPUOffloadPolicy, register_fsdp_forward_method +from typing_extensions import Self, override + + +try: + from transformers.models.glm4_moe_lite import Glm4MoeLiteConfig as HFGlm4MoeLiteConfig +except ImportError: + HFGlm4MoeLiteConfig = None # type: ignore[misc, assignment] +from xtuner.v1.model.moe.moe import BalancingLossConfig, MoEConfig, ZLossConfig +from xtuner.v1.module.attention import DSAMLAConfig, DSAMultiLatentAttention, MLAConfig +from xtuner.v1.module.attention.dsa_topk_sharing import ( + build_dsa_topk_release_plan, + configure_dsa_topk_decoder_lifecycle, +) +from xtuner.v1.module.mtp import MTPConfig, MTPLayer +from xtuner.v1.module.rope import RopeParametersConfig +from xtuner.v1.module.router.noaux_router import NoAuxRouterConfig +from xtuner.v1.utils import default_init_weights + +from .moe import MoE + + +class Glm47Flash(MoE): + """XTuner training model for Hugging Face ``glm4_moe_lite`` checkpoints.""" + + def to_hf_key_list(self, key: str) -> list[str]: + if self.config.tie_word_embeddings and "lm_head" in key: + key = key.replace("lm_head", "embed_tokens") + + if key.startswith("mtp_block."): + match = re.match(r"mtp_block\.layers\.(\d+)\.(.+)", key) + assert match is not None, f"Unexpected GLM-4.7 Flash MTP key: {key}" + mtp_layer_idx = self.config.num_hidden_layers + int(match.group(1)) + key = f"layers.{mtp_layer_idx}.{match.group(2)}" + key = key.replace(".decoder_layer.", ".") + key = re.sub(r"layers\.(\d+)\.final_layernorm\.", r"layers.\1.shared_head.norm.", key) + + if "layers" in key or "embed_tokens" in key: + key = "model." + key + + if "layers" in key: + key = re.sub(r"layers\.(\d+)\.(experts|gate|shared_experts)", r"layers.\1.mlp.\2", key) + + if "fused_w1w3.weight" in key: + expert_prefix = key.removesuffix(".fused_w1w3.weight") + return [ + f"{expert_prefix}.{expert_idx}.{projection}_proj.weight" + for expert_idx in range(self.config.n_routed_experts) + for projection in ("gate", "up") + ] + if "fused_w2.weight" in key: + expert_prefix = key.removesuffix(".fused_w2.weight") + return [ + f"{expert_prefix}.{expert_idx}.down_proj.weight" + for expert_idx in range(self.config.n_routed_experts) + ] + if key.startswith("norm."): + return [key.replace("norm.", "model.norm.")] + if "router.e_score_correction_bias" in key: + return [key.replace("router.e_score_correction_bias", "e_score_correction_bias")] + return [key] + +class Glm47FlashConfig(MoEConfig): + """Configuration for training the 30B-A3B GLM-4.7-Flash checkpoint.""" + + model_type: str = "glm4_moe_lite" + vocab_size: int = 154880 + max_position_embeddings: int = 202752 + pad_token_id: int | None = 154820 + eos_token_id: int = 154820 + hf_eos_token_id: int | list[int] = Field(default_factory=lambda: [154820, 154827, 154829]) + num_hidden_layers: int = 47 + first_k_dense_replace: int = 1 + hidden_size: int = 2048 + intermediate_size: int = 10240 + rms_norm_eps: float = 1e-5 + rope_parameters_cfg: RopeParametersConfig = Field( + default_factory=lambda: RopeParametersConfig(rope_theta=1000000.0) + ) + rope_interleave: bool = True + hidden_act: str = "silu" + attention: MLAConfig = MLAConfig( + kv_lora_rank=512, + q_lora_rank=768, + qk_nope_head_dim=192, + qk_rope_head_dim=64, + v_head_dim=256, + head_dim=64, + num_attention_heads=20, + qkv_bias=False, + o_bias=False, + rope_interleave=True, + ) + hf_head_dim: int = 64 + qk_head_dim: int = 256 + tie_word_embeddings: bool = False + n_routed_experts: int = 64 + n_shared_experts: int = 1 + num_experts_per_tok: int = 4 + hidden_factor: float = 1.0 + moe_intermediate_size: int = 1536 + router: NoAuxRouterConfig = NoAuxRouterConfig( + n_group=1, + topk_group=1, + scoring_func="sigmoid", + norm_topk_prob=True, + router_scaling_factor=1.8, + ) + balancing_loss_cfg: BalancingLossConfig | None = None + z_loss_cfg: ZLossConfig | None = None + mlp_layer_types: list[Literal["dense", "sparse"]] | None = None + num_nextn_predict_layers: int = 1 + mtp_config: MTPConfig | None = MTPConfig(num_layers=1, share_weights=True) + + @computed_field + def num_key_value_heads(self) -> int: + return self.attention.num_attention_heads + + def build(self) -> Glm47Flash: + return Glm47Flash(self) + + @classmethod + def from_hf(cls, hf_path: str | Path) -> Self: + if HFGlm4MoeLiteConfig is None: + raise ImportError("GLM-4.7 Flash requires a Transformers version with glm4_moe_lite support.") + cfg = HFGlm4MoeLiteConfig.from_pretrained(hf_path) + rope_parameters_cfg = RopeParametersConfig.from_hf_config(cfg) + mlp_layer_types = list(cfg.mlp_layer_types) + expected_mlp_types = ["dense"] * cfg.first_k_dense_replace + ["sparse"] * ( + cfg.num_hidden_layers - cfg.first_k_dense_replace + ) + if mlp_layer_types != expected_mlp_types: + raise ValueError("XTuner currently requires GLM-4.7 Flash dense MLP layers to form a prefix.") + + num_nextn_predict_layers = int(getattr(cfg, "num_nextn_predict_layers", 0)) + return cls( + vocab_size=cfg.vocab_size, + max_position_embeddings=cfg.max_position_embeddings, + pad_token_id=getattr(cfg, "pad_token_id", None), + eos_token_id=cfg.eos_token_id[0] if isinstance(cfg.eos_token_id, list) else cfg.eos_token_id, + hf_eos_token_id=cfg.eos_token_id, + num_hidden_layers=cfg.num_hidden_layers, + first_k_dense_replace=cfg.first_k_dense_replace, + hidden_size=cfg.hidden_size, + intermediate_size=cfg.intermediate_size, + rms_norm_eps=cfg.rms_norm_eps, + rope_parameters_cfg=rope_parameters_cfg, + rope_interleave=cfg.rope_interleave, + hidden_act=cfg.hidden_act, + attention=MLAConfig( + kv_lora_rank=cfg.kv_lora_rank, + q_lora_rank=cfg.q_lora_rank, + qk_nope_head_dim=cfg.qk_nope_head_dim, + qk_rope_head_dim=cfg.qk_rope_head_dim, + v_head_dim=cfg.v_head_dim, + head_dim=cfg.qk_rope_head_dim, + num_attention_heads=cfg.num_attention_heads, + qkv_bias=cfg.attention_bias, + o_bias=cfg.attention_bias, + dropout=cfg.attention_dropout, + rope_interleave=cfg.rope_interleave, + ), + hf_head_dim=cfg.qk_rope_head_dim, + qk_head_dim=cfg.qk_head_dim, + tie_word_embeddings=cfg.tie_word_embeddings, + n_routed_experts=cfg.n_routed_experts, + n_shared_experts=cfg.n_shared_experts, + num_experts_per_tok=cfg.num_experts_per_tok, + moe_intermediate_size=cfg.moe_intermediate_size, + router=NoAuxRouterConfig( + n_group=cfg.n_group, + topk_group=cfg.topk_group, + scoring_func="sigmoid", + norm_topk_prob=cfg.norm_topk_prob, + router_scaling_factor=cfg.routed_scaling_factor, + ), + balancing_loss_cfg=None, + mlp_layer_types=mlp_layer_types, + num_nextn_predict_layers=num_nextn_predict_layers, + mtp_config=MTPConfig(num_layers=num_nextn_predict_layers, share_weights=True) + if num_nextn_predict_layers + else None, + ) + + @property + def hf_config(self): + if HFGlm4MoeLiteConfig is None: + return None + return HFGlm4MoeLiteConfig( + architectures=["Glm4MoeLiteForCausalLM"], + vocab_size=self.vocab_size, + max_position_embeddings=self.max_position_embeddings, + pad_token_id=self.pad_token_id, + eos_token_id=self.hf_eos_token_id, + num_hidden_layers=self.num_hidden_layers, + first_k_dense_replace=self.first_k_dense_replace, + mlp_layer_types=self.mlp_layer_types, + hidden_size=self.hidden_size, + intermediate_size=self.intermediate_size, + moe_intermediate_size=self.moe_intermediate_size, + rms_norm_eps=self.rms_norm_eps, + rope_parameters=self.rope_parameters, + rope_interleave=self.rope_interleave, + hidden_act=self.hidden_act, + num_attention_heads=self.attention.num_attention_heads, + num_key_value_heads=self.attention.num_attention_heads, + kv_lora_rank=self.attention.kv_lora_rank, + q_lora_rank=self.attention.q_lora_rank, + qk_nope_head_dim=self.attention.qk_nope_head_dim, + qk_rope_head_dim=self.attention.qk_rope_head_dim, + v_head_dim=self.attention.v_head_dim, + attention_bias=self.attention.qkv_bias or self.attention.o_bias, + attention_dropout=self.attention.dropout, + n_routed_experts=self.n_routed_experts, + n_shared_experts=self.n_shared_experts, + num_experts_per_tok=self.num_experts_per_tok, + n_group=self.router.n_group, + topk_group=self.router.topk_group, + norm_topk_prob=self.router.norm_topk_prob, + routed_scaling_factor=self.router.router_scaling_factor, + tie_word_embeddings=self.tie_word_embeddings, + num_nextn_predict_layers=self.num_nextn_predict_layers, + dtype=torch.bfloat16, + ) + + +class Glm47FlashDSA(Glm47Flash): + """GLM-4.7 Flash with trainable main-stack and physical-MTP indexers.""" + + def _dsa_layers(self) -> list[tuple[torch.nn.Module, DSAMultiLatentAttention]]: + layers: list[tuple[torch.nn.Module, DSAMultiLatentAttention]] = [] + for decoder_layer in self.layers.values(): + self_attn = decoder_layer.self_attn # type: ignore[attr-defined] + if not isinstance(self_attn, DSAMultiLatentAttention): + raise TypeError(f"GLM-4.7 DSA requires DSAMultiLatentAttention, got {type(self_attn).__name__}.") + layers.append((decoder_layer, self_attn)) + + if self.mtp_block is not None and self.config.mtp_config is not None: + num_physical = 1 if self.config.mtp_config.share_weights else self.config.mtp_config.num_layers + for mtp_idx in range(num_physical): + mtp_layer = self.mtp_block.layers[mtp_idx] + if not isinstance(mtp_layer, MTPLayer): + raise TypeError(f"Expected MTPLayer, got {type(mtp_layer).__name__}.") + decoder_layer = mtp_layer.decoder_layer + self_attn = decoder_layer.self_attn # type: ignore[attr-defined] + if not isinstance(self_attn, DSAMultiLatentAttention): + raise TypeError( + f"GLM-4.7 DSA MTP requires DSAMultiLatentAttention, got {type(self_attn).__name__}." + ) + if self_attn.indexer_training is not None and not self_attn.indexer_training.train_mtp_indexer: + self_attn.disable_indexer_training() + layers.append((decoder_layer, self_attn)) + return layers + + @override + def _configure_model_specific_layer_lifecycle(self) -> None: + dsa_layers = self._dsa_layers() + sample_attention = dsa_layers[0][1] + if sample_attention.indexer_training is not None and sample_attention.indexer_training.indexer_only: + self.requires_grad_(False) + for _, self_attn in dsa_layers: + if self_attn.indexer_training is not None and self_attn.source_layer_idx == self_attn.layer_idx: + self_attn.indexer.requires_grad_(True) + + num_physical_mtp = len(dsa_layers) - self.config.num_hidden_layers + release_plan = build_dsa_topk_release_plan( + num_main_layers=self.config.num_hidden_layers, + num_mtp_layers=num_physical_mtp, + indexer_types=sample_attention.indexer_types, + index_skip_topk_offset=sample_attention.index_skip_topk_offset, + index_topk_freq=sample_attention.index_topk_freq, + ) + for decoder_layer, self_attn in dsa_layers: + configure_dsa_topk_decoder_lifecycle( + decoder_layer=decoder_layer, + attention=self_attn, + release_plan=release_plan, + ) + + @override + def _fully_shard_model_specific_submodules(self) -> None: + """Make side-channel indexer loss visible to composable FSDP.""" + + assert self.fsdp_config is not None + assert self.fsdp_mesh is not None + mesh = self.fsdp_mesh if self.hsdp_mesh is None else self.hsdp_mesh + offload_policy = CPUOffloadPolicy() if self.fsdp_config.cpu_offload else None + for _, self_attn in self._dsa_layers(): + if self_attn.source_layer_idx != self_attn.layer_idx or not hasattr(self_attn, "indexer"): + continue + indexer = self_attn.indexer + if not any(parameter.requires_grad for parameter in indexer.parameters()): + continue + self._fully_shard( + mesh=mesh, + mp_policy=self.mp_policy, + reshard_after_forward=self.fsdp_config.reshard_after_forward, + offload_policy=offload_policy, + module=indexer, + ) + register_fsdp_forward_method(indexer, "project_features") + + @override + @torch.no_grad() + def from_hf(self, hf_path: str | Path, strict: bool = True) -> tuple: + """Load official dense weights and initialize only absent indexers.""" + + loaded_keys, unloaded_keys, missing_keys = super().from_hf(hf_path, strict=False) + indexers: list[tuple[str, DSAMultiLatentAttention]] = [] + indexer_names: set[str] = set() + for layer_idx, decoder_layer in self.layers.items(): + indexers.append((f"layers.{layer_idx}.self_attn.indexer", decoder_layer.self_attn)) # type: ignore[attr-defined] + if self.mtp_block is not None and self.config.mtp_config is not None: + num_physical = 1 if self.config.mtp_config.share_weights else self.config.mtp_config.num_layers + for mtp_idx in range(num_physical): + attention = self.mtp_block.layers[mtp_idx].decoder_layer.self_attn + indexers.append((f"mtp_block.layers.{mtp_idx}.decoder_layer.self_attn.indexer", attention)) + + for prefix, attention in indexers: + self_attn = cast(DSAMultiLatentAttention, attention) + if self_attn.source_layer_idx != self_attn.layer_idx: + continue + current_indexer_names = { + self._clean_param_name(name) for name, _ in self_attn.indexer.named_parameters(prefix=prefix) + } + indexer_names.update(current_indexer_names) + loaded_indexer_names = loaded_keys & current_indexer_names + if loaded_indexer_names and loaded_indexer_names != current_indexer_names: + raise RuntimeError( + f"Partially loaded DSA indexer {prefix}: {sorted(current_indexer_names - loaded_indexer_names)}" + ) + if not loaded_indexer_names: + default_init_weights(self_attn.indexer) + + if unloaded_base_keys := unloaded_keys - indexer_names: + raise RuntimeError(f"Failed to load GLM-4.7 base weights: {sorted(unloaded_base_keys)}") + return loaded_keys, unloaded_keys, missing_keys + + +class Glm47FlashDSAConfig(Glm47FlashConfig): + """Two-stage DSA conversion config for official GLM-4.7 Flash.""" + + model_type: str = "glm4_moe_lite_dsa" + attention: DSAMLAConfig = DSAMLAConfig( + kv_lora_rank=512, + q_lora_rank=768, + qk_nope_head_dim=192, + qk_rope_head_dim=64, + v_head_dim=256, + head_dim=64, + num_attention_heads=20, + qkv_bias=False, + o_bias=False, + rope_interleave=True, + index_topk=2048, + index_head_dim=128, + index_n_heads=32, + index_topk_freq=4, + index_skip_topk_offset=1, + indexer_rope_interleave=True, + ) + + def _normalize_indexer_types(self) -> None: + num_physical_mtp = 0 + if self.mtp_config is not None: + num_physical_mtp = 1 if self.mtp_config.share_weights else self.mtp_config.num_layers + expected = [ + "full" if layer_idx % self.attention.index_topk_freq == 0 else "shared" + for layer_idx in range(self.num_hidden_layers) + ] + expected.extend(["full"] * num_physical_mtp) + if self.attention.indexer_types is None: + self.attention.indexer_types = expected + elif self.attention.indexer_types != expected: + raise ValueError(f"GLM-4.7 DSA expects indexer_types={expected}, got {self.attention.indexer_types}.") + + def build(self) -> Glm47FlashDSA: + self._normalize_indexer_types() + return Glm47FlashDSA(self) + + @classmethod + def from_hf(cls, hf_path: str | Path) -> Self: + dense_config = Glm47FlashConfig.from_hf(hf_path) + dense_attention = dense_config.attention + config_values = { + name: getattr(dense_config, name) + for name in Glm47FlashConfig.model_fields + if name not in ("attention", "model_type") + } + config = cls( + **config_values, + attention=DSAMLAConfig( + **dense_attention.model_dump(), + index_topk=2048, + index_head_dim=128, + index_n_heads=32, + index_topk_freq=4, + index_skip_topk_offset=1, + indexer_rope_interleave=dense_attention.rope_interleave, + ), + ) + # The official checkpoint has one physical MTP layer and one training + # prediction depth. No training-time weight-sharing loop is needed. + if config.mtp_config is not None: + config.mtp_config = config.mtp_config.model_copy(update={"share_weights": False}) + config._normalize_indexer_types() + return config diff --git a/xtuner/v1/model/moe/glm52.py b/xtuner/v1/model/moe/glm52.py index 838546677e..42ace4ec75 100644 --- a/xtuner/v1/model/moe/glm52.py +++ b/xtuner/v1/model/moe/glm52.py @@ -84,11 +84,25 @@ def _configure_model_specific_layer_lifecycle(self) -> None: assert isinstance(self_attn, DSAMultiLatentAttention), ( f"GLM-5.2 MTP requires DSAMultiLatentAttention, got {type(self_attn).__name__}." ) + # Keep the checkpoint-backed physical MTP indexer frozen unless + # its training is explicitly requested. This happens before + # optimizer construction and preserves the previous baseline. + if self_attn.indexer_training is not None and not self_attn.indexer_training.train_mtp_indexer: + self_attn.disable_indexer_training() dsa_layers.append((decoder_layer, self_attn)) if mtp_idx == 0: mtp_attention = self_attn sample_attn = dsa_layers[0][1] + if sample_attn.indexer_training is not None and sample_attn.indexer_training.indexer_only: + # Make the sparse-attention teacher stationary for the strict + # overfit gate. Only enabled Full/source indexers are restored to + # trainable; shared layers own no indexer, and MTP remains frozen + # unless ``train_mtp_indexer`` was explicitly enabled. + self.requires_grad_(False) + for _, self_attn in dsa_layers: + if self_attn.indexer_training is not None and self_attn.source_layer_idx == self_attn.layer_idx: + self_attn.indexer.requires_grad_(True) release_plan = build_dsa_topk_release_plan( num_main_layers=self.config.num_hidden_layers, num_mtp_layers=num_physical_mtp_layers, diff --git a/xtuner/v1/model/moe/moe.py b/xtuner/v1/model/moe/moe.py index f27e0a2dbc..a4b35dcee7 100644 --- a/xtuner/v1/model/moe/moe.py +++ b/xtuner/v1/model/moe/moe.py @@ -108,6 +108,7 @@ class MoEModelOutputs(ModelOutputs): z_loss: torch.Tensor | None = None tokens_per_expert_global: torch.Tensor mtp_loss: torch.Tensor | None = None + indexer_loss: torch.Tensor | None = None def free_nongrad_feature(self): """Release large intermediate tensors not needed for backward or @@ -241,6 +242,11 @@ def _maybe_offload_router(self, tensor: torch.Tensor) -> torch.Tensor: def _configure_model_specific_layer_lifecycle(self) -> None: return + def _fully_shard_model_specific_submodules(self) -> None: + """Shard model-specific children before their enclosing layers.""" + + return + def _z_loss_dist_token_count( self, z_ctx: list[ZLossContext] | ZLossContext | None, @@ -483,6 +489,21 @@ def _prepare_seq_ctx_topk_cache(seq_ctx_list: Sequence[SequenceContext]) -> None for offload_slot, seq_ctx in enumerate(seq_ctx_list): seq_ctx.dsa_topk_cache.offload_slot = offload_slot + @staticmethod + def _consume_indexer_losses(seq_ctx_list: Sequence[SequenceContext]) -> torch.Tensor | None: + """Average source-layer indexer losses and release their Python owners.""" + + losses = [loss for seq_ctx in seq_ctx_list for loss in seq_ctx.dsa_topk_cache.indexer_losses] + for seq_ctx in seq_ctx_list: + seq_ctx.dsa_topk_cache.indexer_losses.clear() + if not losses: + return None + return torch.stack(losses).mean() + + def _indexer_only_training(self) -> bool: + training_cfg = getattr(getattr(self.config, "attention", None), "indexer_training", None) + return bool(training_cfg is not None and training_cfg.indexer_only) + def _micro_batch_forward( self, seq_ctx_list: list[SequenceContext], @@ -536,6 +557,7 @@ def _micro_batch_forward( # Initialize output containers output: dict = {} + indexer_only = self._indexer_only_training() # Only the logits side is ever exposed to callers in the micro-batch path; the # weights side is not part of the returned schema, so we never accumulate it. @@ -617,6 +639,7 @@ def _micro_batch_forward( ) assert hidden_states_list, "XTuner Internal Error, found empty hidden states for domino EP" + indexer_loss_contexts = list(seq_ctx_list) if self.mtp_block is not None: assert self.config.mtp_config is not None @@ -645,25 +668,29 @@ def _micro_batch_forward( position_embeddings=position_embeddings_list, seq_ctx=mtp_seq_ctx_list, ) + indexer_loss_contexts.extend(mtp_seq_ctx_list) mtp_losses = torch.tensor(0.0, device=DEVICE) has_mtp_loss = False - for micro_batch_idx, (loss_ctx_dict, mtp_outputs) in enumerate(zip(loss_ctx_list, mtp_outputs_per_mb)): - mtp_loss_ctx_list = loss_ctx_dict.get("mtp") - if mtp_loss_ctx_list is None: - continue - - micro_batch_mtp_losses = torch.tensor(0.0, device=DEVICE) - for mtp_idx, (mtp_hidden, mtp_ctx) in enumerate(zip(mtp_outputs, mtp_loss_ctx_list)): - mtp_hidden_states, mtp_router_results, _, _ = mtp_hidden - mtp_loss, _ = self.lm_head(mtp_hidden_states, cast(MTPLossContext, mtp_ctx)) - micro_batch_mtp_losses += mtp_loss + if not indexer_only: + for micro_batch_idx, (loss_ctx_dict, mtp_outputs) in enumerate( + zip(loss_ctx_list, mtp_outputs_per_mb) + ): + mtp_loss_ctx_list = loss_ctx_dict.get("mtp") + if mtp_loss_ctx_list is None: + continue - if keep_router: - router_logits_list[micro_batch_idx][f"mtp_layer{mtp_idx}"] = mtp_router_results + micro_batch_mtp_losses = torch.tensor(0.0, device=DEVICE) + for mtp_idx, (mtp_hidden, mtp_ctx) in enumerate(zip(mtp_outputs, mtp_loss_ctx_list)): + mtp_hidden_states, mtp_router_results, _, _ = mtp_hidden + mtp_loss, _ = self.lm_head(mtp_hidden_states, cast(MTPLossContext, mtp_ctx)) + micro_batch_mtp_losses += mtp_loss + + if keep_router: + router_logits_list[micro_batch_idx][f"mtp_layer{mtp_idx}"] = mtp_router_results - mtp_losses += micro_batch_mtp_losses / len(mtp_loss_ctx_list) - has_mtp_loss = True + mtp_losses += micro_batch_mtp_losses / len(mtp_loss_ctx_list) + has_mtp_loss = True if has_mtp_loss: # MTP routed experts feed the same balancing / z aux loss as the main MoE layers @@ -707,22 +734,28 @@ def _micro_batch_forward( output["mtp_loss"] = mtp_losses * self.config.mtp_config.loss_scaling_factor - # Apply final norm to all micro-batches - cat_hidden_states = torch.cat(hidden_states_list, dim=1) - cat_hidden_states = self.norm(cat_hidden_states) - - # Process final outputs for each micro-batch - # Extract LM loss context from dict - lm_loss_ctx_list = [loss_ctx_dict["lm"] for loss_ctx_dict in loss_ctx_list] - cat_loss_ctx = type(lm_loss_ctx_list[0]).cat(lm_loss_ctx_list) - loss, (logits, extra_info) = self.lm_head(cat_hidden_states, cast(LMHeadLossContext, cat_loss_ctx)) - - # Aggregate losses (mean across micro-batches) - output["loss"] = loss.sum() - moe_extra_info = ModelForwardExtraLogInfo() - if extra_info: - moe_extra_info.append(extra_info) - output["extra_info"] = moe_extra_info + indexer_loss = self._consume_indexer_losses(indexer_loss_contexts) + if indexer_loss is not None: + output["indexer_loss"] = indexer_loss + + logits = None + if not indexer_only: + # Apply final norm to all micro-batches + cat_hidden_states = torch.cat(hidden_states_list, dim=1) + cat_hidden_states = self.norm(cat_hidden_states) + + # Process final outputs for each micro-batch + # Extract LM loss context from dict + lm_loss_ctx_list = [loss_ctx_dict["lm"] for loss_ctx_dict in loss_ctx_list] + cat_loss_ctx = type(lm_loss_ctx_list[0]).cat(lm_loss_ctx_list) + loss, (logits, extra_info) = self.lm_head(cat_hidden_states, cast(LMHeadLossContext, cat_loss_ctx)) + + # Aggregate losses (mean across micro-batches) + output["loss"] = loss.sum() + moe_extra_info = ModelForwardExtraLogInfo() + if extra_info: + moe_extra_info.append(extra_info) + output["extra_info"] = moe_extra_info split_aux_output = self.aux_loss.finalize( balancing_ctx=balancing_ctx, @@ -779,6 +812,7 @@ def _forward( position_embeddings = self.rotary_emb(hidden_states, position_ids) output: dict = {} # type: ignore + indexer_only = self._indexer_only_training() if self.config.return_hidden_states: output["hidden_states"] = [] @@ -847,21 +881,20 @@ def _forward( output["hidden_states"].append(hidden_states) layer_hidden_states = hidden_states - hidden_states = self.norm(hidden_states) + if not indexer_only: + hidden_states = self.norm(hidden_states) - # Get LM loss context from dict - lm_loss_ctx = loss_ctx["lm"] if loss_ctx is not None else None - loss, (logits, extra_info) = self.lm_head(hidden_states, lm_loss_ctx) # type: ignore - output["loss"] = loss - output["logits"] = logits - output["extra_info"] = extra_info + # Get LM loss context from dict + lm_loss_ctx = loss_ctx["lm"] if loss_ctx is not None else None + loss, (logits, extra_info) = self.lm_head(hidden_states, lm_loss_ctx) # type: ignore + output["loss"] = loss + output["logits"] = logits + output["extra_info"] = extra_info + indexer_loss_contexts = [seq_ctx] # MTP forward pass and loss computation - if ( - self.mtp_block is not None - and loss_ctx is not None - and (mtp_loss_ctx_list := loss_ctx.get("mtp")) is not None - ): + mtp_loss_ctx_list = loss_ctx.get("mtp") if loss_ctx is not None else None + if self.mtp_block is not None and (indexer_only or mtp_loss_ctx_list is not None): mtp_seq_ctx = seq_ctx.copy( input_ids=input_ids.clone() if input_ids is not None else None, position_ids=position_ids.clone(), @@ -884,39 +917,48 @@ def _forward( position_embeddings=position_embeddings, seq_ctx=mtp_seq_ctx, ) + indexer_loss_contexts.append(mtp_seq_ctx) - # Compute MTP losses for each depth - mtp_losses = torch.tensor(0.0, device=DEVICE) - for idx, (mtp_hidden, mtp_ctx) in enumerate(zip(mtp_outputs, mtp_loss_ctx_list)): - mtp_hidden_states, mtp_router_results, mtp_router_weights, mtp_router_topk_ids = mtp_hidden + if not indexer_only: + assert mtp_loss_ctx_list is not None + # Compute MTP losses for each depth + mtp_losses = torch.tensor(0.0, device=DEVICE) + for idx, (mtp_hidden, mtp_ctx) in enumerate(zip(mtp_outputs, mtp_loss_ctx_list)): + mtp_hidden_states, mtp_router_results, mtp_router_weights, mtp_router_topk_ids = mtp_hidden - if keep_router: - output["router_logits"][f"mtp_layer{idx}"] = mtp_router_results - output["router_weights"][f"mtp_layer{idx}"] = mtp_router_weights - # Inject this MTP layer's z-loss before lm_head so backward through mtp_loss - # traverses the AuxLossScaler node and releases this layer's logsumexp activations. - mtp_hidden_states = self.aux_loss.accumulate( - selected_router_weights=mtp_router_weights.index_select(0, mtp_nonpad_indices) - .contiguous() - .float(), - selected_router_logits=mtp_router_results.index_select(0, mtp_nonpad_indices).contiguous().float(), - selected_experts=mtp_router_topk_ids.index_select(0, mtp_nonpad_indices).contiguous(), - hidden_states=mtp_hidden_states, - balancing_ctx=balancing_ctx, - z_ctx=z_ctx, - num_tokens_local=mtp_non_pad_token, - num_tokens_global=mtp_num_tokens_global, - world_size=mtp_z_world_size, - ) - mtp_loss, _ = self.lm_head(mtp_hidden_states, cast(MTPLossContext, mtp_ctx)) - mtp_losses += mtp_loss + if keep_router: + output["router_logits"][f"mtp_layer{idx}"] = mtp_router_results + output["router_weights"][f"mtp_layer{idx}"] = mtp_router_weights + # Inject this MTP layer's z-loss before lm_head so backward through mtp_loss + # traverses the AuxLossScaler node and releases this layer's logsumexp activations. + mtp_hidden_states = self.aux_loss.accumulate( + selected_router_weights=mtp_router_weights.index_select(0, mtp_nonpad_indices) + .contiguous() + .float(), + selected_router_logits=mtp_router_results.index_select(0, mtp_nonpad_indices) + .contiguous() + .float(), + selected_experts=mtp_router_topk_ids.index_select(0, mtp_nonpad_indices).contiguous(), + hidden_states=mtp_hidden_states, + balancing_ctx=balancing_ctx, + z_ctx=z_ctx, + num_tokens_local=mtp_non_pad_token, + num_tokens_global=mtp_num_tokens_global, + world_size=mtp_z_world_size, + ) + mtp_loss, _ = self.lm_head(mtp_hidden_states, cast(MTPLossContext, mtp_ctx)) + mtp_losses += mtp_loss + + # Average MTP losses across depths and scale + mtp_losses = mtp_losses / len(mtp_loss_ctx_list) + scaled_mtp_loss = mtp_losses * self.config.mtp_config.loss_scaling_factor # type: ignore - # Average MTP losses across depths and scale - mtp_losses = mtp_losses / len(mtp_loss_ctx_list) - scaled_mtp_loss = mtp_losses * self.config.mtp_config.loss_scaling_factor # type: ignore + # Add to total loss + output["mtp_loss"] = scaled_mtp_loss - # Add to total loss - output["mtp_loss"] = scaled_mtp_loss + indexer_loss = self._consume_indexer_losses(indexer_loss_contexts) + if indexer_loss is not None: + output["indexer_loss"] = indexer_loss split_aux_output = self.aux_loss.finalize( balancing_ctx=balancing_ctx, @@ -1146,6 +1188,8 @@ def fully_shard( param_dtype=self.fsdp_config.param_dtype, reduce_dtype=fsdp_config.reduce_dtype ) + self._fully_shard_model_specific_submodules() + for layer_idx, layer in tqdm(self.layers.items(), desc="[FSDP Sharding]"): layer_idx = int(layer_idx) if self._should_recompute( diff --git a/xtuner/v1/module/attention/__init__.py b/xtuner/v1/module/attention/__init__.py index dedd2b1451..cb36a3c0a3 100644 --- a/xtuner/v1/module/attention/__init__.py +++ b/xtuner/v1/module/attention/__init__.py @@ -1,6 +1,6 @@ # Copyright (c) OpenMMLab. All rights reserved. from .attn_outputs import AttnOutputs -from .dsa_mla import DSAMLAConfig, DSAMultiLatentAttention +from .dsa_mla import DSAIndexerTrainingConfig, DSAMLAConfig, DSAMultiLatentAttention from .gated_deltanet import GatedDeltaNet, GatedDeltaNetConfig from .mha import MHAConfig, MultiHeadAttention from .mla import MLAConfig, MultiLatentAttention @@ -13,6 +13,7 @@ "MHAConfig", "MLAConfig", "DSAMLAConfig", + "DSAIndexerTrainingConfig", "AttnOutputs", "GatedDeltaNet", "GatedDeltaNetConfig", diff --git a/xtuner/v1/module/attention/dsa_mla.py b/xtuner/v1/module/attention/dsa_mla.py index 6a22b97dbc..22b8f7b8a1 100644 --- a/xtuner/v1/module/attention/dsa_mla.py +++ b/xtuner/v1/module/attention/dsa_mla.py @@ -1,18 +1,21 @@ # Copyright (c) OpenMMLab. All rights reserved. -from typing import Literal, cast +from typing import Literal, NamedTuple, cast import torch +from pydantic import BaseModel, ConfigDict, Field, model_validator from torch import nn from torch.distributed.tensor import DTensor from xtuner.v1.config import GenerateConfig from xtuner.v1.data_proto import SequenceContext from xtuner.v1.float8.config import Float8Config +from xtuner.v1.loss import dense_dsa_indexer_kl_loss from xtuner.v1.module.rope import RopeScalingConfig from xtuner.v1.ops.comm import gather_for_sequence_parallel from xtuner.v1.ops.sparse_mla import ( DSATopKIndicesProtocol, SparseMLAProtocol, + dsa_indexer_kl_loss, ensure_cudnn_dsa_runtime_available, ensure_tilelang_runtime_available, get_dsa_topk_indices, @@ -22,7 +25,7 @@ from ..linear import build_linear from .attn_outputs import AttnOutputs from .dsa_topk_sharing import build_dsa_topk_release_plan, dsa_topk_source_layer, get_dsa_topk_sharing_runtime -from .mla import MLAConfig, MultiLatentAttention, mla_apply_rotary_pos_emb +from .mla import MLAConfig, MultiLatentAttention, mla_apply_rotary_pos_emb, mla_apply_rotary_pos_emb_non_interleaved class LayerNorm(nn.Module): @@ -57,6 +60,48 @@ def extra_repr(self): return f"{self.normalized_shape}, eps={self.eps}" +class DSAIndexerTrainingConfig(BaseModel): + """Two-stage DSA indexer distillation configuration. + + ``None`` on :class:`DSAMLAConfig` remains the strict frozen baseline. This + ``dense_warmup`` keeps the dense attention teacher frozen and trains new + source indexers with a query-blocked PyTorch KL. ``sparse`` runs the + selected sparse backend. Each source indexer is supervised by the attention + layer that owns it; shared consumer layers only reuse its discrete top-k. + + Sequence parallelism and decoder activation checkpoint replay are not + supported while indexer loss is active. ``train_mtp_indexer`` additionally + enables the checkpoint-backed physical MTP source indexer; when one + physical MTP layer is reused for multiple prediction depths, only the + first (index-compute) depth is supervised. + + ``indexer_only`` is a diagnostic overfit mode: GLM freezes the teacher and + every non-indexer parameter, leaving only the selected source indexers + trainable. ``debug_interval`` prints per-source teacher/student + distribution statistics without changing the loss. + """ + + model_config = ConfigDict(extra="forbid") + stage: Literal["dense_warmup", "sparse"] = "sparse" + loss_coeff: float = Field(default=1.0, ge=0.0) + train_mtp_indexer: bool = False + indexer_only: bool = False + dense_query_block_size: int = Field(default=256, ge=1) + debug_interval: int = Field(default=0, ge=0) + + @model_validator(mode="after") + def validate_stage(self) -> "DSAIndexerTrainingConfig": + if self.stage == "dense_warmup" and not self.indexer_only: + raise ValueError("dense_warmup requires indexer_only=True so the full-attention teacher stays frozen.") + return self + + +class DSAIndexerFeatures(NamedTuple): + q: torch.Tensor + k: torch.Tensor + weights: torch.Tensor + + class DSAIndexer(nn.Module): def __init__( self, @@ -67,13 +112,16 @@ def __init__( index_head_dim: int, index_n_heads: int, index_topk: int, + rope_interleave: bool = True, indexer_backend: Literal["torch", "tilelang", "cudnn_dsa"] = "torch", + trainable: bool = False, ): super().__init__() self.qk_rope_head_dim = qk_rope_head_dim self.index_head_dim = index_head_dim self.index_n_heads = index_n_heads self.index_topk = index_topk + self.rope_interleave = rope_interleave self.indexer_backend = indexer_backend self.dsa_topk_indices_func: DSATopKIndicesProtocol = get_dsa_topk_indices(indexer_backend) # wq_b.weight: [index_n_heads * index_head_dim, q_lora_rank] @@ -83,19 +131,20 @@ def __init__( self.k_norm = LayerNorm(index_head_dim, eps=1e-6) # weights_proj.weight: [index_n_heads, hidden_size] self.weights_proj = build_linear(hidden_size, index_n_heads, bias=False) - # The indexer only produces integer DSA top-k IDs under no_grad, so its - # parameters must not be registered with the training optimizer. - self.requires_grad_(False) + # ``trainable=False`` is the historical and strict frozen baseline. + # Top-k selection itself remains no-grad even when sparse KL training is + # enabled; only ``project_features`` participates in autograd. + if not trainable: + self.requires_grad_(False) - @torch.no_grad() - def forward( + def project_features( self, hidden_states: torch.Tensor, q_resid: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], seq_ctx: SequenceContext, - ) -> torch.Tensor: - """Compute DSA top-k indices for each local query token. + ) -> DSAIndexerFeatures: + """Project indexer features while preserving an optional autograd graph. Shapes use ``S`` for the local sequence length and ``S_g`` for the SP-gathered global KV length. Numbers below follow GLM-5.2 defaults @@ -131,9 +180,11 @@ def forward( k_pe, k_nope = torch.split(k, [self.qk_rope_head_dim, self.index_head_dim - self.qk_rope_head_dim], dim=-1) k_pe = k_pe.view(bsz, seq_len, 1, self.qk_rope_head_dim).transpose(1, 2) - # GLM-MoE-DSA applies interleaved RoPE in the indexer, matching HF PR #46842. + # GLM-5.2 uses interleaved RoPE. GLM-4.7 follows its checkpoint + # setting for both the dense teacher and newly initialized indexer. cos, sin = position_embeddings - q_pe, k_pe = mla_apply_rotary_pos_emb(q_pe, k_pe, cos, sin) + rope_fn = mla_apply_rotary_pos_emb if self.rope_interleave else mla_apply_rotary_pos_emb_non_interleaved + q_pe, k_pe = rope_fn(q_pe, k_pe, cos, sin) # q_pe: [bsz, S, Ni, Dr]; k_pe: [bsz, S, Dr] q_pe = q_pe.transpose(1, 2) k_pe = k_pe.transpose(1, 2).squeeze(2) @@ -141,10 +192,8 @@ def forward( # q: [bsz, S, Ni, Di]; k: [bsz, S, Di] q = torch.cat([q_pe, q_nope], dim=-1) k = torch.cat([k_pe, k_nope], dim=-1) - # weights: [bsz, S, Ni] - weights = self.weights_proj(hidden_states).float() * (self.index_n_heads**-0.5) - - # Top-k 索引是整数,不需要梯度,所以整个 indexer 都放在 no_grad 下。 + weights = self.weights_proj(hidden_states) + # Top-k 索引是整数,不需要梯度,所以 selection 始终放在 no_grad 下。 # 这解释了 Case 1 为什么只在 compile 下显错: # eager COMPUTE: indexer 不产生槽位 -> SparseMLA 保存 [A, B, C] # eager REUSE: cache read 不产生槽位 -> SparseMLA 保存 [A, B, C] @@ -156,16 +205,33 @@ def forward( # Index Q 按 query token 保持分片,只有 K 需要全局 gather。 # k: [bsz, S_g, Di] k = gather_for_sequence_parallel(k, dim=1, sp_mesh=seq_ctx.sequence_parallel_mesh) + return DSAIndexerFeatures(q, k, weights) + + @torch.no_grad() + def select_topk(self, features: DSAIndexerFeatures, seq_ctx: SequenceContext) -> torch.Tensor: + """Select integer sparse IDs without retaining the indexer graph.""" + # returns topk_indices: [S, 1, K] return self.dsa_topk_indices_func( - q, - k, - weights, + features.q.detach(), + features.k.detach(), + features.weights.detach().float() * (self.index_n_heads**-0.5), seq_ctx, index_head_dim=self.index_head_dim, index_topk=self.index_topk, ) + @torch.no_grad() + def forward( + self, + hidden_states: torch.Tensor, + q_resid: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + seq_ctx: SequenceContext, + ) -> torch.Tensor: + features = self.project_features(hidden_states, q_resid, position_embeddings, seq_ctx) + return self.select_topk(features, seq_ctx) + class DSAMLAConfig(MLAConfig): index_topk: int @@ -176,6 +242,7 @@ class DSAMLAConfig(MLAConfig): indexer_rope_interleave: bool = True indexer_types: list[str] | None = None sparse_mla_backend: Literal["torch", "tilelang", "cudnn_dsa"] = "torch" + indexer_training: DSAIndexerTrainingConfig | None = None def build( self, @@ -214,6 +281,7 @@ def __init__( indexer_rope_interleave: bool = True, indexer_types: list[str] | None = None, sparse_mla_backend: Literal["torch", "tilelang", "cudnn_dsa"] = "torch", + indexer_training: DSAIndexerTrainingConfig | dict | None = None, **kwargs, ): super().__init__(**kwargs) @@ -240,6 +308,9 @@ def __init__( self.indexer_rope_interleave = indexer_rope_interleave self.indexer_types = indexer_types self.sparse_mla_backend = sparse_mla_backend + self.indexer_training = ( + None if indexer_training is None else DSAIndexerTrainingConfig.model_validate(indexer_training) + ) self.sparse_mla_func: SparseMLAProtocol = get_sparse_mla(sparse_mla_backend) if indexer_types is None: self.dsa_topk_last_use, self.dsa_topk_recompute_release = {}, {} @@ -273,9 +344,17 @@ def __init__( index_head_dim=self.index_head_dim, index_n_heads=self.index_n_heads, index_topk=self.index_topk, + rope_interleave=self.indexer_rope_interleave, indexer_backend=self.sparse_mla_backend, + trainable=self.indexer_training is not None, ) + def disable_indexer_training(self) -> None: + """Restore strict frozen behavior, used by physical MTP layers.""" + + self.indexer_training = None + if hasattr(self, "indexer"): + self.indexer.requires_grad_(False) def get_muon_split_sizes(self) -> dict[nn.Parameter, tuple[int, ...]]: """Return the logical row blocks used by GLM MuonSplit.""" return { @@ -286,11 +365,101 @@ def get_muon_split_sizes(self) -> dict[nn.Parameter, tuple[int, ...]]: * self.num_attention_heads, } + def _validate_indexer_training_runtime(self, seq_ctx: SequenceContext) -> None: + if self.training and not torch.is_grad_enabled(): + raise RuntimeError( + "DSA indexer training does not support activation checkpointing; set recompute_ratio=0." + ) + if seq_ctx.sequence_parallel_mesh is not None and seq_ctx.sequence_parallel_mesh.size() > 1: + raise RuntimeError("DSA indexer training requires sequence parallel size 1.") + + @torch.no_grad() + def _project_dense_teacher_states( + self, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Return Q-LoRA residual and explicit dense teacher Q/K states.""" + + bsz, q_len, _ = hidden_states.shape + assert self.q_lora_rank is not None + q_resid = self.q_a_layernorm(self.q_a_proj(hidden_states)) + q = self.q_b_proj(q_resid).view(bsz, q_len, self.num_attention_heads, self.q_head_dim) + q_nope, q_pe = torch.split(q, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) + + compressed_kv = self.kv_a_proj_with_mqa(hidden_states) + compressed_kv, k_pe = torch.split(compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) + kv = self.kv_b_proj(self.kv_a_layernorm(compressed_kv)).view( + bsz, + q_len, + self.num_attention_heads, + self.qk_nope_head_dim + self.v_head_dim, + ) + k_nope = kv[..., : self.qk_nope_head_dim] + + q_pe = q_pe.transpose(1, 2) + k_pe = k_pe.view(bsz, q_len, 1, self.qk_rope_head_dim).transpose(1, 2) + cos, sin = position_embeddings + rope_fn = mla_apply_rotary_pos_emb if self.rope_interleave else mla_apply_rotary_pos_emb_non_interleaved + q_pe, k_pe = rope_fn(q_pe, k_pe, cos, sin) + q_pe = q_pe.transpose(1, 2) + k_pe = k_pe.transpose(1, 2).expand(-1, -1, self.num_attention_heads, -1) + return q_resid, torch.cat([q_nope, q_pe], dim=-1), torch.cat([k_nope, k_pe], dim=-1) + + def _forward_dense_warmup( + self, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + seq_ctx: SequenceContext, + ) -> AttnOutputs: + """Run original dense MLA while distilling only the DSA indexers.""" + + dense_outputs = MultiLatentAttention.forward(self, hidden_states, position_embeddings, seq_ctx) + training_cfg = self.indexer_training + if training_cfg is None or training_cfg.loss_coeff == 0 or not self.training: + return dense_outputs + if self.source_layer_idx != self.layer_idx: + return dense_outputs + + self._validate_indexer_training_runtime(seq_ctx) + q_resid, teacher_q, teacher_k = self._project_dense_teacher_states(hidden_states.detach(), position_embeddings) + features = self.indexer.project_features( + hidden_states.detach(), + q_resid.detach(), + position_embeddings, + seq_ctx, + ) + + indexer_loss = dense_dsa_indexer_kl_loss( + features.q, + features.k, + features.weights * ((self.index_n_heads * self.index_head_dim) ** -0.5), + teacher_q, + teacher_k, + seq_ctx, + softmax_scale=self.softmax_scale, + loss_coefficient=training_cfg.loss_coeff, + query_block_size=training_cfg.dense_query_block_size, + ) + seq_ctx.dsa_topk_cache.indexer_losses.append(indexer_loss) + return dense_outputs + def forward( self, hidden_states: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], seq_ctx: SequenceContext, + ) -> AttnOutputs: + training_cfg = self.indexer_training + if training_cfg is not None and training_cfg.stage == "dense_warmup": + return self._forward_dense_warmup(hidden_states, position_embeddings, seq_ctx) + return self._forward_sparse(hidden_states, position_embeddings, seq_ctx) + + def _forward_sparse( + self, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + seq_ctx: SequenceContext, ) -> AttnOutputs: """Absorbed DSA-MLA forward for packed training (``bsz == 1``). @@ -336,7 +505,8 @@ def forward( k_pe = k_pe.view(bsz, q_len, 1, self.qk_rope_head_dim).transpose(1, 2) cos, sin = position_embeddings - q_pe, k_pe = mla_apply_rotary_pos_emb(q_pe, k_pe, cos, sin) + rope_fn = mla_apply_rotary_pos_emb if self.rope_interleave else mla_apply_rotary_pos_emb_non_interleaved + q_pe, k_pe = rope_fn(q_pe, k_pe, cos, sin) # kv_b_proj.weight: [N * (Dn + Dv), Rkv] if isinstance(self.kv_b_proj.weight, DTensor): @@ -365,17 +535,51 @@ def forward( # key_states: [S_g, 1, Rkv + Dr] key_states = gather_for_sequence_parallel(key_states, dim=0, sp_mesh=seq_ctx.sequence_parallel_mesh) - # topk_indices: [S, 1, K] - topk_indices = get_dsa_topk_sharing_runtime().get_or_compute( - layer=self, - seq_ctx=seq_ctx, - compute_source_topk=lambda: self.indexer( - hidden_states, - q_resid, + training_cfg = self.indexer_training + indexer_features: DSAIndexerFeatures | None = None + indexer_loss_enabled = ( + training_cfg is not None + and training_cfg.stage == "sparse" + and training_cfg.loss_coeff > 0 + and self.source_layer_idx == self.layer_idx + and self.training + ) + if indexer_loss_enabled: + self._validate_indexer_training_runtime(seq_ctx) + topk_runtime = get_dsa_topk_sharing_runtime() + # GLM-5.2 runs the physical MTP layer once as an index-compute layer, + # then reuses its discrete top-k for later logical MTP depths. Mirror + # that contract during training: later depths must neither rerun the + # indexer projections nor add duplicate KL terms. + reuse_mtp_iteration_topk = topk_runtime.reuses_mtp_iteration_topk(layer=self, seq_ctx=seq_ctx) + train_source_indexer = indexer_loss_enabled and torch.is_grad_enabled() and not reuse_mtp_iteration_topk + if train_source_indexer: + # The indexer learns from attention, but must not inject an extra + # gradient path into the transformer hidden/Q-LoRA activations. + indexer_features = self.indexer.project_features( + hidden_states.detach(), + q_resid.detach(), position_embeddings, seq_ctx, - ), - ) + ) + topk_indices = topk_runtime.get_or_compute( + layer=self, + seq_ctx=seq_ctx, + compute_source_topk=lambda: self.indexer.select_topk(indexer_features, seq_ctx), + ) + else: + # ``loss_coeff=0`` follows this no-grad path so optimizer-visible + # indexer parameters retain ``grad is None`` rather than zero grads. + topk_indices = topk_runtime.get_or_compute( + layer=self, + seq_ctx=seq_ctx, + compute_source_topk=lambda: self.indexer( + hidden_states, + q_resid, + position_embeddings, + seq_ctx, + ), + ) sparse_mla_outputs = self.sparse_mla_func( query_states, key_states, @@ -386,6 +590,44 @@ def forward( # raw_output: [S, N, Rkv]; softmax_lse: [S, N] raw_output = sparse_mla_outputs.raw_output softmax_lse = sparse_mla_outputs.softmax_lse + if train_source_indexer: + assert indexer_features is not None + valid_query_rows = q_len - seq_ctx.num_padding + if valid_query_rows <= 0: + # Packed data may legitimately produce an all-padding shard. + # Keep every indexer parameter in the FSDP backward graph while + # contributing no supervision on that rank. + indexer_loss = ( + indexer_features.q.sum() + indexer_features.k.sum() + indexer_features.weights.sum() + ) * 0.0 + else: + valid_query_mask = torch.arange(q_len, device=query_states.device).unsqueeze(0) < valid_query_rows + attn_q_for_loss = query_states.unsqueeze(0) + attn_lse_for_loss = softmax_lse.unsqueeze(0) + # The cuDNN score-recompute MMA requires an attention-head count + # divisible by 8. Zero Q with +inf LSE contributes exactly zero to + # its head-summed teacher before the existing L1 normalization. + pad_heads = (-self.num_attention_heads) % 8 + if pad_heads: + attn_q_for_loss = torch.nn.functional.pad(attn_q_for_loss, (0, 0, 0, pad_heads)) + attn_lse_for_loss = torch.nn.functional.pad( + attn_lse_for_loss, (0, pad_heads), value=float("inf") + ) + indexer_loss = dsa_indexer_kl_loss( + indexer_features.q, + indexer_features.k, + indexer_features.weights * ((self.index_n_heads * self.index_head_dim) ** -0.5), + attn_q_for_loss, + key_states.squeeze(1).unsqueeze(0), + attn_lse_for_loss, + topk_indices.squeeze(1).unsqueeze(0), + softmax_scale=self.softmax_scale, + row_coefficient=training_cfg.loss_coeff / valid_query_rows, + valid_query_mask=valid_query_mask, + debug_name=f"layer{self.layer_idx}", + debug_interval=training_cfg.debug_interval, + ) + seq_ctx.dsa_topk_cache.indexer_losses.append(indexer_loss) # raw_output: [S, N, Dv] -> [bsz, S, N * Dv] raw_output = torch.einsum("shm,hdm->shd", raw_output, w_vc) raw_output = raw_output.reshape(bsz, q_len, self.num_attention_heads * self.v_head_dim).contiguous() diff --git a/xtuner/v1/module/attention/dsa_topk_sharing.py b/xtuner/v1/module/attention/dsa_topk_sharing.py index 7e0ca4fd27..a1b6869679 100644 --- a/xtuner/v1/module/attention/dsa_topk_sharing.py +++ b/xtuner/v1/module/attention/dsa_topk_sharing.py @@ -261,6 +261,20 @@ def get_or_compute( residency.store_gpu(seq_ctx, layer.layer_idx, topk_indices) return topk_indices + def reuses_mtp_iteration_topk( + self, + *, + layer: DSATopKSharingLayerProtocol, + seq_ctx: SequenceContext, + ) -> bool: + """Whether this logical MTP depth reuses the first depth's top-k.""" + + return self._can_reuse_mtp_iteration_topk( + seq_ctx, + layer.source_layer_idx, + self._residency(), + ) + def after_sparse_mla_use(self, *, layer: DSATopKSharingLayerProtocol, seq_ctx: SequenceContext) -> None: residency = self._residency() cache = seq_ctx.dsa_topk_cache diff --git a/xtuner/v1/module/attention/mla.py b/xtuner/v1/module/attention/mla.py index c17eff2a79..fed2a39067 100644 --- a/xtuner/v1/module/attention/mla.py +++ b/xtuner/v1/module/attention/mla.py @@ -58,6 +58,7 @@ class MLAConfig(BaseModel): qk_rope_head_dim: int qk_nope_head_dim: int v_head_dim: int + rope_interleave: bool = True def build( self, @@ -165,6 +166,24 @@ def mla_apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1): return q_embed, k_embed +def mla_apply_rotary_pos_emb_non_interleaved(q, k, cos, sin, unsqueeze_dim=1): + """Apply half-split RoPE without converting interleaved feature pairs. + + Args: + q (torch.Tensor): Query rotary features. + k (torch.Tensor): Key rotary features. + cos (torch.Tensor): Cosine rotary coefficients. + sin (torch.Tensor): Sine rotary coefficients. + unsqueeze_dim (int): Head dimension used to broadcast the coefficients. + + Returns: + tuple[torch.Tensor, torch.Tensor]: Rotated query and key features. + """ + cos = cos.unsqueeze(unsqueeze_dim) + sin = sin.unsqueeze(unsqueeze_dim) + return (q * cos) + (rotate_half(q) * sin), (k * cos) + (rotate_half(k) * sin) + + def yarn_get_mscale(scale=1.0, mscale=1.0): if scale <= 1: return 1.0 @@ -196,6 +215,7 @@ def __init__( layer_type: Literal["full_attention", "sliding_attention"] | None = None, sliding_window: int = -1, layer_idx: int = 0, + rope_interleave: bool = True, ): super().__init__() self.name = f"layers.{layer_idx}.self_attn" @@ -212,6 +232,7 @@ def __init__( self.qkv_bias = qkv_bias self.o_bias = o_bias self.qk_norm = qk_norm + self.rope_interleave = rope_interleave self.float8_cfg = float8_cfg self.generate_config = generate_config self.q_head_dim = qk_nope_head_dim + qk_rope_head_dim @@ -304,7 +325,8 @@ def forward_training( k_nope, value_states = torch.split(kv, [self.qk_nope_head_dim, self.v_head_dim], dim=-1) cos, sin = position_embeddings - q_pe, k_pe = mla_apply_rotary_pos_emb(q_pe, k_pe, cos, sin) + rope_fn = mla_apply_rotary_pos_emb if self.rope_interleave else mla_apply_rotary_pos_emb_non_interleaved + q_pe, k_pe = rope_fn(q_pe, k_pe, cos, sin) query_states = k_pe.new_empty(bsz, self.num_attention_heads, q_len, self.q_head_dim) query_states[:, :, :, : self.qk_nope_head_dim] = q_nope @@ -383,7 +405,8 @@ def prefilling( cos, sin = position_embeddings q_pe = q_pe.transpose(1, 2) k_pe = k_pe.transpose(1, 2) - q_pe, k_pe = mla_apply_rotary_pos_emb(q_pe, k_pe, cos, sin) + rope_fn = mla_apply_rotary_pos_emb if self.rope_interleave else mla_apply_rotary_pos_emb_non_interleaved + q_pe, k_pe = rope_fn(q_pe, k_pe, cos, sin) q_pe = q_pe.transpose(1, 2) k_pe = k_pe.transpose(1, 2) @@ -502,7 +525,8 @@ def decoding( q_pe = q_pe.transpose(1, 2) k_pe = k_pe.transpose(1, 2) - q_pe, k_pe = mla_apply_rotary_pos_emb(q_pe, k_pe, cos, sin) + rope_fn = mla_apply_rotary_pos_emb if self.rope_interleave else mla_apply_rotary_pos_emb_non_interleaved + q_pe, k_pe = rope_fn(q_pe, k_pe, cos, sin) q_pe = q_pe.transpose(1, 2) k_pe = k_pe.transpose(1, 2) @@ -586,7 +610,8 @@ def forward( cos, sin = position_embeddings # cos = torch.load('cos.pth').cuda() # sin = torch.load('sin.pth').cuda() - q_pe, k_pe = mla_apply_rotary_pos_emb(q_pe, k_pe, cos, sin) + rope_fn = mla_apply_rotary_pos_emb if self.rope_interleave else mla_apply_rotary_pos_emb_non_interleaved + q_pe, k_pe = rope_fn(q_pe, k_pe, cos, sin) query_states = k_pe.new_empty(bsz, self.num_attention_heads, q_len, self.q_head_dim) query_states[:, :, :, : self.qk_nope_head_dim] = q_nope diff --git a/xtuner/v1/ops/sparse_mla/__init__.py b/xtuner/v1/ops/sparse_mla/__init__.py index e76663add6..2bcdabddfc 100644 --- a/xtuner/v1/ops/sparse_mla/__init__.py +++ b/xtuner/v1/ops/sparse_mla/__init__.py @@ -7,6 +7,36 @@ from .pytorch import torch_dsa_topk_indices, torch_sparse_mla +def dsa_indexer_kl_from_distribution(*args, **kwargs): + from .cudnn_dsa_indexer_loss import dsa_indexer_kl_from_distribution as _impl + + return _impl(*args, **kwargs) + + +def dsa_indexer_kl_loss(*args, **kwargs): + from .cudnn_dsa_indexer_loss import dsa_indexer_kl_loss as _impl + + return _impl(*args, **kwargs) + + +def sparse_attention_target(*args, **kwargs): + from .cudnn_dsa_indexer_loss import sparse_attention_target as _impl + + return _impl(*args, **kwargs) + + +def sparse_indexer_predict(*args, **kwargs): + from .cudnn_dsa_indexer_loss import sparse_indexer_predict as _impl + + return _impl(*args, **kwargs) + + +def ensure_cudnn_dsa_indexer_training_available() -> None: + from .cudnn_dsa_indexer_loss import ensure_cudnn_dsa_indexer_training_available as _impl + + return _impl() + + def get_sparse_mla(backend: SparseMLABackend) -> SparseMLAProtocol: if backend == "torch": return torch_sparse_mla @@ -97,7 +127,10 @@ def indexer_fwd_interface(*args, **kwargs): "SparseMLABackend", "SparseMLAOutputs", "SparseMLAProtocol", + "dsa_indexer_kl_from_distribution", + "dsa_indexer_kl_loss", "dsa_topk_indices", + "ensure_cudnn_dsa_indexer_training_available", "ensure_cudnn_dsa_runtime_available", "ensure_tilelang_runtime_available", "get_dsa_topk_indices", @@ -106,6 +139,8 @@ def indexer_fwd_interface(*args, **kwargs): "sparse_mla", "sparse_mla_bwd", "sparse_mla_fwd_interface", + "sparse_attention_target", + "sparse_indexer_predict", "torch_dsa_topk_indices", "torch_sparse_mla", ] diff --git a/xtuner/v1/ops/sparse_mla/cudnn_dsa_indexer_loss.py b/xtuner/v1/ops/sparse_mla/cudnn_dsa_indexer_loss.py new file mode 100644 index 0000000000..1038743662 --- /dev/null +++ b/xtuner/v1/ops/sparse_mla/cudnn_dsa_indexer_loss.py @@ -0,0 +1,602 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""cuDNN sparse DSA indexer distillation loss. + +The public loss uses XTuner-owned sum reduction semantics. cuDNN score +recompute produces the FP32 teacher/prediction distributions, while an opaque +manual-autograd operator routes the KL gradient only to indexer Q/K/weights. +""" + +from __future__ import annotations + +import torch +from torch import Tensor + +from xtuner.v1.utils import log_rank0 + + +_CUDNN_INDEXER_BACKWARD_MIN_HEADS = 64 +_CUDNN_INDEXER_BACKWARD_BLOCK_I = 128 +_INDEXER_LOSS_DEBUG_CALLS: dict[str, int] = {} + + +def _copy_aligned_grad_loss(grad_loss: Tensor, device: torch.device) -> Tensor: + """Copy an autograd scalar into a fresh 16-byte-aligned FP32 allocation.""" + + if grad_loss.numel() != 1: + raise ValueError(f"grad_loss must contain exactly one element, got shape {tuple(grad_loss.shape)}") + + aligned_grad_loss = torch.empty(1, dtype=torch.float32, device=device) + aligned_grad_loss.copy_(grad_loss.detach().to(device=device, dtype=torch.float32).reshape(1)) + return aligned_grad_loss + + +def _aligned_contiguous(tensor: Tensor) -> Tensor: + """Return a contiguous tensor whose actual data pointer is 16-byte aligned.""" + + # The kernel contract is addr(tensor) mod 16 = 0; stride contiguity alone + # does not imply this when the tensor is a storage-offset view. + tensor = tensor.contiguous() + if tensor.data_ptr() % 16 != 0: + tensor = tensor.clone() + return tensor + + +def _pad_indexer_heads_for_cudnn(index_q: Tensor, index_weights: Tensor) -> tuple[Tensor, Tensor, int]: + """Pad sub-64-head indexer inputs without changing their score function.""" + + index_heads = index_q.shape[-2] + if index_weights.shape[-1] != index_heads: + raise ValueError( + "index_q and index_weights must have the same number of index heads, " + f"got {index_heads} and {index_weights.shape[-1]}" + ) + if index_heads >= _CUDNN_INDEXER_BACKWARD_MIN_HEADS: + return index_q, index_weights, index_heads + + # cudnn要求有>=64个head,但是glm 5.2的indexer只有32个,所以要pad到64个 + # For H' = 64, extend q'_h = q_h and w'_h = w_h for h <= H, and set + # q'_h = w'_h = 0 for H < h <= H'. With + # score(q, k, w) = sum_{h=1}^{H} w_h * ReLU(), + # the padded terms are zero, so score(q', k, w') = score(q, k, w). + padded_heads = _CUDNN_INDEXER_BACKWARD_MIN_HEADS - index_heads + return ( + torch.nn.functional.pad(index_q, (0, 0, 0, padded_heads)), + torch.nn.functional.pad(index_weights, (0, padded_heads)), + index_heads, + ) + + +def _prepare_sparse_topk( + topk_indices: Tensor, + topk_length: Tensor | None, +) -> tuple[Tensor, Tensor, Tensor]: + if topk_indices.ndim != 3: + raise ValueError(f"topk_indices must have shape (B, S_q, K), got {tuple(topk_indices.shape)}") + + valid_slots = topk_indices != -1 + if topk_length is None: + topk_length = valid_slots.sum(dim=-1, dtype=torch.int32) + elif topk_length.shape != topk_indices.shape[:2]: + raise ValueError( + f"topk_length must have shape {tuple(topk_indices.shape[:2])}, got {tuple(topk_length.shape)}" + ) + + safe_topk = topk_indices.clamp_min(0).to(dtype=torch.int32).contiguous() + return safe_topk, topk_length.to(dtype=torch.int32).contiguous(), valid_slots + + +def _validate_distribution_shapes( + target: Tensor, + predict: Tensor, + topk_indices: Tensor, +) -> None: + if target.shape != predict.shape: + raise ValueError(f"target and predict must have the same shape, got {target.shape} and {predict.shape}") + if target.shape != topk_indices.shape: + raise ValueError(f"target/predict must match topk_indices shape {topk_indices.shape}, got {target.shape}") + + +def _mask_invalid_query_rows( + target: Tensor, + predict: Tensor, + topk_indices: Tensor, + valid_query_mask: Tensor | None, +) -> tuple[Tensor, Tensor, Tensor]: + """Remove padded query rows from both the KL value and its manual backward. + + Packed SFT batches keep a fixed physical sequence length. Their tail + padding is represented as one or more causal chunks, so the top-k kernel + still returns non-negative indices for those query rows. Consequently, + ``topk_indices != -1`` alone cannot distinguish real queries from padding. + + For a padding row, setting ``target=predict=0`` makes cuDNN's score-gradient + signal zero, while setting ``topk_indices=-1`` makes XTuner's public KL + reduction exclude the same row. Keeping the physical tensor shapes intact + also avoids compiling a new cuDNN kernel for every effective sequence + length. + """ + + if valid_query_mask is None: + return target, predict, topk_indices + expected_shape = topk_indices.shape[:-1] + if valid_query_mask.shape != expected_shape: + raise ValueError( + f"valid_query_mask must have shape {tuple(expected_shape)}, got {tuple(valid_query_mask.shape)}" + ) + + valid_query_mask = valid_query_mask.to(device=topk_indices.device, dtype=torch.bool) + invalid_rows = ~valid_query_mask.unsqueeze(-1) + return ( + target.masked_fill(invalid_rows, 0.0), + predict.masked_fill(invalid_rows, 0.0), + topk_indices.masked_fill(invalid_rows, -1), + ) + + +@torch.no_grad() +def _maybe_log_indexer_loss_diagnostics( + target: Tensor, + predict: Tensor, + topk_indices: Tensor, + *, + debug_name: str | None, + debug_interval: int, +) -> None: + r"""Log compact distribution diagnostics on rank 0 at a fixed interval. + + For every valid query row, the reported quantities include + + .. math:: + + H(p)=-\sum_i p_i\log p_i,\qquad + D_{KL}(p\|q)=\sum_i p_i\log\frac{p_i}{q_i}, + + where ``p`` is the attention teacher and ``q`` is the indexer prediction. + ``top1_match`` is the mean indicator + :math:`\mathbb{1}[\arg\max p=\arg\max q]`. + """ + + if debug_name is None or debug_interval <= 0: + return + if torch.distributed.is_initialized() and torch.distributed.get_rank() != 0: + return + + call = _INDEXER_LOSS_DEBUG_CALLS.get(debug_name, 0) + 1 + _INDEXER_LOSS_DEBUG_CALLS[debug_name] = call + if call != 1 and call % debug_interval != 0: + return + + target_f32 = target.float() + predict_f32 = predict.float() + valid_slots = topk_indices != -1 + valid_rows = valid_slots.any(dim=-1) + num_valid_rows = int(valid_rows.sum().item()) + if num_valid_rows == 0: + log_rank0.info( + f"[DSA_INDEXER_LOSS] name={debug_name} call={call} valid_rows=0 (distribution diagnostics skipped)" + ) + return + target_entropy = -torch.special.xlogy(target_f32, target_f32).sum(dim=-1) + predict_entropy = -torch.special.xlogy(predict_f32, predict_f32).sum(dim=-1) + per_row_kl = torch.special.xlogy(target_f32, target_f32).sum(dim=-1) - torch.special.xlogy( + target_f32, predict_f32 + ).sum(dim=-1) + target_top1 = target_f32.argmax(dim=-1) + predict_top1 = predict_f32.argmax(dim=-1) + target_top1_predict = predict_f32.gather(dim=-1, index=target_top1.unsqueeze(-1)).squeeze(-1) + + valid_target_entropy = target_entropy[valid_rows] + valid_predict_entropy = predict_entropy[valid_rows] + valid_kl = per_row_kl[valid_rows] + valid_target_max = target_f32.amax(dim=-1)[valid_rows] + valid_predict_max = predict_f32.amax(dim=-1)[valid_rows] + valid_target_top1_predict = target_top1_predict[valid_rows] + valid_top1_match = (target_top1 == predict_top1)[valid_rows].float() + mean_topk = valid_slots.sum(dim=-1, dtype=torch.float32)[valid_rows].mean() + + log_rank0.info( + "[DSA_INDEXER_LOSS] " + f"name={debug_name} call={call} valid_rows={num_valid_rows} " + f"mean_topk={mean_topk.item():.2f} kl_mean={valid_kl.mean().item():.6f} " + f"kl_max={valid_kl.max().item():.6f} target_entropy={valid_target_entropy.mean().item():.6f} " + f"predict_entropy={valid_predict_entropy.mean().item():.6f} " + f"target_max={valid_target_max.mean().item():.6f} predict_max={valid_predict_max.mean().item():.6f} " + f"top1_match={valid_top1_match.mean().item():.6f} " + f"predict_at_target_top1={valid_target_top1_predict.mean().item():.6f}" + ) + + +def _standard_kl_loss( + target: Tensor, + predict: Tensor, + topk_indices: Tensor, + *, + row_coefficient: float, +) -> Tensor: + """ + 1. 计算每个 query 的 teacher/student KL。 + 2. 忽略完全无有效 top-k 的 padding query。 + 3. 检查 (p_i>0,q_i=0) 的无限 KL 情况。 + 4. 将所有 query 的 KL 求和并乘 row_coefficient。 + """ + + _validate_distribution_shapes(target, predict, topk_indices) + target_f32 = target.float() + predict_f32 = predict.float() + valid_slots = topk_indices != -1 + valid_rows = valid_slots.any(dim=-1) + + invalid_predict = valid_slots & (target_f32 > 0) & (predict_f32 <= 0) + torch._assert( + ~invalid_predict.any(), + "DSA indexer predict has zero probability where target is positive; standard KL is not finite.", + ) + + target_self_term = torch.special.xlogy(target_f32, target_f32).sum(dim=-1) + target_cross_term = torch.special.xlogy(target_f32, predict_f32).sum(dim=-1) + per_row_kl = (target_self_term - target_cross_term).masked_fill(~valid_rows, 0.0) + return per_row_kl.sum() * float(row_coefficient) + + +@torch.no_grad() +def sparse_attention_target( + attn_q: Tensor, + attn_k: Tensor, + attn_lse: Tensor, + topk_indices: Tensor, + *, + softmax_scale: float, + topk_length: Tensor | None = None, +) -> Tensor: + """Recompute the head-aggregated sparse attention teacher distribution.""" + + from cudnn.deepseek_sparse_attention.score_recompute import sparse_attn_score_recompute_wrapper + + safe_topk, topk_length, valid_slots = _prepare_sparse_topk(topk_indices, topk_length) + outputs = sparse_attn_score_recompute_wrapper( + attn_q.contiguous(), + attn_k.contiguous(), + attn_lse.float().contiguous(), + safe_topk, + softmax_scale=float(softmax_scale), + topk_length=topk_length, + topk_indices_global=False, + ) + return outputs["target"].float().masked_fill(~valid_slots, 0.0) + + +@torch.no_grad() +def sparse_indexer_predict( + index_q: Tensor, + index_k: Tensor, + index_weights: Tensor, + topk_indices: Tensor, + *, + topk_length: Tensor | None = None, +) -> Tensor: + """Recompute the FP32 sparse indexer prediction distribution.""" + + from cudnn.deepseek_sparse_attention.score_recompute import sparse_indexer_score_recompute_wrapper + + safe_topk, topk_length, valid_slots = _prepare_sparse_topk(topk_indices, topk_length) + outputs = sparse_indexer_score_recompute_wrapper( + index_q.contiguous(), + index_k.contiguous(), + index_weights.contiguous(), + safe_topk, + topk_length=topk_length, + topk_indices_global=False, + ) + return outputs["predict"].float().masked_fill(~valid_slots, 0.0) + + +def _xtuner_indexer_backward( + index_q: Tensor, + index_k: Tensor, + index_weights: Tensor, + target: Tensor, + predict: Tensor, + safe_topk_indices: Tensor, + row_coefficient: float, + grad_loss: Tensor, +) -> tuple[Tensor, Tensor, Tensor]: + """ + 1. pad 32 head 到cudnn需要的64head + 2. 检查topk是否满足分块要求 + 3. 转换loss系数,训练目标是1/N \\sum KL_i, cudnn计算的loss自带了一个1/N + """ + + from cudnn.deepseek_sparse_attention.indexer_backward import indexer_backward_wrapper + + index_q_for_cudnn, index_weights_for_cudnn, index_heads = _pad_indexer_heads_for_cudnn(index_q, index_weights) + topk = safe_topk_indices.shape[-1] + # The SM90 kernel tiles the sparse dimension in blocks of I=128, hence + # the admissible top-k sizes satisfy K mod I = 0. + if topk % _CUDNN_INDEXER_BACKWARD_BLOCK_I != 0: + raise ValueError( + "cuDNN sparse indexer backward requires topk to be a multiple of " + f"{_CUDNN_INDEXER_BACKWARD_BLOCK_I}, got {topk}" + ) + + # The current cuDNN SM90 indexer-backward kernel requires at least 64 + # index heads, while GLM-5.2 uses 32. For each query/key pair, + # s_bt = sum_{h=1}^{H} w_bh * ReLU(q_bh^T k_t). + # Extending q_h=w_h=0 for h>H gives s'_bt=s_bt. The padded terms also + # contribute zero to dK; dQ/dW are projected back to their first H heads. + # Keep the already-applied 32-head scaling unchanged. + # + # Let N=B*S be the number of physical rows. The backend computes + # L_backend = c_backend * (1/N) * sum_{i=1}^{N} KL_i, + # while XTuner exposes + # L_xtuner = c_row * sum_{i=1}^{N} KL_i. + # Therefore c_backend = N * c_row makes the two losses and gradients equal. + physical_rows = index_q.shape[0] * index_q.shape[1] + backend_loss_coeff = float(row_coefficient) * physical_rows + # ``grad_loss`` may be an aligned-looking contiguous view into autograd's + # shared scalar buffer whose storage offset makes its actual data pointer + # fail CuTe DSL's 16-byte alignment requirement. cuDNN converts it with + # ``copy=False``, so force a fresh allocation here. ``contiguous()`` alone + # is insufficient because it can return the original contiguous view. + # By the chain rule, dL_outer/dtheta = g * dL_kl/dtheta, where + # g=dL_outer/dL_kl. Relocating the scalar keeps g_aligned=g, so gradients + # are numerically unchanged. + aligned_grad_loss = _copy_aligned_grad_loss(grad_loss, index_q.device) + # The cuDNN wrapper overwrites target/predict while forming score gradients. + # Work on fresh aligned buffers so the custom op remains functionally pure, + # caller-visible distributions stay intact, and retain_graph backward gets a + # new unmodified pair on every invocation. + target_for_cudnn = _aligned_contiguous(target.clone(memory_format=torch.contiguous_format)) + predict_for_cudnn = _aligned_contiguous(predict.clone(memory_format=torch.contiguous_format)) + outputs = indexer_backward_wrapper( + _aligned_contiguous(index_q_for_cudnn), + _aligned_contiguous(index_weights_for_cudnn), + _aligned_contiguous(index_k), + target_for_cudnn, + predict_for_cudnn, + _aligned_contiguous(safe_topk_indices.to(dtype=torch.int32)), + sm_scale=1.0, + loss_coeff=backend_loss_coeff, + grad_loss=aligned_grad_loss, + topk_indices_global=False, + ) + # This is the projection P_H onto the model-owned coordinates: + # dQ = P_H(dQ') = dQ'[..., :H, :] and dW = P_H(dW') = dW'[..., :H]. + return ( + outputs["d_index_q"][..., :index_heads, :].contiguous(), + outputs["d_index_k"], + outputs["d_weights"][..., :index_heads].contiguous(), + ) + + +@torch.library.custom_op( + "sparse_mla::cudnn_dsa_indexer_kl_backward", + mutates_args=(), + device_types="cuda", +) +def _cudnn_dsa_indexer_kl_backward( + index_q: Tensor, + index_k: Tensor, + index_weights: Tensor, + target: Tensor, + predict: Tensor, + safe_topk_indices: Tensor, + row_coefficient: float, + grad_loss: Tensor, +) -> tuple[Tensor, Tensor, Tensor]: + return _xtuner_indexer_backward( + index_q, + index_k, + index_weights, + target, + predict, + safe_topk_indices, + row_coefficient, + grad_loss, + ) + + +@_cudnn_dsa_indexer_kl_backward.register_fake +def _( + index_q: Tensor, + index_k: Tensor, + index_weights: Tensor, + target: Tensor, + predict: Tensor, + safe_topk_indices: Tensor, + row_coefficient: float, + grad_loss: Tensor, +) -> tuple[Tensor, Tensor, Tensor]: + del target, predict, safe_topk_indices, row_coefficient, grad_loss + return torch.empty_like(index_q), torch.empty_like(index_k), torch.empty_like(index_weights) + + +@torch.library.custom_op( + "sparse_mla::cudnn_dsa_indexer_kl_from_distribution", + mutates_args=(), + device_types="cuda", +) +def _cudnn_dsa_indexer_kl_from_distribution( + index_q: Tensor, + index_k: Tensor, + index_weights: Tensor, + target: Tensor, + predict: Tensor, + topk_indices: Tensor, + row_coefficient: float, +) -> Tensor: + del index_q, index_k, index_weights + return _standard_kl_loss( + target, + predict, + topk_indices, + row_coefficient=row_coefficient, + ) + + +@_cudnn_dsa_indexer_kl_from_distribution.register_fake +def _( + index_q: Tensor, + index_k: Tensor, + index_weights: Tensor, + target: Tensor, + predict: Tensor, + topk_indices: Tensor, + row_coefficient: float, +) -> Tensor: + del index_q, index_k, index_weights, predict, topk_indices, row_coefficient + return target.new_empty((), dtype=torch.float32) + + +def _setup_indexer_kl_context(ctx, inputs, output) -> None: + del output + index_q, index_k, index_weights, target, predict, topk_indices, row_coefficient = inputs + safe_topk_indices, _, _ = _prepare_sparse_topk(topk_indices, topk_length=None) + ctx.row_coefficient = row_coefficient + ctx.save_for_backward( + index_q, + index_k, + index_weights, + target, + predict, + safe_topk_indices, + ) + + +def _indexer_kl_backward(ctx, grad_output: Tensor): + index_q, index_k, index_weights, target, predict, safe_topk_indices = ctx.saved_tensors + d_index_q, d_index_k, d_weights = _cudnn_dsa_indexer_kl_backward( + index_q, + index_k, + index_weights, + target, + predict, + safe_topk_indices, + ctx.row_coefficient, + grad_output, + ) + return d_index_q, d_index_k, d_weights, None, None, None, None + + +_cudnn_dsa_indexer_kl_from_distribution.register_autograd( + _indexer_kl_backward, + setup_context=_setup_indexer_kl_context, +) + + +def dsa_indexer_kl_from_distribution( + index_q: Tensor, + index_k: Tensor, + index_weights: Tensor, + target: Tensor, + predict: Tensor, + topk_indices: Tensor, + *, + row_coefficient: float, + valid_query_mask: Tensor | None = None, + debug_name: str | None = None, + debug_interval: int = 0, +) -> Tensor: + """Compute sparse indexer KL with gradients only for indexer features.""" + + if float(row_coefficient) == 0.0: + return target.new_zeros((), dtype=torch.float32) + _validate_distribution_shapes(target, predict, topk_indices) + target, predict, topk_indices = _mask_invalid_query_rows( + target, + predict, + topk_indices, + valid_query_mask, + ) + _maybe_log_indexer_loss_diagnostics( + target, + predict, + topk_indices, + debug_name=debug_name, + debug_interval=debug_interval, + ) + return _cudnn_dsa_indexer_kl_from_distribution( + index_q, + index_k, + index_weights, + target.detach(), + predict.detach(), + topk_indices, + float(row_coefficient), + ) + + +def dsa_indexer_kl_loss( + index_q: Tensor, + index_k: Tensor, + index_weights: Tensor, + attn_q: Tensor, + attn_k: Tensor, + attn_lse: Tensor, + topk_indices: Tensor, + *, + softmax_scale: float, + row_coefficient: float, + topk_length: Tensor | None = None, + valid_query_mask: Tensor | None = None, + debug_name: str | None = None, + debug_interval: int = 0, +) -> Tensor: + """Convenience wrapper for one attention teacher and one indexer.""" + + if float(row_coefficient) == 0.0: + return index_q.new_zeros((), dtype=torch.float32) + target = sparse_attention_target( + attn_q.detach(), + attn_k.detach(), + attn_lse.detach(), + topk_indices, + softmax_scale=softmax_scale, + topk_length=topk_length, + ) + predict = sparse_indexer_predict( + index_q.detach(), + index_k.detach(), + index_weights.detach(), + topk_indices, + topk_length=topk_length, + ) + return dsa_indexer_kl_from_distribution( + index_q, + index_k, + index_weights, + target, + predict, + topk_indices, + row_coefficient=row_coefficient, + valid_query_mask=valid_query_mask, + debug_name=debug_name, + debug_interval=debug_interval, + ) + + +def ensure_cudnn_dsa_indexer_training_available() -> None: + try: + from cudnn.deepseek_sparse_attention.indexer_backward import indexer_backward_wrapper + from cudnn.deepseek_sparse_attention.score_recompute import ( + sparse_attn_score_recompute_wrapper, + sparse_indexer_score_recompute_wrapper, + ) + + _ = ( + indexer_backward_wrapper, + sparse_attn_score_recompute_wrapper, + sparse_indexer_score_recompute_wrapper, + ) + except Exception as exc: + raise RuntimeError( + "cuDNN DSA indexer training requires score-recompute and indexer-backward support." + ) from exc + + +__all__ = [ + "dsa_indexer_kl_from_distribution", + "dsa_indexer_kl_loss", + "ensure_cudnn_dsa_indexer_training_available", + "sparse_attention_target", + "sparse_indexer_predict", +]