diff --git a/tests/engine/test_glm52_moe_train_engine.py b/tests/engine/test_glm52_moe_train_engine.py index baf40d3b6..5695844f0 100644 --- a/tests/engine/test_glm52_moe_train_engine.py +++ b/tests/engine/test_glm52_moe_train_engine.py @@ -204,7 +204,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 +213,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 +228,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..40677a3fc 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_without_sequence_context_cache: 模型通过显式 IDs 完成前反向。 TestGlm52SequenceParallel test_mtp_loss_and_gradients_match_full_sequence: SP2 的 MTP loss 与梯度匹配完整序列。 """ @@ -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/module/attention/test_dsa_mla.py b/tests/module/attention/test_dsa_mla.py index 289a6ec8c..3e255ee50 100644 --- a/tests/module/attention/test_dsa_mla.py +++ b/tests/module/attention/test_dsa_mla.py @@ -23,14 +23,13 @@ import pytest import torch import torch.distributed as dist -import torch.nn as nn from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import CheckpointImpl from xtuner._testing import DeterministicDDPTestCase from xtuner.v1.data_proto import SequenceContext from xtuner.v1.model.utils import checkpoint_wrapper from xtuner.v1.module.attention import DSAMLAConfig -from xtuner.v1.module.attention.dsa_topk_sharing import register_dsa_topk_decoder_lifecycle_hooks +from xtuner.v1.module.decoder_layer.dense_decoder_layer import DenseDecoderLayer from xtuner.v1.ops.sparse_mla import dsa_topk_indices, sparse_mla from xtuner.v1.utils.test_utils import init_data_mesh @@ -103,6 +102,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, @@ -116,26 +119,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) -> DenseDecoderLayer: + return DenseDecoderLayer( + hidden_size=4, + intermediate_size=8, + hidden_act="silu", + attention_config=_tiny_dsa_config(indexer_types), + layer_idx=layer_idx, + ) class TestTorchSparseMLA: @@ -185,15 +179,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) @@ -201,39 +197,51 @@ 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) - def test_reentrant_checkpoint_reuses_and_releases_topk(self): - # 验证真实 source/shared decoder 经 reentrant checkpoint 重算后梯度有限且缓存释放。 + def test_reentrant_checkpoint_preserves_explicit_topk_storage(self): + # 验证真实 decoder 的 flat IDs 输入输出可经 reentrant checkpoint 完成反向。 torch.manual_seed(0) source_block = checkpoint_wrapper( - _TinyDsaDecoderBlock(_tiny_dsa_attention(indexer_types=["full", "shared"], layer_idx=0)), + _tiny_dsa_decoder(["full", "shared"], layer_idx=0), checkpoint_impl=CheckpointImpl.REENTRANT, ) shared_block = checkpoint_wrapper( - _TinyDsaDecoderBlock(_tiny_dsa_attention(indexer_types=["full", "shared"], layer_idx=1)), + _tiny_dsa_decoder(["full", "shared"], layer_idx=1), checkpoint_impl=CheckpointImpl.REENTRANT, ) 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_hidden, source_ids = source_block( + hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + ) + shared_hidden, shared_ids = shared_block( + source_hidden, + source_ids, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + ) + assert shared_ids.untyped_storage().data_ptr() == source_ids.untyped_storage().data_ptr() + shared_hidden.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: @@ -322,12 +330,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) @@ -338,12 +347,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..877bd9027 100644 --- a/tests/module/test_dense_decoder_layer.py +++ b/tests/module/test_dense_decoder_layer.py @@ -87,11 +87,20 @@ def test_batched_inputs_match_independent_forwards(self): ) assert isinstance(outputs, tuple) - for output, reference_output in zip(outputs, reference_outputs): + n = len(hidden_states) + assert len(outputs) == 2 * n + output_hidden = outputs[:n] + output_ids = outputs[n:] + reference_hidden = tuple(result[0] for result in reference_outputs) + reference_ids = tuple(result[1] 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.py b/xtuner/v1/model/moe/glm52.py index fd8bea587..89c7702fa 100644 --- a/xtuner/v1/model/moe/glm52.py +++ b/xtuner/v1/model/moe/glm52.py @@ -1,3 +1,4 @@ +import os import re from pathlib import Path from typing import Literal @@ -7,14 +8,14 @@ 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.loss import BalancingLossContext, ZLossContext 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, + dsa_topk_source_layers, ) from xtuner.v1.module.mtp import MTPConfig, MTPLayer from xtuner.v1.module.rope import RopeParametersConfig @@ -23,10 +24,8 @@ from .moe import MoE -# 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), @@ -59,63 +58,203 @@ 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] + 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 + ) + + @override + 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 GLM layers while carrying DSA top-k IDs as explicit tensors.""" + activation_offload = int(os.getenv("XTUNER_ACTIVATION_OFFLOAD", "0")) == 1 + dsa_topk_offload = int(os.getenv("XTUNER_DSA_TOPK_OFFLOAD", "0")) == 1 + dsa_topk_ids: torch.Tensor | None = None + offload_block_idx = 0 + + for idx, decoder_layer in self.layers.items(): + layer_idx = int(idx) + is_source = self._dsa_topk_source_layers[layer_idx] == layer_idx + if is_source: + dsa_topk_ids = None + else: + assert dsa_topk_ids is not None, f"DSA shared layer {layer_idx} requires dsa_topk_ids." + + layer_inputs = (hidden_states,) if dsa_topk_ids is None else (hidden_states, dsa_topk_ids) + offload_tensors = ( + [hidden_states] if activation_offload and layer_idx >= self.config.first_k_dense_replace else [] ) + if dsa_topk_offload and layer_idx in self._dsa_topk_last_consumers and dsa_topk_ids is not None: + offload_tensors.append(dsa_topk_ids) + + with self._saved_tensors_offload_ctx(offload_block_idx, offload_tensors): + layer_results = decoder_layer( + *layer_inputs, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + ) + if offload_tensors: + offload_block_idx += 1 + + assert isinstance(layer_results, tuple) + if layer_idx < self.config.first_k_dense_replace: + assert len(layer_results) == 2 + hidden_states, dsa_topk_ids = layer_results + else: + assert len(layer_results) == 5 + hidden_states, router_results, router_weights, router_topk_ids, dsa_topk_ids = layer_results + 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.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, + if self.config.return_hidden_states: + output["hidden_states"].append(hidden_states) + + return hidden_states + + @override + 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 GLM layers for micro-batches with flat checkpoint + inputs/results.""" + activation_offload = int(os.getenv("XTUNER_ACTIVATION_OFFLOAD", "0")) == 1 + dsa_topk_offload = int(os.getenv("XTUNER_DSA_TOPK_OFFLOAD", "0")) == 1 + dsa_topk_ids_list: list[torch.Tensor] | None = None + offload_block_idx = 0 + n = len(hidden_states_list) + + for idx, decoder_layer in self.layers.items(): + layer_idx = int(idx) + is_source = self._dsa_topk_source_layers[layer_idx] == layer_idx + if is_source: + dsa_topk_ids_list = None + else: + assert dsa_topk_ids_list is not None, f"DSA shared layer {layer_idx} requires dsa_topk_ids." + + layer_inputs = list(hidden_states_list) + if dsa_topk_ids_list is not None: + layer_inputs.extend(dsa_topk_ids_list) + + offload_tensors = ( + list(hidden_states_list) + if activation_offload and layer_idx >= self.config.first_k_dense_replace + else [] + ) + if dsa_topk_offload and layer_idx in self._dsa_topk_last_consumers and dsa_topk_ids_list is not None: + offload_tensors.extend(dsa_topk_ids_list) + + with self._saved_tensors_offload_ctx(offload_block_idx, offload_tensors): + layer_results = decoder_layer( + *layer_inputs, + position_embeddings=position_embeddings_list, + seq_ctx=seq_ctx_list, + ) + if offload_tensors: + offload_block_idx += 1 + + assert isinstance(layer_results, tuple) + if layer_idx < self.config.first_k_dense_replace: + assert len(layer_results) == 2 * n + hidden_states_list = list(layer_results[:n]) + dsa_topk_ids_list = list(layer_results[n:]) + continue + + assert len(layer_results) == 5 * n + hidden_states = layer_results[:n] + router_logits = layer_results[n : 2 * n] + router_weights = layer_results[2 * n : 3 * n] + router_topk_ids = layer_results[3 * n : 4 * n] + dsa_topk_ids_list = list(layer_results[4 * n :]) + + for micro_batch_idx, hidden_state in enumerate(hidden_states): + hidden_states_list[micro_batch_idx] = hidden_state + if keep_router: + router_logits_list[micro_batch_idx][f"layer{idx}"] = self._maybe_offload_router( + router_logits[micro_batch_idx] + ) + + 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 to_hf_key_list(self, key: str) -> list[str]: if self.config.tie_word_embeddings and "lm_head" in key: key = key.replace("lm_head", "embed_tokens") diff --git a/xtuner/v1/model/moe/moe.py b/xtuner/v1/model/moe/moe.py index 23b2369ed..fed5003ba 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 @@ -23,7 +24,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, @@ -208,7 +209,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] @@ -238,9 +239,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, @@ -540,72 +565,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. - hidden_states_list = list( - decoder_layer( - *hidden_states_list, - position_embeddings=position_embeddings_list, - seq_ctx=seq_ctx_list, - ) - ) - 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 = decoder_layer( - *hidden_states_list, - position_embeddings=position_embeddings_list, - seq_ctx=seq_ctx_list, - ) - else: - layer_results = decoder_layer( - *hidden_states_list, - position_embeddings=position_embeddings_list, - seq_ctx=seq_ctx_list, - ) - hidden_states = layer_results[: len(hidden_states_list)] - router_logits = layer_results[len(hidden_states_list) : len(hidden_states_list) * 2] - router_weights = layer_results[len(hidden_states_list) * 2 : len(hidden_states_list) * 3] - router_topk_ids = layer_results[len(hidden_states_list) * 3 :] - - # 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(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" @@ -624,7 +596,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(), ) ) @@ -741,6 +712,79 @@ 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.""" + activation_offload = int(os.getenv("XTUNER_ACTIVATION_OFFLOAD", "0")) == 1 + offload_block_idx = 0 + + 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. + hidden_states_list = list( + decoder_layer( + *hidden_states_list, + position_embeddings=position_embeddings_list, + seq_ctx=seq_ctx_list, + ) + ) + continue + + offload_tensors = list(hidden_states_list) if activation_offload else [] + with self._saved_tensors_offload_ctx(offload_block_idx, offload_tensors): + layer_results = decoder_layer( + *hidden_states_list, + position_embeddings=position_embeddings_list, + seq_ctx=seq_ctx_list, + ) + if offload_tensors: + offload_block_idx += 1 + + n = len(hidden_states_list) + hidden_states = layer_results[:n] + router_logits = layer_results[n : 2 * n] + router_weights = layer_results[2 * n : 3 * n] + router_topk_ids = layer_results[3 * n :] + + # 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 @@ -782,57 +826,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: - hidden_states = decoder_layer( - hidden_states, - position_embeddings=position_embeddings, - seq_ctx=seq_ctx, - ) - 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 = decoder_layer( - hidden_states, - position_embeddings=position_embeddings, - seq_ctx=seq_ctx, - ) - - else: - layer_results = decoder_layer( - hidden_states, - position_embeddings=position_embeddings, - seq_ctx=seq_ctx, - ) - hidden_states, router_results, router_weights, router_topk_ids = layer_results - 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) + 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) @@ -854,7 +867,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] @@ -923,6 +935,65 @@ 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.""" + activation_offload = int(os.getenv("XTUNER_ACTIVATION_OFFLOAD", "0")) == 1 + offload_block_idx = 0 + + for idx, decoder_layer in self.layers.items(): + layer_idx = int(idx) + if layer_idx < self.config.first_k_dense_replace: + hidden_states = decoder_layer( + hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + ) + else: + offload_tensors = [hidden_states] if activation_offload else [] + with self._saved_tensors_offload_ctx(offload_block_idx, offload_tensors): + layer_results = decoder_layer( + hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + ) + if offload_tensors: + offload_block_idx += 1 + + hidden_states, router_results, router_weights, router_topk_ids = layer_results + 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 build_embeddings(self, config: MoEConfig): return nn.Embedding(config.vocab_size, config.hidden_size, config.pad_token_id) @@ -1190,27 +1261,13 @@ def fully_shard( if self._should_recompute(None, mtp_idx=mtp_idx) or ( self.config.mtp_config is not None and self.config.mtp_config.share_weights ): # share mtp head must recompute - # MTP 默认使用 reentrant 的原因: - # Case 1:最小触发条件是 compile, topk offload, MTP share weights and depth > 1. - # 多个 logical depth 共用 top-k cache。reentrant 的 original - # 关闭 grad、replay 开启 grad,DSA 能据此正确更新 cache 计数。 - # original 不建立内部图,所以 replay 可以安全复用离散 top-k。 - # non-reentrant 的两次执行都开启 grad,却仍沿用该复用策略, - # 因而出现 original=COMPUTE、replay=REUSE,无法重建相同清单。 - # - # indexer 本身始终 no_grad。不开 compile 时,多执行/少执行一次 - # indexer 不会改变 eager autograd 的保存清单;开启 compile 后, - # COMPUTE/REUSE 经过不同 graph break 和 compiled block,才可能让 - # checkpoint 保存槽位错位并报 different metadata。例如 original - # 保存 [A, B, C]、replay 保存 [A, X, C] 时,槽位 1 的 metadata - # 不同。后续若显式记录 ORIGINAL/REPLAY phase,可再让 - # non-reentrant 正确推进 cache 状态。 + # DSA IDs are flat checkpoint inputs/results, so shared MTP + # logical depths keep the same storage without cache state. + # A source depth recomputes its no-grad indexer during replay. # - # 使用 reentrant 时还必须用 pytree_reentrant_checkpoint: - # Case 2:触发条件是 EP > 1, intra-layer micro-batch > 1(例如 micro2). - # micro2 传入 [embedding_0, embedding_1];pytree 把 list 内 Tensor - # 展开后,checkpoint 才能在 replay 前逐个 detach,并在 backward - # 中把梯度交回原始 embedding graph。 + # Reentrant microbatch MTP still needs the pytree adapter: + # future_embeddings remains a list keyword, and each tensor + # must be detached/reconnected by CheckpointFunction. use_reentrant = self.fsdp_config.mtp_checkpoint_use_reentrant if use_reentrant: mtp_layer = checkpoint_wrapper( diff --git a/xtuner/v1/module/attention/attn_outputs.py b/xtuner/v1/module/attention/attn_outputs.py index e78cf2841..9c90a7c12 100644 --- a/xtuner/v1/module/attention/attn_outputs.py +++ b/xtuner/v1/module/attention/attn_outputs.py @@ -8,3 +8,4 @@ class AttnOutputs(TypedDict, total=False): raw_output: torch.Tensor softmax_lse: torch.Tensor | None attn_logits: torch.Tensor | None + dsa_topk_ids: torch.Tensor diff --git a/xtuner/v1/module/attention/dsa_mla.py b/xtuner/v1/module/attention/dsa_mla.py index 23e0f9f68..d125ed8cf 100644 --- a/xtuner/v1/module/attention/dsa_mla.py +++ b/xtuner/v1/module/attention/dsa_mla.py @@ -4,6 +4,7 @@ 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 @@ -21,7 +22,7 @@ from ..linear import build_linear from .attn_outputs import AttnOutputs -from .dsa_topk_sharing import build_dsa_topk_release_plan, dsa_topk_source_layer, get_dsa_topk_sharing_runtime +from .dsa_topk_sharing import dsa_topk_source_layer from .mla import MLAConfig, MultiLatentAttention, mla_apply_rotary_pos_emb @@ -141,15 +142,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 +233,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,6 +261,7 @@ def forward( hidden_states: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], seq_ctx: SequenceContext, + dsa_topk_ids: torch.Tensor | None = None, ) -> AttnOutputs: """Absorbed DSA-MLA forward for packed training (``bsz == 1``). @@ -352,21 +336,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 +374,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, + ) -> AttnOutputs: ... + + __call__ = nn.Module.__call__ diff --git a/xtuner/v1/module/attention/dsa_topk_sharing.py b/xtuner/v1/module/attention/dsa_topk_sharing.py index ce6e2f574..4e6053040 100644 --- a/xtuner/v1/module/attention/dsa_topk_sharing.py +++ b/xtuner/v1/module/attention/dsa_topk_sharing.py @@ -1,30 +1,4 @@ # 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( @@ -34,7 +8,7 @@ def dsa_topk_source_layer( index_skip_topk_offset: int, index_topk_freq: int, ) -> int: - """Resolve the physical indexer source for one logical DSA layer.""" + """Resolve the 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 @@ -52,467 +26,20 @@ def dsa_topk_source_layer( 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( +def dsa_topk_source_layers( *, - num_main_layers: int, - num_mtp_layers: int, + num_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( +) -> 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, ) - 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() - }, + for layer_idx in range(num_layers) ) - - -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. Reentrant checkpoint replay invokes the decoder - # module and these hooks again, 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 426e353b9..4a74a11e8 100644 --- a/xtuner/v1/module/decoder_layer/dense_decoder_layer.py +++ b/xtuner/v1/module/decoder_layer/dense_decoder_layer.py @@ -1,4 +1,4 @@ -from typing import Literal +from typing import Literal, cast import torch import torch.nn as nn @@ -6,7 +6,14 @@ from xtuner.v1.config import GenerateConfig from xtuner.v1.data_proto import SequenceContext from xtuner.v1.float8.config import Float8Config -from xtuner.v1.module import AttnOutputs, GatedDeltaNetConfig, MHAConfig, MLAConfig, RMSNorm +from xtuner.v1.module import ( + AttnOutputs, + DSAMultiLatentAttention, + GatedDeltaNetConfig, + MHAConfig, + MLAConfig, + RMSNorm, +) from xtuner.v1.module.rope import RopeScalingConfig from xtuner.v1.ops.act_fn import get_act_fn from xtuner.v1.utils import ForwardState @@ -74,7 +81,7 @@ def __init__( def forward( self, - *hidden_states: torch.Tensor, + *layer_inputs: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], seq_ctx: SequenceContext | list[SequenceContext], ) -> torch.Tensor | tuple[torch.Tensor, ...]: @@ -82,45 +89,77 @@ def forward( Keeping the micro-batch loop inside the decoder layer lets outer FSDP and checkpoint wrappers materialize the layer only once, while each - attention call keeps its own ``SequenceContext``. + attention call keeps its own ``SequenceContext``. DSA callers append + explicit IDs after the hidden-state inputs; DSA results use the flat + ``(hidden..., dsa_topk_ids...)`` layout required by reentrant checkpoint. """ - if len(hidden_states) == 1: + if isinstance(seq_ctx, SequenceContext): + assert len(layer_inputs) in (1, 2), ( + "Single-microbatch DenseDecoderLayer expects hidden_states and optional dsa_topk_ids." + ) assert isinstance(position_embeddings, tuple) and len(position_embeddings) == 2 - assert isinstance(seq_ctx, SequenceContext) return self._forward( - hidden_states=hidden_states[0], + hidden_states=layer_inputs[0], position_embeddings=position_embeddings, seq_ctx=seq_ctx, + dsa_topk_ids=layer_inputs[1] if len(layer_inputs) == 2 else None, ) - assert isinstance(position_embeddings, list) and len(position_embeddings) == len(hidden_states) - assert isinstance(seq_ctx, list) and len(seq_ctx) == len(hidden_states) + n = len(seq_ctx) + assert len(layer_inputs) in (n, 2 * n), ( + f"Multi-microbatch DenseDecoderLayer expects {n} hidden states and optional {n} dsa_topk_ids." + ) + hidden_states = layer_inputs[:n] + dsa_topk_ids: tuple[torch.Tensor | None, ...] + if len(layer_inputs) == n: + dsa_topk_ids = (None,) * n + else: + dsa_topk_ids = layer_inputs[n:] + + assert isinstance(position_embeddings, list) and len(position_embeddings) == n assert all(hidden.shape == hidden_states[0].shape for hidden in hidden_states) - return tuple( + layer_results = tuple( self._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, position_embeddings, seq_ctx ) - for hidden, position_embedding, context in zip(hidden_states, position_embeddings, seq_ctx) ) + if isinstance(layer_results[0], tuple): + return tuple(result[0] for result in layer_results) + tuple(result[1] for result in layer_results) # type: ignore[index] + return layer_results # type: ignore[return-value] def _forward( self, hidden_states: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], seq_ctx: SequenceContext, - ) -> torch.Tensor: + dsa_topk_ids: torch.Tensor | None, + ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: residual = hidden_states hidden_states = self.input_layernorm(hidden_states) # Self Attention - attn_outputs: AttnOutputs = self.self_attn( - hidden_states=hidden_states, - position_embeddings=position_embeddings, - seq_ctx=seq_ctx, - ) + if dsa_topk_ids is None: + attn_outputs: AttnOutputs = self.self_attn( + hidden_states=hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + ) + else: + dsa_attn = cast(DSAMultiLatentAttention, self.self_attn) + attn_outputs = dsa_attn( + hidden_states=hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + dsa_topk_ids=dsa_topk_ids, + ) + dsa_topk_ids = attn_outputs.get("dsa_topk_ids") hidden_states = attn_outputs["projected_output"] hidden_states = residual + hidden_states @@ -130,7 +169,9 @@ def _forward( hidden_states = self.mlp(hidden_states) hidden_states = residual + hidden_states - return hidden_states + if dsa_topk_ids is None: + return hidden_states + return hidden_states, dsa_topk_ids def prefilling( self, diff --git a/xtuner/v1/module/decoder_layer/moe_decoder_layer.py b/xtuner/v1/module/decoder_layer/moe_decoder_layer.py index 00e5d6c27..fb0e8ab38 100644 --- a/xtuner/v1/module/decoder_layer/moe_decoder_layer.py +++ b/xtuner/v1/module/decoder_layer/moe_decoder_layer.py @@ -14,6 +14,7 @@ from xtuner.v1.float8 import Float8Config from xtuner.v1.module import ( AttnOutputs, + DSAMultiLatentAttention, GatedDeltaNet, GatedDeltaNetConfig, GreedyRouterConfig, @@ -291,7 +292,7 @@ def __init__( def forward( self, - *hidden_states: torch.Tensor, + *layer_inputs: torch.Tensor, seq_ctx: SequenceContext | list[SequenceContext], position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]] | None = None, ) -> tuple[HiddenStates, RouterLogits, RouterWeights, RouterTopKIds] | tuple[torch.Tensor, ...]: @@ -305,33 +306,43 @@ def forward( Returns: tuple: Output hidden states, router logits, router weights, and the - expert IDs selected by the router. + expert IDs selected by the router. DSA layers append + ``dsa_topk_ids``; multi-microbatch results keep each category + contiguous in a flat ``4 * N`` or ``5 * N`` tuple. """ - if len(hidden_states) == 1: - assert isinstance(seq_ctx, SequenceContext), ( - f"seq_ctx should be a SequenceContext instance but got {seq_ctx}" + if isinstance(seq_ctx, SequenceContext): + assert len(layer_inputs) in (1, 2), ( + "Single-microbatch MoEDecoderLayer expects hidden_states and optional dsa_topk_ids." ) assert isinstance(position_embeddings, tuple) and len(position_embeddings) == 2, ( "position_embeddings should be a tuple of two tensors (position_ids, position_embeds)" ) return self._forward( - hidden_states=hidden_states[0], + hidden_states=layer_inputs[0], seq_ctx=seq_ctx, position_embeddings=position_embeddings, + dsa_topk_ids=layer_inputs[1] if len(layer_inputs) == 2 else None, ) + + n = len(seq_ctx) + assert len(layer_inputs) in (n, 2 * n), ( + f"Multi-microbatch MoEDecoderLayer expects {n} hidden states and optional {n} dsa_topk_ids." + ) + assert isinstance(position_embeddings, list) and len(position_embeddings) == n, ( + "position_embeddings should be a list of tuples with the same length as seq_ctx" + ) + dsa_topk_ids_list: list[torch.Tensor | None] + if len(layer_inputs) == n: + dsa_topk_ids_list = [None] * n 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" - ) + dsa_topk_ids_list = list(layer_inputs[n:]) - return self._micro_batch_forward( - hidden_states_list=list(hidden_states), - seq_ctx_list=seq_ctx, - position_embeddings_list=position_embeddings, - ) + return self._micro_batch_forward( + hidden_states_list=list(layer_inputs[:n]), + dsa_topk_ids_list=dsa_topk_ids_list, + 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 @@ -376,12 +387,14 @@ def _forward( hidden_states: torch.Tensor, seq_ctx: SequenceContext, position_embeddings: tuple[torch.Tensor, torch.Tensor], - ) -> tuple[HiddenStates, RouterLogits, RouterWeights, RouterTopKIds]: - residual, hidden_states, router_results = self._pre_moe_forward( + dsa_topk_ids: torch.Tensor | None, + ) -> tuple[torch.Tensor, ...]: + residual, hidden_states, router_results, dsa_topk_ids = self._pre_moe_forward( hidden_states=hidden_states, seq_ctx=seq_ctx, position_embeddings=position_embeddings, state=ForwardState.TRAINING, + dsa_topk_ids=dsa_topk_ids, ) origin_shape = hidden_states.shape @@ -461,16 +474,20 @@ def _forward( residual=residual, shared_experts_out=shared_experts_out, ) - return ( + layer_results = ( hidden_states, router_results["logits"], router_results["router_weights"], router_results["topk_ids"], ) + if dsa_topk_ids is None: + return layer_results + return (*layer_results, dsa_topk_ids) def _micro_batch_forward( self, hidden_states_list: list[torch.Tensor], + dsa_topk_ids_list: list[torch.Tensor | None], seq_ctx_list: list[SequenceContext], position_embeddings_list: list[tuple[torch.Tensor, torch.Tensor]], ) -> tuple[torch.Tensor, ...]: @@ -481,6 +498,7 @@ def _micro_batch_forward( intra_layer_micro_batch = len(hidden_states_list) residual_list: list[torch.Tensor] = [] router_results_list: list[RouterResults] = [] + dsa_topk_ids_out: list[torch.Tensor | None] = [] pre_dispatched_list: list[PreDispatchResult] = [] dispatched_list: list[DispatchResult] = [] @@ -489,18 +507,21 @@ def _micro_batch_forward( # Attention + gate + pre-dispatch for ( hidden_states, + dsa_topk_ids, seq_ctx, position_embeddings, ) in zip( hidden_states_list, + dsa_topk_ids_list, seq_ctx_list, position_embeddings_list, ): - residual, hidden_states, router_results = self._pre_moe_forward( + residual, hidden_states, router_results, dsa_topk_ids = self._pre_moe_forward( hidden_states=hidden_states, seq_ctx=seq_ctx, position_embeddings=position_embeddings, state=ForwardState.TRAINING, + dsa_topk_ids=dsa_topk_ids, ) pre_moe_forward_out_list.append(hidden_states) hidden_states = hidden_states.view(-1, hidden_states.shape[-1]) @@ -512,6 +533,7 @@ def _micro_batch_forward( pre_dispatched_list.append(pre_dispatched) residual_list.append(residual) router_results_list.append(router_results) + dsa_topk_ids_out.append(dsa_topk_ids) post_dispatched_list: list[PostDispatchResult] = [] experts_out_list: list[torch.Tensor] = [] @@ -601,7 +623,11 @@ def _micro_batch_forward( 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] - return tuple(hidden_states_out_list + router_logits + router_weights + router_topk_ids) + layer_results = hidden_states_out_list + router_logits + router_weights + router_topk_ids + if all(dsa_topk_ids is None for dsa_topk_ids in dsa_topk_ids_out): + return tuple(layer_results) + assert all(dsa_topk_ids is not None for dsa_topk_ids in dsa_topk_ids_out) + return tuple(layer_results + cast(list[torch.Tensor], dsa_topk_ids_out)) def _pre_moe_forward( self, @@ -610,7 +636,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]: + dsa_topk_ids: torch.Tensor | None = None, + ) -> tuple[torch.Tensor, torch.Tensor, RouterResults, torch.Tensor | None]: # 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 @@ -618,11 +645,21 @@ def _pre_moe_forward( # Self Attention if state == ForwardState.TRAINING: - attn_outputs: AttnOutputs = self.self_attn( - hidden_states=hidden_states, - position_embeddings=position_embeddings, - seq_ctx=seq_ctx, - ) + if dsa_topk_ids is None: + attn_outputs: AttnOutputs = self.self_attn( + hidden_states=hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + ) + else: + dsa_attn = cast(DSAMultiLatentAttention, self.self_attn) + attn_outputs = dsa_attn( + hidden_states=hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + dsa_topk_ids=dsa_topk_ids, + ) + dsa_topk_ids = attn_outputs.get("dsa_topk_ids") hidden_states = attn_outputs["projected_output"] elif state == ForwardState.PREFILLING: assert past_key_values is not None, "past_key_values should be provided in pre-filling state" @@ -656,7 +693,7 @@ def _pre_moe_forward( else: rollout_routed_experts = None router_results: RouterResults = self.gate(hidden_states, rollout_routed_experts) - return residual, hidden_states, router_results + return residual, hidden_states, router_results, dsa_topk_ids def _shared_experts_forward( self, diff --git a/xtuner/v1/module/mtp/mtp_block.py b/xtuner/v1/module/mtp/mtp_block.py index 9d43f685e..1057de6c0 100644 --- a/xtuner/v1/module/mtp/mtp_block.py +++ b/xtuner/v1/module/mtp/mtp_block.py @@ -173,6 +173,7 @@ def _forward( """ mtp_outputs: list[MTPDepthOutput] = [] current_hidden_states = hidden_states.detach() if self.mtp_config.detach_mtp_inputs else hidden_states + current_dsa_topk_ids: torch.Tensor | None = None current_seq_ctx = seq_ctx num_steps = self.mtp_config.num_layers @@ -186,12 +187,20 @@ def _forward( if self.mtp_config.detach_mtp_inputs: future_embeddings = future_embeddings.detach() - current_hidden_states, router_logits, router_weights, router_topk_ids = layer( - current_hidden_states, + layer_inputs = ( + (current_hidden_states,) + if current_dsa_topk_ids is None + else (current_hidden_states, current_dsa_topk_ids) + ) + layer_results = layer( + *layer_inputs, future_embeddings=future_embeddings, position_embeddings=position_embeddings, seq_ctx=current_seq_ctx, ) + assert len(layer_results) in (4, 5) + current_hidden_states, router_logits, router_weights, router_topk_ids = layer_results[:4] + current_dsa_topk_ids = layer_results[4] if len(layer_results) == 5 else None mtp_outputs.append((current_hidden_states, router_logits, router_weights, router_topk_ids)) return mtp_outputs @@ -209,6 +218,7 @@ def _micro_batch_forward( # outputs_per_mb[mb_idx][depth_idx] to match the single-microbatch API shape. outputs_per_mb: list[list[MTPDepthOutput]] = [[] for _ in range(n)] current_hidden_states_list = list(hidden_states_list) + current_dsa_topk_ids_list: list[torch.Tensor] | None = None current_seq_ctx_list = list(seq_ctx_list) num_steps = self.mtp_config.num_layers @@ -218,20 +228,24 @@ def _micro_batch_forward( 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_inputs = list(current_hidden_states_list) + if current_dsa_topk_ids_list is not None: + layer_inputs.extend(current_dsa_topk_ids_list) layer_results = layer( - *current_hidden_states_list, + *layer_inputs, future_embeddings=future_embeddings_list, position_embeddings=position_embeddings_list, seq_ctx=current_seq_ctx_list, ) - assert isinstance(layer_results, tuple) and len(layer_results) == 4 * n, ( - f"MTPLayer multi-microbatch forward should return a flat tuple of length {4 * n}, " + assert isinstance(layer_results, tuple) and len(layer_results) in (4 * n, 5 * n), ( + f"MTPLayer multi-microbatch forward should return a flat tuple of length {4 * n} or {5 * n}, " f"got {len(layer_results) if isinstance(layer_results, tuple) else type(layer_results)}" ) new_hidden = list(layer_results[:n]) router_logits = list(layer_results[n : 2 * n]) router_weights = list(layer_results[2 * n : 3 * n]) - router_topk_ids = list(layer_results[3 * n :]) + router_topk_ids = list(layer_results[3 * n : 4 * n]) + current_dsa_topk_ids_list = list(layer_results[4 * n :]) if len(layer_results) == 5 * n else None 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 711c7f4b2..ca01ab2a4 100644 --- a/xtuner/v1/module/mtp/mtp_layer.py +++ b/xtuner/v1/module/mtp/mtp_layer.py @@ -81,20 +81,17 @@ def __init__( def forward( self, - *hidden_states: torch.Tensor, + *layer_inputs: 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], ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] | tuple[torch.Tensor, ...]: """Forward pass through the MTP layer. - Mirrors :meth:`MoEDecoderLayer.forward`: when a single ``hidden_states`` tensor is - provided, the layer runs the regular single-microbatch path and returns a 4-tuple - ``(hidden, router_logits, router_weights, router_topk_ids)``. When ``N`` hidden states are provided - (intra-layer micro-batching / domino EP), ``future_embeddings``, ``position_embeddings`` - and ``seq_ctx`` must be lists of length ``N``; the per-microbatch preprocessing - (enorm/hnorm/eh_proj) is run independently and a single underlying decoder forward - is issued so the inner MoE EP communication can be overlapped across micro-batches. + Mirrors :meth:`MoEDecoderLayer.forward`: DSA calls append explicit + ``dsa_topk_ids`` to both the flat inputs and results. The enclosing + :class:`MTPBlock` consumes that extra category internally and keeps the + public per-depth output unchanged. Args: hidden_states (torch.Tensor): One or more hidden state tensors. A single tensor @@ -107,41 +104,46 @@ def forward( seq_ctx (SequenceContext | list[SequenceContext]): Sequence context per micro-batch. Returns: - tuple: For single-microbatch input, a 4-tuple - ``(hidden_states, router_logits, router_weights, router_topk_ids)``. - For ``N`` micro-batches, a flat tuple of length ``4 * N`` matching the - convention used by :meth:`MoEDecoderLayer._micro_batch_forward`: - ``(hidden_0, ..., hidden_{N-1}, router_logits_0, ..., - router_weights_{N-1}, router_topk_ids_0, ..., router_topk_ids_{N-1})``. + tuple: A flat ``4 * N`` result for non-DSA decoders or ``5 * N`` + result for DSA decoders, with ``dsa_topk_ids`` as the final + category. """ - if len(hidden_states) == 1: + if isinstance(seq_ctx, SequenceContext): + assert len(layer_inputs) in (1, 2), ( + "Single-microbatch MTPLayer expects hidden_states and optional dsa_topk_ids." + ) assert isinstance(future_embeddings, torch.Tensor), ( "future_embeddings should be a Tensor in single-microbatch mode" ) - assert isinstance(seq_ctx, SequenceContext), ( - "seq_ctx should be a SequenceContext instance in single-microbatch mode" - ) assert isinstance(position_embeddings, tuple) and len(position_embeddings) == 2, ( "position_embeddings should be a (cos, sin) tuple in single-microbatch mode" ) return self._forward( - hidden_states=hidden_states[0], + hidden_states=layer_inputs[0], + dsa_topk_ids=layer_inputs[1] if len(layer_inputs) == 2 else None, future_embeddings=future_embeddings, position_embeddings=position_embeddings, seq_ctx=seq_ctx, ) + n = len(seq_ctx) + assert len(layer_inputs) in (n, 2 * n), ( + f"Multi-microbatch MTPLayer expects {n} hidden states and optional {n} dsa_topk_ids." + ) assert isinstance(future_embeddings, list), ( "future_embeddings should be a list aligned with hidden_states in multi-microbatch mode" ) - assert isinstance(seq_ctx, list), ( - "seq_ctx should be a list aligned with hidden_states in multi-microbatch mode" - ) assert isinstance(position_embeddings, list), ( "position_embeddings should be a list aligned with hidden_states in multi-microbatch mode" ) + dsa_topk_ids_list: list[torch.Tensor | None] + if len(layer_inputs) == n: + dsa_topk_ids_list = [None] * n + else: + dsa_topk_ids_list = list(layer_inputs[n:]) return self._micro_batch_forward( - hidden_states_list=list(hidden_states), + hidden_states_list=list(layer_inputs[:n]), + dsa_topk_ids_list=dsa_topk_ids_list, future_embeddings_list=future_embeddings, position_embeddings_list=position_embeddings, seq_ctx_list=seq_ctx, @@ -150,25 +152,33 @@ def forward( def _forward( self, hidden_states: torch.Tensor, + dsa_topk_ids: torch.Tensor | None, future_embeddings: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], seq_ctx: SequenceContext, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + ) -> tuple[torch.Tensor, ...]: projected = self._preprocess(hidden_states=hidden_states, future_embeddings=future_embeddings) - hidden_states, router_results, router_weights, router_topk_ids = self.decoder_layer( - projected, + decoder_inputs = (projected,) if dsa_topk_ids is None else (projected, dsa_topk_ids) + layer_results = self.decoder_layer( + *decoder_inputs, position_embeddings=position_embeddings, seq_ctx=seq_ctx, ) + assert len(layer_results) in (4, 5) + hidden_states, router_results, router_weights, router_topk_ids = layer_results[:4] hidden_states = self.final_layernorm(hidden_states) - return hidden_states, router_results, router_weights, router_topk_ids + mtp_results = (hidden_states, router_results, router_weights, router_topk_ids) + if len(layer_results) == 4: + return mtp_results + return (*mtp_results, layer_results[4]) def _micro_batch_forward( self, *, hidden_states_list: list[torch.Tensor], + dsa_topk_ids_list: list[torch.Tensor | None], future_embeddings_list: list[torch.Tensor], position_embeddings_list: list[tuple[torch.Tensor, torch.Tensor]], seq_ctx_list: list[SequenceContext], @@ -185,22 +195,30 @@ def _micro_batch_forward( for h, e in zip(hidden_states_list, future_embeddings_list) ] + decoder_inputs = list(projected_list) + if any(dsa_topk_ids is not None for dsa_topk_ids in dsa_topk_ids_list): + assert all(dsa_topk_ids is not None for dsa_topk_ids in dsa_topk_ids_list) + decoder_inputs.extend(dsa_topk_ids for dsa_topk_ids in dsa_topk_ids_list if dsa_topk_ids is not None) layer_results = self.decoder_layer( - *projected_list, + *decoder_inputs, position_embeddings=position_embeddings_list, seq_ctx=seq_ctx_list, ) - assert isinstance(layer_results, tuple) and len(layer_results) == 4 * n, ( + assert isinstance(layer_results, tuple) and len(layer_results) in (4 * n, 5 * n), ( "Multi-microbatch MTP requires the wrapped decoder layer to return a flat " - f"(hidden..., router_logits..., router_weights..., router_topk_ids...) tuple of length {4 * n}; " + "(hidden..., router_logits..., router_weights..., router_topk_ids..., optional dsa_topk_ids...) " + f"tuple of length {4 * n} or {5 * n}; " f"got length {len(layer_results) if isinstance(layer_results, tuple) else type(layer_results)}" ) hidden_out = [self.final_layernorm(h) for h in layer_results[:n]] router_logits = list(layer_results[n : 2 * n]) router_weights = list(layer_results[2 * n : 3 * n]) - router_topk_ids = list(layer_results[3 * n :]) - return tuple(hidden_out + router_logits + router_weights + router_topk_ids) + router_topk_ids = list(layer_results[3 * n : 4 * n]) + mtp_results = hidden_out + router_logits + router_weights + router_topk_ids + if len(layer_results) == 4 * n: + return tuple(mtp_results) + return tuple(mtp_results + list(layer_results[4 * n :])) def _preprocess( self, 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