From 790a7a8ef582aa6b4c1d53f74f6767ed81d30321 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=96=BD=E5=98=89=E9=98=B3?= Date: Fri, 14 Aug 2026 14:34:18 +0800 Subject: [PATCH] [Feature] Support GLM-5.2 source-layer DSA indexer training --- .gitignore | 5 + examples/v1/config/sft_glm5p2.py | 30 +- examples/v1/scripts/train_glm52_indexer.sh | 91 +++ tests/model/test_glm52_moe.py | 32 +- tests/module/attention/test_dsa_mla.py | 202 +++++- tests/ops/test_cudnn_dsa_indexer_loss.py | 561 ++++++++++++++++ xtuner/v1/data_proto/sequence_context.py | 3 + xtuner/v1/model/moe/glm52.py | 13 + xtuner/v1/model/moe/moe.py | 20 + xtuner/v1/module/attention/__init__.py | 3 +- xtuner/v1/module/attention/dsa_mla.py | 168 ++++- xtuner/v1/ops/sparse_mla/__init__.py | 35 + .../ops/sparse_mla/cudnn_dsa_indexer_loss.py | 614 ++++++++++++++++++ 13 files changed, 1748 insertions(+), 29 deletions(-) create mode 100755 examples/v1/scripts/train_glm52_indexer.sh create mode 100644 tests/ops/test_cudnn_dsa_indexer_loss.py create mode 100644 xtuner/v1/ops/sparse_mla/cudnn_dsa_indexer_loss.py 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/examples/v1/config/sft_glm5p2.py b/examples/v1/config/sft_glm5p2.py index cab158ec2f..82817cd526 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,15 @@ 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")), + 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 +109,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 +148,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 +159,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/model/test_glm52_moe.py b/tests/model/test_glm52_moe.py index 5dd0c77365..ce97e7c3fa 100644 --- a/tests/model/test_glm52_moe.py +++ b/tests/model/test_glm52_moe.py @@ -27,7 +27,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 @@ -131,6 +131,36 @@ 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_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] + @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..cc25498f73 100644 --- a/tests/module/attention/test_dsa_mla.py +++ b/tests/module/attention/test_dsa_mla.py @@ -3,6 +3,7 @@ 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_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,8 +29,10 @@ 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 import DSAIndexerTrainingConfig, DSAMLAConfig +from xtuner.v1.module.attention import dsa_mla as dsa_mla_module from xtuner.v1.module.attention.dsa_topk_sharing import 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 +105,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 +119,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 +172,201 @@ 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_training_weights_preserve_existing_topk_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") + + features = attention.indexer.project_features(hidden_states, q_resid, position_embeddings, seq_ctx) + 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.selection_weights, + ) + training_logits = torch.einsum("bsht,bsh->bst", torch.relu(dot), features.training_weights.float()) + + 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_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..7795c46a5c --- /dev/null +++ b/tests/ops/test_cudnn_dsa_indexer_loss.py @@ -0,0 +1,561 @@ +"""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, + target_xlogx: torch.Tensor | None = None, +) -> torch.Tensor: + valid_rows = (topk_indices != -1).any(dim=-1) + if target_xlogx is None: + 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_target_xlogx_recovers_strict_average_multi_layer_kl(self): + target_0 = torch.tensor([[[0.7, 0.2, 0.1]]], dtype=torch.float32) + target_1 = torch.tensor([[[0.1, 0.3, 0.6]]], dtype=torch.float32) + predict = torch.tensor([[[0.2, 0.5, 0.3]]], dtype=torch.float32) + topk_indices = torch.tensor([[[0, 1, 2]]], dtype=torch.int32) + target_mean = (target_0 + target_1) / 2 + mean_xlogx = ( + torch.special.xlogy(target_0, target_0).sum(dim=-1) + torch.special.xlogy(target_1, target_1).sum(dim=-1) + ) / 2 + + actual = _standard_kl_loss( + target_mean, + predict, + topk_indices, + row_coefficient=0.4, + target_xlogx=mean_xlogx, + ) + expected = ( + _explicit_standard_kl(target_0, predict, topk_indices, row_coefficient=0.4) + + _explicit_standard_kl(target_1, predict, topk_indices, row_coefficient=0.4) + ) / 2 + + torch.testing.assert_close(actual, expected) + + 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/model/moe/glm52.py b/xtuner/v1/model/moe/glm52.py index 768537823b..4d7378c292 100644 --- a/xtuner/v1/model/moe/glm52.py +++ b/xtuner/v1/model/moe/glm52.py @@ -84,11 +84,24 @@ 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__}." ) + # Train only main-stack Full/source indexers. + # Freeze before optimizer construction, so MTP parameters are + # neither updated nor used to build an auxiliary loss graph. + if self_attn.indexer_training is not None: + 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 main-stack Full/source indexers are restored + # to trainable; shared layers own no indexer and MTP stays frozen. + self.requires_grad_(False) + for _, self_attn in dsa_layers[: self.config.num_hidden_layers]: + if 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..32d4e12b0d 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 @@ -483,6 +484,17 @@ 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 _micro_batch_forward( self, seq_ctx_list: list[SequenceContext], @@ -618,6 +630,10 @@ def _micro_batch_forward( assert hidden_states_list, "XTuner Internal Error, found empty hidden states for domino EP" + indexer_loss = self._consume_indexer_losses(seq_ctx_list) + if indexer_loss is not None: + output["indexer_loss"] = indexer_loss + if self.mtp_block is not None: assert self.config.mtp_config is not None @@ -846,6 +862,10 @@ def _forward( if self.config.return_hidden_states: output["hidden_states"].append(hidden_states) + indexer_loss = self._consume_indexer_losses([seq_ctx]) + if indexer_loss is not None: + output["indexer_loss"] = indexer_loss + layer_hidden_states = hidden_states hidden_states = self.norm(hidden_states) 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 063e740b73..db0671ef22 100644 --- a/xtuner/v1/module/attention/dsa_mla.py +++ b/xtuner/v1/module/attention/dsa_mla.py @@ -1,7 +1,8 @@ # Copyright (c) OpenMMLab. All rights reserved. -from typing import Literal +from typing import Literal, NamedTuple import torch +from pydantic import BaseModel, ConfigDict, Field from torch import nn from torch.distributed.tensor import DTensor @@ -13,6 +14,7 @@ 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, @@ -57,6 +59,34 @@ def extra_repr(self): return f"{self.normalized_shape}, eps={self.eps}" +class DSAIndexerTrainingConfig(BaseModel): + """Source-layer-only sparse indexer training configuration. + + ``None`` on :class:`DSAMLAConfig` remains the strict frozen baseline. This + implementation excludes IndexShare supervision, sequence parallelism, + activation checkpoint replay, and MTP indexers. + + ``indexer_only`` is a diagnostic overfit mode: GLM freezes the teacher and + every non-indexer parameter, leaving only main-stack source indexers + trainable. ``debug_interval`` prints per-source teacher/student + distribution statistics without changing the loss. + """ + + model_config = ConfigDict(extra="forbid") + loss_coeff: float = Field(default=1.0, ge=0.0) + supervision: Literal["source_layer"] = "source_layer" + train_mtp_indexer: Literal[False] = False + indexer_only: bool = False + debug_interval: int = Field(default=0, ge=0) + + +class DSAIndexerFeatures(NamedTuple): + q: torch.Tensor + k: torch.Tensor + selection_weights: torch.Tensor + training_weights: torch.Tensor + + class DSAIndexer(nn.Module): def __init__( self, @@ -68,6 +98,7 @@ def __init__( index_n_heads: int, index_topk: int, indexer_backend: Literal["torch", "tilelang", "cudnn_dsa"] = "torch", + trainable: bool = False, ): super().__init__() self.qk_rope_head_dim = qk_rope_head_dim @@ -83,19 +114,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 @@ -141,10 +173,14 @@ 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 下。 + raw_weights = self.weights_proj(hidden_states) + # The selection backends apply ``index_head_dim**-0.5`` internally. + # cuDNN indexer backward uses ``sm_scale=1`` and therefore consumes the + # complete effective scaling in its BF16 weights tensor. + selection_weights = raw_weights.float() * (self.index_n_heads**-0.5) + training_weights = raw_weights * ((self.index_n_heads * self.index_head_dim) ** -0.5) + + # 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 +192,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, selection_weights, training_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.selection_weights.detach(), 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 +229,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 +268,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 +295,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 = {}, {} @@ -274,8 +332,16 @@ def __init__( index_n_heads=self.index_n_heads, index_topk=self.index_topk, 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 forward( self, hidden_states: torch.Tensor, @@ -355,17 +421,50 @@ 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, + indexer_features: DSAIndexerFeatures | None = None + indexer_loss_enabled = ( + self.indexer_training is not None + and self.indexer_training.loss_coeff > 0 + and self.source_layer_idx == self.layer_idx + ) + if indexer_loss_enabled and self.training and not torch.is_grad_enabled(): + raise RuntimeError( + "DSA indexer training does not support activation checkpointing; set recompute_ratio=0." + ) + if ( + indexer_loss_enabled + and 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.") + train_source_indexer = indexer_loss_enabled and torch.is_grad_enabled() + 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 = get_dsa_topk_sharing_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 = get_dsa_topk_sharing_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, @@ -376,6 +475,27 @@ 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: + raise ValueError("DSA indexer training requires at least one attention-valid query row.") + valid_query_mask = torch.arange(q_len, device=query_states.device).unsqueeze(0) < valid_query_rows + indexer_loss = dsa_indexer_kl_loss( + indexer_features.q, + indexer_features.k, + indexer_features.training_weights, + query_states.unsqueeze(0), + key_states.squeeze(1).unsqueeze(0), + softmax_lse.unsqueeze(0), + topk_indices.squeeze(1).unsqueeze(0), + softmax_scale=self.softmax_scale, + row_coefficient=self.indexer_training.loss_coeff / valid_query_rows, + valid_query_mask=valid_query_mask, + debug_name=f"layer{self.layer_idx}", + debug_interval=self.indexer_training.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/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..7fb66ae510 --- /dev/null +++ b/xtuner/v1/ops/sparse_mla/cudnn_dsa_indexer_loss.py @@ -0,0 +1,614 @@ +# 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, + target_xlogx: Tensor | None, +) -> 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}") + if target_xlogx is not None and target_xlogx.shape != target.shape[:-1]: + raise ValueError(f"target_xlogx must have shape {target.shape[:-1]}, got {tuple(target_xlogx.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, + target_xlogx: Tensor | None = None, +) -> 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_xlogx) + 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.", + ) + + if target_xlogx is None: + target_self_term = torch.special.xlogy(target_f32, target_f32).sum(dim=-1) + else: + target_self_term = target_xlogx.float() + 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() + predict_for_cudnn = _aligned_contiguous(predict).clone() + 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, + target_xlogx: Tensor | None, +) -> Tensor: + del index_q, index_k, index_weights + return _standard_kl_loss( + target, + predict, + topk_indices, + row_coefficient=row_coefficient, + target_xlogx=target_xlogx, + ) + + +@_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, + target_xlogx: Tensor | None, +) -> Tensor: + del index_q, index_k, index_weights, predict, topk_indices, row_coefficient, target_xlogx + 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, 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, + target_xlogx: Tensor | None = None, + 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_xlogx) + 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), + target_xlogx.detach() if target_xlogx is not None else None, + ) + + +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", +]