Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 1 addition & 7 deletions tests/engine/test_glm52_moe_train_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,8 +32,7 @@
from xtuner.v1.loss.ce_loss import CELossConfig
from xtuner.v1.model import get_model_config_from_hf
from xtuner.v1.model.base import ModelItem
from xtuner.v1.model.moe.glm52 import Glm52MoEConfig
from xtuner.v1.module.attention import DSAMLAConfig
from xtuner.v1.model.moe.glm52 import DSAMLAConfig, Glm52MoEConfig
from xtuner.v1.module.mtp import MTPConfig
from xtuner.v1.module.router.noaux_router import NoAuxRouter, NoAuxRouterConfig
from xtuner.v1.utils import pad_to_max_length
Expand Down Expand Up @@ -204,7 +203,6 @@ def test_sp2_ep4_micro2_compile_offload_train_step(self):
engine.init_model_weights()
sp_mesh = init_data_mesh(str(DEVICE), sp_size=2)["sp"]
data_batches = []
seq_ctx_list = []

try:
for micro_batch_idx in range(4):
Expand All @@ -214,7 +212,6 @@ def test_sp2_ep4_micro2_compile_offload_train_step(self):
data = {"seq_ctx": full_seq_ctx, "shifted_labels": input_ids[:, 1:]}
loss_ctx = engine.model.build_loss_ctx_batch([data], sp_mesh=sp_mesh)[0]
seq_ctx = full_seq_ctx.split(sp_mesh)
seq_ctx_list.append(seq_ctx)
data_batches.append(ModelItem(seq_ctx=seq_ctx, loss_ctx=loss_ctx))

with mock.patch.dict(
Expand All @@ -230,9 +227,6 @@ def test_sp2_ep4_micro2_compile_offload_train_step(self):
self.assertTrue(math.isfinite(step_info["logs_info"]["reduced_mtp_loss"]))
self.assertTrue(math.isfinite(float(grad_norm)))
self.assertTrue(engine.optimizer.state)
for seq_ctx in seq_ctx_list:
self.assertEqual(seq_ctx.dsa_topk_cache.indices, {})
self.assertEqual(seq_ctx.dsa_topk_cache.offloaded, {})
finally:
del engine
torch.cuda.empty_cache()
Expand Down
26 changes: 25 additions & 1 deletion tests/model/test_glm52_moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,8 @@
TestGlm52RouterBias
test_scratch_init_zeroes_main_and_mtp_biases: 从头初始化清零主干与 MTP router bias。
test_update_bias_handles_main_and_shared_mtp_loads: bias 更新覆盖主干并聚合共享 MTP 深度。
TestGlm52ExplicitDsaDataflow
test_model_forward_backward_with_explicit_dsa_dataflow: 模型通过显式 IDs 完成前反向。
TestGlm52SequenceParallel
test_mtp_loss_and_gradients_match_full_sequence: SP2 的 MTP loss 与梯度匹配完整序列。
"""
Expand All @@ -27,7 +29,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.model.moe.glm52 import 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
Expand Down Expand Up @@ -222,6 +224,28 @@ def test_update_bias_handles_main_and_shared_mtp_loads(self):
)


@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
class TestGlm52ExplicitDsaDataflow:
def test_model_forward_backward_with_explicit_dsa_dataflow(self):
# 验证 GLM public forward/backward 经显式 DSA IDs 数据流产生有限 loss 和梯度。
config = _tiny_glm52_config()
config.mtp_config = None
model = config.build().to(device="cuda", dtype=torch.bfloat16)
model.init_weights()

input_ids = torch.tensor([[2, 3, 4, 5]], device="cuda")
shifted_labels = torch.tensor([[3, 4, 5, 6]], device="cuda")
seq_ctx = SequenceContext.from_input_ids((input_ids,), device="cuda")
data = {"seq_ctx": seq_ctx, "shifted_labels": shifted_labels}
loss_ctx = model.build_loss_ctx_batch([data], sp_mesh=None)[0]

output = model(seq_ctx=seq_ctx, loss_ctx=loss_ctx)
output["loss"].backward()

assert torch.isfinite(output["loss"])
assert any(parameter.grad is not None for parameter in model.parameters())


@unittest.skipUnless(torch.cuda.device_count() >= 2, "requires 2 CUDA devices")
class TestGlm52SequenceParallel(DeterministicDDPTestCase):
def test_mtp_loss_and_gradients_match_full_sequence(self):
Expand Down
15 changes: 2 additions & 13 deletions tests/model/test_glm52_mtp_checkpoint_repro.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,7 @@
"""GLM-5.2 MTP checkpoint 的真实训练回归测试。

TestGlm52CompiledMTPCheckpoint
test_shared_mtp_depths_train_with_compile_and_topk_offload: 共享 MTP 深度可在 compile/offload 下训练
(GLM-5.2 兼容待做,暂 xfail)。
test_shared_mtp_depths_train_with_compile_and_topk_offload: 共享 MTP 深度可在 compile/offload 下训练。
TestGlm52MicroBatchMTPCheckpoint
test_nested_micro_batch_inputs_preserve_gradients: EP2 micro2 的嵌套 embedding 梯度可正确反传。
"""
Expand All @@ -12,7 +11,6 @@
import unittest
from unittest import mock

import pytest
import torch

from xtuner._testing import DeterministicDDPTestCase
Expand All @@ -21,8 +19,7 @@
from xtuner.v1.engine.train_engine import TrainEngine
from xtuner.v1.loss.ce_loss import CELossConfig
from xtuner.v1.model.base import ModelItem
from xtuner.v1.model.moe.glm52 import Glm52MoEConfig
from xtuner.v1.module.attention import DSAMLAConfig
from xtuner.v1.model.moe.glm52 import DSAMLAConfig, Glm52MoEConfig
from xtuner.v1.module.mtp import MTPConfig
from xtuner.v1.module.router.noaux_router import NoAuxRouterConfig

Expand Down Expand Up @@ -105,14 +102,6 @@ def _model_item(engine: TrainEngine, start: int) -> ModelItem:

@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
class TestGlm52CompiledMTPCheckpoint(DeterministicDDPTestCase):
@pytest.mark.xfail(
reason="Shared MTP depths share one DSA top-k cache whose phase detection still relies on "
"`torch.is_grad_enabled()`, which only held for the removed reentrant implementation. Under "
"compile the resulting COMPUTE/REUSE divergence trips torch's "
"'Recomputed values ... have different metadata' check. Part of the pending GLM-5.2 "
"compatibility work; MTP gradients themselves are healthy (19/19 non-zero, matching base).",
strict=False,
)
def test_shared_mtp_depths_train_with_compile_and_topk_offload(self):
# 验证共享 MTP 深度可在 compile/offload 下训练且 loss 有限。
self.create_pg("cuda")
Expand Down
130 changes: 69 additions & 61 deletions tests/module/attention/test_dsa_mla.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,8 @@
test_padded_indices_support_int32_and_backward: PyTorch 后端处理 padding、int32 和反向传播。
TestDSAAttention
test_packed_inputs_respect_causal_boundaries_and_backward: packed attention 遵守分段因果边界并可反传。
test_shared_layers_reuse_topk_without_cross_context_leak: shared layer 复用当前样本 top-k 且不跨样本泄漏
test_checkpoint_reuses_and_releases_topk: checkpoint 重算复用并最终释放 top-k(GLM-5.2 兼容待做,暂 xfail)
test_shared_layer_consumes_explicit_topk_ids: shared layer 复用显式 top-k IDs,漏传时立即报错
test_selective_checkpoint_preserves_explicit_topk_storage: 显式 IDs 穿过 non-reentrant selective checkpoint
TestAcceleratedSparseMLA
test_tilelang_forward_backward_matches_torch: TileLang 前反向数值与 PyTorch 后端一致。
test_compiled_cudnn_backward_matches_tilelang: 编译后的 cuDNN DSA 前反向与 TileLang 一致。
Expand All @@ -23,13 +23,12 @@
import pytest
import torch
import torch.distributed as dist
import torch.nn as nn

from xtuner._testing import DeterministicDDPTestCase
from xtuner.v1.data_proto import SequenceContext
from xtuner.v1.model.utils import apply_gradient_checkpointing
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.model.moe.glm52 import DSAMLAConfig
from xtuner.v1.model.moe.glm52.decoder_layer import GLM52DenseDecoderLayer
from xtuner.v1.model.utils import apply_selective_checkpointing
from xtuner.v1.ops.sparse_mla import dsa_topk_indices, sparse_mla
from xtuner.v1.utils.test_utils import init_data_mesh

Expand Down Expand Up @@ -102,6 +101,10 @@ def _tiny_dsa_attention(
indexer_types: list[str] | None = None,
layer_idx: int = 0,
):
return _tiny_dsa_config(indexer_types).build(hidden_size=4, layer_idx=layer_idx)


def _tiny_dsa_config(indexer_types: list[str] | None = None) -> DSAMLAConfig:
return DSAMLAConfig(
num_attention_heads=2,
head_dim=2,
Expand All @@ -115,26 +118,17 @@ def _tiny_dsa_attention(
index_n_heads=2,
indexer_types=indexer_types,
sparse_mla_backend="torch",
).build(hidden_size=4, layer_idx=layer_idx)


class _TinyDsaDecoderBlock(nn.Module):
def __init__(self, attention: nn.Module) -> None:
super().__init__()
self.self_attn = attention
register_dsa_topk_decoder_lifecycle_hooks(self)

def forward(
self,
hidden_states: torch.Tensor,
position_embeddings: tuple[torch.Tensor, torch.Tensor],
seq_ctx: SequenceContext,
) -> torch.Tensor:
return self.self_attn(
hidden_states=hidden_states,
position_embeddings=position_embeddings,
seq_ctx=seq_ctx,
)["projected_output"]
)


def _tiny_dsa_decoder(indexer_types: list[str], layer_idx: int) -> GLM52DenseDecoderLayer:
return GLM52DenseDecoderLayer(
hidden_size=4,
intermediate_size=8,
hidden_act="silu",
attention_config=_tiny_dsa_config(indexer_types),
layer_idx=layer_idx,
)


class TestTorchSparseMLA:
Expand Down Expand Up @@ -184,59 +178,71 @@ def test_packed_inputs_respect_causal_boundaries_and_backward(self):
assert outputs["raw_output"].shape == (1, 5, 6)
assert torch.isfinite(outputs["projected_output"]).all()
assert torch.isfinite(hidden_states.grad).all()
topk = seq_ctx.dsa_topk_cache.indices[0]
topk = outputs["dsa_topk_ids"]
assert topk.dtype == torch.int32
assert topk.is_contiguous()
for token_idx, seq_start in [(0, 0), (1, 0), (2, 2), (3, 2), (4, 2)]:
valid_indices = topk[token_idx, 0][topk[token_idx, 0] != -1]
assert valid_indices.numel() == token_idx - seq_start + 1
assert valid_indices.min().item() >= seq_start
assert valid_indices.max().item() <= token_idx

def test_shared_layers_reuse_topk_without_cross_context_leak(self):
# 验证 shared attention 复用同一 SequenceContext 的 source top-k,其他 context 保持独立
def test_shared_layer_consumes_explicit_topk_ids(self):
# 验证 shared attention 复用显式 IDs,并在漏传时立即报错
torch.manual_seed(0)
source_attention = _tiny_dsa_attention(indexer_types=["full", "shared"], layer_idx=0)
shared_attention = _tiny_dsa_attention(indexer_types=["full", "shared"], layer_idx=1)
position_embeddings = (torch.ones(1, 4, 2), torch.zeros(1, 4, 2))
hidden_states = torch.randn(1, 4, 4)
seq_ctx = SequenceContext.from_input_ids((torch.tensor([[1, 2, 3, 4]]),), device="cpu")

source_attention(hidden_states, position_embeddings, seq_ctx)
source_topk = seq_ctx.dsa_topk_cache.indices[0]
shared_output = shared_attention(hidden_states, position_embeddings, seq_ctx)["projected_output"]

other_seq_ctx = SequenceContext.from_input_ids((torch.tensor([[5, 6, 7, 8]]),), device="cpu")
source_attention(torch.randn(1, 4, 4), position_embeddings, other_seq_ctx)
source_outputs = source_attention(hidden_states, position_embeddings, seq_ctx)
dsa_topk_ids = source_outputs["dsa_topk_ids"]
shared_outputs = shared_attention(
hidden_states,
position_embeddings,
seq_ctx,
dsa_topk_ids=dsa_topk_ids,
)

assert torch.isfinite(shared_output).all()
assert seq_ctx.dsa_topk_cache.indices[0] is source_topk
assert other_seq_ctx.dsa_topk_cache.indices[0] is not source_topk
assert torch.isfinite(shared_outputs["projected_output"]).all()
assert shared_outputs["dsa_topk_ids"] is dsa_topk_ids
with pytest.raises(RuntimeError, match="requires dsa_topk_ids"):
shared_attention(hidden_states, position_embeddings, seq_ctx)

@pytest.mark.xfail(
reason="DSA top-k lifecycle still infers the checkpoint phase from `torch.is_grad_enabled()`, "
"which only held for the removed reentrant implementation, so the shared cache is never "
"released. Restoring this is part of the pending GLM-5.2 compatibility work.",
strict=False,
)
def test_checkpoint_reuses_and_releases_topk(self):
# 验证真实 source/shared decoder 经 checkpoint 重算后梯度有限且缓存释放。
def test_selective_checkpoint_preserves_explicit_topk_storage(self):
# 验证显式 IDs 可穿过 non-reentrant selective checkpoint 并完成反向。
torch.manual_seed(0)
source_block = apply_gradient_checkpointing(
_TinyDsaDecoderBlock(_tiny_dsa_attention(indexer_types=["full", "shared"], layer_idx=0))
source_block = apply_selective_checkpointing(
_tiny_dsa_decoder(["full", "shared"], layer_idx=0),
[("mlp.begin", "mlp.end")],
)
shared_block = apply_gradient_checkpointing(
_TinyDsaDecoderBlock(_tiny_dsa_attention(indexer_types=["full", "shared"], layer_idx=1))
shared_block = apply_selective_checkpointing(
_tiny_dsa_decoder(["full", "shared"], layer_idx=1),
[("mlp.begin", "mlp.end")],
)
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 = source_block(hidden_states, position_embeddings=position_embeddings, seq_ctx=seq_ctx)
output = shared_block(output, position_embeddings=position_embeddings, seq_ctx=seq_ctx)
output.square().mean().backward()
source_outputs = source_block(
hidden_states,
position_embeddings=position_embeddings,
seq_ctx=seq_ctx,
)
source_ids = source_outputs["dsa_topk_ids"]
shared_outputs = shared_block(
source_outputs["hidden_states"],
position_embeddings=position_embeddings,
seq_ctx=seq_ctx,
dsa_topk_ids=source_ids,
)
shared_ids = shared_outputs["dsa_topk_ids"]
assert shared_ids.untyped_storage().data_ptr() == source_ids.untyped_storage().data_ptr()
shared_outputs["hidden_states"].square().mean().backward()

assert torch.isfinite(hidden_states.grad).all()
assert seq_ctx.dsa_topk_cache.indices == {}
assert seq_ctx.dsa_topk_cache.offloaded == {}
assert source_ids.dtype == torch.int32


class TestAcceleratedSparseMLA:
Expand Down Expand Up @@ -325,12 +331,13 @@ def test_packed_attention_matches_full_sequence(self):
full_output_grad = torch.randn(1, 8, 4, device="cuda")

full_seq_ctx = SequenceContext.from_input_ids(packed_input_ids, device="cuda")
expected_output = attention(
expected_outputs = attention(
full_hidden_states,
position_embeddings=full_position_embeddings,
seq_ctx=full_seq_ctx,
)["projected_output"]
expected_topk = full_seq_ctx.dsa_topk_cache.indices[0].clone()
)
expected_output = expected_outputs["projected_output"]
expected_topk = expected_outputs["dsa_topk_ids"].clone()
expected_output.backward(full_output_grad)
expected_input_grad = full_hidden_states.grad.clone()
attention.zero_grad(set_to_none=True)
Expand All @@ -341,12 +348,13 @@ def test_packed_attention_matches_full_sequence(self):
shard_start = sp_seq_ctx.sp_rank * shard_size
shard_end = shard_start + shard_size
local_hidden_states = full_hidden_states.detach()[:, shard_start:shard_end].clone().requires_grad_()
local_output = attention(
local_outputs = attention(
local_hidden_states,
position_embeddings=tuple(x[:, shard_start:shard_end] for x in full_position_embeddings),
seq_ctx=sp_seq_ctx,
)["projected_output"]
local_topk = sp_seq_ctx.dsa_topk_cache.indices[0]
)
local_output = local_outputs["projected_output"]
local_topk = local_outputs["dsa_topk_ids"]
local_output.backward(full_output_grad[:, shard_start:shard_end])

gathered_output = [torch.empty_like(local_output) for _ in range(2)]
Expand Down
Loading
Loading