diff --git a/tests/engine/test_glm52_moe_train_engine.py b/tests/engine/test_glm52_moe_train_engine.py index baf40d3b6..997b0ad2e 100644 --- a/tests/engine/test_glm52_moe_train_engine.py +++ b/tests/engine/test_glm52_moe_train_engine.py @@ -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 @@ -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): @@ -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( @@ -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() diff --git a/tests/model/test_glm52_moe.py b/tests/model/test_glm52_moe.py index 5dd0c7736..1e5e7acd0 100644 --- a/tests/model/test_glm52_moe.py +++ b/tests/model/test_glm52_moe.py @@ -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 与梯度匹配完整序列。 """ @@ -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 @@ -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): diff --git a/tests/model/test_glm52_mtp_checkpoint_repro.py b/tests/model/test_glm52_mtp_checkpoint_repro.py index 758f14bb1..fba6d5127 100644 --- a/tests/model/test_glm52_mtp_checkpoint_repro.py +++ b/tests/model/test_glm52_mtp_checkpoint_repro.py @@ -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 梯度可正确反传。 """ @@ -12,7 +11,6 @@ import unittest from unittest import mock -import pytest import torch from xtuner._testing import DeterministicDDPTestCase @@ -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 @@ -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") diff --git a/tests/module/attention/test_dsa_mla.py b/tests/module/attention/test_dsa_mla.py index 6951bf518..02dc77c04 100644 --- a/tests/module/attention/test_dsa_mla.py +++ b/tests/module/attention/test_dsa_mla.py @@ -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 一致。 @@ -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 @@ -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, @@ -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: @@ -184,15 +178,17 @@ 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) @@ -200,43 +196,53 @@ def test_shared_layers_reuse_topk_without_cross_context_leak(self): 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: @@ -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) @@ -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)] diff --git a/tests/module/test_dense_decoder_layer.py b/tests/module/test_dense_decoder_layer.py index 032f8f32d..62fa56b54 100644 --- a/tests/module/test_dense_decoder_layer.py +++ b/tests/module/test_dense_decoder_layer.py @@ -1,6 +1,6 @@ -"""DenseDecoderLayer 多 micro-batch 行为测试。 +"""GLM52DenseDecoderLayer 多 micro-batch 行为测试。 -TestDenseDecoderLayerMicroBatch +TestGLM52DenseDecoderLayerMicroBatch test_batched_inputs_match_independent_forwards: 等长 micro-batch 的输出与梯度等价于独立调用。 """ @@ -9,12 +9,12 @@ import torch from xtuner.v1.data_proto import SequenceContext -from xtuner.v1.module.attention import DSAMLAConfig -from xtuner.v1.module.decoder_layer.dense_decoder_layer import DenseDecoderLayer +from xtuner.v1.model.moe.glm52 import DSAMLAConfig +from xtuner.v1.model.moe.glm52.decoder_layer import GLM52DenseDecoderLayer -def _build_dense_dsa_layer() -> DenseDecoderLayer: - return DenseDecoderLayer( +def _build_dense_dsa_layer() -> GLM52DenseDecoderLayer: + return GLM52DenseDecoderLayer( hidden_size=4, intermediate_size=8, hidden_act="silu", @@ -55,7 +55,7 @@ def _build_inputs() -> tuple[ return hidden_states, position_embeddings, seq_ctx -class TestDenseDecoderLayerMicroBatch: +class TestGLM52DenseDecoderLayerMicroBatch: def test_batched_inputs_match_independent_forwards(self): # 验证一次多输入调用与逐 micro-batch 调用产生相同输出、输入梯度和参数梯度。 torch.manual_seed(0) @@ -69,7 +69,7 @@ def test_batched_inputs_match_independent_forwards(self): ] outputs = layer( - *hidden_states, + hidden_states, position_embeddings=position_embeddings, seq_ctx=seq_ctx, ) @@ -86,12 +86,18 @@ def test_batched_inputs_match_independent_forwards(self): ) ) - assert isinstance(outputs, tuple) - for output, reference_output in zip(outputs, reference_outputs): + output_hidden = outputs["hidden_states"] + output_ids = outputs["dsa_topk_ids"] + reference_hidden = tuple(result["hidden_states"] for result in reference_outputs) + reference_ids = tuple(result["dsa_topk_ids"] for result in reference_outputs) + for output, reference_output in zip(output_hidden, reference_hidden): torch.testing.assert_close(output, reference_output) + for dsa_topk_ids, reference_dsa_topk_ids in zip(output_ids, reference_ids): + torch.testing.assert_close(dsa_topk_ids, reference_dsa_topk_ids) + assert dsa_topk_ids.dtype == torch.int32 - sum(output.sum() for output in outputs).backward() - sum(output.sum() for output in reference_outputs).backward() + sum(output.sum() for output in output_hidden).backward() + sum(output.sum() for output in reference_hidden).backward() for hidden, reference_hidden in zip(hidden_states, reference_hidden_states): torch.testing.assert_close(hidden.grad, reference_hidden.grad) diff --git a/xtuner/v1/data_proto/__init__.py b/xtuner/v1/data_proto/__init__.py index 6194971cb..c30af9de4 100644 --- a/xtuner/v1/data_proto/__init__.py +++ b/xtuner/v1/data_proto/__init__.py @@ -1,7 +1,6 @@ -from .sequence_context import DSATopKCacheState, SequenceContext +from .sequence_context import SequenceContext __all__ = [ - "DSATopKCacheState", "SequenceContext", ] diff --git a/xtuner/v1/data_proto/sequence_context.py b/xtuner/v1/data_proto/sequence_context.py index 6860ee362..e17a1efad 100644 --- a/xtuner/v1/data_proto/sequence_context.py +++ b/xtuner/v1/data_proto/sequence_context.py @@ -1,5 +1,4 @@ # Copyright (c) OpenMMLab. All rights reserved. -import itertools from typing import cast import torch @@ -9,52 +8,6 @@ from .utils import gather_for_sequence_parallel, pad_to_multiple_of, split_for_sequence_parallel -_DSA_TOPK_CONTEXT_IDS = itertools.count() - - -class DSATopKCacheState: - """Mutable DSA cross-layer top-k cache, scoped to one microbatch. - - For example, if source layer 2 provides top-k indices to layers 2, 3, and 4, - its original forward stores ``indices[2]``. After layer 4's no-grad - checkpoint forward, ``checkpoint_active`` becomes true and top-k offload may - replace that entry with ``offloaded[2]``. Backward then replays layers 4, 3, - and 2; layer 2 removes the cache and adds 2 to ``released_sources``. If one - physical MTP source is reused at two logical depths, both MTP counters start - at 2 so the cache is transferred and released only after the second use in - each phase. - """ - - indices: dict[int, torch.Tensor] # GPU-resident top-k, keyed by source layer. - offloaded: dict[int, str] # OffloadManager key for each CPU-resident source. - released_sources: set[int] # Sources whose backward replay lifetime has ended. - checkpoint_active: bool # Whether checkpoint forward retained this cache for replay. - context_id: int # Process-local identifier used to make offload keys unique. - 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. - - def __init__( - self, - *, - indices: dict[int, torch.Tensor] | None = None, - offloaded: dict[int, str] | None = None, - released_sources: set[int] | None = None, - checkpoint_active: bool = False, - context_id: int | None = None, - mtp_forward_uses_remaining: dict[int, int] | None = None, - mtp_replays_remaining: dict[int, int] | None = None, - ) -> None: - # topk_indices format: {source_layer_idx: [seq_len, kv_group, topk]}. - # Invalid/padded sparse slots are represented by -1. - self.indices = {} if indices is None else indices - self.offloaded = {} if offloaded is None else offloaded - self.released_sources = set() if released_sources is None else released_sources - self.checkpoint_active = checkpoint_active - self.context_id = next(_DSA_TOPK_CONTEXT_IDS) if context_id is None else context_id - 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 - - # Avoid using dataclass decorator here to get rid of extra ops called in pytorch 2.8 and above # The extra ops is introduced by function _apply_to_tensors in # https://github.com/pytorch/pytorch/blob/v2.8.0/torch/distributed/fsdp/_fully_shard/_fsdp_state.py @@ -97,7 +50,6 @@ class SequenceContext: # moe routed_experts rollout_routed_experts: torch.Tensor | None offload_rollout_routed_experts: bool - dsa_topk_cache: DSATopKCacheState # Private backing attributes for SP shard reconstruction _raw_input_ids: torch.LongTensor | None @@ -127,7 +79,6 @@ def __init__( num_img_tokens: list[list[int]] | None = None, rollout_routed_experts: torch.Tensor | None = None, offload_rollout_routed_experts: bool = False, - dsa_topk_cache: DSATopKCacheState | None = None, # SP shard metadata: private, accessed via properties below raw_input_ids: torch.LongTensor | None = None, raw_inputs_embeds: torch.FloatTensor | None = None, @@ -162,7 +113,6 @@ def __init__( self.num_img_tokens = num_img_tokens self.rollout_routed_experts = rollout_routed_experts self.offload_rollout_routed_experts = offload_rollout_routed_experts - self.dsa_topk_cache = DSATopKCacheState() if dsa_topk_cache is None else dsa_topk_cache self._raw_input_ids = raw_input_ids self._raw_inputs_embeds = raw_inputs_embeds self._shard_start = shard_start @@ -551,7 +501,6 @@ def copy(self, **overrides) -> Self: offload_rollout_routed_experts=overrides.get( "offload_rollout_routed_experts", self.offload_rollout_routed_experts ), - dsa_topk_cache=overrides.get("dsa_topk_cache", self.dsa_topk_cache), raw_input_ids=overrides.get("raw_input_ids", self._raw_input_ids), raw_inputs_embeds=overrides.get("raw_inputs_embeds", self._raw_inputs_embeds), shard_start=overrides.get("shard_start", self._shard_start), @@ -643,5 +592,4 @@ def data(self) -> dict: "num_img_tokens": self.num_img_tokens, "rollout_routed_experts": self.rollout_routed_experts, "offload_rollout_routed_experts": self.offload_rollout_routed_experts, - "dsa_topk_cache": self.dsa_topk_cache, } diff --git a/xtuner/v1/model/moe/glm52/__init__.py b/xtuner/v1/model/moe/glm52/__init__.py new file mode 100644 index 000000000..1e1c3c668 --- /dev/null +++ b/xtuner/v1/model/moe/glm52/__init__.py @@ -0,0 +1,10 @@ +from .dsa_mla import DSAMLAConfig, DSAMultiLatentAttention +from .glm52 import Glm52MoE, Glm52MoEConfig + + +__all__ = [ + "DSAMLAConfig", + "DSAMultiLatentAttention", + "Glm52MoE", + "Glm52MoEConfig", +] diff --git a/xtuner/v1/model/moe/glm52/decoder_layer.py b/xtuner/v1/model/moe/glm52/decoder_layer.py new file mode 100644 index 000000000..2330389ee --- /dev/null +++ b/xtuner/v1/model/moe/glm52/decoder_layer.py @@ -0,0 +1,208 @@ +from typing import cast + +import torch +from typing_extensions import override + +from xtuner.v1.data_proto import SequenceContext +from xtuner.v1.module import AttnOutputs, RouterResults +from xtuner.v1.module.decoder_layer.dense_decoder_layer import ( + DenseDecoderLayer, + DenseDecoderLayerMicroBatchOutput, + DenseDecoderLayerOutput, +) +from xtuner.v1.module.decoder_layer.moe_decoder_layer import ( + MoEDecoderLayer, + MoEDecoderLayerMicroBatchOutput, + MoEDecoderLayerOutput, +) +from xtuner.v1.utils import checkpoint_record + +from .dsa_mla import DSAMultiLatentAttention, GLM52AttnOutputs + + +class GLM52DenseDecoderLayerOutput(DenseDecoderLayerOutput): + """GLM-5.2 dense-layer output with explicit DSA IDs.""" + + dsa_topk_ids: torch.Tensor + + +class GLM52DenseDecoderLayerMicroBatchOutput(DenseDecoderLayerMicroBatchOutput): + """GLM-5.2 dense-layer outputs for intra-layer micro-batches.""" + + dsa_topk_ids: list[torch.Tensor] + + +class GLM52DenseDecoderLayer(DenseDecoderLayer): + """Dense decoder layer that threads GLM-5.2 DSA IDs explicitly.""" + + @override + def forward( + self, + hidden_states: torch.Tensor | list[torch.Tensor], + *, + position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], + seq_ctx: SequenceContext | list[SequenceContext], + dsa_topk_ids: torch.Tensor | list[torch.Tensor] | None = None, + ) -> GLM52DenseDecoderLayerOutput | GLM52DenseDecoderLayerMicroBatchOutput: + if not isinstance(hidden_states, list): + assert isinstance(position_embeddings, tuple) and len(position_embeddings) == 2 + assert isinstance(seq_ctx, SequenceContext) + assert dsa_topk_ids is None or isinstance(dsa_topk_ids, torch.Tensor) + return self._glm52_forward( + hidden_states=hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + dsa_topk_ids=dsa_topk_ids, + ) + + n = len(hidden_states) + assert isinstance(position_embeddings, list) and len(position_embeddings) == n + assert isinstance(seq_ctx, list) and len(seq_ctx) == n + assert all(hidden.shape == hidden_states[0].shape for hidden in hidden_states) + if dsa_topk_ids is None: + dsa_topk_ids_list: list[torch.Tensor | None] = [None] * n + else: + assert isinstance(dsa_topk_ids, list) and len(dsa_topk_ids) == n + dsa_topk_ids_list = list(dsa_topk_ids) + + layer_results = [ + self._glm52_forward( + hidden_states=hidden, + position_embeddings=position_embedding, + seq_ctx=context, + dsa_topk_ids=topk_ids, + ) + for hidden, topk_ids, position_embedding, context in zip( + hidden_states, dsa_topk_ids_list, position_embeddings, seq_ctx + ) + ] + return { + "hidden_states": [result["hidden_states"] for result in layer_results], + "dsa_topk_ids": [result["dsa_topk_ids"] for result in layer_results], + } + + def _glm52_forward( + self, + *, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + seq_ctx: SequenceContext, + dsa_topk_ids: torch.Tensor | None, + ) -> GLM52DenseDecoderLayerOutput: + residual = hidden_states + hidden_states = self.input_layernorm(hidden_states) + + checkpoint_record("attn.begin") + attention = cast(DSAMultiLatentAttention, self.self_attn) + attn_outputs = attention( + hidden_states=hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + dsa_topk_ids=dsa_topk_ids, + ) + checkpoint_record("attn.end") + hidden_states = residual + attn_outputs["projected_output"] + + residual = hidden_states + hidden_states = self.post_attention_layernorm(hidden_states) + checkpoint_record("mlp.begin") + hidden_states = self.mlp(hidden_states) + checkpoint_record("mlp.end") + hidden_states = residual + hidden_states + + return { + "hidden_states": hidden_states, + "dsa_topk_ids": attn_outputs["dsa_topk_ids"], + } + + +class GLM52MoEDecoderLayerOutput(MoEDecoderLayerOutput): + """GLM-5.2 MoE-layer output with explicit DSA IDs.""" + + dsa_topk_ids: torch.Tensor + + +class GLM52MoEDecoderLayerMicroBatchOutput(MoEDecoderLayerMicroBatchOutput): + """GLM-5.2 MoE-layer outputs for intra-layer micro-batches.""" + + dsa_topk_ids: list[torch.Tensor] + + +class GLM52MoEDecoderLayer(MoEDecoderLayer): + """MoE decoder layer that threads GLM-5.2 DSA IDs explicitly.""" + + @override + def forward( + self, + hidden_states: torch.Tensor | list[torch.Tensor], + *, + seq_ctx: SequenceContext | list[SequenceContext], + position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], + dsa_topk_ids: torch.Tensor | list[torch.Tensor] | None = None, + ) -> GLM52MoEDecoderLayerOutput | GLM52MoEDecoderLayerMicroBatchOutput: + if not isinstance(hidden_states, list): + assert isinstance(seq_ctx, SequenceContext) + assert isinstance(position_embeddings, tuple) and len(position_embeddings) == 2 + assert dsa_topk_ids is None or isinstance(dsa_topk_ids, torch.Tensor) + return cast( + GLM52MoEDecoderLayerOutput, + self._forward( + hidden_states=hidden_states, + seq_ctx=seq_ctx, + position_embeddings=position_embeddings, + attention_kwargs={"dsa_topk_ids": dsa_topk_ids}, + ), + ) + + n = len(hidden_states) + assert isinstance(seq_ctx, list) and len(seq_ctx) == n + assert isinstance(position_embeddings, list) and len(position_embeddings) == n + if dsa_topk_ids is None: + dsa_topk_ids_list: list[torch.Tensor | None] = [None] * n + else: + assert isinstance(dsa_topk_ids, list) and len(dsa_topk_ids) == n + dsa_topk_ids_list = list(dsa_topk_ids) + + return cast( + GLM52MoEDecoderLayerMicroBatchOutput, + self._micro_batch_forward( + hidden_states_list=hidden_states, + seq_ctx_list=seq_ctx, + position_embeddings_list=position_embeddings, + attention_kwargs_list=[{"dsa_topk_ids": topk_ids} for topk_ids in dsa_topk_ids_list], + ), + ) + + @override + def _build_output( + self, + *, + hidden_states: torch.Tensor, + router_results: RouterResults, + attn_outputs: AttnOutputs, + ) -> GLM52MoEDecoderLayerOutput: + glm_attn_outputs = cast(GLM52AttnOutputs, attn_outputs) + return { + "hidden_states": hidden_states, + "router_logits": router_results["logits"], + "router_weights": router_results["router_weights"], + "router_topk_ids": router_results["topk_ids"], + "dsa_topk_ids": glm_attn_outputs["dsa_topk_ids"], + } + + @override + def _build_micro_batch_output( + self, + *, + hidden_states_list: list[torch.Tensor], + router_results_list: list[RouterResults], + attn_outputs_list: list[AttnOutputs], + ) -> GLM52MoEDecoderLayerMicroBatchOutput: + glm_attn_outputs = [cast(GLM52AttnOutputs, output) for output in attn_outputs_list] + return { + "hidden_states": hidden_states_list, + "router_logits": [result["logits"] for result in router_results_list], + "router_weights": [result["router_weights"] for result in router_results_list], + "router_topk_ids": [result["topk_ids"] for result in router_results_list], + "dsa_topk_ids": [output["dsa_topk_ids"] for output in glm_attn_outputs], + } diff --git a/xtuner/v1/module/attention/dsa_mla.py b/xtuner/v1/model/moe/glm52/dsa_mla.py similarity index 88% rename from xtuner/v1/module/attention/dsa_mla.py rename to xtuner/v1/model/moe/glm52/dsa_mla.py index 23e0f9f68..348007502 100644 --- a/xtuner/v1/module/attention/dsa_mla.py +++ b/xtuner/v1/model/moe/glm52/dsa_mla.py @@ -4,10 +4,14 @@ import torch from torch import nn from torch.distributed.tensor import DTensor +from typing_extensions import overload from xtuner.v1.config import GenerateConfig from xtuner.v1.data_proto import SequenceContext from xtuner.v1.float8.config import Float8Config +from xtuner.v1.module.attention.attn_outputs import AttnOutputs +from xtuner.v1.module.attention.mla import MLAConfig, MultiLatentAttention, mla_apply_rotary_pos_emb +from xtuner.v1.module.linear import build_linear from xtuner.v1.module.rope import RopeScalingConfig from xtuner.v1.ops.comm import gather_for_sequence_parallel from xtuner.v1.ops.sparse_mla import ( @@ -19,10 +23,13 @@ get_sparse_mla, ) -from ..linear import build_linear -from .attn_outputs import AttnOutputs -from .dsa_topk_sharing import build_dsa_topk_release_plan, dsa_topk_source_layer, get_dsa_topk_sharing_runtime -from .mla import MLAConfig, MultiLatentAttention, mla_apply_rotary_pos_emb +from .dsa_topk_sharing import dsa_topk_source_layer + + +class GLM52AttnOutputs(AttnOutputs): + """GLM-5.2 attention outputs with explicit cross-layer DSA IDs.""" + + dsa_topk_ids: torch.Tensor class LayerNorm(nn.Module): @@ -141,15 +148,9 @@ def forward( # weights: [bsz, S, Ni] weights = self.weights_proj(hidden_states).float() * (self.index_n_heads**-0.5) - # Top-k 索引是整数,不需要梯度,所以整个 indexer 都放在 no_grad 下。 - # 这解释了 Case 1 为什么只在 compile 下显错: - # eager COMPUTE: indexer 不产生槽位 -> SparseMLA 保存 [A, B, C] - # eager REUSE: cache read 不产生槽位 -> SparseMLA 保存 [A, B, C] - # original/replay 虽然走了不同分支,但 checkpoint 看到的保存清单仍能对齐。 - # compile 会把 indexer 周围的可求导计算按 compiled block 打包;COMPUTE 与 - # REUSE 经过不同 graph break 后,可能分别保存 [A, B, C, D] 和 - # [A, X, C, D],同一槽位的 metadata 不同才触发 CheckpointError。 - # 这里的字母只表示保存槽位,不表示真实变量或 Tensor 数值。 + # IDs are discrete and never need gradients. In the first explicit- + # dataflow implementation, a checkpointed source layer reruns this + # no-grad region during backward replay. # 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) @@ -238,18 +239,6 @@ def __init__( self.indexer_types = indexer_types self.sparse_mla_backend = sparse_mla_backend 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 = {}, {} - else: - release_plan = build_dsa_topk_release_plan( - num_main_layers=len(indexer_types), - num_mtp_layers=0, - indexer_types=indexer_types, - index_skip_topk_offset=index_skip_topk_offset, - index_topk_freq=index_topk_freq, - ) - self.dsa_topk_last_use = release_plan.forward_last_use - self.dsa_topk_recompute_release = release_plan.recompute_release if self.q_lora_rank is None: raise ValueError("DSA MLA requires q_lora_rank because the indexer consumes q_a_layernorm output.") @@ -278,7 +267,8 @@ def forward( hidden_states: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], seq_ctx: SequenceContext, - ) -> AttnOutputs: + dsa_topk_ids: torch.Tensor | None = None, + ) -> GLM52AttnOutputs: """Absorbed DSA-MLA forward for packed training (``bsz == 1``). Shapes use ``S`` for the local sequence length and ``S_g`` for the @@ -352,21 +342,28 @@ 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, - position_embeddings, - seq_ctx, - ), - ) + # A source layer computes IDs once; shared layers receive the same + # explicit tensor reference from the GLM decoder stack. + if dsa_topk_ids is None: + if not hasattr(self, "indexer"): + raise RuntimeError(f"DSA shared layer {self.layer_idx} requires dsa_topk_ids.") + dsa_topk_ids = ( + self.indexer( + hidden_states, + q_resid, + position_embeddings, + seq_ctx, + ) + .to(torch.int32) + .contiguous() + ) + elif dsa_topk_ids.dtype != torch.int32 or not dsa_topk_ids.is_contiguous(): + raise RuntimeError("dsa_topk_ids must be a contiguous torch.int32 tensor.") + sparse_mla_outputs = self.sparse_mla_func( query_states, key_states, - topk_indices, + dsa_topk_ids, self.softmax_scale, value_dim=self.kv_lora_rank, ) @@ -383,4 +380,16 @@ def forward( "raw_output": raw_output, "projected_output": projected_output, "softmax_lse": softmax_lse, + "dsa_topk_ids": dsa_topk_ids, } + + @overload # type: ignore + def __call__( # type: ignore + self, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + seq_ctx: SequenceContext, + dsa_topk_ids: torch.Tensor | None = None, + ) -> GLM52AttnOutputs: ... + + __call__ = nn.Module.__call__ diff --git a/xtuner/v1/model/moe/glm52/dsa_topk_sharing.py b/xtuner/v1/model/moe/glm52/dsa_topk_sharing.py new file mode 100644 index 000000000..baee96931 --- /dev/null +++ b/xtuner/v1/model/moe/glm52/dsa_topk_sharing.py @@ -0,0 +1,46 @@ +# Copyright (c) OpenMMLab. All rights reserved. + + +def dsa_topk_source_layer( + *, + layer_idx: int, + indexer_types: list[str] | None, + index_skip_topk_offset: int, + index_topk_freq: int, +) -> int: + """Resolve the GLM-5.2 source layer whose DSA top-k IDs a layer + consumes.""" + if indexer_types is not None: + if layer_idx < len(indexer_types) and indexer_types[layer_idx] == "full": + return layer_idx + for source_layer_idx in range(min(layer_idx, len(indexer_types) - 1), -1, -1): + if indexer_types[source_layer_idx] == "full": + return source_layer_idx + raise ValueError(f"DSA layer {layer_idx} has no preceding full indexer layer.") + + if index_topk_freq <= 1: + return layer_idx + + source_layer_idx = layer_idx + while (max(source_layer_idx + 1 - index_skip_topk_offset, 0) % index_topk_freq) != 0: + source_layer_idx -= 1 + return source_layer_idx + + +def dsa_topk_source_layers( + *, + num_layers: int, + indexer_types: list[str] | None, + index_skip_topk_offset: int, + index_topk_freq: int, +) -> tuple[int, ...]: + """Return the source-layer index for every layer in one decoder stack.""" + return tuple( + dsa_topk_source_layer( + layer_idx=layer_idx, + indexer_types=indexer_types, + index_skip_topk_offset=index_skip_topk_offset, + index_topk_freq=index_topk_freq, + ) + for layer_idx in range(num_layers) + ) diff --git a/xtuner/v1/model/moe/glm52.py b/xtuner/v1/model/moe/glm52/glm52.py similarity index 72% rename from xtuner/v1/model/moe/glm52.py rename to xtuner/v1/model/moe/glm52/glm52.py index fd8bea587..401d5c44c 100644 --- a/xtuner/v1/model/moe/glm52.py +++ b/xtuner/v1/model/moe/glm52/glm52.py @@ -1,56 +1,70 @@ +import os import re from pathlib import Path -from typing import Literal +from typing import Callable, Literal, cast import torch from pydantic import Field, computed_field from typing_extensions import Self, override from transformers.models.glm_moe_dsa import GlmMoeDsaConfig as HFGlmMoeDsaConfig +from xtuner.v1.data_proto import SequenceContext from xtuner.v1.model.base import DEFAULT_FLOAT8_CFG, TorchCompileOption -from xtuner.v1.model.moe.moe import BalancingLossConfig, MoEConfig, ZLossConfig -from xtuner.v1.module.attention import DSAMLAConfig, DSAMultiLatentAttention -from xtuner.v1.module.attention.dsa_topk_sharing import ( - build_dsa_topk_release_plan, - configure_dsa_mtp_iteration_lifecycle, - configure_dsa_topk_decoder_lifecycle, - dsa_topk_source_layer, +from xtuner.v1.model.moe.moe import BalancingLossConfig, MoE, MoEConfig, ZLossConfig +from xtuner.v1.module.decoder_layer.dense_decoder_layer import ( + DenseDecoderLayerMicroBatchOutput, + DenseDecoderLayerOutput, ) -from xtuner.v1.module.mtp import MTPConfig, MTPLayer +from xtuner.v1.module.decoder_layer.moe_decoder_layer import ( + MoEDecoderLayerMicroBatchOutput, + MoEDecoderLayerOutput, +) +from xtuner.v1.module.mtp import MTPConfig from xtuner.v1.module.rope import RopeParametersConfig from xtuner.v1.module.router.noaux_router import NoAuxRouterConfig -from .moe import MoE +from .decoder_layer import ( + GLM52DenseDecoderLayer, + GLM52DenseDecoderLayerMicroBatchOutput, + GLM52DenseDecoderLayerOutput, + GLM52MoEDecoderLayer, + GLM52MoEDecoderLayerMicroBatchOutput, + GLM52MoEDecoderLayerOutput, +) +from .dsa_mla import DSAMLAConfig, DSAMultiLatentAttention +from .dsa_topk_sharing import dsa_topk_source_layer, dsa_topk_source_layers +from .mtp import GLM52MTPBlock, GLM52MTPLayer -# GLM DSA attention records cross-layer top-k indices in SequenceContext. -# That Python-side cache mutation is intentionally kept out of strict fullgraph -# regions, so decoder/pre-attn/DSA/dense boundaries allow graph breaks while -# pure tensor MoE expert sub-stages stay fullgraph. +# Keep the existing graph boundaries while explicit DSA top-k tensor inputs and +# results are validated. Each boundary can be tightened independently later. MOE_NON_EP_COMPILE_CFG: dict[str, TorchCompileOption] = { "xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEBlock.forward": TorchCompileOption(fullgraph=True), - "xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEDecoderLayer.forward": TorchCompileOption(fullgraph=False), + "xtuner.v1.model.moe.glm52.decoder_layer.GLM52MoEDecoderLayer.forward": TorchCompileOption(fullgraph=False), "xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEDecoderLayer._pre_moe_forward": TorchCompileOption( fullgraph=False ), - "xtuner.v1.module.attention.dsa_mla.DSAMultiLatentAttention.forward": TorchCompileOption(fullgraph=False), + "xtuner.v1.model.moe.glm52.dsa_mla.DSAMultiLatentAttention.forward": TorchCompileOption(fullgraph=False), "xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEDecoderLayer._shared_experts_forward": TorchCompileOption( fullgraph=True ), "xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEDecoderLayer._post_moe_forward": TorchCompileOption( fullgraph=True ), - "xtuner.v1.module.decoder_layer.dense_decoder_layer.DenseDecoderLayer.forward": TorchCompileOption( - fullgraph=False - ), + "xtuner.v1.model.moe.glm52.decoder_layer.GLM52DenseDecoderLayer.forward": TorchCompileOption(fullgraph=False), **DEFAULT_FLOAT8_CFG, } MOE_EP_COMPILE_CFG = MOE_NON_EP_COMPILE_CFG.copy() -MOE_EP_COMPILE_CFG.pop("xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEDecoderLayer.forward") +MOE_EP_COMPILE_CFG.pop("xtuner.v1.model.moe.glm52.decoder_layer.GLM52MoEDecoderLayer.forward") class Glm52MoE(MoE): + dense_decoder_layer_cls = GLM52DenseDecoderLayer + moe_decoder_layer_cls = GLM52MoEDecoderLayer + mtp_layer_cls = GLM52MTPLayer + mtp_block_cls = GLM52MTPBlock + @property @override def default_compile_cfg(self) -> dict[str, TorchCompileOption]: @@ -59,62 +73,121 @@ def default_compile_cfg(self) -> dict[str, TorchCompileOption]: return MOE_NON_EP_COMPILE_CFG @override - def _configure_model_specific_layer_lifecycle(self) -> None: - dsa_layers: list[tuple[torch.nn.Module, DSAMultiLatentAttention]] = [] - mtp_attention: DSAMultiLatentAttention | None = None + def _configure_model_specific_layers(self) -> None: + dsa_layers: list[DSAMultiLatentAttention] = [] for decoder_layer in self.layers.values(): self_attn = decoder_layer.self_attn # type: ignore[attr-defined] assert isinstance(self_attn, DSAMultiLatentAttention), ( f"GLM-5.2 requires DSAMultiLatentAttention, got {type(self_attn).__name__}." ) - dsa_layers.append((decoder_layer, self_attn)) + dsa_layers.append(self_attn) - num_physical_mtp_layers = 0 if self.mtp_block is not None and self.config.mtp_config is not None: num_physical_mtp_layers = 1 if self.config.mtp_config.share_weights else self.config.mtp_config.num_layers for mtp_idx in range(num_physical_mtp_layers): mtp_layer = self.mtp_block.layers[mtp_idx] - assert isinstance(mtp_layer, MTPLayer) - decoder_layer = mtp_layer.decoder_layer - self_attn = decoder_layer.self_attn # type: ignore[attr-defined] + assert isinstance(mtp_layer, GLM52MTPLayer) + self_attn = mtp_layer.decoder_layer.self_attn # type: ignore[attr-defined] assert isinstance(self_attn, DSAMultiLatentAttention), ( f"GLM-5.2 MTP requires DSAMultiLatentAttention, got {type(self_attn).__name__}." ) - dsa_layers.append((decoder_layer, self_attn)) - if mtp_idx == 0: - mtp_attention = self_attn - - sample_attn = dsa_layers[0][1] - release_plan = build_dsa_topk_release_plan( - num_main_layers=self.config.num_hidden_layers, - num_mtp_layers=num_physical_mtp_layers, + + sample_attn = dsa_layers[0] + self._dsa_topk_source_layers = dsa_topk_source_layers( + num_layers=self.config.num_hidden_layers, indexer_types=sample_attn.indexer_types, index_skip_topk_offset=sample_attn.index_skip_topk_offset, index_topk_freq=sample_attn.index_topk_freq, ) - for decoder_layer, self_attn in dsa_layers: - # DSA top-k sharing spans dense prefix, sparse MoE layers, and the - # optional MTP layer. The attention-local default release maps only - # see the main-stack indexer_types, so GLM-5.2 injects a model-level - # plan with the full physical layer topology. - configure_dsa_topk_decoder_lifecycle( - decoder_layer=decoder_layer, - attention=self_attn, - release_plan=release_plan, - ) + self._dsa_topk_last_consumers = frozenset( + layer_idx + for layer_idx, source_layer_idx in enumerate(self._dsa_topk_source_layers) + if layer_idx == self.config.num_hidden_layers - 1 + or self._dsa_topk_source_layers[layer_idx + 1] != source_layer_idx + ) - if ( - self.mtp_block is not None - and self.config.mtp_config is not None - and self.config.mtp_config.share_weights - and self.config.index_share_for_mtp_iteration - ): - assert mtp_attention is not None - configure_dsa_mtp_iteration_lifecycle( - mtp_block=self.mtp_block, - attention=mtp_attention, - num_iterations=self.config.mtp_config.num_layers, + @override + def _call_decoder_layer( + self, + *, + decoder_layer: torch.nn.Module, + layer_idx: int, + hidden_states: torch.Tensor | list[torch.Tensor], + position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], + seq_ctx: SequenceContext | list[SequenceContext], + previous_layer_results: ( + DenseDecoderLayerOutput + | DenseDecoderLayerMicroBatchOutput + | MoEDecoderLayerOutput + | MoEDecoderLayerMicroBatchOutput + | None + ), + ) -> ( + DenseDecoderLayerOutput + | DenseDecoderLayerMicroBatchOutput + | MoEDecoderLayerOutput + | MoEDecoderLayerMicroBatchOutput + ): + """Arrange GLM-5.2 DSA dataflow and offload around one decoder + layer.""" + is_micro_batch = isinstance(hidden_states, list) + is_source_layer = self._dsa_topk_source_layers[layer_idx] == layer_idx + if is_source_layer: + dsa_topk_ids: torch.Tensor | list[torch.Tensor] | None = None + else: + previous_results = cast( + GLM52DenseDecoderLayerOutput + | GLM52DenseDecoderLayerMicroBatchOutput + | GLM52MoEDecoderLayerOutput + | GLM52MoEDecoderLayerMicroBatchOutput, + previous_layer_results, + ) + dsa_topk_ids = previous_results["dsa_topk_ids"] + + activation_offload = int(os.getenv("XTUNER_ACTIVATION_OFFLOAD", "0")) == 1 + dsa_topk_offload = int(os.getenv("XTUNER_DSA_TOPK_OFFLOAD", "0")) == 1 + offload_tensors: list[torch.Tensor] = [] + if activation_offload and layer_idx >= self.config.first_k_dense_replace: + offload_tensors = list(hidden_states) if is_micro_batch else [cast(torch.Tensor, hidden_states)] + if dsa_topk_offload and layer_idx in self._dsa_topk_last_consumers and dsa_topk_ids is not None: + offload_tensors.extend(dsa_topk_ids if isinstance(dsa_topk_ids, list) else [dsa_topk_ids]) + + # The offload context expects a dense zero-based block index across every + # active activation or DSA-ID window, including dense GLM layers. + offload_block_idx = sum( + (activation_offload and previous_idx >= self.config.first_k_dense_replace) + or ( + dsa_topk_offload + and previous_idx in self._dsa_topk_last_consumers + and self._dsa_topk_source_layers[previous_idx] != previous_idx + ) + for previous_idx in (int(idx) for idx in self.layers) + if previous_idx < layer_idx + ) + decoder_forward = cast( + Callable[ + ..., + DenseDecoderLayerOutput + | DenseDecoderLayerMicroBatchOutput + | MoEDecoderLayerOutput + | MoEDecoderLayerMicroBatchOutput, + ], + decoder_layer, + ) + with self._saved_tensors_offload_ctx(offload_block_idx, offload_tensors): + layer_results = decoder_forward( + hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + dsa_topk_ids=dsa_topk_ids, ) + if layer_idx < self.config.first_k_dense_replace: + if is_micro_batch: + return cast(GLM52DenseDecoderLayerMicroBatchOutput, layer_results) + return cast(GLM52DenseDecoderLayerOutput, layer_results) + if is_micro_batch: + return cast(GLM52MoEDecoderLayerMicroBatchOutput, layer_results) + return cast(GLM52MoEDecoderLayerOutput, layer_results) def to_hf_key_list(self, key: str) -> list[str]: if self.config.tie_word_embeddings and "lm_head" in key: diff --git a/xtuner/v1/model/moe/glm52/mtp.py b/xtuner/v1/model/moe/glm52/mtp.py new file mode 100644 index 000000000..12adf480f --- /dev/null +++ b/xtuner/v1/model/moe/glm52/mtp.py @@ -0,0 +1,118 @@ +from typing import cast + +import torch +from typing_extensions import override + +from xtuner.v1.data_proto import SequenceContext +from xtuner.v1.module.decoder_layer.moe_decoder_layer import ( + MoEDecoderLayerMicroBatchOutput, + MoEDecoderLayerOutput, +) +from xtuner.v1.module.mtp import MTPBlock, MTPLayer +from xtuner.v1.module.mtp.mtp_block import MTPInternalOutput + +from .decoder_layer import ( + GLM52MoEDecoderLayerMicroBatchOutput, + GLM52MoEDecoderLayerOutput, +) + + +class GLM52MTPLayer(MTPLayer): + """MTP layer whose wrapped GLM-5.2 decoder consumes explicit DSA IDs.""" + + @override + def forward( + self, + hidden_states: torch.Tensor | list[torch.Tensor], + *, + future_embeddings: torch.Tensor | list[torch.Tensor], + position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], + seq_ctx: SequenceContext | list[SequenceContext], + dsa_topk_ids: torch.Tensor | list[torch.Tensor] | None = None, + ) -> GLM52MoEDecoderLayerOutput | GLM52MoEDecoderLayerMicroBatchOutput: + if not isinstance(hidden_states, list): + assert isinstance(future_embeddings, torch.Tensor) + assert isinstance(position_embeddings, tuple) and len(position_embeddings) == 2 + assert isinstance(seq_ctx, SequenceContext) + assert dsa_topk_ids is None or isinstance(dsa_topk_ids, torch.Tensor) + projected = self._preprocess(hidden_states=hidden_states, future_embeddings=future_embeddings) + layer_results = cast( + GLM52MoEDecoderLayerOutput, + self.decoder_layer( + projected, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + dsa_topk_ids=dsa_topk_ids, + ), + ) + return { + "hidden_states": self.final_layernorm(layer_results["hidden_states"]), + "router_logits": layer_results["router_logits"], + "router_weights": layer_results["router_weights"], + "router_topk_ids": layer_results["router_topk_ids"], + "dsa_topk_ids": layer_results["dsa_topk_ids"], + } + + n = len(hidden_states) + assert isinstance(future_embeddings, list) and len(future_embeddings) == n + assert isinstance(position_embeddings, list) and len(position_embeddings) == n + assert isinstance(seq_ctx, list) and len(seq_ctx) == n + if dsa_topk_ids is None: + decoder_topk_ids = None + else: + assert isinstance(dsa_topk_ids, list) and len(dsa_topk_ids) == n + decoder_topk_ids = dsa_topk_ids + + projected_list = [ + self._preprocess(hidden_states=hidden, future_embeddings=future) + for hidden, future in zip(hidden_states, future_embeddings) + ] + micro_batch_results = cast( + GLM52MoEDecoderLayerMicroBatchOutput, + self.decoder_layer( + projected_list, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + dsa_topk_ids=decoder_topk_ids, + ), + ) + return { + "hidden_states": [self.final_layernorm(hidden) for hidden in micro_batch_results["hidden_states"]], + "router_logits": micro_batch_results["router_logits"], + "router_weights": micro_batch_results["router_weights"], + "router_topk_ids": micro_batch_results["router_topk_ids"], + "dsa_topk_ids": micro_batch_results["dsa_topk_ids"], + } + + +class GLM52MTPBlock(MTPBlock): + """MTP block that keeps DSA sharing private to GLM-5.2.""" + + @override + def _call_decoder_layer( + self, + layer: MTPLayer, + hidden_states: torch.Tensor | list[torch.Tensor], + *, + previous_layer_results: MTPInternalOutput | None, + future_embeddings: torch.Tensor | list[torch.Tensor], + position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], + seq_ctx: SequenceContext | list[SequenceContext], + ) -> MTPInternalOutput: + glm_layer = cast(GLM52MTPLayer, layer) + previous_results = cast( + GLM52MoEDecoderLayerOutput | GLM52MoEDecoderLayerMicroBatchOutput | None, + previous_layer_results, + ) + dsa_topk_ids = None if previous_results is None else previous_results["dsa_topk_ids"] + + return cast( + MoEDecoderLayerOutput | MoEDecoderLayerMicroBatchOutput, + glm_layer( + hidden_states, + future_embeddings=future_embeddings, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + dsa_topk_ids=dsa_topk_ids, + ), + ) diff --git a/xtuner/v1/model/moe/moe.py b/xtuner/v1/model/moe/moe.py index 2e18dc771..e91c74459 100644 --- a/xtuner/v1/model/moe/moe.py +++ b/xtuner/v1/model/moe/moe.py @@ -1,4 +1,5 @@ # Copyright (c) OpenMMLab. All rights reserved. +import contextlib import os import types from pathlib import Path @@ -22,7 +23,7 @@ from typing_extensions import overload, override from xtuner.v1.config import FSDPConfig -from xtuner.v1.data_proto import DSATopKCacheState, SequenceContext +from xtuner.v1.data_proto import SequenceContext from xtuner.v1.float8.float8_handler import Float8Handler from xtuner.v1.loss import ( AuxLossConfig, @@ -225,6 +226,10 @@ class MoE(BaseModel): config: MoEConfig ep_mesh: DeviceMesh | None = None + dense_decoder_layer_cls = DenseDecoderLayer + moe_decoder_layer_cls = MoEDecoderLayer + mtp_layer_cls = MTPLayer + mtp_block_cls = MTPBlock def __init__(self, config: MoEConfig): super().__init__(config) @@ -245,7 +250,7 @@ def __init__(self, config: MoEConfig): self.rotary_emb = self.build_rotary_embedding(config) self.embed_tokens = self.build_embeddings(config) self.mtp_block = self.build_mtp_block(config) if config.mtp_config is not None else None - self._configure_model_specific_layer_lifecycle() + self._configure_model_specific_layers() self.fp32_layers = [self.rotary_emb] @@ -275,9 +280,33 @@ def _maybe_offload_router(self, tensor: torch.Tensor) -> torch.Tensor: return async_offload_to_cpu(tensor, self.offload_stream) return tensor - def _configure_model_specific_layer_lifecycle(self) -> None: + def _configure_model_specific_layers(self) -> None: return + def _saved_tensors_offload_ctx( + self, + block_idx: int, + tensors: list[torch.Tensor], + ) -> contextlib.AbstractContextManager: + """Build one policy-neutral saved-tensor offload window. + + The decoder-stack caller decides which tensors belong to the current + window and advances ``block_idx`` only when the list is non-empty. + """ + if not tensors: + return contextlib.nullcontext() + + storage_ptrs = {tensor.untyped_storage().data_ptr() for tensor in tensors} + return async_save_on_cpu( + h2d_stream=self.offload_stream, + d2h_stream=self.offload_stream, + block_idx=block_idx, + group="text", + custom_check_fn=lambda tensor: tensor.untyped_storage().data_ptr() in storage_ptrs, + prefetch=True, + reserve_pin_memory=True, + ) + def _z_loss_dist_token_count( self, z_ctx: list[ZLossContext] | ZLossContext | None, @@ -577,79 +606,19 @@ def _micro_batch_forward( for seq_ctx in seq_ctx_list: self._mark_dynamic(seq_ctx) - for idx, decoder_layer in self.layers.items(): - layer_idx = int(idx) - - if layer_idx < self.config.first_k_dense_replace: - # Keep each micro-batch in its own SequenceContext while issuing - # one outer layer call, so FSDP materializes dense weights once. - dense_results = cast( - DenseDecoderLayerMicroBatchOutput, - decoder_layer( - hidden_states_list, - position_embeddings=position_embeddings_list, - seq_ctx=seq_ctx_list, - ), - ) - hidden_states_list = dense_results["hidden_states"] - else: - if int(os.getenv("XTUNER_ACTIVATION_OFFLOAD", "0")) == 1: - with async_save_on_cpu( - h2d_stream=self.offload_stream, - d2h_stream=self.offload_stream, - block_idx=layer_idx - self.config.first_k_dense_replace, - group="text", - custom_check_fn=lambda x: x.data_ptr() - in [hidden_states.data_ptr() for hidden_states in hidden_states_list], - prefetch=True, - reserve_pin_memory=True, - ): - layer_results = cast( - MoEDecoderLayerMicroBatchOutput, - decoder_layer( - hidden_states_list, - position_embeddings=position_embeddings_list, - seq_ctx=seq_ctx_list, - ), - ) - else: - layer_results = cast( - MoEDecoderLayerMicroBatchOutput, - decoder_layer( - hidden_states_list, - position_embeddings=position_embeddings_list, - seq_ctx=seq_ctx_list, - ), - ) - router_logits = layer_results["router_logits"] - router_weights = layer_results["router_weights"] - router_topk_ids = layer_results["router_topk_ids"] - - # Update hidden states and (optionally) collect router logits. - # router_weights are only consumed by aux_loss.accumulate below, so we - # never stash them per-MB the way we do for logits. - for i, hidden_states in enumerate(layer_results["hidden_states"]): - hidden_states_list[i] = hidden_states - if keep_router: - router_logits_list[i][f"layer{idx}"] = self._maybe_offload_router(router_logits[i]) - - cat_router_weights = torch.cat(router_weights, dim=0) - cat_router_logits = torch.cat(router_logits, dim=0) - cat_router_topk_ids = torch.cat(router_topk_ids, dim=0) - # Pin the per-layer z-loss to MB0's hidden_states stream. With multiple MBs, only - # one carrier may be chosen — all MBs converge into the same total_loss backward, - # so MB0's path traverses every aux-loss node exactly once. - hidden_states_list[0] = self.aux_loss.accumulate( - selected_router_weights=cat_router_weights.index_select(0, nonpad_indices).contiguous().float(), - selected_router_logits=cat_router_logits.index_select(0, nonpad_indices).contiguous().float(), - selected_experts=cat_router_topk_ids.index_select(0, nonpad_indices).contiguous(), - hidden_states=hidden_states_list[0], - balancing_ctx=balancing_ctx, - z_ctx=z_ctx, - num_tokens_local=non_pad_token, - num_tokens_global=num_tokens_global, - world_size=z_world_size, - ) + hidden_states_list = self._micro_batch_decoder_stack( + hidden_states_list=hidden_states_list, + position_embeddings_list=position_embeddings_list, + seq_ctx_list=seq_ctx_list, + router_logits_list=router_logits_list, + keep_router=keep_router, + balancing_ctx=balancing_ctx, + z_ctx=z_ctx, + nonpad_indices=nonpad_indices, + non_pad_token=non_pad_token, + num_tokens_global=num_tokens_global, + z_world_size=z_world_size, + ) assert hidden_states_list, "XTuner Internal Error, found empty hidden states for domino EP" @@ -668,7 +637,6 @@ def _micro_batch_forward( input_ids=seq_ctx.input_ids.clone() if seq_ctx.input_ids is not None else None, position_ids=seq_ctx.position_ids.clone(), inputs_embeds=seq_ctx.inputs_embeds.clone() if seq_ctx.inputs_embeds is not None else None, - dsa_topk_cache=DSATopKCacheState(), ) ) @@ -783,6 +751,75 @@ def _micro_batch_forward( return MoEModelOutputs(**output, logits=logits) + def _micro_batch_decoder_stack( + self, + *, + hidden_states_list: list[torch.Tensor], + position_embeddings_list: list[tuple[torch.Tensor, torch.Tensor]], + seq_ctx_list: list[SequenceContext], + router_logits_list: list[dict[str, torch.Tensor]], + keep_router: bool, + balancing_ctx: list[BalancingLossContext] | BalancingLossContext | None, + z_ctx: list[ZLossContext] | ZLossContext | None, + nonpad_indices: torch.Tensor, + non_pad_token: int, + num_tokens_global: torch.Tensor | None, + z_world_size: int, + ) -> list[torch.Tensor]: + """Run the main decoder stack for intra-layer micro-batches.""" + previous_layer_results: DenseDecoderLayerMicroBatchOutput | MoEDecoderLayerMicroBatchOutput | None = None + + for idx, decoder_layer in self.layers.items(): + layer_idx = int(idx) + layer_results = self._call_decoder_layer( + decoder_layer=decoder_layer, + layer_idx=layer_idx, + hidden_states=hidden_states_list, + position_embeddings=position_embeddings_list, + seq_ctx=seq_ctx_list, + previous_layer_results=previous_layer_results, + ) + previous_layer_results = cast( + DenseDecoderLayerMicroBatchOutput | MoEDecoderLayerMicroBatchOutput, + layer_results, + ) + if layer_idx < self.config.first_k_dense_replace: + # Keep each micro-batch in its own SequenceContext while issuing + # one outer layer call, so FSDP materializes dense weights once. + dense_results = cast(DenseDecoderLayerMicroBatchOutput, layer_results) + hidden_states_list = dense_results["hidden_states"] + continue + + layer_results = cast(MoEDecoderLayerMicroBatchOutput, layer_results) + hidden_states = layer_results["hidden_states"] + router_logits = layer_results["router_logits"] + router_weights = layer_results["router_weights"] + router_topk_ids = layer_results["router_topk_ids"] + + # Router weights are consumed immediately by aux loss; only logits + # requested by the caller are retained per micro-batch. + for i, hidden_state in enumerate(hidden_states): + hidden_states_list[i] = hidden_state + if keep_router: + router_logits_list[i][f"layer{idx}"] = self._maybe_offload_router(router_logits[i]) + + cat_router_weights = torch.cat(router_weights, dim=0) + cat_router_logits = torch.cat(router_logits, dim=0) + cat_router_topk_ids = torch.cat(router_topk_ids, dim=0) + hidden_states_list[0] = self.aux_loss.accumulate( + selected_router_weights=cat_router_weights.index_select(0, nonpad_indices).contiguous().float(), + selected_router_logits=cat_router_logits.index_select(0, nonpad_indices).contiguous().float(), + selected_experts=cat_router_topk_ids.index_select(0, nonpad_indices).contiguous(), + hidden_states=hidden_states_list[0], + balancing_ctx=balancing_ctx, + z_ctx=z_ctx, + num_tokens_local=non_pad_token, + num_tokens_global=num_tokens_global, + world_size=z_world_size, + ) + + return hidden_states_list + def _forward( self, seq_ctx: SequenceContext, # todo(@yehaochen): support intra layer micro-batch @@ -824,70 +861,26 @@ def _forward( output["router_weights"] = None self._mark_dynamic(seq_ctx) balancing_ctx, z_ctx = self._extract_aux_loss_ctx(loss_ctx) + balancing_ctx = cast(BalancingLossContext | None, balancing_ctx) + z_ctx = cast(ZLossContext | None, z_ctx) # Hoisted out of the per-layer accumulate path: mask is constant across layers. nonpad_indices = torch.nonzero(seq_ctx.mask, as_tuple=True)[1] non_pad_token = nonpad_indices.numel() num_tokens_global, z_world_size = self._z_loss_dist_token_count(z_ctx, non_pad_token, seq_ctx.mask.device) - for idx, decoder_layer in self.layers.items(): - if int(idx) < self.config.first_k_dense_replace: - dense_results = cast( - DenseDecoderLayerOutput, - decoder_layer( - hidden_states, - position_embeddings=position_embeddings, - seq_ctx=seq_ctx, - ), - ) - hidden_states = dense_results["hidden_states"] - else: - if int(os.getenv("XTUNER_ACTIVATION_OFFLOAD", "0")) == 1: - with async_save_on_cpu( - h2d_stream=self.offload_stream, - d2h_stream=self.offload_stream, - block_idx=int(idx), - group="text", - custom_check_fn=lambda x: x.data_ptr() == hidden_states.data_ptr(), - ): - layer_results = cast( - MoEDecoderLayerOutput, - decoder_layer( - hidden_states, - position_embeddings=position_embeddings, - seq_ctx=seq_ctx, - ), - ) - - else: - layer_results = cast( - MoEDecoderLayerOutput, - decoder_layer( - hidden_states, - position_embeddings=position_embeddings, - seq_ctx=seq_ctx, - ), - ) - hidden_states = layer_results["hidden_states"] - router_logits = layer_results["router_logits"] - router_weights = layer_results["router_weights"] - router_topk_ids = layer_results["router_topk_ids"] - if keep_router: - output["router_logits"][f"layer{idx}"] = self._maybe_offload_router(router_logits) - output["router_weights"][f"layer{idx}"] = self._maybe_offload_router(router_weights) - hidden_states = self.aux_loss.accumulate( - selected_router_weights=router_weights.index_select(0, nonpad_indices).contiguous().float(), - selected_router_logits=router_logits.index_select(0, nonpad_indices).contiguous().float(), - selected_experts=router_topk_ids.index_select(0, nonpad_indices).contiguous(), - hidden_states=hidden_states, - balancing_ctx=balancing_ctx, - z_ctx=z_ctx, - num_tokens_local=non_pad_token, - num_tokens_global=num_tokens_global, - world_size=z_world_size, - ) - - if self.config.return_hidden_states: - output["hidden_states"].append(hidden_states) + hidden_states = self._decoder_stack( + hidden_states=hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + output=output, + keep_router=keep_router, + balancing_ctx=balancing_ctx, + z_ctx=z_ctx, + nonpad_indices=nonpad_indices, + non_pad_token=non_pad_token, + num_tokens_global=num_tokens_global, + z_world_size=z_world_size, + ) layer_hidden_states = hidden_states hidden_states = self.norm(hidden_states) @@ -909,7 +902,6 @@ def _forward( input_ids=input_ids.clone() if input_ids is not None else None, position_ids=position_ids.clone(), inputs_embeds=seq_ctx.inputs_embeds.clone() if seq_ctx.inputs_embeds is not None else None, - dsa_topk_cache=DSATopKCacheState(), ) # MTP uses its own mask; main mask's non-pad indices do not apply. mtp_nonpad_indices = torch.nonzero(mtp_seq_ctx.mask, as_tuple=True)[1] @@ -981,6 +973,113 @@ def _forward( return MoEModelOutputs(**output) + def _decoder_stack( + self, + *, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + seq_ctx: SequenceContext, + output: dict, + keep_router: bool, + balancing_ctx: BalancingLossContext | None, + z_ctx: ZLossContext | None, + nonpad_indices: torch.Tensor, + non_pad_token: int, + num_tokens_global: torch.Tensor | None, + z_world_size: int, + ) -> torch.Tensor: + """Run the main decoder stack for one sequence context.""" + previous_layer_results: DenseDecoderLayerOutput | MoEDecoderLayerOutput | None = None + + for idx, decoder_layer in self.layers.items(): + layer_idx = int(idx) + layer_results = self._call_decoder_layer( + decoder_layer=decoder_layer, + layer_idx=layer_idx, + hidden_states=hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + previous_layer_results=previous_layer_results, + ) + previous_layer_results = cast(DenseDecoderLayerOutput | MoEDecoderLayerOutput, layer_results) + if layer_idx < self.config.first_k_dense_replace: + dense_results = cast(DenseDecoderLayerOutput, layer_results) + hidden_states = dense_results["hidden_states"] + else: + layer_results = cast(MoEDecoderLayerOutput, layer_results) + hidden_states = layer_results["hidden_states"] + router_results = layer_results["router_logits"] + router_weights = layer_results["router_weights"] + router_topk_ids = layer_results["router_topk_ids"] + if keep_router: + output["router_logits"][f"layer{idx}"] = self._maybe_offload_router(router_results) + output["router_weights"][f"layer{idx}"] = self._maybe_offload_router(router_weights) + hidden_states = self.aux_loss.accumulate( + selected_router_weights=router_weights.index_select(0, nonpad_indices).contiguous().float(), + selected_router_logits=router_results.index_select(0, nonpad_indices).contiguous().float(), + selected_experts=router_topk_ids.index_select(0, nonpad_indices).contiguous(), + hidden_states=hidden_states, + balancing_ctx=balancing_ctx, + z_ctx=z_ctx, + num_tokens_local=non_pad_token, + num_tokens_global=num_tokens_global, + world_size=z_world_size, + ) + + if self.config.return_hidden_states: + output["hidden_states"].append(hidden_states) + + return hidden_states + + def _call_decoder_layer( + self, + *, + decoder_layer: nn.Module, + layer_idx: int, + hidden_states: torch.Tensor | list[torch.Tensor], + position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], + seq_ctx: SequenceContext | list[SequenceContext], + previous_layer_results: ( + DenseDecoderLayerOutput + | DenseDecoderLayerMicroBatchOutput + | MoEDecoderLayerOutput + | MoEDecoderLayerMicroBatchOutput + | None + ), + ) -> ( + DenseDecoderLayerOutput + | DenseDecoderLayerMicroBatchOutput + | MoEDecoderLayerOutput + | MoEDecoderLayerMicroBatchOutput + ): + """Call one decoder layer and contain its activation-offload window. + + ``previous_layer_results`` is intentionally unused by the generic model; subclasses may + consume private cross-layer fields without widening the common decoder API. + """ + if layer_idx < self.config.first_k_dense_replace: + return decoder_layer( + hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + ) + + offload_tensors = list(hidden_states) if isinstance(hidden_states, list) else [hidden_states] + if int(os.getenv("XTUNER_ACTIVATION_OFFLOAD", "0")) != 1: + offload_tensors = [] + offload_block_idx = sum( + self.config.first_k_dense_replace <= int(previous_idx) < layer_idx for previous_idx in self.layers + ) + with self._saved_tensors_offload_ctx( + offload_block_idx, + offload_tensors, + ): + return decoder_layer( + hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + ) + def build_embeddings(self, config: MoEConfig): return nn.Embedding(config.vocab_size, config.hidden_size, config.pad_token_id) @@ -1003,7 +1102,7 @@ def build_layers(self, config: MoEConfig) -> nn.ModuleDict: ) if layer_idx < config.first_k_dense_replace: - layers[str(layer_idx)] = DenseDecoderLayer( + layers[str(layer_idx)] = self.dense_decoder_layer_cls( hidden_size=config.hidden_size, intermediate_size=config.intermediate_size, mlp_bias=config.mlp_bias, @@ -1018,7 +1117,7 @@ def build_layers(self, config: MoEConfig) -> nn.ModuleDict: layer_idx=layer_idx, ) else: - layers[str(layer_idx)] = MoEDecoderLayer( + layers[str(layer_idx)] = self.moe_decoder_layer_cls( hidden_size=config.hidden_size, intermediate_size=config.intermediate_size, moe_intermediate_size=config.moe_intermediate_size, @@ -1083,7 +1182,7 @@ def build_mtp_block(self, config: MoEConfig) -> MTPBlock: num_physical_layer = 1 if mtp_config.share_weights else mtp_config.num_layers for i in range(num_physical_layer): # Build MoE decoder layer for MTP - decoder_layer = MoEDecoderLayer( + decoder_layer = self.moe_decoder_layer_cls( hidden_size=config.hidden_size, intermediate_size=config.intermediate_size, moe_intermediate_size=config.moe_intermediate_size, @@ -1112,7 +1211,7 @@ def build_mtp_block(self, config: MoEConfig) -> MTPBlock: ) # Wrap decoder layer in MTPLayer - mtp_layer = MTPLayer( + mtp_layer = self.mtp_layer_cls( hidden_size=config.hidden_size, rms_norm_eps=config.rms_norm_eps, rms_norm_type=config.rms_norm_type, @@ -1121,7 +1220,7 @@ def build_mtp_block(self, config: MoEConfig) -> MTPBlock: ) mtp_layers.append(mtp_layer) - return MTPBlock( + return self.mtp_block_cls( mtp_config=mtp_config, mtp_layers=mtp_layers, ) diff --git a/xtuner/v1/module/__init__.py b/xtuner/v1/module/__init__.py index 8dbecfc75..2f1d59e67 100644 --- a/xtuner/v1/module/__init__.py +++ b/xtuner/v1/module/__init__.py @@ -1,7 +1,5 @@ from .attention import ( AttnOutputs, - DSAMLAConfig, - DSAMultiLatentAttention, GatedDeltaNet, GatedDeltaNetConfig, MHAConfig, @@ -28,10 +26,8 @@ "RMSNorm", "MultiHeadAttention", "MultiLatentAttention", - "DSAMultiLatentAttention", "MHAConfig", "MLAConfig", - "DSAMLAConfig", "GatedDeltaNetConfig", "GatedDeltaNet", "AttnOutputs", diff --git a/xtuner/v1/module/attention/__init__.py b/xtuner/v1/module/attention/__init__.py index dedd2b145..d4594014a 100644 --- a/xtuner/v1/module/attention/__init__.py +++ b/xtuner/v1/module/attention/__init__.py @@ -1,6 +1,5 @@ # Copyright (c) OpenMMLab. All rights reserved. from .attn_outputs import AttnOutputs -from .dsa_mla import DSAMLAConfig, DSAMultiLatentAttention from .gated_deltanet import GatedDeltaNet, GatedDeltaNetConfig from .mha import MHAConfig, MultiHeadAttention from .mla import MLAConfig, MultiLatentAttention @@ -8,11 +7,9 @@ __all__ = [ "MultiLatentAttention", - "DSAMultiLatentAttention", "MultiHeadAttention", "MHAConfig", "MLAConfig", - "DSAMLAConfig", "AttnOutputs", "GatedDeltaNet", "GatedDeltaNetConfig", diff --git a/xtuner/v1/module/attention/dsa_topk_sharing.py b/xtuner/v1/module/attention/dsa_topk_sharing.py deleted file mode 100644 index 628675ae0..000000000 --- a/xtuner/v1/module/attention/dsa_topk_sharing.py +++ /dev/null @@ -1,518 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -import os -from dataclasses import dataclass -from functools import partial -from typing import Any, Callable, Protocol, cast - -import torch - -from xtuner.v1.data_proto import SequenceContext -from xtuner.v1.utils.activation_offload import OffloadManager, SwapTensor - - -class DSATopKSharingLayerProtocol(Protocol): - layer_idx: int - source_layer_idx: int - training: bool - indexer_types: list[str] | None - index_skip_topk_offset: int - index_topk_freq: int - dsa_topk_last_use: dict[int, int] - dsa_topk_recompute_release: dict[int, int] - - -@dataclass(frozen=True) -class DSATopKReleasePlan: - forward_last_use: dict[int, int] - recompute_release: dict[int, int] - - -def dsa_topk_source_layer( - *, - layer_idx: int, - indexer_types: list[str] | None, - index_skip_topk_offset: int, - index_topk_freq: int, -) -> int: - """Resolve the physical indexer source for one logical DSA layer.""" - if indexer_types is not None: - if layer_idx < len(indexer_types) and indexer_types[layer_idx] == "full": - return layer_idx - for source_layer_idx in range(min(layer_idx, len(indexer_types) - 1), -1, -1): - if indexer_types[source_layer_idx] == "full": - return source_layer_idx - raise ValueError(f"DSA layer {layer_idx} has no preceding full indexer layer.") - - if index_topk_freq <= 1: - return layer_idx - - source_layer_idx = layer_idx - while (max(source_layer_idx + 1 - index_skip_topk_offset, 0) % index_topk_freq) != 0: - source_layer_idx -= 1 - return source_layer_idx - - -def _dsa_topk_offload_enabled() -> bool: - override = os.getenv("XTUNER_DSA_TOPK_OFFLOAD") - if override is not None: - return int(override) == 1 - # DSA top-k cache is consumed by SparseMLA backward. Keep this offload path - # opt-in instead of coupling it to hidden-state activation offload. - return False - - -def build_dsa_topk_release_plan( - *, - num_main_layers: int, - num_mtp_layers: int, - indexer_types: list[str] | None, - index_skip_topk_offset: int, - index_topk_freq: int, -) -> DSATopKReleasePlan: - consumers: dict[int, list[int]] = {} - for layer_idx in range(num_main_layers + num_mtp_layers): - source_layer_idx = dsa_topk_source_layer( - layer_idx=layer_idx, - indexer_types=indexer_types, - index_skip_topk_offset=index_skip_topk_offset, - index_topk_freq=index_topk_freq, - ) - consumers.setdefault(source_layer_idx, []).append(layer_idx) - - return DSATopKReleasePlan( - forward_last_use={ - source_layer_idx: max(consumer_layers) for source_layer_idx, consumer_layers in consumers.items() - }, - recompute_release={ - source_layer_idx: min(consumer_layers) for source_layer_idx, consumer_layers in consumers.items() - }, - ) - - -class GpuTopKResidency: - def has_cache(self, seq_ctx: SequenceContext, source_layer_idx: int) -> bool: - return source_layer_idx in seq_ctx.dsa_topk_cache.indices - - def store_gpu(self, seq_ctx: SequenceContext, source_layer_idx: int, topk_indices: torch.Tensor) -> None: - seq_ctx.dsa_topk_cache.indices[source_layer_idx] = topk_indices - - def read(self, seq_ctx: SequenceContext, source_layer_idx: int) -> torch.Tensor: - return seq_ctx.dsa_topk_cache.indices[source_layer_idx] - - def after_original_forward_last_use(self, seq_ctx: SequenceContext, source_layer_idx: int) -> None: - return - - def after_recompute_release(self, seq_ctx: SequenceContext, source_layer_idx: int) -> None: - seq_ctx.dsa_topk_cache.indices.pop(source_layer_idx, None) - - def _offload_key(self, seq_ctx: SequenceContext, source_layer_idx: int) -> str: - return f"dsa_topk_{seq_ctx.dsa_topk_cache.context_id}_{source_layer_idx}" - - -class ActivationOffloadedTopKResidency(GpuTopKResidency): - def __init__(self) -> None: - self._streams: dict[int, torch.cuda.Stream] = {} - self._prefetched: dict[tuple[int, int], SwapTensor] = {} - - def has_cache(self, seq_ctx: SequenceContext, source_layer_idx: int) -> bool: - cache = seq_ctx.dsa_topk_cache - return source_layer_idx in cache.indices or source_layer_idx in cache.offloaded - - def read(self, seq_ctx: SequenceContext, source_layer_idx: int) -> torch.Tensor: - cache = seq_ctx.dsa_topk_cache - if source_layer_idx in cache.indices: - self._wait_prefetched(seq_ctx, source_layer_idx) - return cache.indices[source_layer_idx] - return self._read_offloaded(seq_ctx, source_layer_idx) - - # Pinned CPU buffers and stream-side effects must stay outside Inductor graphs. - @torch.compiler.disable - def prefetch(self, seq_ctx: SequenceContext, source_layer_idx: int) -> None: - cache = seq_ctx.dsa_topk_cache - if source_layer_idx in cache.indices or source_layer_idx not in cache.offloaded: - return - - key = cache.offloaded[source_layer_idx] - swap_tensor = OffloadManager().get(key) - stream = self._stream_for_device(swap_tensor.tensor.device) - # Decoder pre-hook runs before the compiled layer body. Launch H2D here - # and wait only when SparseMLA actually consumes top-k in read(). - swap_tensor.prefetch_launch_h2d(stream, True) - cache.indices[source_layer_idx] = swap_tensor.tensor - self._prefetched[self._prefetch_key(seq_ctx, source_layer_idx)] = swap_tensor - - # Pinned CPU buffers and stream-side effects must stay outside Inductor graphs. - @torch.compiler.disable - def _read_offloaded(self, seq_ctx: SequenceContext, source_layer_idx: int) -> torch.Tensor: - cache = seq_ctx.dsa_topk_cache - key = cache.offloaded[source_layer_idx] - swap_tensor = OffloadManager().get(key) - stream = self._stream_for_device(swap_tensor.tensor.device) - working_stream = torch.cuda.current_stream(swap_tensor.tensor.device) - - # DSA top-k cache is not captured by saved_tensors_hooks, so this mirrors - # activation offload's explicit H2D choreography for manual cache state. - stream.wait_stream(working_stream) - with torch.cuda.stream(stream): - swap_tensor.launch_h2d(stream, True, stream) - working_stream.wait_stream(stream) - - cache.indices[source_layer_idx] = swap_tensor.tensor - return swap_tensor.tensor - - # Pinned CPU buffers and stream-side effects must stay outside Inductor graphs. - @torch.compiler.disable - def _wait_prefetched(self, seq_ctx: SequenceContext, source_layer_idx: int) -> None: - swap_tensor = self._prefetched.pop(self._prefetch_key(seq_ctx, source_layer_idx), None) - if swap_tensor is None: - return - swap_tensor.wait_h2d_finished() - - # Pinned CPU buffers and stream-side effects must stay outside Inductor graphs. - @torch.compiler.disable - def after_original_forward_last_use(self, seq_ctx: SequenceContext, source_layer_idx: int) -> None: - cache = seq_ctx.dsa_topk_cache - topk_indices = cache.indices.pop(source_layer_idx) - if not topk_indices.is_cuda: - cache.indices[source_layer_idx] = topk_indices - return - - key = self._offload_key(seq_ctx, source_layer_idx) - cpu_buffer = OffloadManager().get_or_create_pin_memory(key, topk_indices.shape, topk_indices.dtype) - swap_tensor = SwapTensor(topk_indices, key, tensor_cpu=cpu_buffer) - stream = self._stream_for_device(topk_indices.device) - stream.wait_stream(torch.cuda.current_stream(topk_indices.device)) - swap_tensor.launch_d2h(stream) - swap_tensor.wait_d2h_finished(stream, True) - OffloadManager().put(key, swap_tensor) - cache.offloaded[source_layer_idx] = key - - # Pinned CPU buffers and stream-side effects must stay outside Inductor graphs. - @torch.compiler.disable - def after_recompute_release(self, seq_ctx: SequenceContext, source_layer_idx: int) -> None: - cache = seq_ctx.dsa_topk_cache - self._wait_prefetched(seq_ctx, source_layer_idx) - super().after_recompute_release(seq_ctx, source_layer_idx) - key = cache.offloaded.pop(source_layer_idx, None) - if key is None: - return - - stream = self._stream_for_current_device() - OffloadManager().del_may_npu_tensor(key, stream) - if OffloadManager().exist(key): - OffloadManager().clear(key) - - def _stream_for_current_device(self) -> torch.cuda.Stream: - return self._stream_for_device(torch.device("cuda", torch.cuda.current_device())) - - def _stream_for_device(self, device: torch.device) -> torch.cuda.Stream: - device_idx = torch.cuda.current_device() if device.index is None else device.index - if device_idx not in self._streams: - self._streams[device_idx] = torch.cuda.Stream(device=device_idx) - return self._streams[device_idx] - - def _prefetch_key(self, seq_ctx: SequenceContext, source_layer_idx: int) -> tuple[int, int]: - return id(seq_ctx.dsa_topk_cache), source_layer_idx - - -class CrossLayerTopKSharingRuntime: - def __init__(self) -> None: - self._gpu_residency = GpuTopKResidency() - self._offloaded_residency = ActivationOffloadedTopKResidency() - - def get_or_compute( - self, - *, - layer: DSATopKSharingLayerProtocol, - seq_ctx: SequenceContext, - compute_source_topk: Callable[[], torch.Tensor], - ) -> torch.Tensor: - residency = self._residency() - cache = seq_ctx.dsa_topk_cache - source_layer_idx = layer.source_layer_idx - - if source_layer_idx != layer.layer_idx: - self._assert_source_present(layer, seq_ctx, residency) - return residency.read(seq_ctx, source_layer_idx) - - if ( - self._is_checkpoint_recompute(seq_ctx) - and layer.layer_idx not in cache.released_sources - and residency.has_cache(seq_ctx, source_layer_idx) - ): - # Top-k indices are discrete and need no autograd graph. Reentrant - # replay can reuse the original forward cache without rerunning the indexer. - return residency.read(seq_ctx, source_layer_idx) - - if self._can_reuse_mtp_iteration_topk(seq_ctx, source_layer_idx, residency): - return residency.read(seq_ctx, source_layer_idx) - - topk_indices = compute_source_topk() - if layer.layer_idx not in cache.released_sources: - residency.store_gpu(seq_ctx, layer.layer_idx, topk_indices) - return topk_indices - - def after_sparse_mla_use(self, *, layer: DSATopKSharingLayerProtocol, seq_ctx: SequenceContext) -> None: - residency = self._residency() - cache = seq_ctx.dsa_topk_cache - source_layer_idx = layer.source_layer_idx - if self._is_checkpoint_original_forward(layer): - if layer.dsa_topk_last_use.get(source_layer_idx) == layer.layer_idx: - if not self._is_last_mtp_forward_use(seq_ctx, source_layer_idx): - return - # Reentrant checkpoint original forward runs under no_grad, so - # SparseMLA has no autograd ctx. Keep/offload source top-k for - # backward recompute, then release after source replay consumes it. - cache.checkpoint_active = True - residency.after_original_forward_last_use(seq_ctx, source_layer_idx) - return - - if not self._is_checkpoint_recompute(seq_ctx): - return - - release_layer_idx = layer.dsa_topk_recompute_release.get(source_layer_idx) - if release_layer_idx != layer.layer_idx: - return - - if not self._should_release_after_mtp_iteration_recompute(seq_ctx, source_layer_idx): - return - - residency.after_recompute_release(seq_ctx, source_layer_idx) - cache.released_sources.add(source_layer_idx) - - def register_mtp_iteration_topk_sharing( - self, - *, - seq_ctx: SequenceContext, - source_layer_idx: int, - num_iterations: int, - ) -> None: - if num_iterations <= 1: - return - - cache = seq_ctx.dsa_topk_cache - cache.mtp_forward_uses_remaining[source_layer_idx] = num_iterations - cache.mtp_replays_remaining[source_layer_idx] = num_iterations - - def before_layer_forward(self, *, layer: DSATopKSharingLayerProtocol, seq_ctx: SequenceContext) -> None: - if not isinstance(self._residency(), ActivationOffloadedTopKResidency): - return - source_layer_idx = layer.source_layer_idx - if source_layer_idx not in seq_ctx.dsa_topk_cache.offloaded: - return - self._offloaded_residency.prefetch(seq_ctx, source_layer_idx) - - def _residency(self) -> GpuTopKResidency: - if _dsa_topk_offload_enabled() and torch.cuda.is_available(): - return self._offloaded_residency - return self._gpu_residency - - def _is_checkpoint_original_forward(self, layer: DSATopKSharingLayerProtocol) -> bool: - # 这里通过 grad 是否开启来判断当前阶段: - # reentrant: original=False,replay=True,可以区分; - # non-reentrant: original=True, replay=True,无法区分。 - # 例如 MTP depth2 需要在两次 original 和两次 replay 中分别更新 cache 计数; - # non-reentrant 识别不到 original,计数没有正确更新,depth1 replay 就会 - # 沿用仅适合 reentrant 的 cache-reuse 路径。reentrant original 不建内部图, - # replay 复用离散 top-k 是安全的;non-reentrant 则必须重建相同保存清单。 - # compile 只会把 COMPUTE/REUSE 的分支差异暴露为 saved-tensor metadata - # mismatch;即使关闭 compile 不报错,这里的 cache 状态仍然是错误的。 - return layer.training and not torch.is_grad_enabled() - - def _is_checkpoint_recompute(self, seq_ctx: SequenceContext) -> bool: - return seq_ctx.dsa_topk_cache.checkpoint_active and torch.is_grad_enabled() - - def _can_reuse_mtp_iteration_topk( - self, - seq_ctx: SequenceContext, - source_layer_idx: int, - residency: GpuTopKResidency, - ) -> bool: - return source_layer_idx in seq_ctx.dsa_topk_cache.mtp_replays_remaining and residency.has_cache( - seq_ctx, source_layer_idx - ) - - def _is_last_mtp_forward_use(self, seq_ctx: SequenceContext, source_layer_idx: int) -> bool: - cache = seq_ctx.dsa_topk_cache - remaining = cache.mtp_forward_uses_remaining.get(source_layer_idx) - if remaining is None: - return True - - remaining -= 1 - if remaining == 0: - cache.mtp_forward_uses_remaining.pop(source_layer_idx) - return True - - cache.mtp_forward_uses_remaining[source_layer_idx] = remaining - return False - - def _should_release_after_mtp_iteration_recompute( - self, - seq_ctx: SequenceContext, - source_layer_idx: int, - ) -> bool: - remaining = seq_ctx.dsa_topk_cache.mtp_replays_remaining.get(source_layer_idx) - if remaining is None: - return True - - remaining -= 1 - if remaining == 0: - seq_ctx.dsa_topk_cache.mtp_replays_remaining.pop(source_layer_idx) - return True - - seq_ctx.dsa_topk_cache.mtp_replays_remaining[source_layer_idx] = remaining - return False - - def _assert_source_present( - self, - layer: DSATopKSharingLayerProtocol, - seq_ctx: SequenceContext, - residency: GpuTopKResidency, - ) -> None: - if residency.has_cache(seq_ctx, layer.source_layer_idx): - return - raise AssertionError( - "DSA index-share: skip layer " - f"{layer.layer_idx} needs source layer {layer.source_layer_idx} top-k, " - "but it is not present in this microbatch SequenceContext. " - "Cross-pipeline top-k sharing is not supported." - ) - - -_DSA_TOPK_SHARING_RUNTIME = CrossLayerTopKSharingRuntime() - - -def get_dsa_topk_sharing_runtime() -> CrossLayerTopKSharingRuntime: - return _DSA_TOPK_SHARING_RUNTIME - - -def configure_dsa_topk_decoder_lifecycle( - *, - decoder_layer: torch.nn.Module, - attention: DSATopKSharingLayerProtocol, - release_plan: DSATopKReleasePlan, -) -> None: - # The release maps and decoder hooks are one lifecycle contract: source - # caches are kept/offloaded until the planned consumer layer runs. - attention.dsa_topk_last_use = release_plan.forward_last_use - attention.dsa_topk_recompute_release = release_plan.recompute_release - register_dsa_topk_decoder_lifecycle_hooks(decoder_layer) - - -def configure_dsa_mtp_iteration_lifecycle( - *, - mtp_block: torch.nn.Module, - attention: DSATopKSharingLayerProtocol, - num_iterations: int, -) -> None: - if num_iterations <= 1: - return - - # The outer MTP block runs once per model forward, while its checkpointed - # physical layer replays once per logical depth during backward. Register - # the shared cache ownership before either sequence starts. - mtp_block.register_forward_pre_hook( - partial( - _dsa_mtp_iteration_lifecycle_pre_hook, - source_layer_idx=attention.source_layer_idx, - num_iterations=num_iterations, - ), - with_kwargs=True, - ) - - -@torch.compiler.disable -def before_dsa_topk_decoder_forward(attention: object, seq_ctx: SequenceContext | list[SequenceContext]) -> None: - assert hasattr(attention, "dsa_topk_last_use"), "DSA top-k lifecycle requires a DSA attention module." - - runtime = get_dsa_topk_sharing_runtime() - for ctx in seq_ctx if isinstance(seq_ctx, list) else [seq_ctx]: - runtime.before_layer_forward(layer=cast(DSATopKSharingLayerProtocol, attention), seq_ctx=ctx) - - -@torch.compiler.disable -def after_dsa_topk_decoder_forward(attention: object, seq_ctx: SequenceContext | list[SequenceContext]) -> None: - assert hasattr(attention, "dsa_topk_last_use"), "DSA top-k lifecycle requires a DSA attention module." - - runtime = get_dsa_topk_sharing_runtime() - for ctx in seq_ctx if isinstance(seq_ctx, list) else [seq_ctx]: - runtime.after_sparse_mla_use(layer=cast(DSATopKSharingLayerProtocol, attention), seq_ctx=ctx) - - -def _get_seq_ctx_from_forward( - args: tuple[Any, ...], - kwargs: dict[str, Any], -) -> SequenceContext | list[SequenceContext]: - seq_ctx = kwargs.get("seq_ctx") - if seq_ctx is None and len(args) >= 3: - seq_ctx = args[2] - assert seq_ctx is not None, "DSA top-k lifecycle requires seq_ctx in decoder forward." - assert isinstance(seq_ctx, SequenceContext | list), ( - f"DSA top-k lifecycle expected SequenceContext or list, got {type(seq_ctx).__name__}." - ) - return seq_ctx - - -def _dsa_topk_decoder_lifecycle_pre_hook( - module: torch.nn.Module, - args: tuple[Any, ...], - kwargs: dict[str, Any], -) -> None: - seq_ctx = _get_seq_ctx_from_forward(args, kwargs) - before_dsa_topk_decoder_forward(module.self_attn, seq_ctx) # type: ignore[attr-defined] - - -def _dsa_topk_decoder_lifecycle_post_hook( - module: torch.nn.Module, - args: tuple[Any, ...], - kwargs: dict[str, Any], - _output: Any, -) -> None: - seq_ctx = _get_seq_ctx_from_forward(args, kwargs) - after_dsa_topk_decoder_forward(module.self_attn, seq_ctx) # type: ignore[attr-defined] - - -@torch.compiler.disable -def _dsa_mtp_iteration_lifecycle_pre_hook( - _module: torch.nn.Module, - args: tuple[Any, ...], - kwargs: dict[str, Any], - *, - source_layer_idx: int, - num_iterations: int, -) -> None: - seq_ctx = _get_seq_ctx_from_forward(args, kwargs) - - runtime = get_dsa_topk_sharing_runtime() - for ctx in seq_ctx if isinstance(seq_ctx, list) else [seq_ctx]: - runtime.register_mtp_iteration_topk_sharing( - seq_ctx=ctx, - source_layer_idx=source_layer_idx, - num_iterations=num_iterations, - ) - - -def register_dsa_topk_decoder_lifecycle_hooks(decoder_layer: torch.nn.Module) -> None: - if getattr(decoder_layer, "_dsa_topk_decoder_lifecycle_hooks_registered", False): - return - assert hasattr(decoder_layer, "self_attn"), "DSA top-k lifecycle requires decoder_layer.self_attn." - assert hasattr(decoder_layer.self_attn, "dsa_topk_last_use"), ( # type: ignore[attr-defined] - "DSA top-k lifecycle requires a DSA attention module." - ) - - # Pinned-memory, CUDA-stream and OffloadManager side effects cannot run in - # an Inductor graph. The previous in-attention implementation therefore - # recorded only pending actions and flushed them later. Remove that - # transient state by keeping the entire residency transition at the decoder - # boundary: the pre-hook launches H2D and the post-hook directly runs - # after_sparse_mla_use. Checkpoint recompute re-invokes the decoder module - # and these hooks along with it, so main, micro-batch and MTP callers do not - # need separate lifecycle handling. - # - # This deliberately delays eager D2H until the decoder returns, losing its - # overlap with attention projection and MoE compute; lifecycle is also no - # longer adjacent to SparseMLA's exact last use. Direct attention callers - # must therefore run through a decoder with these hooks registered. - decoder_layer.register_forward_pre_hook(_dsa_topk_decoder_lifecycle_pre_hook, with_kwargs=True) - decoder_layer.register_forward_hook(_dsa_topk_decoder_lifecycle_post_hook, with_kwargs=True) - object.__setattr__(decoder_layer, "_dsa_topk_decoder_lifecycle_hooks_registered", True) diff --git a/xtuner/v1/module/decoder_layer/dense_decoder_layer.py b/xtuner/v1/module/decoder_layer/dense_decoder_layer.py index 94d1a00aa..0fc874382 100644 --- a/xtuner/v1/module/decoder_layer/dense_decoder_layer.py +++ b/xtuner/v1/module/decoder_layer/dense_decoder_layer.py @@ -113,7 +113,6 @@ def forward( Rotary position embeddings ``(cos, sin)``, aligned with ``hidden_states``. seq_ctx (SequenceContext | list[SequenceContext]): Sequence context, aligned with ``hidden_states``. - Returns: DenseDecoderLayerOutput | DenseDecoderLayerMicroBatchOutput: Output hidden states. A single tensor for a single ``hidden_states`` tensor, a per-micro-batch list for a list diff --git a/xtuner/v1/module/decoder_layer/moe_decoder_layer.py b/xtuner/v1/module/decoder_layer/moe_decoder_layer.py index 25d1e2e29..927b1f9ae 100644 --- a/xtuner/v1/module/decoder_layer/moe_decoder_layer.py +++ b/xtuner/v1/module/decoder_layer/moe_decoder_layer.py @@ -1,5 +1,5 @@ from functools import partial -from typing import Literal, Protocol, TypeAlias, TypedDict, cast +from typing import Callable, Literal, Protocol, TypeAlias, TypedDict, cast import torch import torch.nn as nn @@ -349,19 +349,19 @@ def forward( seq_ctx=seq_ctx, position_embeddings=position_embeddings, ) - else: - assert isinstance(seq_ctx, list) and len(seq_ctx) == len(hidden_states), ( - "seq_ctx should be a list of SequenceContext instances with the same length as hidden_states" - ) - assert isinstance(position_embeddings, list) and len(position_embeddings) == len(hidden_states), ( - "position_embeddings should be a list of tuples with the same length as hidden_states" - ) - return self._micro_batch_forward( - hidden_states_list=hidden_states, - seq_ctx_list=seq_ctx, - position_embeddings_list=position_embeddings, - ) + assert isinstance(seq_ctx, list) and len(seq_ctx) == len(hidden_states), ( + "seq_ctx should be a list of SequenceContext instances with the same length as hidden_states" + ) + assert isinstance(position_embeddings, list) and len(position_embeddings) == len(hidden_states), ( + "position_embeddings should be a list of tuples with the same length as hidden_states" + ) + + return self._micro_batch_forward( + hidden_states_list=hidden_states, + seq_ctx_list=seq_ctx, + position_embeddings_list=position_embeddings, + ) def _hf_expert_forward_for_debug(self, hidden_states: torch.Tensor, router_results: RouterResults, origin_shape): # xtuner: num_experts * 2 * expert_dim, hidden_size @@ -406,12 +406,14 @@ def _forward( hidden_states: torch.Tensor, seq_ctx: SequenceContext, position_embeddings: tuple[torch.Tensor, torch.Tensor], + attention_kwargs: dict[str, object] | None = None, ) -> MoEDecoderLayerOutput: - residual, hidden_states, router_results = self._pre_moe_forward( + residual, hidden_states, router_results, attn_outputs = self._pre_moe_forward( hidden_states=hidden_states, seq_ctx=seq_ctx, position_embeddings=position_embeddings, state=ForwardState.TRAINING, + attention_kwargs=attention_kwargs, ) origin_shape = hidden_states.shape @@ -501,6 +503,20 @@ def _forward( residual=residual, shared_experts_out=shared_experts_out, ) + return self._build_output( + hidden_states=hidden_states, + router_results=router_results, + attn_outputs=attn_outputs, + ) + + def _build_output( + self, + *, + hidden_states: torch.Tensor, + router_results: RouterResults, + attn_outputs: AttnOutputs, + ) -> MoEDecoderLayerOutput: + """Build the public output; model-specific decoders may extend it.""" return { "hidden_states": hidden_states, "router_logits": router_results["logits"], @@ -513,6 +529,7 @@ def _micro_batch_forward( hidden_states_list: list[torch.Tensor], seq_ctx_list: list[SequenceContext], position_embeddings_list: list[tuple[torch.Tensor, torch.Tensor]], + attention_kwargs_list: list[dict[str, object]] | None = None, ) -> MoEDecoderLayerMicroBatchOutput: origin_shape = hidden_states_list[0].shape assert all(hidden_states.shape == origin_shape for hidden_states in hidden_states_list), ( @@ -521,6 +538,10 @@ def _micro_batch_forward( intra_layer_micro_batch = len(hidden_states_list) residual_list: list[torch.Tensor] = [] router_results_list: list[RouterResults] = [] + attn_outputs_list: list[AttnOutputs] = [] + if attention_kwargs_list is None: + attention_kwargs_list = [{} for _ in hidden_states_list] + assert len(attention_kwargs_list) == intra_layer_micro_batch pre_dispatched_list: list[PreDispatchResult] = [] dispatched_list: list[DispatchResult] = [] @@ -529,18 +550,21 @@ def _micro_batch_forward( # Attention + gate + pre-dispatch for ( hidden_states, + attention_kwargs, seq_ctx, position_embeddings, ) in zip( hidden_states_list, + attention_kwargs_list, seq_ctx_list, position_embeddings_list, ): - residual, hidden_states, router_results = self._pre_moe_forward( + residual, hidden_states, router_results, attn_outputs = self._pre_moe_forward( hidden_states=hidden_states, seq_ctx=seq_ctx, position_embeddings=position_embeddings, state=ForwardState.TRAINING, + attention_kwargs=attention_kwargs, ) pre_moe_forward_out_list.append(hidden_states) hidden_states = hidden_states.view(-1, hidden_states.shape[-1]) @@ -554,6 +578,7 @@ def _micro_batch_forward( pre_dispatched_list.append(pre_dispatched) residual_list.append(residual) router_results_list.append(router_results) + attn_outputs_list.append(attn_outputs) post_dispatched_list: list[PostDispatchResult] = [] experts_out_list: list[torch.Tensor] = [] @@ -655,8 +680,23 @@ def _micro_batch_forward( ) hidden_states_out_list.append(hidden_states) + return self._build_micro_batch_output( + hidden_states_list=hidden_states_out_list, + router_results_list=router_results_list, + attn_outputs_list=attn_outputs_list, + ) + + def _build_micro_batch_output( + self, + *, + hidden_states_list: list[torch.Tensor], + router_results_list: list[RouterResults], + attn_outputs_list: list[AttnOutputs], + ) -> MoEDecoderLayerMicroBatchOutput: + """Build the public micro-batch output; model-specific decoders may + extend it.""" return { - "hidden_states": hidden_states_out_list, + "hidden_states": hidden_states_list, "router_logits": [router_results["logits"] for router_results in router_results_list], "router_weights": [router_results["router_weights"] for router_results in router_results_list], "router_topk_ids": [router_results["topk_ids"] for router_results in router_results_list], @@ -669,7 +709,8 @@ def _pre_moe_forward( position_embeddings: tuple[torch.Tensor, torch.Tensor], state: ForwardState, past_key_values: list[list[torch.Tensor]] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor, RouterResults]: + attention_kwargs: dict[str, object] | None = None, + ) -> tuple[torch.Tensor, torch.Tensor, RouterResults, AttnOutputs]: # NOTE: In order to allow `torch.compile` to compile the ops before and after attention as much as possible, # attention, post-layernorm and gate are implemented in one function residual = hidden_states @@ -678,13 +719,16 @@ def _pre_moe_forward( # Self Attention checkpoint_record("attn.begin") if state == ForwardState.TRAINING: - attn_outputs: AttnOutputs = self.self_attn( + attention_forward = cast(Callable[..., AttnOutputs], self.self_attn) + attn_outputs = attention_forward( hidden_states=hidden_states, position_embeddings=position_embeddings, seq_ctx=seq_ctx, + **(attention_kwargs or {}), ) hidden_states = attn_outputs["projected_output"] elif state == ForwardState.PREFILLING: + attn_outputs = {} assert past_key_values is not None, "past_key_values should be provided in pre-filling state" hidden_states = self.self_attn.prefilling( # type: ignore hidden_states=hidden_states, @@ -693,6 +737,7 @@ def _pre_moe_forward( past_key_values=past_key_values, ) elif state == ForwardState.DECODING: + attn_outputs = {} assert past_key_values is not None, "past_key_values should be provided in decoding state" hidden_states = self.self_attn.decoding( # type: ignore hidden_states=hidden_states, @@ -719,7 +764,7 @@ def _pre_moe_forward( checkpoint_record("moe.gate.begin") router_results: RouterResults = self.gate(hidden_states, rollout_routed_experts) checkpoint_record("moe.gate.end") - return residual, hidden_states, router_results + return residual, hidden_states, router_results, attn_outputs def _shared_experts_forward( self, diff --git a/xtuner/v1/module/mtp/mtp_block.py b/xtuner/v1/module/mtp/mtp_block.py index 8a0d585b8..f13f1c223 100644 --- a/xtuner/v1/module/mtp/mtp_block.py +++ b/xtuner/v1/module/mtp/mtp_block.py @@ -1,6 +1,6 @@ """Multi-Token Prediction (MTP) Block implementation.""" -from typing import Callable +from typing import Callable, cast import torch import torch.nn as nn @@ -20,6 +20,8 @@ """One MTP depth produces the same keyed outputs as the decoder layer it wraps.""" +MTPInternalOutput = MoEDecoderLayerOutput | MoEDecoderLayerMicroBatchOutput + class MTPBlock(nn.Module): """Multi-Token Prediction (MTP) block containing multiple MTP layers. @@ -175,10 +177,11 @@ def _forward( mtp_outputs: list[MTPDepthOutput] = [] current_hidden_states = hidden_states.detach() if self.mtp_config.detach_mtp_inputs else hidden_states current_seq_ctx = seq_ctx + previous_layer_results: MTPInternalOutput | None = None num_steps = self.mtp_config.num_layers for step in range(num_steps): - layer = self.layers[0] if self.mtp_config.share_weights else self.layers[step] + layer = cast(MTPLayer, self.layers[0] if self.mtp_config.share_weights else self.layers[step]) # Roll each packed sequence independently so we get the (i+k)-th token while # respecting per-sequence boundaries inside the packed batch. current_seq_ctx = roll_sequence_context(current_seq_ctx, shifts=-1) @@ -187,17 +190,52 @@ def _forward( if self.mtp_config.detach_mtp_inputs: future_embeddings = future_embeddings.detach() - layer_results: MTPDepthOutput = layer( - current_hidden_states, - future_embeddings=future_embeddings, - position_embeddings=position_embeddings, - seq_ctx=current_seq_ctx, + layer_results = cast( + MoEDecoderLayerOutput, + self._call_decoder_layer( + layer, + current_hidden_states, + previous_layer_results=previous_layer_results, + future_embeddings=future_embeddings, + position_embeddings=position_embeddings, + seq_ctx=current_seq_ctx, + ), ) + previous_layer_results = layer_results current_hidden_states = layer_results["hidden_states"] - mtp_outputs.append(layer_results) + mtp_outputs.append( + { + "hidden_states": current_hidden_states, + "router_logits": layer_results["router_logits"], + "router_weights": layer_results["router_weights"], + "router_topk_ids": layer_results["router_topk_ids"], + } + ) return mtp_outputs + def _call_decoder_layer( + self, + layer: MTPLayer, + hidden_states: torch.Tensor | list[torch.Tensor], + *, + previous_layer_results: MTPInternalOutput | None, + future_embeddings: torch.Tensor | list[torch.Tensor], + position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], + seq_ctx: SequenceContext | list[SequenceContext], + ) -> MTPInternalOutput: + """Call one MTP decoder layer. + + Subclasses can consume model-specific fields from ``previous_layer_results`` while the + generic MTP input and public output stay model-agnostic. + """ + return layer( + hidden_states, + future_embeddings=future_embeddings, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + ) + def _micro_batch_forward( self, *, @@ -212,20 +250,27 @@ def _micro_batch_forward( outputs_per_mb: list[list[MTPDepthOutput]] = [[] for _ in range(n)] current_hidden_states_list = list(hidden_states_list) current_seq_ctx_list = list(seq_ctx_list) + previous_layer_results: MTPInternalOutput | None = None num_steps = self.mtp_config.num_layers for step in range(num_steps): - layer = self.layers[0] if self.mtp_config.share_weights else self.layers[step] + layer = cast(MTPLayer, self.layers[0] if self.mtp_config.share_weights else self.layers[step]) current_seq_ctx_list = [roll_sequence_context(ctx, shifts=-1) for ctx in current_seq_ctx_list] future_embeddings_list = [self._embed_future(ctx, embed_tokens_fn) for ctx in current_seq_ctx_list] - layer_results: MoEDecoderLayerMicroBatchOutput = layer( - current_hidden_states_list, - future_embeddings=future_embeddings_list, - position_embeddings=position_embeddings_list, - seq_ctx=current_seq_ctx_list, + layer_results = cast( + MoEDecoderLayerMicroBatchOutput, + self._call_decoder_layer( + layer, + current_hidden_states_list, + previous_layer_results=previous_layer_results, + future_embeddings=future_embeddings_list, + position_embeddings=position_embeddings_list, + seq_ctx=current_seq_ctx_list, + ), ) + previous_layer_results = layer_results for mb_idx in range(n): outputs_per_mb[mb_idx].append( diff --git a/xtuner/v1/module/mtp/mtp_layer.py b/xtuner/v1/module/mtp/mtp_layer.py index 620500cc9..410691275 100644 --- a/xtuner/v1/module/mtp/mtp_layer.py +++ b/xtuner/v1/module/mtp/mtp_layer.py @@ -107,7 +107,6 @@ def forward( Rotary position embeddings ``(cos, sin)``, aligned with ``hidden_states``. seq_ctx (SequenceContext | list[SequenceContext]): Sequence context, aligned with ``hidden_states``. - Returns: MoEDecoderLayerOutput | MoEDecoderLayerMicroBatchOutput: The wrapped decoder layer's outputs with the MTP final layernorm applied to the hidden states. @@ -185,7 +184,6 @@ def _micro_batch_forward( self._preprocess(hidden_states=h, future_embeddings=e) for h, e in zip(hidden_states_list, future_embeddings_list) ] - layer_results: MoEDecoderLayerMicroBatchOutput = self.decoder_layer( projected_list, position_embeddings=position_embeddings_list, diff --git a/xtuner/v1/ops/sparse_mla/pytorch.py b/xtuner/v1/ops/sparse_mla/pytorch.py index 7813973cc..0a0392bde 100644 --- a/xtuner/v1/ops/sparse_mla/pytorch.py +++ b/xtuner/v1/ops/sparse_mla/pytorch.py @@ -68,7 +68,7 @@ def torch_dsa_topk_indices( topk = min(index_topk, kv_len) topk_scores, topk_indices = index_scores.topk(topk, dim=-1) topk_indices = topk_indices.masked_fill(topk_scores == -torch.inf, -1) - return topk_indices.squeeze(0).unsqueeze(1) + return topk_indices.squeeze(0).unsqueeze(1).to(torch.int32) def _packed_causal_mask(seq_ctx: SequenceContext, query_len: int, kv_len: int, device: torch.device) -> torch.Tensor: diff --git a/xtuner/v1/ops/sparse_mla/tilelang.py b/xtuner/v1/ops/sparse_mla/tilelang.py index 603e75a4f..d186942fe 100644 --- a/xtuner/v1/ops/sparse_mla/tilelang.py +++ b/xtuner/v1/ops/sparse_mla/tilelang.py @@ -170,7 +170,7 @@ def _tilelang_dsa_topk_indices_from_ranges( topk = min(index_topk, k.shape[0]) topk_scores, topk_indices = logits.topk(topk, dim=-1) topk_indices = topk_indices.masked_fill(topk_scores == -torch.inf, -1) - return topk_indices.to(torch.int64).unsqueeze(1) + return topk_indices.to(torch.int32).unsqueeze(1) @_tilelang_dsa_topk_indices_from_ranges.register_fake @@ -183,7 +183,7 @@ def _( index_topk: int, ) -> Tensor: topk = min(index_topk, k.shape[0]) - return torch.empty((q.shape[0], 1, topk), device=q.device, dtype=torch.int64) + return torch.empty((q.shape[0], 1, topk), device=q.device, dtype=torch.int32) @functools.cache