diff --git a/tests/engine/test_glm52_moe_train_engine.py b/tests/engine/test_glm52_moe_train_engine.py index 88dd9386d9..efb94f7100 100644 --- a/tests/engine/test_glm52_moe_train_engine.py +++ b/tests/engine/test_glm52_moe_train_engine.py @@ -36,8 +36,7 @@ from xtuner.v1.loss.ce_loss import CELossConfig from xtuner.v1.model import get_model_config_from_hf from xtuner.v1.model.base import ModelItem -from xtuner.v1.model.moe.glm52 import Glm52MoEConfig -from xtuner.v1.module.attention import DSAMLAConfig +from xtuner.v1.model.moe.glm52 import DSAMLAConfig, Glm52MoEConfig from xtuner.v1.module.mtp import MTPConfig from xtuner.v1.module.router.noaux_router import NoAuxRouter, NoAuxRouterConfig from xtuner.v1.utils import pad_to_max_length @@ -208,7 +207,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): @@ -218,7 +216,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( @@ -234,9 +231,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/engine/test_moe_train_engine_float8.py b/tests/engine/test_moe_train_engine_float8.py index ec6df8ba9c..4d6f8dfe4c 100644 --- a/tests/engine/test_moe_train_engine_float8.py +++ b/tests/engine/test_moe_train_engine_float8.py @@ -20,7 +20,7 @@ from xtuner.v1.utils.device import get_device from xtuner.v1.model.base import ModelItem from xtuner.v1.loss.ce_loss import CELossConfig -from xtuner.v1.model.moe.moe import BalancingLossConfig +from xtuner.v1.model.moe.moe import MOE_BLOCK_FORWARD, BalancingLossConfig @@ -35,9 +35,9 @@ class TestMoEEngineFloat8(DeterministicDDPTestCase): "device,ep_size,hsdp_sharding_size,sim_tol,rtol", [ ("cuda", 1, int(os.getenv("XTUNER_TEST_WORLD_SIZE", "8")), 0.01, 0.01), - # ep8 is a smoke/trend coverage for the FSDP shard-mesh-size-1 FP8 path. - # It shares the ep1 reference below, but is not expected to align step-by-step - # because EP changes routing/collective order and accumulates FP8 numeric drift. + # EP8 covers checkpoint replay across layer-varying routed-token shapes while MoEBlock + # remains fullgraph-compiled. It shares the EP1 reference below, but EP changes routing + # and collective order, so the two loss curves need not align step-by-step. # Observed 10-step loss: # [2.4714, 2.4714, 1.8044, 1.5210, 0.9570, 0.6952, 0.4370, 0.3123, 0.1714, 0.1100] ("cuda", 8, int(os.getenv("XTUNER_TEST_WORLD_SIZE", "8")), 0.01, 0.15), @@ -66,6 +66,9 @@ def test_tile_wise_fp8(self, device, ep_size, hsdp_sharding_size, sim_tol, rtol) optim_cfg=optim_cfg, fsdp_cfg=fsdp_cfg, ) + if ep_size > 1: + # Checkpoint replay must remain correct while the EP expert block stays fullgraph. + self.assertEqual(engine.model.compile_cfg.get(MOE_BLOCK_FORWARD), {"fullgraph": True}) engine.from_hf(hf_path=QWEN3_MOE_PATH) loss_cfg = CELossConfig() diff --git a/tests/model/test_fsdp_checkpoint.py b/tests/model/test_fsdp_checkpoint.py new file mode 100644 index 0000000000..a2ebf5e0fe --- /dev/null +++ b/tests/model/test_fsdp_checkpoint.py @@ -0,0 +1,174 @@ +import torch + +from xtuner._testing import DeterministicDDPTestCase +from xtuner.v1.config import FSDPConfig +from xtuner.v1.data_proto import SequenceContext +from xtuner.v1.model.dense.qwen3 import Qwen3DenseConfig +from xtuner.v1.module.attention import MHAConfig +from xtuner.v1.utils.compile import is_compiled_function + + +class TestFSDPCheckpoint(DeterministicDDPTestCase): + @property + def world_size(self) -> int: + return 2 + + def test_reentrant_checkpoint_keeps_fsdp_outside_recompute(self): + self.create_pg("cuda") + config = Qwen3DenseConfig( + vocab_size=64, + max_position_embeddings=64, + eos_token_id=2, + bos_token_id=1, + num_hidden_layers=2, + hidden_size=32, + intermediate_size=64, + rms_norm_eps=1e-6, + hidden_act="silu", + attention=MHAConfig( + num_attention_heads=4, + num_key_value_heads=2, + head_dim=8, + qk_norm=True, + ), + compile_cfg=False, + ) + model = config.build().cuda() + grad_modes: list[bool] = [] + original_layer = model.layers["0"] + original_layer.register_forward_pre_hook(lambda _module, _inputs: grad_modes.append(torch.is_grad_enabled())) + + model.fully_shard( + FSDPConfig( + param_dtype=torch.bfloat16, + reduce_dtype=torch.bfloat16, + torch_compile=False, + ) + ) + checkpoint_calls = 0 + + def record_checkpoint_call(_module, _inputs, _output): + nonlocal checkpoint_calls + checkpoint_calls += 1 + + model.layers["0"].register_forward_hook(record_checkpoint_call) + input_ids = torch.randint(0, config.vocab_size, (1, 8), device="cuda") + output = model(SequenceContext.from_input_ids((input_ids,))) + assert output.logits is not None + output.logits.sum().backward() + + # The original layer and its lifecycle hooks must be replayed, while + # the outer FSDP/checkpoint boundary is one logical forward only. + assert grad_modes == [False, True] + assert checkpoint_calls == 1 + + def test_qwen3_vl_checkpoint_compile_allows_pytree_boundary(self): + self.create_pg("cuda") + from xtuner.v1.model.compose.qwen3_vl.qwen3_vl_config import Qwen3VLVisionConfig + + compile_target = "xtuner.v1.model.compose.qwen3_vl.modeling_vision.Qwen3VLVisionLayer.forward" + config = Qwen3VLVisionConfig( + depth=1, + hidden_size=32, + intermediate_size=64, + num_attention_heads=4, + patch_size=2, + temporal_patch_size=1, + spatial_merge_size=1, + num_position_embeddings=4, + deepstack_visual_indexes=[], + attn_impl="eager_attention", + compile_cfg={compile_target: {"fullgraph": True}}, + ) + model = config.build().cuda() + model.fully_shard(FSDPConfig(vision_recompute_ratio=1.0)) + assert is_compiled_function(model.blocks[0].forward) + + hidden_states = torch.randn(4, 32, device="cuda", dtype=torch.bfloat16, requires_grad=True) + cu_seqlens = torch.tensor([0, 4], device="cuda", dtype=torch.int32) + cos = torch.ones(4, 8, device="cuda", dtype=torch.bfloat16) + sin = torch.zeros_like(cos) + output = model.blocks[0](hidden_states, cu_seqlens, 4, (cos, sin)) + output.square().sum().backward() + + assert hidden_states.grad is not None + assert torch.isfinite(hidden_states.grad).all() + + def test_intern_s1_checkpoint_compile_allows_pytree_boundary(self): + self.create_pg("cuda") + from xtuner.v1.model.compose.intern_s1.intern_s1_config import InternS1VisionConfig + + compile_target = "xtuner.v1.model.compose.intern_s1.modeling_vision.InternS1VisionLayer.forward" + config = InternS1VisionConfig( + image_size=(4, 4), + patch_size=(2, 2), + num_hidden_layers=1, + hidden_size=32, + intermediate_size=64, + num_attention_heads=4, + attn_impl="eager_attention", + compile_cfg={compile_target: {"fullgraph": True}}, + ) + model = config.build().cuda() + model.fully_shard(FSDPConfig(vision_recompute_ratio=1.0)) + assert is_compiled_function(model.encoder.layer[0].forward) + + hidden_states = torch.randn(1, 4, 32, device="cuda", dtype=torch.bfloat16, requires_grad=True) + output = model.encoder.layer[0](hidden_states) + output.square().sum().backward() + + assert hidden_states.grad is not None + assert torch.isfinite(hidden_states.grad).all() + + def test_mixed_dense_checkpoint_compile_allows_pytree_boundary(self): + self.create_pg("cuda") + from xtuner.v1.loss.ce_loss import CELossConfig + from xtuner.v1.model.dense.qwen3_5_text import Qwen3_5_VLTextDenseConfig + from xtuner.v1.module.attention import GatedDeltaNetConfig + + config = Qwen3_5_VLTextDenseConfig( + vocab_size=64, + max_position_embeddings=64, + eos_token_id=2, + num_hidden_layers=4, + hidden_size=128, + intermediate_size=256, + rms_norm_eps=1e-6, + hidden_act="silu", + attention=MHAConfig( + with_gate=True, + num_attention_heads=4, + num_key_value_heads=2, + head_dim=32, + qk_norm=True, + rms_norm_eps=1e-6, + rms_norm_type="zero_centered", + ), + linear_attention=GatedDeltaNetConfig( + num_value_heads=4, + num_key_heads=4, + key_head_dim=16, + value_head_dim=16, + conv_kernel_dim=4, + hidden_act="silu", + rms_norm_eps=1e-6, + ), + ) + model = config.build().cuda() + model.fully_shard(FSDPConfig(recompute_ratio=1.0, torch_compile=True)) + assert all(is_compiled_function(layer.forward) for layer in model.layers.values()) + + input_ids = torch.randint(0, config.vocab_size, (1, 16), device="cuda") + seq_ctx = SequenceContext.from_input_ids((input_ids[:, :-1],)) + loss_config = CELossConfig(mode="eager") + loss_ctx = loss_config.build( + data={"shifted_labels": input_ids[:, 1:]}, + sp_mesh=None, + ) + loss_ctx = loss_config.loss_ctx_cls.build_batches([loss_ctx])[0] + output = model(seq_ctx, {"lm": loss_ctx}) + assert output.loss is not None + output.loss.backward() + + assert torch.isfinite(output.loss) + assert any(parameter.grad is not None for parameter in model.parameters()) diff --git a/tests/model/test_glm52_moe.py b/tests/model/test_glm52_moe.py index a09aeca5e0..b9b82800ee 100644 --- a/tests/model/test_glm52_moe.py +++ b/tests/model/test_glm52_moe.py @@ -9,6 +9,8 @@ TestGlm52RouterBias test_scratch_init_zeroes_main_and_mtp_biases: 从头初始化清零主干与 MTP router bias。 test_update_bias_handles_main_and_shared_mtp_loads: bias 更新覆盖主干并聚合共享 MTP 深度。 +TestGlm52ExplicitDsaDataflow + test_model_forward_backward_with_explicit_dsa_dataflow: 模型通过显式 IDs 完成前反向。 TestGlm52SequenceParallel test_mtp_loss_and_gradients_match_full_sequence: SP2 的 MTP loss 与梯度匹配完整序列。 """ @@ -29,7 +31,7 @@ from xtuner.v1.data_proto import SequenceContext from xtuner.v1.loss.ce_loss import CELossConfig from xtuner.v1.model import Glm52MoEConfig, get_model_config, get_model_config_from_hf -from xtuner.v1.module.attention import DSAMLAConfig +from xtuner.v1.model.moe.glm52 import DSAMLAConfig from xtuner.v1.module.mtp import MTPConfig from xtuner.v1.module.router.noaux_router import NoAuxRouterConfig from xtuner.v1.utils.test_utils import init_data_mesh @@ -273,6 +275,28 @@ def test_update_bias_handles_main_and_shared_mtp_loads(self): ) +@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA") +class TestGlm52ExplicitDsaDataflow: + def test_model_forward_backward_with_explicit_dsa_dataflow(self): + # 验证 GLM public forward/backward 经显式 DSA IDs 数据流产生有限 loss 和梯度。 + config = _tiny_glm52_config() + config.mtp_config = None + model = config.build().to(device="cuda", dtype=torch.bfloat16) + model.init_weights() + + input_ids = torch.tensor([[2, 3, 4, 5]], device="cuda") + shifted_labels = torch.tensor([[3, 4, 5, 6]], device="cuda") + seq_ctx = SequenceContext.from_input_ids((input_ids,), device="cuda") + data = {"seq_ctx": seq_ctx, "shifted_labels": shifted_labels} + loss_ctx = model.build_loss_ctx_batch([data], sp_mesh=None)[0] + + output = model(seq_ctx=seq_ctx, loss_ctx=loss_ctx) + output["loss"].backward() + + assert torch.isfinite(output["loss"]) + assert any(parameter.grad is not None for parameter in model.parameters()) + + @unittest.skipUnless(torch.cuda.device_count() >= 2, "requires 2 CUDA devices") class TestGlm52SequenceParallel(DeterministicDDPTestCase): def test_mtp_loss_and_gradients_match_full_sequence(self): diff --git a/tests/model/test_glm52_mtp_checkpoint_repro.py b/tests/model/test_glm52_mtp_checkpoint_repro.py index 87863fcfb3..17a8048afc 100644 --- a/tests/model/test_glm52_mtp_checkpoint_repro.py +++ b/tests/model/test_glm52_mtp_checkpoint_repro.py @@ -1,7 +1,8 @@ -"""GLM-5.2 MTP reentrant checkpoint 的真实训练回归测试。 +"""GLM-5.2 MTP checkpoint 的真实训练回归测试。 TestGlm52CompiledMTPCheckpoint - test_shared_mtp_depths_train_with_compile_and_topk_offload: 全 detach 的共享 MTP 可在 compile/offload 下训练。 + test_shared_mtp_depths_train_with_compile_and_topk_offload: 全 detach 的共享 MTP 可训练且 source 不重算。 + test_topk_offload_uses_pinned_memory_and_restores_ids: top-k IDs 真实 D2H 并在 backward 恢复。 TestGlm52MicroBatchMTPCheckpoint test_nested_micro_batch_inputs_preserve_gradients: EP2 micro2 的嵌套 embedding 梯度可正确反传。 """ @@ -20,10 +21,10 @@ from xtuner.v1.engine.train_engine import TrainEngine from xtuner.v1.loss.ce_loss import CELossConfig from xtuner.v1.model.base import ModelItem -from xtuner.v1.model.moe.glm52 import Glm52MoEConfig -from xtuner.v1.module.attention import DSAMLAConfig +from xtuner.v1.model.moe.glm52 import DSAMLAConfig, Glm52MoEConfig from xtuner.v1.module.mtp import MTPConfig from xtuner.v1.module.router.noaux_router import NoAuxRouterConfig +from xtuner.v1.utils.activation_offload import OffloadManager def _tiny_mtp_config( @@ -92,6 +93,7 @@ def _build_engine( compile_model: bool, detach_mtp_inputs: bool = False, detach_mtp_lm_head_weight: bool = False, + recompute_ratio: float = 0.0, ) -> TrainEngine: engine = TrainEngine( model_cfg=_tiny_mtp_config( @@ -105,7 +107,7 @@ def _build_engine( fsdp_cfg=FSDPConfig( ep_size=ep_size, cpu_offload=False, - recompute_ratio=0.0, + recompute_ratio=recompute_ratio, torch_compile=compile_model, ), intra_layer_micro_batch=intra_layer_micro_batch, @@ -125,7 +127,7 @@ def _model_item(engine: TrainEngine, start: int) -> ModelItem: @unittest.skipUnless(torch.cuda.is_available(), "requires CUDA") class TestGlm52CompiledMTPCheckpoint(DeterministicDDPTestCase): def test_shared_mtp_depths_train_with_compile_and_topk_offload(self): - # 验证全 detach 输入仍会触发默认 reentrant replay,并为共享 MTP 参数生成梯度。 + # 验证全 detach 的共享 MTP 在 compile/offload 下训练,且首个 logical depth 的 source 不重算。 self.create_pg("cuda") engine = _build_engine( intra_layer_micro_batch=1, @@ -135,6 +137,15 @@ def test_shared_mtp_depths_train_with_compile_and_topk_offload(self): detach_mtp_inputs=True, detach_mtp_lm_head_weight=True, ) + mtp_indexer_calls = 0 + + def count_mtp_indexer(*_args): + nonlocal mtp_indexer_calls + mtp_indexer_calls += 1 + + mtp_hook = engine.model.mtp_block.layers[0].decoder_layer.self_attn.indexer.register_forward_hook( + count_mtp_indexer + ) try: with mock.patch.dict( os.environ, @@ -155,7 +166,51 @@ def test_shared_mtp_depths_train_with_compile_and_topk_offload(self): local_grad = projection_grad.to_local() if isinstance(projection_grad, DTensor) else projection_grad assert torch.isfinite(local_grad).all() assert local_grad.norm() > 0 + assert mtp_indexer_calls == 1 + finally: + mtp_hook.remove() + del engine + torch.cuda.empty_cache() + + def test_topk_offload_uses_pinned_memory_and_restores_ids(self): + # 主层 checkpoint 会将显式 IDs 作为 SavedVariable;验证 D2H 真实命中并在 backward 恢复。 + self.create_pg("cuda") + engine = _build_engine( + intra_layer_micro_batch=1, + ep_size=1, + mtp_num_layers=1, + compile_model=False, + recompute_ratio=1.0, + ) + offload_manager = OffloadManager() + offload_manager.clear(group="text", clear_pin_memory_cache=True) + source_ids: list[torch.Tensor] = [] + + def record_source_ids(_module, _inputs, output): + source_ids.append(output["dsa_topk_ids"]) + + hook = engine.model.layers["0"].self_attn.indexer.register_forward_hook(record_source_ids) + try: + with mock.patch.dict( + os.environ, + {"XTUNER_ACTIVATION_OFFLOAD": "0", "XTUNER_DSA_TOPK_OFFLOAD": "1"}, + ): + step_info = engine.train_step([_model_item(engine, 2)]) + + assert math.isfinite(step_info["total_loss"]) + assert len(source_ids) == 1 + torch.cuda.synchronize() + expected_ids = source_ids[0].detach().cpu() + pinned_ids = [ + tensor + for key, tensor in offload_manager.pin_memory_cache.items() + if key.startswith("text_") and tensor.dtype == torch.int32 + ] + assert pinned_ids and all(tensor.is_pinned() for tensor in pinned_ids) + assert any(torch.equal(tensor, expected_ids) for tensor in pinned_ids) finally: + hook.remove() + offload_manager.clear(group="text", clear_pin_memory_cache=True) del engine torch.cuda.empty_cache() @@ -175,6 +230,15 @@ def test_nested_micro_batch_inputs_preserve_gradients(self): mtp_num_layers=1, compile_model=False, ) + mtp_indexer_calls = 0 + + def count_mtp_indexer(*_args): + nonlocal mtp_indexer_calls + mtp_indexer_calls += 1 + + hook = engine.model.mtp_block.layers[0].decoder_layer.self_attn.indexer.register_forward_hook( + count_mtp_indexer + ) try: with mock.patch.dict( os.environ, @@ -189,7 +253,9 @@ def test_nested_micro_batch_inputs_preserve_gradients(self): assert math.isfinite(step_info["total_loss"]) assert math.isfinite(step_info["logs_info"]["reduced_mtp_loss"]) + assert mtp_indexer_calls == 2 finally: + hook.remove() del engine torch.cuda.empty_cache() diff --git a/tests/model/test_qwen3_5_dense.py b/tests/model/test_qwen3_5_dense.py index 181fadb0ab..4be05fe156 100644 --- a/tests/model/test_qwen3_5_dense.py +++ b/tests/model/test_qwen3_5_dense.py @@ -99,7 +99,7 @@ def test_decoder_layer_bitwise_parity(self, device, layer_idx): loss_hf.backward() x_xt = base.clone().requires_grad_(True) - o_xt = xt_layer(x_xt, position_embeddings=(cos, sin), seq_ctx=seq_ctx) + o_xt = xt_layer(x_xt, position_embeddings=(cos, sin), seq_ctx=seq_ctx)["hidden_states"] loss_xt = F.cross_entropy(F.linear(model.norm(o_xt), model.lm_head.weight).reshape(-1, cfg.vocab_size), labels) loss_xt.backward() diff --git a/tests/model/test_recompute.py b/tests/model/test_recompute.py new file mode 100644 index 0000000000..84636bada0 --- /dev/null +++ b/tests/model/test_recompute.py @@ -0,0 +1,263 @@ +"""Activation checkpoint public-behavior regression tests.""" + +import pytest +import torch +from torch import nn +from torch.autograd.graph import saved_tensors_hooks + +from xtuner.v1.model.utils import apply_activation_checkpointing, reuse_during_recompute +from xtuner.v1.utils import clean_param_name + + +class _KeywordOnlyBlock(nn.Module): + """A forward shape that requires pytree adaptation with reentrant checkpointing. + + Tensors arrive nested in a dict and behind a keyword-only argument, and the result is returned + as a dict rather than a tensor or a tuple of tensors. + """ + + def __init__(self) -> None: + super().__init__() + self.linear = nn.Linear(4, 4) + self.tag = "block" + + def forward(self, inputs: dict[str, torch.Tensor], *, scale: float) -> dict[str, torch.Tensor]: + return {"out": self.linear(inputs["x"]) * scale} + + +class _FlexibleBlock(nn.Module): + """接受任意摆放的输入:位置的容器、字典、关键字参数,用来覆盖各种嵌套形状。""" + + def __init__(self) -> None: + super().__init__() + # 输入 4 维、输出 6 维:输出与输入形状不同,断言才不会把输出误当成输入。 + self.linear = nn.Linear(4, 6) + + def forward(self, inputs, *, scale: float, extra: torch.Tensor | None = None) -> dict[str, torch.Tensor]: + tensors = list(inputs.values()) if isinstance(inputs, dict) else list(inputs) + if extra is not None: + tensors.append(extra) + return {"out": sum(self.linear(t) * scale for t in tensors)} + + +class _GradModeBlock(nn.Module): + def __init__(self) -> None: + super().__init__() + self.linear = nn.Linear(4, 4) + self.grad_modes: list[bool] = [] + + def forward(self, x: torch.Tensor) -> torch.Tensor: + self.grad_modes.append(torch.is_grad_enabled()) + return self.linear(x) + + +class _ParameterOnlyBlock(nn.Module): + def __init__(self) -> None: + super().__init__() + self.weight = nn.Parameter(torch.tensor([2.0])) + + def forward(self, inputs: torch.Tensor) -> torch.Tensor: + return self.weight * inputs + + +class _ReservedNamesBlock(nn.Module): + def forward( + self, + x: torch.Tensor, + *, + function: float, + preserve_rng_state: float, + context_fn: float, + ) -> dict[str, torch.Tensor]: + return {"out": x * function * preserve_rng_state * context_fn} + + +class _ChangingOutputBlock(nn.Module): + def forward(self, x: torch.Tensor) -> dict[str, torch.Tensor] | list[torch.Tensor]: + output = x.square() + return [output] if torch.is_grad_enabled() else {"out": output} + + +class _ReusableMetadata(nn.Module): + def __init__(self) -> None: + super().__init__() + self.calls: list[tuple[float, int]] = [] + + def forward(self, x: torch.Tensor, *, offset: int) -> dict[str, torch.Tensor]: + self.calls.append((x.item(), offset)) + return {"ids": (x.detach() * 10 + offset).to(torch.int64)} + + +class _ReuseTwiceBlock(nn.Module): + def __init__(self, metadata: _ReusableMetadata) -> None: + super().__init__() + self.metadata = metadata + + def forward(self, x: torch.Tensor) -> dict[str, torch.Tensor]: + first = reuse_during_recompute(self.metadata, x, offset=1)["ids"] + second = reuse_during_recompute(self.metadata, x, offset=2)["ids"] + return {"out": x * (first + 2 * second)} + + +class TestCheckpointWrapper: + def test_wrapper_is_transparent_to_state_dict_and_attributes(self): + # 包裹层不能出现在参数名里,否则 checkpoint 的存/取与非重算模型不兼容。 + plain = _KeywordOnlyBlock() + wrapped = apply_activation_checkpointing(_KeywordOnlyBlock()) + wrapped.load_state_dict(plain.state_dict()) + + assert sorted(wrapped.state_dict()) == sorted(plain.state_dict()) + assert sorted(name for name, _ in wrapped.named_parameters()) == sorted( + name for name, _ in plain.named_parameters() + ) + assert torch.equal(wrapped.state_dict()["linear.weight"], plain.state_dict()["linear.weight"]) + assert wrapped.tag == "block" + + def test_checkpoint_is_fixed_reentrant(self): + wrapped = apply_activation_checkpointing(_GradModeBlock()) + wrapped(torch.randn(2, 4, requires_grad=True)).sum().backward() + + # Reentrant checkpoint runs the original pass without a graph, then replays it with grad. + assert wrapped.grad_modes == [False, True] + + def test_detached_inputs_still_update_module_parameters(self): + # MTP may detach every backbone input, but replay must still produce module gradients. + wrapped = apply_activation_checkpointing(_ParameterOnlyBlock()) + detached_inputs = torch.tensor([3.0]) + + wrapped(detached_inputs).sum().backward() + + assert wrapped.weight.grad is not None + torch.testing.assert_close(wrapped.weight.grad, torch.tensor([3.0])) + assert detached_inputs.grad is None + + def test_frozen_module_with_detached_inputs_does_not_force_replay(self): + # A frozen vision block has neither grad inputs nor trainable parameters. Its output is a + # constant for downstream training, so checkpoint must not create an unusable replay edge. + wrapped = apply_activation_checkpointing(_ParameterOnlyBlock().requires_grad_(False)) + detached_inputs = torch.tensor([3.0]) + downstream_weight = nn.Parameter(torch.tensor([4.0])) + + output = wrapped(detached_inputs) + (output * downstream_weight).sum().backward() + + assert wrapped.weight.grad is None + torch.testing.assert_close(downstream_weight.grad, output.detach()) + + def test_checkpoint_option_names_are_forwarded_to_the_module(self): + x = torch.tensor(2.0, requires_grad=True) + wrapped = apply_activation_checkpointing(_ReservedNamesBlock()) + + output = wrapped( + x, + function=3.0, + preserve_rng_state=5.0, + context_fn=7.0, + )["out"] + output.backward() + + assert output.item() == 210.0 + assert x.grad.item() == 105.0 + + def test_preserve_rng_state_replays_the_same_dropout_mask(self): + wrapped = apply_activation_checkpointing(nn.Dropout(p=0.5), preserve_rng_state=True) + x = torch.ones(32, requires_grad=True) + + torch.manual_seed(0) + output = wrapped(x) + output.sum().backward() + + torch.testing.assert_close(x.grad, output) + + def test_non_tensor_signature_preserves_gradients(self): + # 非 tensor 签名下梯度必须与不重算完全一致。 + torch.manual_seed(0) + plain = _KeywordOnlyBlock() + wrapped = apply_activation_checkpointing(_KeywordOnlyBlock()) + wrapped.load_state_dict(plain.state_dict()) + + x = torch.randn(2, 4, requires_grad=True) + plain({"x": x}, scale=2.0)["out"].square().sum().backward() + baseline_input_grad, x.grad = x.grad.clone(), None + + wrapped({"x": x}, scale=2.0)["out"].square().sum().backward() + + assert torch.equal(x.grad, baseline_input_grad) + assert torch.equal(wrapped.linear.weight.grad, plain.linear.weight.grad) + + def test_root_parameter_names_can_be_normalized(self): + root = nn.Module() + root.block = apply_activation_checkpointing(_KeywordOnlyBlock()) + + names = {clean_param_name(name) for name, _ in root.named_parameters()} + + assert names == {"block.linear.weight", "block.linear.bias"} + + def test_output_structure_must_match_during_replay(self): + wrapped = apply_activation_checkpointing(_ChangingOutputBlock()) + x = torch.randn(2, requires_grad=True) + + with pytest.raises(RuntimeError, match="different output PyTree structure"): + wrapped(x)["out"].sum().backward() + + @pytest.mark.parametrize( + "make_call", + [ + pytest.param(lambda block, x: block([x], scale=2.0), id="nested-in-list"), + pytest.param(lambda block, x: block({"x": x}, scale=2.0), id="nested-in-dict"), + pytest.param(lambda block, x: block([], scale=2.0, extra=x), id="passed-by-keyword"), + ], + ) + def test_input_tensors_reach_the_ambient_saved_tensor_hooks(self, make_call): + # 激活 offload 是靠外层 saved_tensors_hooks 拿到层输入的,而 checkpoint 只把**顶层** + # tensor 参数包成 SavedVariable(构造它才会触发 hook)。所以嵌套在容器里、或走关键字 + # 传进来的 tensor 会一个 hook 都不经过——offload 静默空转,梯度却完全正确,没有任何 + # 现象能暴露它。这里直接断言 hook 收得到。 + packed: list[int] = [] + + class _Record(saved_tensors_hooks): + # 按 data_ptr 认张量,不按 shape:区域的输出很容易和输入同形, + # 按 shape 断言会把输出当成输入,测试变成恒绿。 + def __init__(self) -> None: + super().__init__(lambda t: (packed.append(t.data_ptr()), t)[1], lambda t: t) + + wrapped = apply_activation_checkpointing(_FlexibleBlock()) + x = torch.randn(2, 4, requires_grad=True) + + with _Record(): + make_call(wrapped, x)["out"].square().sum().backward() + + assert x.data_ptr() in packed + + +class TestReuseDuringRecompute: + def test_outside_checkpoint_calls_the_callable_each_time(self): + metadata = _ReusableMetadata() + x = torch.tensor(1.0) + + reuse_during_recompute(metadata, x, offset=1) + reuse_during_recompute(metadata, x, offset=2) + + assert metadata.calls == [(1.0, 1), (1.0, 2)] + + def test_fifo_calls_are_isolated_between_checkpoint_invocations(self): + # The two graphs replay in reverse forward order. Each graph must retain + # its own FIFO even though both use the same callable object twice. + metadata = _ReusableMetadata() + block = apply_activation_checkpointing(_ReuseTwiceBlock(metadata)) + first = torch.tensor(1.0, requires_grad=True) + second = torch.tensor(2.0, requires_grad=True) + + first_output = block(first)["out"] + second_output = block(second)["out"] + second_output.backward() + first_output.backward() + + assert metadata.calls == [ + (1.0, 1), + (1.0, 2), + (2.0, 1), + (2.0, 2), + ] + assert first.grad.item() == 35.0 + assert second.grad.item() == 65.0 diff --git a/tests/module/attention/test_dsa_mla.py b/tests/module/attention/test_dsa_mla.py index 3f3fde1b79..94a1ee57a4 100644 --- a/tests/module/attention/test_dsa_mla.py +++ b/tests/module/attention/test_dsa_mla.py @@ -4,8 +4,9 @@ test_padded_indices_support_int32_and_backward: PyTorch 后端处理 padding、int32 和反向传播。 TestDSAAttention test_packed_inputs_respect_causal_boundaries_and_backward: packed attention 遵守分段因果边界并可反传。 - test_shared_layers_reuse_topk_without_cross_context_leak: shared layer 复用当前样本 top-k 且不跨样本泄漏。 - test_reentrant_checkpoint_reuses_and_releases_topk: checkpoint 重算复用并最终释放 top-k。 + test_indexer_freeze_contract: indexer 默认冻结,未实现的可训练模式在 build 时拒绝。 + test_shared_layer_consumes_explicit_topk_ids: shared layer 复用显式 top-k IDs,漏传时立即报错。 + test_checkpoint_reuses_source_topk_storage: 显式 IDs 穿过 checkpoint,source indexer 不重算。 TestAcceleratedSparseMLA test_tilelang_forward_backward_matches_torch: TileLang 前反向数值与 PyTorch 后端一致。 test_compiled_cudnn_backward_matches_tilelang: 编译后的 cuDNN DSA 前反向与 TileLang 一致。 @@ -23,14 +24,12 @@ 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.model.moe.glm52 import DSAMLAConfig +from xtuner.v1.model.moe.glm52.decoder_layer import GLM52DenseDecoderLayer +from xtuner.v1.model.utils import apply_activation_checkpointing 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) -> GLM52DenseDecoderLayer: + return GLM52DenseDecoderLayer( + hidden_size=4, + intermediate_size=8, + hidden_act="silu", + attention_config=_tiny_dsa_config(indexer_types), + layer_idx=layer_idx, + ) class TestTorchSparseMLA: @@ -185,15 +179,27 @@ 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_indexer_freeze_contract(self): + attention = _tiny_dsa_attention(indexer_types=["full"], layer_idx=0) + + assert all(not parameter.requires_grad for parameter in attention.indexer.parameters()) + + config = _tiny_dsa_config(["full"]) + config.freeze_dsa_indexer = False + with pytest.raises(ValueError, match="freeze_dsa_indexer=False"): + config.build(hidden_size=4, layer_idx=0) + + 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 +207,62 @@ 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_checkpoint_reuses_source_topk_storage(self): + # 验证显式 IDs 穿过 checkpoint 且真实 indexer backend 只执行一次。 torch.manual_seed(0) - source_block = checkpoint_wrapper( - _TinyDsaDecoderBlock(_tiny_dsa_attention(indexer_types=["full", "shared"], layer_idx=0)), - checkpoint_impl=CheckpointImpl.REENTRANT, + source_block = apply_activation_checkpointing( + _tiny_dsa_decoder(["full", "shared"], layer_idx=0) ) - shared_block = checkpoint_wrapper( - _TinyDsaDecoderBlock(_tiny_dsa_attention(indexer_types=["full", "shared"], layer_idx=1)), - checkpoint_impl=CheckpointImpl.REENTRANT, + shared_block = apply_activation_checkpointing( + _tiny_dsa_decoder(["full", "shared"], layer_idx=1) ) 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() + indexer_calls = 0 + + def count_indexer_call(*_args): + nonlocal indexer_calls + indexer_calls += 1 + + hook = source_block.self_attn.indexer.register_forward_hook(count_indexer_call) + + try: + source_outputs = source_block( + hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + ) + source_ids = source_outputs["dsa_topk_ids"] + shared_outputs = shared_block( + source_outputs["hidden_states"], + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + dsa_topk_ids=source_ids, + ) + shared_ids = shared_outputs["dsa_topk_ids"] + assert shared_ids.untyped_storage().data_ptr() == source_ids.untyped_storage().data_ptr() + shared_outputs["hidden_states"].square().mean().backward() + finally: + hook.remove() 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 + assert indexer_calls == 1 # The multiprocess cases must run before TileLang JIT is initialized in the @@ -257,12 +286,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) @@ -273,12 +303,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 032f8f32d3..1b02c6b3be 100644 --- a/tests/module/test_dense_decoder_layer.py +++ b/tests/module/test_dense_decoder_layer.py @@ -1,6 +1,6 @@ -"""DenseDecoderLayer 多 micro-batch 行为测试。 +"""GLM52DenseDecoderLayer 多 micro-batch 行为测试。 -TestDenseDecoderLayerMicroBatch +TestGLM52DenseDecoderLayerMicroBatch test_batched_inputs_match_independent_forwards: 等长 micro-batch 的输出与梯度等价于独立调用。 """ @@ -9,12 +9,12 @@ import torch from xtuner.v1.data_proto import SequenceContext -from xtuner.v1.module.attention import DSAMLAConfig -from xtuner.v1.module.decoder_layer.dense_decoder_layer import DenseDecoderLayer +from xtuner.v1.model.moe.glm52 import DSAMLAConfig +from xtuner.v1.model.moe.glm52.decoder_layer import GLM52DenseDecoderLayer -def _build_dense_dsa_layer() -> DenseDecoderLayer: - return DenseDecoderLayer( +def _build_dense_dsa_layer() -> GLM52DenseDecoderLayer: + return GLM52DenseDecoderLayer( hidden_size=4, intermediate_size=8, hidden_act="silu", @@ -55,7 +55,7 @@ def _build_inputs() -> tuple[ return hidden_states, position_embeddings, seq_ctx -class TestDenseDecoderLayerMicroBatch: +class TestGLM52DenseDecoderLayerMicroBatch: def test_batched_inputs_match_independent_forwards(self): # 验证一次多输入调用与逐 micro-batch 调用产生相同输出、输入梯度和参数梯度。 torch.manual_seed(0) @@ -69,11 +69,11 @@ def test_batched_inputs_match_independent_forwards(self): ] outputs = layer( - *hidden_states, + hidden_states, position_embeddings=position_embeddings, seq_ctx=seq_ctx, ) - reference_outputs = tuple( + reference_outputs = [ reference_layer( hidden, position_embeddings=position_embedding, @@ -84,14 +84,20 @@ def test_batched_inputs_match_independent_forwards(self): position_embeddings, reference_seq_ctx, ) - ) + ] - assert isinstance(outputs, tuple) - for output, reference_output in zip(outputs, reference_outputs): + output_hidden = outputs["hidden_states"] + output_ids = outputs["dsa_topk_ids"] + reference_hidden = tuple(result["hidden_states"] for result in reference_outputs) + reference_ids = tuple(result["dsa_topk_ids"] for result in reference_outputs) + for output, reference_output in zip(output_hidden, reference_hidden): torch.testing.assert_close(output, reference_output) + for dsa_topk_ids, reference_dsa_topk_ids in zip(output_ids, reference_ids): + torch.testing.assert_close(dsa_topk_ids, reference_dsa_topk_ids) + assert dsa_topk_ids.dtype == torch.int32 - sum(output.sum() for output in outputs).backward() - sum(output.sum() for output in reference_outputs).backward() + sum(output.sum() for output in output_hidden).backward() + sum(output.sum() for output in reference_hidden).backward() for hidden, reference_hidden in zip(hidden_states, reference_hidden_states): torch.testing.assert_close(hidden.grad, reference_hidden.grad) diff --git a/tests/utils/test_checkpoint_wrapper_checker.py b/tests/utils/test_checkpoint_wrapper_checker.py deleted file mode 100644 index 53cbf4da85..0000000000 --- a/tests/utils/test_checkpoint_wrapper_checker.py +++ /dev/null @@ -1,74 +0,0 @@ -from torch._prims_common import check -from xtuner.v1.model.utils import checkpoint_wrapper -import torch.nn as nn -import torch -import pytest - - -# Missing typehints -class ErrorDecoderLayer1(nn.Module): - def forward(self, x): - return x - - -# Inputs args missing raw tensor -class ErrorDecoderLayer2(nn.Module): - def forward(self, x: list[torch.Tensor], y: tuple[torch.Tensor], z: dict[str, torch.Tensor]) -> torch.Tensor: - ... - - -# Missing return type -class ErrorDecoderLayer3(nn.Module): - def forward(self, x: torch.Tensor, y: tuple[torch.Tensor], z: dict[str, torch.Tensor]): - ... - -# Missing raw tensor in return type -class ErrorDecoderLayer4(nn.Module): - def forward(self, x: torch.Tensor, y: tuple[torch.Tensor], z: dict[str, torch.Tensor]) -> tuple[list[torch.Tensor], int]: - ... - - -# return type must be a tuple -class ErrorDecoderLayer5(nn.Module): - def forward(self, x: torch.Tensor, y: tuple[torch.Tensor], z: dict[str, torch.Tensor]) -> list[torch.Tensor]: - ... - - -class DecoderLayer1(nn.Module): - def forward(self, x: torch.Tensor, y: tuple[torch.Tensor], z: dict[str, torch.Tensor]) -> torch.Tensor: - ... - - -class DecoderLayer2(nn.Module): - def forward(self, x: torch.Tensor, y: tuple[torch.Tensor], z: dict[str, torch.Tensor]) -> tuple[torch.Tensor, int]: - ... - - -class DecoderLayer3(nn.Module): - def forward( - self, x: torch.Tensor, y: tuple[torch.Tensor], z: dict[str, torch.Tensor] - ) -> tuple[torch.Tensor, int] | torch.Tensor: - ... - - -def test_checkpoint_wrapper_checker(): - with pytest.raises(TypeError): - checkpoint_wrapper(ErrorDecoderLayer1()) - - with pytest.raises(TypeError): - checkpoint_wrapper(ErrorDecoderLayer2()) - - with pytest.raises(TypeError): - checkpoint_wrapper(ErrorDecoderLayer3()) - - with pytest.raises(TypeError): - checkpoint_wrapper(ErrorDecoderLayer4()) - - with pytest.raises(TypeError): - checkpoint_wrapper(ErrorDecoderLayer5()) - - # Correct cases - checkpoint_wrapper(DecoderLayer1()) - checkpoint_wrapper(DecoderLayer2()) - checkpoint_wrapper(DecoderLayer3()) - diff --git a/tests/utils/test_pytree_reentrant_checkpoint.py b/tests/utils/test_pytree_reentrant_checkpoint.py deleted file mode 100644 index 63e83ead29..0000000000 --- a/tests/utils/test_pytree_reentrant_checkpoint.py +++ /dev/null @@ -1,65 +0,0 @@ -"""Pytree reentrant checkpoint 的梯度行为测试。 - -TestPytreeReentrantCheckpoint - test_nested_inputs_preserve_both_gradient_paths: 嵌套输入在 checkpoint 内外复用时梯度正确汇合。 - test_detached_inputs_still_update_module_parameters: 全 detached 输入仍可触发重算并更新模块参数。 -""" - -import torch -from torch import nn -from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import CheckpointImpl - -from xtuner.v1.model.utils import checkpoint_wrapper, pytree_reentrant_checkpoint - - -class NestedTensorBlock(nn.Module): - def forward(self, direct: torch.Tensor, nested: list[torch.Tensor]) -> torch.Tensor: - return direct * nested[0] - - -class ParameterOnlyBlock(nn.Module): - def __init__(self) -> None: - super().__init__() - self.weight = nn.Parameter(torch.tensor([2.0])) - - def forward(self, inputs: torch.Tensor) -> torch.Tensor: - return self.weight * inputs - - -class TestPytreeReentrantCheckpoint: - def test_nested_inputs_preserve_both_gradient_paths(self): - # 验证嵌套 Tensor 在 checkpoint 内外同时使用时不会重复反传旧 graph,且梯度正确相加。 - direct_source = torch.tensor([2.0], requires_grad=True) - nested_source = torch.tensor([5.0], requires_grad=True) - direct = direct_source * 2 - nested = nested_source * 3 - block = checkpoint_wrapper( - NestedTensorBlock(), - checkpoint_impl=CheckpointImpl.REENTRANT, - checkpoint_fn=pytree_reentrant_checkpoint, - ) - - loss = block(direct, nested=[nested]).sum() + nested.square().sum() - loss.backward() - - torch.testing.assert_close(direct_source.grad, torch.tensor([30.0])) - torch.testing.assert_close(nested_source.grad, torch.tensor([102.0])) - - def test_detached_inputs_still_update_module_parameters(self): - # MTP 会刻意 detach backbone 输入,但 checkpoint 仍须重算模块以生成其参数梯度。 - module = ParameterOnlyBlock() - block = checkpoint_wrapper( - module, - checkpoint_impl=CheckpointImpl.REENTRANT, - checkpoint_fn=pytree_reentrant_checkpoint, - ) - detached_inputs = torch.tensor([3.0]) - unrelated_parameter = nn.Parameter(torch.tensor([1.0])) - - loss = block(detached_inputs).sum() + unrelated_parameter.sum() * 0.0 - loss.backward() - - assert module.weight.grad is not None - torch.testing.assert_close(module.weight.grad, torch.tensor([3.0])) - assert detached_inputs.grad is None - torch.testing.assert_close(unrelated_parameter.grad, torch.tensor([0.0])) diff --git a/xtuner/v1/config/fsdp.py b/xtuner/v1/config/fsdp.py index 278295deec..7335d6d7a5 100644 --- a/xtuner/v1/config/fsdp.py +++ b/xtuner/v1/config/fsdp.py @@ -18,10 +18,6 @@ class FSDPConfig(BaseModel): recompute_ratio: Annotated[float, Parameter(help="Gradient checkpointing ratio for memory optimization")] = 1.0 vision_recompute_ratio: Annotated[float, Parameter(help="Recompute ratio for vision modules")] = 1.0 checkpoint_preserve_rng_state: Annotated[bool, Parameter(help="Preserve RNG state during checkpointing")] = True - mtp_checkpoint_use_reentrant: Annotated[ - bool, - Parameter(help="Use reentrant checkpointing for MTP layers"), - ] = True # Training-time FSDP CPU offload is version-sensitive for XTuner model configs # that keep selected fp32 trainable parameters outside FSDP via # fp32_keys_pattern. The Qwen3.5-VL MoE RL path was verified to run on Torch diff --git a/xtuner/v1/data_proto/__init__.py b/xtuner/v1/data_proto/__init__.py index 6194971cb2..c30af9de46 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 fa1829a91c..e17a1efad3 100644 --- a/xtuner/v1/data_proto/sequence_context.py +++ b/xtuner/v1/data_proto/sequence_context.py @@ -8,49 +8,6 @@ from .utils import gather_for_sequence_parallel, pad_to_multiple_of, split_for_sequence_parallel -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. - offload_slot: int # Stable offload slot among concurrently active microbatches. - mtp_forward_uses_remaining: dict[int, int] # Original-forward MTP uses left per shared source. - mtp_replays_remaining: dict[int, int] # Backward MTP replays left per shared source. - - 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, - offload_slot: int = 0, - 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.offload_slot = offload_slot - self.mtp_forward_uses_remaining = {} if mtp_forward_uses_remaining is None else mtp_forward_uses_remaining - self.mtp_replays_remaining = {} if mtp_replays_remaining is None else mtp_replays_remaining - - # 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 @@ -93,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 @@ -123,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, @@ -158,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 @@ -547,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), @@ -639,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/compose/intern_s1/modeling_vision.py b/xtuner/v1/model/compose/intern_s1/modeling_vision.py index 71d9cd50c0..48a5f1d629 100644 --- a/xtuner/v1/model/compose/intern_s1/modeling_vision.py +++ b/xtuner/v1/model/compose/intern_s1/modeling_vision.py @@ -34,8 +34,7 @@ fully_shard, ) from xtuner.v1.ops.attn_imp import attn_impl_mapping, AttnOpOutputs -from xtuner.v1.model.utils.checkpointing import checkpoint_wrapper -from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import CheckpointImpl +from xtuner.v1.model.utils.checkpointing import apply_activation_checkpointing from xtuner.v1.module import RMSNorm from xtuner.v1.ops.others import Dropout from xtuner.v1.ops.act_fn import get_act_fn @@ -408,11 +407,15 @@ def fully_shard( layer = self.encoder.layer[layer_idx] if layer_idx < num_recompute_layers: - layer = checkpoint_wrapper(layer, - preserve_rng_state=checkpoint_preserve_rng_state, - checkpoint_impl=CheckpointImpl.REENTRANT) + layer = apply_activation_checkpointing( + layer, + preserve_rng_state=checkpoint_preserve_rng_state, + ) if self.config.drop_path_rate == 0.0 and self.compile_cfg: - layer.forward = torch.compile(layer.forward, fullgraph=True) + # The PyTree checkpoint adapter is an intentional graph break; model compute + # inside it keeps the independently configured full-graph compilation. + compiled_forward = torch.compile(type(layer).forward, fullgraph=False) + layer.forward = compiled_forward.__get__(layer, type(layer)) self.encoder.layer[layer_idx] = layer diff --git a/xtuner/v1/model/compose/qwen3_vl/modeling_vision.py b/xtuner/v1/model/compose/qwen3_vl/modeling_vision.py index c61c1cb021..59d745559c 100644 --- a/xtuner/v1/model/compose/qwen3_vl/modeling_vision.py +++ b/xtuner/v1/model/compose/qwen3_vl/modeling_vision.py @@ -24,9 +24,8 @@ from torch.distributed.device_mesh import init_device_mesh import torch.distributed as dist from xtuner.v1.utils.compile import maybe_compile -from xtuner.v1.model.utils.checkpointing import checkpoint_wrapper +from xtuner.v1.model.utils.checkpointing import apply_activation_checkpointing from xtuner.v1.module import AttnOutputs -from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import CheckpointImpl from torch.distributed.device_mesh import DeviceMesh from tqdm import tqdm from xtuner.v1.ops.comm.all_to_all import ulysses_all_to_all @@ -343,11 +342,15 @@ def fully_shard( layer = self.blocks[layer_idx] if layer_idx < num_recompute_layers: - layer = checkpoint_wrapper(layer, - preserve_rng_state=checkpoint_preserve_rng_state, - checkpoint_impl=CheckpointImpl.REENTRANT) + layer = apply_activation_checkpointing( + layer, + preserve_rng_state=checkpoint_preserve_rng_state, + ) if self.compile_cfg: - layer.forward = torch.compile(layer.forward, fullgraph=True) + # The PyTree checkpoint adapter is an intentional graph break; model compute + # inside it keeps the independently configured full-graph compilation. + compiled_forward = torch.compile(type(layer).forward, fullgraph=False) + layer.forward = compiled_forward.__get__(layer, type(layer)) self.blocks[layer_idx] = layer diff --git a/xtuner/v1/model/dense/dense.py b/xtuner/v1/model/dense/dense.py index 47952f0c6d..519cf5d757 100644 --- a/xtuner/v1/model/dense/dense.py +++ b/xtuner/v1/model/dense/dense.py @@ -6,7 +6,6 @@ import torch.distributed as dist import torch.nn.functional as F from torch import nn -from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import CheckpointImpl from torch.distributed.device_mesh import DeviceMesh, init_device_mesh from torch.distributed.fsdp import ( CPUOffloadPolicy, @@ -27,7 +26,7 @@ TorchCompileOption, TransformerConfig, ) -from xtuner.v1.model.utils import checkpoint_wrapper +from xtuner.v1.model.utils import apply_activation_checkpointing from xtuner.v1.module import ( GatedDeltaNetConfig, LMHead, @@ -35,7 +34,7 @@ MLAConfig, RMSNorm, ) -from xtuner.v1.module.decoder_layer.dense_decoder_layer import DenseDecoderLayer +from xtuner.v1.module.decoder_layer.dense_decoder_layer import DenseDecoderLayer, DenseDecoderLayerOutput from xtuner.v1.utils import ( get_device, get_logger, @@ -98,11 +97,12 @@ def forward( self._mark_dynamic(seq_ctx) for idx, decoder_layer in self.layers.items(): - hidden_states = decoder_layer( + layer_results: DenseDecoderLayerOutput = decoder_layer( hidden_states, position_embeddings=position_embeddings, seq_ctx=seq_ctx, ) + hidden_states = layer_results["hidden_states"] if self.config.return_hidden_states: output["hidden_states"].append(hidden_states) @@ -235,18 +235,16 @@ def fully_shard( layer = self.layers[str(int(layer_idx))] layer_idx = int(layer_idx) if layer_idx < num_recompute_layers: - layer = checkpoint_wrapper( - layer, preserve_rng_state=checkpoint_preserve_rng_state, checkpoint_impl=CheckpointImpl.REENTRANT + layer = apply_activation_checkpointing( + layer, + preserve_rng_state=checkpoint_preserve_rng_state, ) - # __class__ without self attribute - - # Linear-attention (GatedDeltaNet) layers write ``seq_ctx.seq_idx`` inside the - # checkpoint region; compiling the wrapped layer with ``fullgraph=True`` turns the - # checkpoint into a HigherOrderOperator that rejects that side effect. Such layers are - # still compiled, but with ``fullgraph=False`` so the write can graph-break. if self.compile_cfg: - fullgraph = self.config.layers_type[layer_idx] != "linear_attention" - layer.forward = torch.compile(layer.forward, fullgraph=fullgraph) + # Every checkpointed layer crosses the eager PyTree adapter, so the wrapper + # must allow a graph break regardless of its attention type. Model compute + # inside it keeps the independently configured full-graph compilation. + compiled_forward = torch.compile(type(layer).forward, fullgraph=False) + layer.forward = compiled_forward.__get__(layer, type(layer)) self.layers[str(layer_idx)] = layer self._fully_shard( diff --git a/xtuner/v1/model/dense/qwen3vl_text.py b/xtuner/v1/model/dense/qwen3vl_text.py index f41b760ca6..b10f6ef57f 100644 --- a/xtuner/v1/model/dense/qwen3vl_text.py +++ b/xtuner/v1/model/dense/qwen3vl_text.py @@ -6,6 +6,7 @@ from xtuner.v1.data_proto import SequenceContext from xtuner.v1.loss import BaseLossContext from xtuner.v1.model.base import ModelOutputs +from xtuner.v1.module.decoder_layer.dense_decoder_layer import DenseDecoderLayerOutput from .qwen3 import Qwen3Dense, Qwen3Dense4BConfig, Qwen3Dense8BConfig @@ -64,11 +65,12 @@ def forward( # type: ignore[override] # ===================================================== for idx, decoder_layer in self.layers.items(): - hidden_states = decoder_layer( + layer_results: DenseDecoderLayerOutput = decoder_layer( hidden_states, position_embeddings=position_embeddings, seq_ctx=seq_ctx, ) + hidden_states = layer_results["hidden_states"] if deepstack_visual_embeds is not None and ((idx := int(idx)) in range(len(deepstack_visual_embeds))): assert visual_pos_masks is not None diff --git a/xtuner/v1/model/moe/glm52/__init__.py b/xtuner/v1/model/moe/glm52/__init__.py new file mode 100644 index 0000000000..1e1c3c668a --- /dev/null +++ b/xtuner/v1/model/moe/glm52/__init__.py @@ -0,0 +1,10 @@ +from .dsa_mla import DSAMLAConfig, DSAMultiLatentAttention +from .glm52 import Glm52MoE, Glm52MoEConfig + + +__all__ = [ + "DSAMLAConfig", + "DSAMultiLatentAttention", + "Glm52MoE", + "Glm52MoEConfig", +] diff --git a/xtuner/v1/model/moe/glm52/decoder_layer.py b/xtuner/v1/model/moe/glm52/decoder_layer.py new file mode 100644 index 0000000000..dcafe056c8 --- /dev/null +++ b/xtuner/v1/model/moe/glm52/decoder_layer.py @@ -0,0 +1,203 @@ +from typing import cast + +import torch +from typing_extensions import override + +from xtuner.v1.data_proto import SequenceContext +from xtuner.v1.module import AttnOutputs, RouterResults +from xtuner.v1.module.decoder_layer.dense_decoder_layer import ( + DenseDecoderLayer, + DenseDecoderLayerMicroBatchOutput, + DenseDecoderLayerOutput, +) +from xtuner.v1.module.decoder_layer.moe_decoder_layer import ( + MoEDecoderLayer, + MoEDecoderLayerMicroBatchOutput, + MoEDecoderLayerOutput, +) + +from .dsa_mla import DSAMultiLatentAttention, GLM52AttnOutputs + + +class GLM52DenseDecoderLayerOutput(DenseDecoderLayerOutput): + """GLM-5.2 dense-layer output with explicit DSA IDs.""" + + dsa_topk_ids: torch.Tensor + + +class GLM52DenseDecoderLayerMicroBatchOutput(DenseDecoderLayerMicroBatchOutput): + """GLM-5.2 dense-layer outputs for intra-layer micro-batches.""" + + dsa_topk_ids: list[torch.Tensor] + + +class GLM52DenseDecoderLayer(DenseDecoderLayer): + """Dense decoder layer that threads GLM-5.2 DSA IDs explicitly.""" + + @override + def forward( + self, + hidden_states: torch.Tensor | list[torch.Tensor], + *, + position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], + seq_ctx: SequenceContext | list[SequenceContext], + dsa_topk_ids: torch.Tensor | list[torch.Tensor] | None = None, + ) -> GLM52DenseDecoderLayerOutput | GLM52DenseDecoderLayerMicroBatchOutput: + if not isinstance(hidden_states, list): + assert isinstance(position_embeddings, tuple) and len(position_embeddings) == 2 + assert isinstance(seq_ctx, SequenceContext) + assert dsa_topk_ids is None or isinstance(dsa_topk_ids, torch.Tensor) + return self._glm52_forward( + hidden_states=hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + dsa_topk_ids=dsa_topk_ids, + ) + + n = len(hidden_states) + assert isinstance(position_embeddings, list) and len(position_embeddings) == n + assert isinstance(seq_ctx, list) and len(seq_ctx) == n + assert all(hidden.shape == hidden_states[0].shape for hidden in hidden_states) + if dsa_topk_ids is None: + dsa_topk_ids_list: list[torch.Tensor | None] = [None] * n + else: + assert isinstance(dsa_topk_ids, list) and len(dsa_topk_ids) == n + dsa_topk_ids_list = list(dsa_topk_ids) + + layer_results = [ + self._glm52_forward( + hidden_states=hidden, + position_embeddings=position_embedding, + seq_ctx=context, + dsa_topk_ids=topk_ids, + ) + for hidden, topk_ids, position_embedding, context in zip( + hidden_states, dsa_topk_ids_list, position_embeddings, seq_ctx + ) + ] + return { + "hidden_states": [result["hidden_states"] for result in layer_results], + "dsa_topk_ids": [result["dsa_topk_ids"] for result in layer_results], + } + + def _glm52_forward( + self, + *, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + seq_ctx: SequenceContext, + dsa_topk_ids: torch.Tensor | None, + ) -> GLM52DenseDecoderLayerOutput: + residual = hidden_states + hidden_states = self.input_layernorm(hidden_states) + + attention = cast(DSAMultiLatentAttention, self.self_attn) + attn_outputs = attention( + hidden_states=hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + dsa_topk_ids=dsa_topk_ids, + ) + hidden_states = residual + attn_outputs["projected_output"] + + residual = hidden_states + hidden_states = self.post_attention_layernorm(hidden_states) + hidden_states = self.mlp(hidden_states) + hidden_states = residual + hidden_states + + return { + "hidden_states": hidden_states, + "dsa_topk_ids": attn_outputs["dsa_topk_ids"], + } + + +class GLM52MoEDecoderLayerOutput(MoEDecoderLayerOutput): + """GLM-5.2 MoE-layer output with explicit DSA IDs.""" + + dsa_topk_ids: torch.Tensor + + +class GLM52MoEDecoderLayerMicroBatchOutput(MoEDecoderLayerMicroBatchOutput): + """GLM-5.2 MoE-layer outputs for intra-layer micro-batches.""" + + dsa_topk_ids: list[torch.Tensor] + + +class GLM52MoEDecoderLayer(MoEDecoderLayer): + """MoE decoder layer that threads GLM-5.2 DSA IDs explicitly.""" + + @override + def forward( + self, + hidden_states: torch.Tensor | list[torch.Tensor], + *, + seq_ctx: SequenceContext | list[SequenceContext], + position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], + dsa_topk_ids: torch.Tensor | list[torch.Tensor] | None = None, + ) -> GLM52MoEDecoderLayerOutput | GLM52MoEDecoderLayerMicroBatchOutput: + if not isinstance(hidden_states, list): + assert isinstance(seq_ctx, SequenceContext) + assert isinstance(position_embeddings, tuple) and len(position_embeddings) == 2 + assert dsa_topk_ids is None or isinstance(dsa_topk_ids, torch.Tensor) + return cast( + GLM52MoEDecoderLayerOutput, + self._forward( + hidden_states=hidden_states, + seq_ctx=seq_ctx, + position_embeddings=position_embeddings, + attention_kwargs={"dsa_topk_ids": dsa_topk_ids}, + ), + ) + + n = len(hidden_states) + assert isinstance(seq_ctx, list) and len(seq_ctx) == n + assert isinstance(position_embeddings, list) and len(position_embeddings) == n + if dsa_topk_ids is None: + dsa_topk_ids_list: list[torch.Tensor | None] = [None] * n + else: + assert isinstance(dsa_topk_ids, list) and len(dsa_topk_ids) == n + dsa_topk_ids_list = list(dsa_topk_ids) + + return cast( + GLM52MoEDecoderLayerMicroBatchOutput, + self._micro_batch_forward( + hidden_states_list=hidden_states, + seq_ctx_list=seq_ctx, + position_embeddings_list=position_embeddings, + attention_kwargs_list=[{"dsa_topk_ids": topk_ids} for topk_ids in dsa_topk_ids_list], + ), + ) + + @override + def _build_output( + self, + *, + hidden_states: torch.Tensor, + router_results: RouterResults, + attn_outputs: AttnOutputs, + ) -> GLM52MoEDecoderLayerOutput: + glm_attn_outputs = cast(GLM52AttnOutputs, attn_outputs) + return { + "hidden_states": hidden_states, + "router_logits": router_results["logits"], + "router_weights": router_results["router_weights"], + "router_topk_ids": router_results["topk_ids"], + "dsa_topk_ids": glm_attn_outputs["dsa_topk_ids"], + } + + @override + def _build_micro_batch_output( + self, + *, + hidden_states_list: list[torch.Tensor], + router_results_list: list[RouterResults], + attn_outputs_list: list[AttnOutputs], + ) -> GLM52MoEDecoderLayerMicroBatchOutput: + glm_attn_outputs = [cast(GLM52AttnOutputs, output) for output in attn_outputs_list] + return { + "hidden_states": hidden_states_list, + "router_logits": [result["logits"] for result in router_results_list], + "router_weights": [result["router_weights"] for result in router_results_list], + "router_topk_ids": [result["topk_ids"] for result in router_results_list], + "dsa_topk_ids": [output["dsa_topk_ids"] for output in glm_attn_outputs], + } diff --git a/xtuner/v1/module/attention/dsa_mla.py b/xtuner/v1/model/moe/glm52/dsa_mla.py similarity index 82% rename from xtuner/v1/module/attention/dsa_mla.py rename to xtuner/v1/model/moe/glm52/dsa_mla.py index 6a22b97dbc..d75e0dce5d 100644 --- a/xtuner/v1/module/attention/dsa_mla.py +++ b/xtuner/v1/model/moe/glm52/dsa_mla.py @@ -1,13 +1,18 @@ # Copyright (c) OpenMMLab. All rights reserved. -from typing import Literal, cast +from typing import Literal, TypedDict, cast import torch from torch import nn from torch.distributed.tensor import DTensor +from typing_extensions import NotRequired, overload from xtuner.v1.config import GenerateConfig from xtuner.v1.data_proto import SequenceContext from xtuner.v1.float8.config import Float8Config +from xtuner.v1.model.utils import reuse_during_recompute +from xtuner.v1.module.attention.attn_outputs import AttnOutputs +from xtuner.v1.module.attention.mla import MLAConfig, MultiLatentAttention, mla_apply_rotary_pos_emb +from xtuner.v1.module.linear import build_linear from xtuner.v1.module.rope import RopeScalingConfig from xtuner.v1.ops.comm import gather_for_sequence_parallel from xtuner.v1.ops.sparse_mla import ( @@ -19,10 +24,21 @@ get_sparse_mla, ) -from ..linear import build_linear -from .attn_outputs import AttnOutputs -from .dsa_topk_sharing import build_dsa_topk_release_plan, dsa_topk_source_layer, get_dsa_topk_sharing_runtime -from .mla import MLAConfig, MultiLatentAttention, mla_apply_rotary_pos_emb +from .dsa_topk_sharing import dsa_topk_source_layer + + +class GLM52AttnOutputs(AttnOutputs): + """GLM-5.2 attention outputs with explicit cross-layer DSA IDs.""" + + dsa_topk_ids: torch.Tensor + + +class DSAIndexerOutput(TypedDict): + """Frozen indexer result; logits are reserved for a future trainable + path.""" + + dsa_topk_ids: torch.Tensor + dsa_topk_logits: NotRequired[torch.Tensor] class LayerNorm(nn.Module): @@ -83,18 +99,14 @@ def __init__( self.k_norm = LayerNorm(index_head_dim, eps=1e-6) # weights_proj.weight: [index_n_heads, hidden_size] self.weights_proj = build_linear(hidden_size, index_n_heads, bias=False) - # The indexer only produces integer DSA top-k IDs under no_grad, so its - # parameters must not be registered with the training optimizer. - self.requires_grad_(False) - @torch.no_grad() def forward( self, hidden_states: torch.Tensor, q_resid: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], seq_ctx: SequenceContext, - ) -> torch.Tensor: + ) -> DSAIndexerOutput: """Compute DSA top-k indices for each local query token. Shapes use ``S`` for the local sequence length and ``S_g`` for the @@ -144,27 +156,19 @@ 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 数值。 # 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) - # returns topk_indices: [S, 1, K] - return self.dsa_topk_indices_func( - q, - k, - weights, - seq_ctx, - index_head_dim=self.index_head_dim, - index_topk=self.index_topk, + # The indexer contract owns dtype/layout so checkpoint reuse never needs + # to clone or convert the storage shared with downstream layers. + dsa_topk_ids = ( + self.dsa_topk_indices_func( + q, k, weights, seq_ctx, index_head_dim=self.index_head_dim, index_topk=self.index_topk + ) + .to(torch.int32) + .contiguous() ) + return {"dsa_topk_ids": dsa_topk_ids} class DSAMLAConfig(MLAConfig): @@ -176,6 +180,7 @@ class DSAMLAConfig(MLAConfig): indexer_rope_interleave: bool = True indexer_types: list[str] | None = None sparse_mla_backend: Literal["torch", "tilelang", "cudnn_dsa"] = "torch" + freeze_dsa_indexer: bool = True def build( self, @@ -186,6 +191,8 @@ def build( generate_config: GenerateConfig | None = None, float8_cfg: Float8Config | None = None, ) -> "DSAMultiLatentAttention": + if not self.freeze_dsa_indexer: + raise ValueError("freeze_dsa_indexer=False is not supported until the indexer has a differentiable output") if self.sparse_mla_backend in ("tilelang", "cudnn_dsa"): ensure_tilelang_runtime_available() if self.sparse_mla_backend == "cudnn_dsa": @@ -214,6 +221,7 @@ def __init__( indexer_rope_interleave: bool = True, indexer_types: list[str] | None = None, sparse_mla_backend: Literal["torch", "tilelang", "cudnn_dsa"] = "torch", + freeze_dsa_indexer: bool = True, **kwargs, ): super().__init__(**kwargs) @@ -240,19 +248,8 @@ def __init__( self.indexer_rope_interleave = indexer_rope_interleave self.indexer_types = indexer_types self.sparse_mla_backend = sparse_mla_backend + self.freeze_dsa_indexer = freeze_dsa_indexer 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.") @@ -275,6 +272,8 @@ def __init__( index_topk=self.index_topk, indexer_backend=self.sparse_mla_backend, ) + if self.freeze_dsa_indexer: + self.indexer.requires_grad_(False) def get_muon_split_sizes(self) -> dict[nn.Parameter, tuple[int, ...]]: """Return the logical row blocks used by GLM MuonSplit.""" @@ -291,7 +290,8 @@ def forward( hidden_states: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], seq_ctx: SequenceContext, - ) -> AttnOutputs: + dsa_topk_ids: torch.Tensor | None = None, + ) -> GLM52AttnOutputs: """Absorbed DSA-MLA forward for packed training (``bsz == 1``). Shapes use ``S`` for the local sequence length and ``S_g`` for the @@ -365,21 +365,35 @@ 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.") + if self.freeze_dsa_indexer: + with torch.no_grad(): + indexer_output = reuse_during_recompute( + self.indexer, + hidden_states, + q_resid, + position_embeddings, + seq_ctx, + ) + else: + indexer_output = self.indexer( + hidden_states, + q_resid, + position_embeddings, + seq_ctx, + ) + dsa_topk_ids = indexer_output["dsa_topk_ids"] + 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, ) @@ -396,4 +410,16 @@ def forward( "raw_output": raw_output, "projected_output": projected_output, "softmax_lse": softmax_lse, + "dsa_topk_ids": dsa_topk_ids, } + + @overload # type: ignore + def __call__( # type: ignore + self, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + seq_ctx: SequenceContext, + dsa_topk_ids: torch.Tensor | None = None, + ) -> GLM52AttnOutputs: ... + + __call__ = nn.Module.__call__ diff --git a/xtuner/v1/model/moe/glm52/dsa_topk_sharing.py b/xtuner/v1/model/moe/glm52/dsa_topk_sharing.py new file mode 100644 index 0000000000..baee969313 --- /dev/null +++ b/xtuner/v1/model/moe/glm52/dsa_topk_sharing.py @@ -0,0 +1,46 @@ +# Copyright (c) OpenMMLab. All rights reserved. + + +def dsa_topk_source_layer( + *, + layer_idx: int, + indexer_types: list[str] | None, + index_skip_topk_offset: int, + index_topk_freq: int, +) -> int: + """Resolve the GLM-5.2 source layer whose DSA top-k IDs a layer + consumes.""" + if indexer_types is not None: + if layer_idx < len(indexer_types) and indexer_types[layer_idx] == "full": + return layer_idx + for source_layer_idx in range(min(layer_idx, len(indexer_types) - 1), -1, -1): + if indexer_types[source_layer_idx] == "full": + return source_layer_idx + raise ValueError(f"DSA layer {layer_idx} has no preceding full indexer layer.") + + if index_topk_freq <= 1: + return layer_idx + + source_layer_idx = layer_idx + while (max(source_layer_idx + 1 - index_skip_topk_offset, 0) % index_topk_freq) != 0: + source_layer_idx -= 1 + return source_layer_idx + + +def dsa_topk_source_layers( + *, + num_layers: int, + indexer_types: list[str] | None, + index_skip_topk_offset: int, + index_topk_freq: int, +) -> tuple[int, ...]: + """Return the source-layer index for every layer in one decoder stack.""" + return tuple( + dsa_topk_source_layer( + layer_idx=layer_idx, + indexer_types=indexer_types, + index_skip_topk_offset=index_skip_topk_offset, + index_topk_freq=index_topk_freq, + ) + for layer_idx in range(num_layers) + ) diff --git a/xtuner/v1/model/moe/glm52.py b/xtuner/v1/model/moe/glm52/glm52.py similarity index 72% rename from xtuner/v1/model/moe/glm52.py rename to xtuner/v1/model/moe/glm52/glm52.py index 1c71eb4c54..fe1cf80a8e 100644 --- a/xtuner/v1/model/moe/glm52.py +++ b/xtuner/v1/model/moe/glm52/glm52.py @@ -1,6 +1,7 @@ +import os import re from pathlib import Path -from typing import Literal +from typing import Callable, Literal, cast import torch from pydantic import Field, computed_field @@ -11,50 +12,64 @@ from transformers.models.glm_moe_dsa import GlmMoeDsaConfig as HFGlmMoeDsaConfig except ImportError: HFGlmMoeDsaConfig = None # type: ignore[misc, assignment] + +from xtuner.v1.data_proto import SequenceContext from xtuner.v1.model.base import DEFAULT_FLOAT8_CFG, TorchCompileOption -from xtuner.v1.model.moe.moe import BalancingLossConfig, MoEConfig, ZLossConfig -from xtuner.v1.module.attention import DSAMLAConfig, DSAMultiLatentAttention -from xtuner.v1.module.attention.dsa_topk_sharing import ( - build_dsa_topk_release_plan, - configure_dsa_mtp_iteration_lifecycle, - configure_dsa_topk_decoder_lifecycle, - dsa_topk_source_layer, +from xtuner.v1.model.moe.moe import BalancingLossConfig, MoE, MoEConfig, ZLossConfig +from xtuner.v1.module.decoder_layer.dense_decoder_layer import ( + DenseDecoderLayerMicroBatchOutput, + DenseDecoderLayerOutput, +) +from xtuner.v1.module.decoder_layer.moe_decoder_layer import ( + MoEDecoderLayerMicroBatchOutput, + MoEDecoderLayerOutput, ) -from xtuner.v1.module.mtp import MTPConfig, MTPLayer +from xtuner.v1.module.mtp import MTPConfig from xtuner.v1.module.rope import RopeParametersConfig from xtuner.v1.module.router.noaux_router import NoAuxRouterConfig -from .moe import MoE +from .decoder_layer import ( + GLM52DenseDecoderLayer, + GLM52DenseDecoderLayerMicroBatchOutput, + GLM52DenseDecoderLayerOutput, + GLM52MoEDecoderLayer, + GLM52MoEDecoderLayerMicroBatchOutput, + GLM52MoEDecoderLayerOutput, +) +from .dsa_mla import DSAMLAConfig, DSAMultiLatentAttention +from .dsa_topk_sharing import dsa_topk_source_layer, dsa_topk_source_layers +from .mtp import GLM52MTPBlock, GLM52MTPLayer -# GLM DSA attention records cross-layer top-k indices in SequenceContext. -# That Python-side cache mutation is intentionally kept out of strict fullgraph -# regions, so decoder/pre-attn/DSA/dense boundaries allow graph breaks while -# pure tensor MoE expert sub-stages stay fullgraph. +# Keep the existing graph boundaries while explicit DSA top-k tensor inputs and +# results are validated. Each boundary can be tightened independently later. MOE_NON_EP_COMPILE_CFG: dict[str, TorchCompileOption] = { "xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEBlock.forward": TorchCompileOption(fullgraph=True), - "xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEDecoderLayer.forward": TorchCompileOption(fullgraph=False), + "xtuner.v1.model.moe.glm52.decoder_layer.GLM52MoEDecoderLayer.forward": TorchCompileOption(fullgraph=False), "xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEDecoderLayer._pre_moe_forward": TorchCompileOption( fullgraph=False ), - "xtuner.v1.module.attention.dsa_mla.DSAMultiLatentAttention.forward": TorchCompileOption(fullgraph=False), + "xtuner.v1.model.moe.glm52.dsa_mla.DSAMultiLatentAttention.forward": TorchCompileOption(fullgraph=False), "xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEDecoderLayer._shared_experts_forward": TorchCompileOption( fullgraph=True ), "xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEDecoderLayer._post_moe_forward": TorchCompileOption( fullgraph=True ), - "xtuner.v1.module.decoder_layer.dense_decoder_layer.DenseDecoderLayer.forward": TorchCompileOption( - fullgraph=False - ), + "xtuner.v1.model.moe.glm52.decoder_layer.GLM52DenseDecoderLayer.forward": TorchCompileOption(fullgraph=False), **DEFAULT_FLOAT8_CFG, } MOE_EP_COMPILE_CFG = MOE_NON_EP_COMPILE_CFG.copy() -MOE_EP_COMPILE_CFG.pop("xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEDecoderLayer.forward") +MOE_EP_COMPILE_CFG.pop("xtuner.v1.model.moe.glm52.decoder_layer.GLM52MoEDecoderLayer.forward") class Glm52MoE(MoE): + dense_decoder_layer_cls = GLM52DenseDecoderLayer + moe_decoder_layer_cls = GLM52MoEDecoderLayer + mtp_layer_cls = GLM52MTPLayer + mtp_block_cls = GLM52MTPBlock + @property @override def default_compile_cfg(self) -> dict[str, TorchCompileOption]: @@ -63,62 +78,121 @@ def default_compile_cfg(self) -> dict[str, TorchCompileOption]: return MOE_NON_EP_COMPILE_CFG @override - def _configure_model_specific_layer_lifecycle(self) -> None: - dsa_layers: list[tuple[torch.nn.Module, DSAMultiLatentAttention]] = [] - mtp_attention: DSAMultiLatentAttention | None = None + def _configure_model_specific_layers(self) -> None: + dsa_layers: list[DSAMultiLatentAttention] = [] for decoder_layer in self.layers.values(): self_attn = decoder_layer.self_attn # type: ignore[attr-defined] assert isinstance(self_attn, DSAMultiLatentAttention), ( f"GLM-5.2 requires DSAMultiLatentAttention, got {type(self_attn).__name__}." ) - dsa_layers.append((decoder_layer, self_attn)) + dsa_layers.append(self_attn) - num_physical_mtp_layers = 0 if self.mtp_block is not None and self.config.mtp_config is not None: num_physical_mtp_layers = 1 if self.config.mtp_config.share_weights else self.config.mtp_config.num_layers for mtp_idx in range(num_physical_mtp_layers): mtp_layer = self.mtp_block.layers[mtp_idx] - assert isinstance(mtp_layer, MTPLayer) - decoder_layer = mtp_layer.decoder_layer - self_attn = decoder_layer.self_attn # type: ignore[attr-defined] + assert isinstance(mtp_layer, GLM52MTPLayer) + self_attn = mtp_layer.decoder_layer.self_attn # type: ignore[attr-defined] assert isinstance(self_attn, DSAMultiLatentAttention), ( f"GLM-5.2 MTP requires DSAMultiLatentAttention, got {type(self_attn).__name__}." ) - dsa_layers.append((decoder_layer, self_attn)) - if mtp_idx == 0: - mtp_attention = self_attn - - sample_attn = dsa_layers[0][1] - release_plan = build_dsa_topk_release_plan( - num_main_layers=self.config.num_hidden_layers, - num_mtp_layers=num_physical_mtp_layers, + + sample_attn = dsa_layers[0] + self._dsa_topk_source_layers = dsa_topk_source_layers( + num_layers=self.config.num_hidden_layers, indexer_types=sample_attn.indexer_types, index_skip_topk_offset=sample_attn.index_skip_topk_offset, index_topk_freq=sample_attn.index_topk_freq, ) - for decoder_layer, self_attn in dsa_layers: - # DSA top-k sharing spans dense prefix, sparse MoE layers, and the - # optional MTP layer. The attention-local default release maps only - # see the main-stack indexer_types, so GLM-5.2 injects a model-level - # plan with the full physical layer topology. - configure_dsa_topk_decoder_lifecycle( - decoder_layer=decoder_layer, - attention=self_attn, - release_plan=release_plan, - ) + self._dsa_topk_last_consumers = frozenset( + layer_idx + for layer_idx, source_layer_idx in enumerate(self._dsa_topk_source_layers) + if layer_idx == self.config.num_hidden_layers - 1 + or self._dsa_topk_source_layers[layer_idx + 1] != source_layer_idx + ) - if ( - self.mtp_block is not None - and self.config.mtp_config is not None - and self.config.mtp_config.share_weights - and self.config.index_share_for_mtp_iteration - ): - assert mtp_attention is not None - configure_dsa_mtp_iteration_lifecycle( - mtp_block=self.mtp_block, - attention=mtp_attention, - num_iterations=self.config.mtp_config.num_layers, + @override + def _call_decoder_layer( + self, + *, + decoder_layer: torch.nn.Module, + layer_idx: int, + hidden_states: torch.Tensor | list[torch.Tensor], + position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], + seq_ctx: SequenceContext | list[SequenceContext], + previous_layer_results: ( + DenseDecoderLayerOutput + | DenseDecoderLayerMicroBatchOutput + | MoEDecoderLayerOutput + | MoEDecoderLayerMicroBatchOutput + | None + ), + ) -> ( + DenseDecoderLayerOutput + | DenseDecoderLayerMicroBatchOutput + | MoEDecoderLayerOutput + | MoEDecoderLayerMicroBatchOutput + ): + """Arrange GLM-5.2 DSA dataflow and offload around one decoder + layer.""" + is_micro_batch = isinstance(hidden_states, list) + is_source_layer = self._dsa_topk_source_layers[layer_idx] == layer_idx + if is_source_layer: + dsa_topk_ids: torch.Tensor | list[torch.Tensor] | None = None + else: + previous_results = cast( + GLM52DenseDecoderLayerOutput + | GLM52DenseDecoderLayerMicroBatchOutput + | GLM52MoEDecoderLayerOutput + | GLM52MoEDecoderLayerMicroBatchOutput, + previous_layer_results, + ) + dsa_topk_ids = previous_results["dsa_topk_ids"] + + activation_offload = int(os.getenv("XTUNER_ACTIVATION_OFFLOAD", "0")) == 1 + dsa_topk_offload = int(os.getenv("XTUNER_DSA_TOPK_OFFLOAD", "0")) == 1 + offload_tensors: list[torch.Tensor] = [] + if activation_offload and layer_idx >= self.config.first_k_dense_replace: + offload_tensors = list(hidden_states) if is_micro_batch else [cast(torch.Tensor, hidden_states)] + if dsa_topk_offload and layer_idx in self._dsa_topk_last_consumers and dsa_topk_ids is not None: + offload_tensors.extend(dsa_topk_ids if isinstance(dsa_topk_ids, list) else [dsa_topk_ids]) + + # The offload context expects a dense zero-based block index across every + # active activation or DSA-ID window, including dense GLM layers. + offload_block_idx = sum( + (activation_offload and previous_idx >= self.config.first_k_dense_replace) + or ( + dsa_topk_offload + and previous_idx in self._dsa_topk_last_consumers + and self._dsa_topk_source_layers[previous_idx] != previous_idx + ) + for previous_idx in (int(idx) for idx in self.layers) + if previous_idx < layer_idx + ) + decoder_forward = cast( + Callable[ + ..., + DenseDecoderLayerOutput + | DenseDecoderLayerMicroBatchOutput + | MoEDecoderLayerOutput + | MoEDecoderLayerMicroBatchOutput, + ], + decoder_layer, + ) + with self._saved_tensors_offload_ctx(offload_block_idx, offload_tensors): + layer_results = decoder_forward( + hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + dsa_topk_ids=dsa_topk_ids, ) + if layer_idx < self.config.first_k_dense_replace: + if is_micro_batch: + return cast(GLM52DenseDecoderLayerMicroBatchOutput, layer_results) + return cast(GLM52DenseDecoderLayerOutput, layer_results) + if is_micro_batch: + return cast(GLM52MoEDecoderLayerMicroBatchOutput, layer_results) + return cast(GLM52MoEDecoderLayerOutput, layer_results) def to_hf_key_list(self, key: str) -> list[str]: if self.config.tie_word_embeddings and "lm_head" in key: diff --git a/xtuner/v1/model/moe/glm52/mtp.py b/xtuner/v1/model/moe/glm52/mtp.py new file mode 100644 index 0000000000..12adf480ff --- /dev/null +++ b/xtuner/v1/model/moe/glm52/mtp.py @@ -0,0 +1,118 @@ +from typing import cast + +import torch +from typing_extensions import override + +from xtuner.v1.data_proto import SequenceContext +from xtuner.v1.module.decoder_layer.moe_decoder_layer import ( + MoEDecoderLayerMicroBatchOutput, + MoEDecoderLayerOutput, +) +from xtuner.v1.module.mtp import MTPBlock, MTPLayer +from xtuner.v1.module.mtp.mtp_block import MTPInternalOutput + +from .decoder_layer import ( + GLM52MoEDecoderLayerMicroBatchOutput, + GLM52MoEDecoderLayerOutput, +) + + +class GLM52MTPLayer(MTPLayer): + """MTP layer whose wrapped GLM-5.2 decoder consumes explicit DSA IDs.""" + + @override + def forward( + self, + hidden_states: torch.Tensor | list[torch.Tensor], + *, + future_embeddings: torch.Tensor | list[torch.Tensor], + position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], + seq_ctx: SequenceContext | list[SequenceContext], + dsa_topk_ids: torch.Tensor | list[torch.Tensor] | None = None, + ) -> GLM52MoEDecoderLayerOutput | GLM52MoEDecoderLayerMicroBatchOutput: + if not isinstance(hidden_states, list): + assert isinstance(future_embeddings, torch.Tensor) + assert isinstance(position_embeddings, tuple) and len(position_embeddings) == 2 + assert isinstance(seq_ctx, SequenceContext) + assert dsa_topk_ids is None or isinstance(dsa_topk_ids, torch.Tensor) + projected = self._preprocess(hidden_states=hidden_states, future_embeddings=future_embeddings) + layer_results = cast( + GLM52MoEDecoderLayerOutput, + self.decoder_layer( + projected, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + dsa_topk_ids=dsa_topk_ids, + ), + ) + return { + "hidden_states": self.final_layernorm(layer_results["hidden_states"]), + "router_logits": layer_results["router_logits"], + "router_weights": layer_results["router_weights"], + "router_topk_ids": layer_results["router_topk_ids"], + "dsa_topk_ids": layer_results["dsa_topk_ids"], + } + + n = len(hidden_states) + assert isinstance(future_embeddings, list) and len(future_embeddings) == n + assert isinstance(position_embeddings, list) and len(position_embeddings) == n + assert isinstance(seq_ctx, list) and len(seq_ctx) == n + if dsa_topk_ids is None: + decoder_topk_ids = None + else: + assert isinstance(dsa_topk_ids, list) and len(dsa_topk_ids) == n + decoder_topk_ids = dsa_topk_ids + + projected_list = [ + self._preprocess(hidden_states=hidden, future_embeddings=future) + for hidden, future in zip(hidden_states, future_embeddings) + ] + micro_batch_results = cast( + GLM52MoEDecoderLayerMicroBatchOutput, + self.decoder_layer( + projected_list, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + dsa_topk_ids=decoder_topk_ids, + ), + ) + return { + "hidden_states": [self.final_layernorm(hidden) for hidden in micro_batch_results["hidden_states"]], + "router_logits": micro_batch_results["router_logits"], + "router_weights": micro_batch_results["router_weights"], + "router_topk_ids": micro_batch_results["router_topk_ids"], + "dsa_topk_ids": micro_batch_results["dsa_topk_ids"], + } + + +class GLM52MTPBlock(MTPBlock): + """MTP block that keeps DSA sharing private to GLM-5.2.""" + + @override + def _call_decoder_layer( + self, + layer: MTPLayer, + hidden_states: torch.Tensor | list[torch.Tensor], + *, + previous_layer_results: MTPInternalOutput | None, + future_embeddings: torch.Tensor | list[torch.Tensor], + position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], + seq_ctx: SequenceContext | list[SequenceContext], + ) -> MTPInternalOutput: + glm_layer = cast(GLM52MTPLayer, layer) + previous_results = cast( + GLM52MoEDecoderLayerOutput | GLM52MoEDecoderLayerMicroBatchOutput | None, + previous_layer_results, + ) + dsa_topk_ids = None if previous_results is None else previous_results["dsa_topk_ids"] + + return cast( + MoEDecoderLayerOutput | MoEDecoderLayerMicroBatchOutput, + glm_layer( + hidden_states, + future_embeddings=future_embeddings, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + dsa_topk_ids=dsa_topk_ids, + ), + ) diff --git a/xtuner/v1/model/moe/moe.py b/xtuner/v1/model/moe/moe.py index b0111afa73..15dc0d628e 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 @@ -11,7 +12,6 @@ from pydantic import ConfigDict from torch import nn from torch.distributed._functional_collectives import all_reduce -from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import CheckpointImpl from torch.distributed.device_mesh import DeviceMesh, init_device_mesh from torch.distributed.distributed_c10d import ReduceOp from torch.distributed.fsdp import ( @@ -23,7 +23,7 @@ from typing_extensions import overload, override from xtuner.v1.config import FSDPConfig -from xtuner.v1.data_proto import DSATopKCacheState, SequenceContext +from xtuner.v1.data_proto import SequenceContext from xtuner.v1.float8.float8_handler import Float8Handler from xtuner.v1.loss import ( AuxLossConfig, @@ -47,9 +47,8 @@ ) from xtuner.v1.model.utils import ( ModelForwardExtraLogInfo, - checkpoint_wrapper, + apply_activation_checkpointing, module_dict_repr, - pytree_reentrant_checkpoint, ) from xtuner.v1.module import ( GatedDeltaNetConfig, @@ -61,8 +60,19 @@ NoAuxRouterConfig, RMSNorm, ) -from xtuner.v1.module.decoder_layer.dense_decoder_layer import DenseDecoderLayer -from xtuner.v1.module.decoder_layer.moe_decoder_layer import MoEActFnConfig, MoEBlock, MoEDecoderLayer, MoEGate +from xtuner.v1.module.decoder_layer.dense_decoder_layer import ( + DenseDecoderLayer, + DenseDecoderLayerMicroBatchOutput, + DenseDecoderLayerOutput, +) +from xtuner.v1.module.decoder_layer.moe_decoder_layer import ( + MoEActFnConfig, + MoEBlock, + MoEDecoderLayer, + MoEDecoderLayerMicroBatchOutput, + MoEDecoderLayerOutput, + MoEGate, +) from xtuner.v1.module.mtp import MTPBlock, MTPConfig, MTPLayer from xtuner.v1.utils import ( get_device, @@ -81,8 +91,9 @@ logger = get_logger() +MOE_BLOCK_FORWARD = "xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEBlock.forward" MOE_NON_EP_COMPILE_CFG: dict[str, TorchCompileOption] = { - "xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEBlock.forward": TorchCompileOption(fullgraph=True), + MOE_BLOCK_FORWARD: TorchCompileOption(fullgraph=True), "xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEDecoderLayer.forward": TorchCompileOption(fullgraph=True), "xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEDecoderLayer._pre_moe_forward": TorchCompileOption( fullgraph=True @@ -190,6 +201,10 @@ class MoE(BaseModel): ep_mesh: DeviceMesh | None = None expert_tp_mesh: DeviceMesh | None = None ep_tp_mesh: DeviceMesh | None = None + dense_decoder_layer_cls = DenseDecoderLayer + moe_decoder_layer_cls = MoEDecoderLayer + mtp_layer_cls = MTPLayer + mtp_block_cls = MTPBlock def __init__(self, config: MoEConfig): # Concrete MoE configs override build(), so validate dispatcher support @@ -246,7 +261,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] @@ -276,9 +291,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() + + data_ptrs = {tensor.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.data_ptr() in data_ptrs, + prefetch=True, + reserve_pin_memory=True, + ) + def _z_loss_dist_token_count( self, z_ctx: list[ZLossContext] | ZLossContext | None, @@ -513,14 +552,6 @@ def post_micro_batch_forward(self, batch_outputs: Sequence[MoEModelOutputs]) -> moe_info = cast(MoEBatchForwardInfo, base_info) return moe_info - @staticmethod - def _prepare_seq_ctx_topk_cache(seq_ctx_list: Sequence[SequenceContext]) -> None: - # A slot owns both runtime offload state and pinned storage. TrainEngine - # finishes backward before the next accumulation group, so only contexts - # in one model call can be live concurrently and need distinct slots. - for offload_slot, seq_ctx in enumerate(seq_ctx_list): - seq_ctx.dsa_topk_cache.offload_slot = offload_slot - def _micro_batch_forward( self, seq_ctx_list: list[SequenceContext], @@ -532,7 +563,6 @@ def _micro_batch_forward( This method processes multiple micro-batches in parallel, similar to how MoEDecoderLayer handles micro-batching at the layer level. """ - self._prepare_seq_ctx_topk_cache(seq_ctx_list) if self.config.return_hidden_states: raise NotImplementedError @@ -587,72 +617,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" @@ -671,14 +648,11 @@ 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( - offload_slot=seq_ctx.dsa_topk_cache.offload_slot, - ), ) ) mtp_outputs_per_mb = self.mtp_block( - *hidden_states_list, + hidden_states_list, embed_tokens_fn=self.embed_tokens, position_embeddings=position_embeddings_list, seq_ctx=mtp_seq_ctx_list, @@ -693,12 +667,11 @@ def _micro_batch_forward( micro_batch_mtp_losses = torch.tensor(0.0, device=DEVICE) for mtp_idx, (mtp_hidden, mtp_ctx) in enumerate(zip(mtp_outputs, mtp_loss_ctx_list)): - mtp_hidden_states, mtp_router_results, _, _ = mtp_hidden - mtp_loss, _ = self.lm_head(mtp_hidden_states, cast(MTPLossContext, mtp_ctx)) + mtp_loss, _ = self.lm_head(mtp_hidden["hidden_states"], cast(MTPLossContext, mtp_ctx)) micro_batch_mtp_losses += mtp_loss if keep_router: - router_logits_list[micro_batch_idx][f"mtp_layer{mtp_idx}"] = mtp_router_results + router_logits_list[micro_batch_idx][f"mtp_layer{mtp_idx}"] = mtp_hidden["router_logits"] mtp_losses += micro_batch_mtp_losses / len(mtp_loss_ctx_list) has_mtp_loss = True @@ -719,13 +692,13 @@ def _micro_batch_forward( # loss already rides on, so backward traverses each MTP aux node exactly once. for mtp_idx in range(self.config.mtp_config.num_layers): cat_mtp_router_weights = torch.cat( - [mb_outputs[mtp_idx][2] for mb_outputs in mtp_outputs_per_mb], dim=0 + [mb_outputs[mtp_idx]["router_weights"] for mb_outputs in mtp_outputs_per_mb], dim=0 ) cat_mtp_router_logits = torch.cat( - [mb_outputs[mtp_idx][1] for mb_outputs in mtp_outputs_per_mb], dim=0 + [mb_outputs[mtp_idx]["router_logits"] for mb_outputs in mtp_outputs_per_mb], dim=0 ) cat_mtp_router_topk_ids = torch.cat( - [mb_outputs[mtp_idx][3] for mb_outputs in mtp_outputs_per_mb], dim=0 + [mb_outputs[mtp_idx]["router_topk_ids"] for mb_outputs in mtp_outputs_per_mb], dim=0 ) hidden_states_list[0] = self.aux_loss.accumulate( selected_router_weights=cat_mtp_router_weights.index_select(0, nonpad_indices) @@ -783,20 +756,87 @@ def _micro_batch_forward( layer_router_logits_list: list[torch.Tensor] = [] for micro_batch_idx in range(len(seq_ctx_list)): layer_router_logits_list.append(router_logits_list[micro_batch_idx][layer_name].detach()) - router_logits = torch.stack(layer_router_logits_list, dim=0).unsqueeze(0) - router_logits_dict[layer_name] = router_logits + router_logits_dict[layer_name] = torch.stack(layer_router_logits_list, dim=0).unsqueeze(0) output["router_logits"] = router_logits_dict return MoEModelOutputs(**output, logits=logits) + def _micro_batch_decoder_stack( + self, + *, + hidden_states_list: list[torch.Tensor], + position_embeddings_list: list[tuple[torch.Tensor, torch.Tensor]], + seq_ctx_list: list[SequenceContext], + router_logits_list: list[dict[str, torch.Tensor]], + keep_router: bool, + balancing_ctx: list[BalancingLossContext] | BalancingLossContext | None, + z_ctx: list[ZLossContext] | ZLossContext | None, + nonpad_indices: torch.Tensor, + non_pad_token: int, + num_tokens_global: torch.Tensor | None, + z_world_size: int, + ) -> list[torch.Tensor]: + """Run the main decoder stack for intra-layer micro-batches.""" + previous_layer_results: DenseDecoderLayerMicroBatchOutput | MoEDecoderLayerMicroBatchOutput | None = None + + for idx, decoder_layer in self.layers.items(): + layer_idx = int(idx) + layer_results = self._call_decoder_layer( + decoder_layer=decoder_layer, + layer_idx=layer_idx, + hidden_states=hidden_states_list, + position_embeddings=position_embeddings_list, + seq_ctx=seq_ctx_list, + previous_layer_results=previous_layer_results, + ) + previous_layer_results = cast( + DenseDecoderLayerMicroBatchOutput | MoEDecoderLayerMicroBatchOutput, + layer_results, + ) + if layer_idx < self.config.first_k_dense_replace: + # Keep each micro-batch in its own SequenceContext while issuing + # one outer layer call, so FSDP materializes dense weights once. + dense_results = cast(DenseDecoderLayerMicroBatchOutput, layer_results) + hidden_states_list = dense_results["hidden_states"] + continue + + layer_results = cast(MoEDecoderLayerMicroBatchOutput, layer_results) + hidden_states = layer_results["hidden_states"] + router_logits = layer_results["router_logits"] + router_weights = layer_results["router_weights"] + router_topk_ids = layer_results["router_topk_ids"] + + # Router weights are consumed immediately by aux loss; only logits + # requested by the caller are retained per micro-batch. + for i, hidden_state in enumerate(hidden_states): + hidden_states_list[i] = hidden_state + if keep_router: + router_logits_list[i][f"layer{idx}"] = self._maybe_offload_router(router_logits[i]) + + cat_router_weights = torch.cat(router_weights, dim=0) + cat_router_logits = torch.cat(router_logits, dim=0) + cat_router_topk_ids = torch.cat(router_topk_ids, dim=0) + hidden_states_list[0] = self.aux_loss.accumulate( + selected_router_weights=cat_router_weights.index_select(0, nonpad_indices).contiguous().float(), + selected_router_logits=cat_router_logits.index_select(0, nonpad_indices).contiguous().float(), + selected_experts=cat_router_topk_ids.index_select(0, nonpad_indices).contiguous(), + hidden_states=hidden_states_list[0], + balancing_ctx=balancing_ctx, + z_ctx=z_ctx, + num_tokens_local=non_pad_token, + num_tokens_global=num_tokens_global, + world_size=z_world_size, + ) + + return hidden_states_list + def _forward( self, seq_ctx: SequenceContext, # todo(@yehaochen): support intra layer micro-batch loss_ctx: MoELossContextDict | None, return_router_logits: bool = False, ) -> MoEModelOutputs: - self._prepare_seq_ctx_topk_cache([seq_ctx]) input_ids = seq_ctx.input_ids position_ids = seq_ctx.position_ids @@ -832,57 +872,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) @@ -904,9 +913,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( - offload_slot=seq_ctx.dsa_topk_cache.offload_slot, - ), ) # 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] @@ -926,10 +932,13 @@ def _forward( # Compute MTP losses for each depth mtp_losses = torch.tensor(0.0, device=DEVICE) for idx, (mtp_hidden, mtp_ctx) in enumerate(zip(mtp_outputs, mtp_loss_ctx_list)): - mtp_hidden_states, mtp_router_results, mtp_router_weights, mtp_router_topk_ids = mtp_hidden + mtp_hidden_states = mtp_hidden["hidden_states"] + mtp_router_logits = mtp_hidden["router_logits"] + mtp_router_weights = mtp_hidden["router_weights"] + mtp_router_topk_ids = mtp_hidden["router_topk_ids"] if keep_router: - output["router_logits"][f"mtp_layer{idx}"] = mtp_router_results + output["router_logits"][f"mtp_layer{idx}"] = mtp_router_logits output["router_weights"][f"mtp_layer{idx}"] = mtp_router_weights # Inject this MTP layer's z-loss before lm_head so backward through mtp_loss # traverses the AuxLossScaler node and releases this layer's logsumexp activations. @@ -937,7 +946,7 @@ def _forward( selected_router_weights=mtp_router_weights.index_select(0, mtp_nonpad_indices) .contiguous() .float(), - selected_router_logits=mtp_router_results.index_select(0, mtp_nonpad_indices).contiguous().float(), + selected_router_logits=mtp_router_logits.index_select(0, mtp_nonpad_indices).contiguous().float(), selected_experts=mtp_router_topk_ids.index_select(0, mtp_nonpad_indices).contiguous(), hidden_states=mtp_hidden_states, balancing_ctx=balancing_ctx, @@ -975,6 +984,113 @@ def _forward( return MoEModelOutputs(**output) + def _decoder_stack( + self, + *, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + seq_ctx: SequenceContext, + output: dict, + keep_router: bool, + balancing_ctx: BalancingLossContext | None, + z_ctx: ZLossContext | None, + nonpad_indices: torch.Tensor, + non_pad_token: int, + num_tokens_global: torch.Tensor | None, + z_world_size: int, + ) -> torch.Tensor: + """Run the main decoder stack for one sequence context.""" + previous_layer_results: DenseDecoderLayerOutput | MoEDecoderLayerOutput | None = None + + for idx, decoder_layer in self.layers.items(): + layer_idx = int(idx) + layer_results = self._call_decoder_layer( + decoder_layer=decoder_layer, + layer_idx=layer_idx, + hidden_states=hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + previous_layer_results=previous_layer_results, + ) + previous_layer_results = cast(DenseDecoderLayerOutput | MoEDecoderLayerOutput, layer_results) + if layer_idx < self.config.first_k_dense_replace: + dense_results = cast(DenseDecoderLayerOutput, layer_results) + hidden_states = dense_results["hidden_states"] + else: + layer_results = cast(MoEDecoderLayerOutput, layer_results) + hidden_states = layer_results["hidden_states"] + router_results = layer_results["router_logits"] + router_weights = layer_results["router_weights"] + router_topk_ids = layer_results["router_topk_ids"] + if keep_router: + output["router_logits"][f"layer{idx}"] = self._maybe_offload_router(router_results) + output["router_weights"][f"layer{idx}"] = self._maybe_offload_router(router_weights) + hidden_states = self.aux_loss.accumulate( + selected_router_weights=router_weights.index_select(0, nonpad_indices).contiguous().float(), + selected_router_logits=router_results.index_select(0, nonpad_indices).contiguous().float(), + selected_experts=router_topk_ids.index_select(0, nonpad_indices).contiguous(), + hidden_states=hidden_states, + balancing_ctx=balancing_ctx, + z_ctx=z_ctx, + num_tokens_local=non_pad_token, + num_tokens_global=num_tokens_global, + world_size=z_world_size, + ) + + if self.config.return_hidden_states: + output["hidden_states"].append(hidden_states) + + return hidden_states + + def _call_decoder_layer( + self, + *, + decoder_layer: nn.Module, + layer_idx: int, + hidden_states: torch.Tensor | list[torch.Tensor], + position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], + seq_ctx: SequenceContext | list[SequenceContext], + previous_layer_results: ( + DenseDecoderLayerOutput + | DenseDecoderLayerMicroBatchOutput + | MoEDecoderLayerOutput + | MoEDecoderLayerMicroBatchOutput + | None + ), + ) -> ( + DenseDecoderLayerOutput + | DenseDecoderLayerMicroBatchOutput + | MoEDecoderLayerOutput + | MoEDecoderLayerMicroBatchOutput + ): + """Call one decoder layer and contain its activation-offload window. + + ``previous_layer_results`` is intentionally unused by the generic model; subclasses may + consume private cross-layer fields without widening the common decoder API. + """ + if layer_idx < self.config.first_k_dense_replace: + return decoder_layer( + hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + ) + + offload_tensors = list(hidden_states) if isinstance(hidden_states, list) else [hidden_states] + if int(os.getenv("XTUNER_ACTIVATION_OFFLOAD", "0")) != 1: + offload_tensors = [] + offload_block_idx = sum( + self.config.first_k_dense_replace <= int(previous_idx) < layer_idx for previous_idx in self.layers + ) + with self._saved_tensors_offload_ctx( + offload_block_idx, + offload_tensors, + ): + return decoder_layer( + hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + ) + def build_embeddings(self, config: MoEConfig): return nn.Embedding(config.vocab_size, config.hidden_size, config.pad_token_id) @@ -997,7 +1113,7 @@ def build_layers(self, config: MoEConfig) -> nn.ModuleDict: ) if layer_idx < config.first_k_dense_replace: - layers[str(layer_idx)] = DenseDecoderLayer( + layers[str(layer_idx)] = self.dense_decoder_layer_cls( hidden_size=config.hidden_size, intermediate_size=config.intermediate_size, mlp_bias=config.mlp_bias, @@ -1012,7 +1128,7 @@ def build_layers(self, config: MoEConfig) -> nn.ModuleDict: layer_idx=layer_idx, ) else: - layers[str(layer_idx)] = MoEDecoderLayer( + layers[str(layer_idx)] = self.moe_decoder_layer_cls( hidden_size=config.hidden_size, intermediate_size=config.intermediate_size, moe_intermediate_size=config.moe_intermediate_size, @@ -1079,7 +1195,7 @@ def build_mtp_block(self, config: MoEConfig) -> MTPBlock: num_physical_layer = 1 if mtp_config.share_weights else mtp_config.num_layers for i in range(num_physical_layer): # Build MoE decoder layer for MTP - decoder_layer = MoEDecoderLayer( + decoder_layer = self.moe_decoder_layer_cls( hidden_size=config.hidden_size, intermediate_size=config.intermediate_size, moe_intermediate_size=config.moe_intermediate_size, @@ -1110,7 +1226,7 @@ def build_mtp_block(self, config: MoEConfig) -> MTPBlock: ) # Wrap decoder layer in MTPLayer - mtp_layer = MTPLayer( + mtp_layer = self.mtp_layer_cls( hidden_size=config.hidden_size, rms_norm_eps=config.rms_norm_eps, rms_norm_type=config.rms_norm_type, @@ -1119,7 +1235,7 @@ def build_mtp_block(self, config: MoEConfig) -> MTPBlock: ) mtp_layers.append(mtp_layer) - return MTPBlock( + return self.mtp_block_cls( mtp_config=mtp_config, mtp_layers=mtp_layers, ) @@ -1193,6 +1309,7 @@ def fully_shard( mp_policy = MixedPrecisionPolicy( param_dtype=self.fsdp_config.param_dtype, reduce_dtype=fsdp_config.reduce_dtype ) + checkpoint_preserve_rng_state = fsdp_config.checkpoint_preserve_rng_state for layer_idx, layer in tqdm(self.layers.items(), desc="[FSDP Sharding]"): layer_idx = int(layer_idx) @@ -1200,7 +1317,10 @@ def fully_shard( layer_idx=layer_idx, mtp_idx=None, ): - layer = checkpoint_wrapper(layer, checkpoint_impl=CheckpointImpl.REENTRANT) + layer = apply_activation_checkpointing( + layer, + preserve_rng_state=checkpoint_preserve_rng_state, + ) self.layers[str(layer_idx)] = layer if layer_idx >= len(self.layers) - 1 and self.mtp_block is None: @@ -1252,39 +1372,10 @@ 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 状态。 - # - # 使用 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。 - use_reentrant = self.fsdp_config.mtp_checkpoint_use_reentrant - if use_reentrant: - mtp_layer = checkpoint_wrapper( - mtp_layer, - checkpoint_impl=CheckpointImpl.REENTRANT, - checkpoint_fn=pytree_reentrant_checkpoint, - ) - else: - mtp_layer = checkpoint_wrapper( - mtp_layer, - checkpoint_impl=CheckpointImpl.NO_REENTRANT, - ) + mtp_layer = apply_activation_checkpointing( + mtp_layer, + preserve_rng_state=checkpoint_preserve_rng_state, + ) self.mtp_block.layers[mtp_idx] = mtp_layer reshard_after_forward = mtp_idx != len(self.mtp_block.layers) - 1 @@ -1326,8 +1417,7 @@ def fully_shard( def default_compile_cfg(self) -> dict[str, TorchCompileOption]: if use_moe_ep_compile_cfg(self.config): return MOE_EP_COMPILE_CFG - else: - return MOE_NON_EP_COMPILE_CFG + return MOE_NON_EP_COMPILE_CFG @property def need_update_bias(self) -> bool: diff --git a/xtuner/v1/model/moe/qwen3vl_text.py b/xtuner/v1/model/moe/qwen3vl_text.py index 45742d8527..33e952029a 100644 --- a/xtuner/v1/model/moe/qwen3vl_text.py +++ b/xtuner/v1/model/moe/qwen3vl_text.py @@ -4,6 +4,8 @@ import torch from xtuner.v1.data_proto import SequenceContext +from xtuner.v1.module.decoder_layer.dense_decoder_layer import DenseDecoderLayerOutput +from xtuner.v1.module.decoder_layer.moe_decoder_layer import MoEDecoderLayerOutput from xtuner.v1.utils.activation_offload import async_save_on_cpu from .moe import MoELossContextDict, MoEModelOutputs @@ -130,11 +132,12 @@ def _forward( for idx, decoder_layer in self.layers.items(): if int(idx) < self.config.first_k_dense_replace: - hidden_states = decoder_layer( + dense_results: DenseDecoderLayerOutput = decoder_layer( hidden_states, position_embeddings=position_embeddings, seq_ctx=seq_ctx, ) + hidden_states = dense_results["hidden_states"] else: if int(os.getenv("XTUNER_ACTIVATION_OFFLOAD", "0")) == 1: offload_stream = decoder_layer._get_fsdp_state()._comm_ctx.all_gather_stream @@ -145,25 +148,29 @@ def _forward( depth=len(self.layers), custom_check_fn=lambda x: x.data_ptr() == hidden_states.data_ptr(), ): - hidden_states, router_results, router_weights, router_topk_ids = decoder_layer( + layer_results: MoEDecoderLayerOutput = decoder_layer( hidden_states, position_embeddings=position_embeddings, seq_ctx=seq_ctx, ) else: - hidden_states, router_results, router_weights, router_topk_ids = decoder_layer( + layer_results = decoder_layer( hidden_states, position_embeddings=position_embeddings, seq_ctx=seq_ctx, ) + hidden_states = layer_results["hidden_states"] + router_logits = layer_results["router_logits"] + router_weights = layer_results["router_weights"] + router_topk_ids = layer_results["router_topk_ids"] if keep_router: - output["router_logits"][f"layer{idx}"] = router_results + output["router_logits"][f"layer{idx}"] = router_logits output["router_weights"][f"layer{idx}"] = 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_router_logits=router_logits.index_select(0, nonpad_indices).contiguous().float(), selected_experts=router_topk_ids.index_select(0, nonpad_indices).contiguous(), hidden_states=hidden_states, balancing_ctx=balancing_ctx, diff --git a/xtuner/v1/model/utils/__init__.py b/xtuner/v1/model/utils/__init__.py index ba77e7c3c0..111b47f172 100644 --- a/xtuner/v1/model/utils/__init__.py +++ b/xtuner/v1/model/utils/__init__.py @@ -1,5 +1,10 @@ -from .checkpointing import checkpoint_wrapper, pytree_reentrant_checkpoint +from .checkpointing import apply_activation_checkpointing, reuse_during_recompute from .misc import ModelForwardExtraLogInfo, module_dict_repr -__all__ = ["checkpoint_wrapper", "pytree_reentrant_checkpoint", "module_dict_repr", "ModelForwardExtraLogInfo"] +__all__ = [ + "apply_activation_checkpointing", + "reuse_during_recompute", + "module_dict_repr", + "ModelForwardExtraLogInfo", +] diff --git a/xtuner/v1/model/utils/checkpointing.py b/xtuner/v1/model/utils/checkpointing.py index 0cc9282793..76249c544c 100644 --- a/xtuner/v1/model/utils/checkpointing.py +++ b/xtuner/v1/model/utils/checkpointing.py @@ -1,131 +1,154 @@ -import inspect -from types import UnionType -from typing import Any, Callable, Union, get_args, get_origin +"""Activation checkpointing entry points.""" + +from collections import deque +from contextvars import ContextVar +from typing import Any, Callable import torch import torch.nn as nn -from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import ( - checkpoint_wrapper as ptd_checkpoint_wrapper, -) -from torch.utils._pytree import tree_flatten, tree_unflatten +from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import CheckpointImpl, checkpoint_wrapper +from torch.utils._pytree import TreeSpec, tree_flatten, tree_unflatten from torch.utils.checkpoint import checkpoint -from xtuner.v1.utils import copy_signature - - -# TODO: Currently xtuner uses the internal, outdated `torch.distributed.algorithms._checkpoint.checkpoint_wrapper` interface -# We should look for opportunities to use the public, updated interface in the future -# NOTE: -# PyTorch's `torch.distributed.algorithms._checkpoint.checkpoint_wrapper` has some limitations. Modules decorated with `checkpoint_wrapper` -# must have forward interfaces that conform to the specifications of `torch.autograd.function.Function`. -# Specifically, for input parameters, the `forward` interface must explicitly accept parameters of type `torch.Tensor` to ensure proper gradient backpropagation. -# For return values, the `forward` interface must return either `torch.Tensor` or tuple[torch.Tensor, ...]. -# For example If the forward interface is declared as: -# def forward(self, x: tuple[torch.Tensor], y: list[torch.Tensor]) -> torch.Tensor | tuple[torch.Tensor, ...]: -# This interface will break the gradient graph because the inputs don't meet the requirements. For instance, x and y are not of type `torch.Tensor` -# `_check_signature_of_forward` will check (not exhaustively) whether the signature of the `forward` interface meets the requirements to identify issues early. +__all__ = ["apply_activation_checkpointing", "reuse_during_recompute"] -def _check_signature_of_forward(module: nn.Module): - def _is_tensor_or_tuple_tensor(arg_type: type): - if arg_type is torch.Tensor: - return True +class _ActivationCheckpointFrame: + """Reusable outputs owned by one checkpoint invocation.""" - origin_type = get_origin(arg_type) + def __init__(self) -> None: + self.outputs: dict[Callable[..., Any], deque[Any]] = {} - if not origin_type or origin_type not in (tuple, UnionType, Union): - return False + def save(self, function: Callable[..., Any], output: Any) -> None: + self.outputs.setdefault(function, deque()).append(output) - if origin_type in [UnionType, Union]: - type_list = get_args(arg_type) - return any(_is_tensor_or_tuple_tensor(t) for t in type_list) + def replay(self, function: Callable[..., Any]) -> Any: + queue = self.outputs.get(function) + if not queue: + raise RuntimeError("Checkpoint replay has no matching reusable output from the original forward") + return queue.popleft() - else: - type_list = get_args(arg_type) - return any(t is torch.Tensor for t in type_list) + def assert_consumed(self) -> None: + if any(self.outputs.values()): + raise RuntimeError("Checkpoint replay did not consume every reusable output from the original forward") - def _has_missing_type(arg_type: type): - if arg_type is inspect._empty: - return True - origin_arg = get_origin(arg_type) - return any(_has_missing_type(t) for t in get_args(origin_arg)) - input_type = inspect.signature(module.forward).parameters - ret_type = inspect.signature(module.forward).return_annotation - - for name, arg_type in input_type.items(): - if _has_missing_type(arg_type.annotation): - raise TypeError( - f"The type of argument '{name}' of {module.__class__.__name__}.forward must be annotated, but got " - f"{name} unannotated." - ) - - if _has_missing_type(ret_type): - raise TypeError( - f"The return type of {module.__class__.__name__}.forward must be annotated, but got {ret_type}" - ) +# The bool is true only while a reentrant checkpoint invocation is replaying. +_CURRENT_ACTIVATION_CHECKPOINT: ContextVar[tuple[_ActivationCheckpointFrame, bool] | None] = ContextVar( + "xtuner_current_activation_checkpoint", + default=None, +) - for arg_type in input_type.values(): - origin_arg = get_origin(arg_type.annotation) - # Union[Tensor, None] or Optional[Tensor] is legal - if origin_arg: - if torch.Tensor in origin_arg: - break - else: - if arg_type.annotation is torch.Tensor: - break - else: - raise TypeError( - f"The type of all arguments of the {module.__class__.__name__}.forward must be torch.Tensor, but got " - f"{input_type}" - ) - if not _is_tensor_or_tuple_tensor(ret_type): - raise TypeError( - f"The return type of {module.__class__.__name__}.forward must be torch.Tensor or tuple of torch.Tensor, " - f"but got {ret_type}" +@torch.compiler.disable(recursive=False) +def reuse_during_recompute(function: Callable[..., Any], /, *args: Any, **kwargs: Any) -> Any: + """Run a no-grad callable once and reuse its output during checkpoint + replay. + + Outside :func:`apply_activation_checkpointing`, this is a direct call. Inside a + checkpoint invocation, calls to the same stable callable are matched in FIFO order. + """ + state = _CURRENT_ACTIVATION_CHECKPOINT.get() + if state is None: + return function(*args, **kwargs) + + frame, is_replay = state + if is_replay: + return frame.replay(function) + + output = function(*args, **kwargs) + flat_output, _ = tree_flatten(output) + if any(isinstance(leaf, torch.Tensor) and leaf.requires_grad for leaf in flat_output): + raise RuntimeError("reuse_during_recompute only supports Tensor outputs that do not require gradients") + frame.save(function, output) + return output + + +def apply_activation_checkpointing( + module: nn.Module, + *, + preserve_rng_state: bool = True, +) -> nn.Module: + """Wrap ``module`` in fixed reentrant activation checkpointing. + + The PyTree bridge keeps nested positional and keyword inputs visible to autograd and saved-tensor hooks, and + restores structured model outputs. + """ + module_has_trainable_parameters = any(parameter.requires_grad for parameter in module.parameters()) + + def checkpoint_fn(function: Callable[..., Any], /, *args: Any, **kwargs: Any) -> Any: + return _checkpoint_pytree( + function, + call_args=args, + call_kwargs=kwargs, + preserve_rng_state=preserve_rng_state, + module_has_trainable_parameters=module_has_trainable_parameters, ) - -@copy_signature(ptd_checkpoint_wrapper) -def checkpoint_wrapper(module: nn.Module, *args, **kwargs): - _check_signature_of_forward(module) - return ptd_checkpoint_wrapper(module, *args, **kwargs) - - -def pytree_reentrant_checkpoint( - function: Callable[..., torch.Tensor | tuple[torch.Tensor, ...]], - *args: Any, - **kwargs: Any, -) -> torch.Tensor | tuple[torch.Tensor, ...]: - """让嵌套 Tensor 也成为 reentrant checkpoint 的 autograd 输入。""" - # CheckpointWrapper 只打包一层。例如: - # future_embeddings=[embedding_0, embedding_1] - # 对原生 CheckpointFunction 来说只是“一个 list 参数”,它看不到 list 里的 - # 两个 Tensor,也就不会 detach 它们。这可能会造成反向传播的错误,因为这两个 - # Tensor 的梯度应该交由 CheckpointFunction.backward 的返回值交回原始 Tensor, - # 而不是由他们自己来传递梯度。 - # tree_flatten 会把输入变成近似: - # hidden, embedding_0, embedding_1 - # 这样 checkpoint 能逐个 detach;tree_unflatten 再在 replay 前把 list 还原。 - flat_inputs, input_spec = tree_flatten((args, kwargs)) - + # FSDP must wrap this checkpoint boundary. Its output hook then unshards + # parameters before replay; wrapping FSDP itself would replay an FSDP forward. + return checkpoint_wrapper( + module, + checkpoint_impl=CheckpointImpl.REENTRANT, + checkpoint_fn=checkpoint_fn, + ) + + +@torch.compiler.disable(recursive=False) +def _checkpoint_pytree( + function: Callable[..., Any], + /, + *, + call_args: tuple[Any, ...], + call_kwargs: dict[str, Any], + preserve_rng_state: bool, + module_has_trainable_parameters: bool, +) -> Any: + """Adapt a structured module call to reentrant checkpoint's flat + boundary.""" + flat_inputs, input_spec = tree_flatten((call_args, call_kwargs)) has_grad_input = any(isinstance(value, torch.Tensor) and value.requires_grad for value in flat_inputs) checkpoint_inputs = flat_inputs - if not has_grad_input: - # Reentrant checkpoint 没有 grad 输入时不会重算。这个零元素 leaf 只负责 - # 保留 backward 入口,不参与模块计算,也不会把梯度接回已 detach 的原图。 + needs_grad_entry = not has_grad_input and module_has_trainable_parameters + if needs_grad_entry: + # MTP can detach every model input while keeping trainable parameters. The empty leaf + # preserves its replay entry, but frozen modules must remain detached and skip replay. first_tensor = next(value for value in flat_inputs if isinstance(value, torch.Tensor)) checkpoint_entry = torch.empty(0, device=first_tensor.device, requires_grad=True) checkpoint_inputs = [checkpoint_entry, *flat_inputs] - - def run_function(*replayed_flat_inputs: Any) -> torch.Tensor | tuple[torch.Tensor, ...]: - # 这里只还原参数结构,不会把 detached Tensor 重新连接到旧 graph;梯度由 - # CheckpointFunction.backward 的返回值交回原始 Tensor。 - if not has_grad_input: - replayed_flat_inputs = replayed_flat_inputs[1:] - replayed_args, replayed_kwargs = tree_unflatten(list(replayed_flat_inputs), input_spec) - return function(*replayed_args, **replayed_kwargs) - - return checkpoint(run_function, *checkpoint_inputs, use_reentrant=True) + output_spec: TreeSpec | None = None + frame = _ActivationCheckpointFrame() + + @torch.compiler.disable(recursive=False) + def call_with_original_signature(*replayed_inputs: Any) -> tuple[Any, ...]: + nonlocal output_spec + if needs_grad_entry: + replayed_inputs = replayed_inputs[1:] + replayed_args, replayed_kwargs = tree_unflatten(list(replayed_inputs), input_spec) + is_replay = torch.is_grad_enabled() + token = _CURRENT_ACTIVATION_CHECKPOINT.set((frame, is_replay)) + try: + output = function(*replayed_args, **replayed_kwargs) + finally: + _CURRENT_ACTIVATION_CHECKPOINT.reset(token) + + flat_outputs, current_output_spec = tree_flatten(output) + if output_spec is None: + output_spec = current_output_spec + elif current_output_spec != output_spec: + raise RuntimeError("Checkpoint replay returned a different output PyTree structure") + if is_replay: + frame.assert_consumed() + return tuple(flat_outputs) + + flat_outputs = checkpoint( + call_with_original_signature, + *checkpoint_inputs, + use_reentrant=True, + preserve_rng_state=preserve_rng_state, + ) + assert output_spec is not None, "XTuner internal error: checkpoint did not run the function" + if not isinstance(flat_outputs, tuple): + flat_outputs = (flat_outputs,) + return tree_unflatten(list(flat_outputs), output_spec) diff --git a/xtuner/v1/module/__init__.py b/xtuner/v1/module/__init__.py index 8dbecfc75c..2f1d59e674 100644 --- a/xtuner/v1/module/__init__.py +++ b/xtuner/v1/module/__init__.py @@ -1,7 +1,5 @@ from .attention import ( AttnOutputs, - DSAMLAConfig, - DSAMultiLatentAttention, GatedDeltaNet, GatedDeltaNetConfig, MHAConfig, @@ -28,10 +26,8 @@ "RMSNorm", "MultiHeadAttention", "MultiLatentAttention", - "DSAMultiLatentAttention", "MHAConfig", "MLAConfig", - "DSAMLAConfig", "GatedDeltaNetConfig", "GatedDeltaNet", "AttnOutputs", diff --git a/xtuner/v1/module/attention/__init__.py b/xtuner/v1/module/attention/__init__.py index dedd2b1451..d4594014a3 100644 --- a/xtuner/v1/module/attention/__init__.py +++ b/xtuner/v1/module/attention/__init__.py @@ -1,6 +1,5 @@ # Copyright (c) OpenMMLab. All rights reserved. from .attn_outputs import AttnOutputs -from .dsa_mla import DSAMLAConfig, DSAMultiLatentAttention from .gated_deltanet import GatedDeltaNet, GatedDeltaNetConfig from .mha import MHAConfig, MultiHeadAttention from .mla import MLAConfig, MultiLatentAttention @@ -8,11 +7,9 @@ __all__ = [ "MultiLatentAttention", - "DSAMultiLatentAttention", "MultiHeadAttention", "MHAConfig", "MLAConfig", - "DSAMLAConfig", "AttnOutputs", "GatedDeltaNet", "GatedDeltaNetConfig", diff --git a/xtuner/v1/module/attention/dsa_topk_sharing.py b/xtuner/v1/module/attention/dsa_topk_sharing.py deleted file mode 100644 index 7e0ca4fd27..0000000000 --- a/xtuner/v1/module/attention/dsa_topk_sharing.py +++ /dev/null @@ -1,527 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -import os -from dataclasses import dataclass -from functools import partial -from typing import Any, Callable, Protocol, cast - -import torch - -from xtuner.v1.data_proto import SequenceContext -from xtuner.v1.utils.activation_offload import OffloadManager, SwapTensor - - -class DSATopKSharingLayerProtocol(Protocol): - layer_idx: int - source_layer_idx: int - training: bool - indexer_types: list[str] | None - index_skip_topk_offset: int - index_topk_freq: int - dsa_topk_last_use: dict[int, int] - dsa_topk_recompute_release: dict[int, int] - - -@dataclass(frozen=True) -class DSATopKReleasePlan: - forward_last_use: dict[int, int] - recompute_release: dict[int, int] - - -def dsa_topk_source_layer( - *, - layer_idx: int, - indexer_types: list[str] | None, - index_skip_topk_offset: int, - index_topk_freq: int, -) -> int: - """Resolve the physical indexer source for one logical DSA layer.""" - if indexer_types is not None: - if layer_idx < len(indexer_types) and indexer_types[layer_idx] == "full": - return layer_idx - for source_layer_idx in range(min(layer_idx, len(indexer_types) - 1), -1, -1): - if indexer_types[source_layer_idx] == "full": - return source_layer_idx - raise ValueError(f"DSA layer {layer_idx} has no preceding full indexer layer.") - - if index_topk_freq <= 1: - return layer_idx - - source_layer_idx = layer_idx - while (max(source_layer_idx + 1 - index_skip_topk_offset, 0) % index_topk_freq) != 0: - source_layer_idx -= 1 - return source_layer_idx - - -def _dsa_topk_offload_enabled() -> bool: - override = os.getenv("XTUNER_DSA_TOPK_OFFLOAD") - if override is not None: - return int(override) == 1 - # DSA top-k cache is consumed by SparseMLA backward. Keep this offload path - # opt-in instead of coupling it to hidden-state activation offload. - return False - - -def build_dsa_topk_release_plan( - *, - num_main_layers: int, - num_mtp_layers: int, - indexer_types: list[str] | None, - index_skip_topk_offset: int, - index_topk_freq: int, -) -> DSATopKReleasePlan: - consumers: dict[int, list[int]] = {} - for layer_idx in range(num_main_layers + num_mtp_layers): - source_layer_idx = dsa_topk_source_layer( - layer_idx=layer_idx, - indexer_types=indexer_types, - index_skip_topk_offset=index_skip_topk_offset, - index_topk_freq=index_topk_freq, - ) - consumers.setdefault(source_layer_idx, []).append(layer_idx) - - return DSATopKReleasePlan( - forward_last_use={ - source_layer_idx: max(consumer_layers) for source_layer_idx, consumer_layers in consumers.items() - }, - recompute_release={ - source_layer_idx: min(consumer_layers) for source_layer_idx, consumer_layers in consumers.items() - }, - ) - - -class GpuTopKResidency: - def has_cache(self, seq_ctx: SequenceContext, source_layer_idx: int) -> bool: - return source_layer_idx in seq_ctx.dsa_topk_cache.indices - - def store_gpu(self, seq_ctx: SequenceContext, source_layer_idx: int, topk_indices: torch.Tensor) -> None: - seq_ctx.dsa_topk_cache.indices[source_layer_idx] = topk_indices - - def read(self, seq_ctx: SequenceContext, source_layer_idx: int) -> torch.Tensor: - return seq_ctx.dsa_topk_cache.indices[source_layer_idx] - - def after_original_forward_last_use(self, seq_ctx: SequenceContext, source_layer_idx: int) -> None: - return - - def after_recompute_release(self, seq_ctx: SequenceContext, source_layer_idx: int) -> None: - seq_ctx.dsa_topk_cache.indices.pop(source_layer_idx, None) - - def _offload_key(self, seq_ctx: SequenceContext, source_layer_idx: int) -> str: - return f"dsa_topk_{seq_ctx.dsa_topk_cache.offload_slot}_{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) - # A slot/source pair owns both the runtime entry and its reusable pinned - # storage. Reusing it while still live would overwrite the earlier D2H - # result, so fail loudly if the scheduling contract is violated. - if OffloadManager().has_runtime_key(key): - raise RuntimeError(f"DSA top-k offload slot is still active: {key}") - 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 426e353b92..d38f13abb8 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, TypedDict import torch import torch.nn as nn @@ -14,6 +14,27 @@ from ..linear import build_linear +class DenseDecoderLayerOutput(TypedDict): + """Per-micro-batch outputs of one :class:`DenseDecoderLayer` forward. + + A dense layer only produces hidden states, but it reports them through the same keyed contract + as :class:`~xtuner.v1.module.decoder_layer.moe_decoder_layer.MoEDecoderLayer` so that the two + layer families stay interchangeable to their callers. + """ + + hidden_states: torch.Tensor + + +class DenseDecoderLayerMicroBatchOutput(TypedDict): + """Outputs of one :class:`DenseDecoderLayer` forward over several micro- + batches. + + Each field holds one entry per micro-batch, in input order. + """ + + hidden_states: list[torch.Tensor] + + class DenseMLP(nn.Module): def __init__( self, @@ -74,36 +95,53 @@ def __init__( def forward( self, - *hidden_states: torch.Tensor, + hidden_states: torch.Tensor | list[torch.Tensor], + *, position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], seq_ctx: SequenceContext | list[SequenceContext], - ) -> torch.Tensor | tuple[torch.Tensor, ...]: + ) -> DenseDecoderLayerOutput | DenseDecoderLayerMicroBatchOutput: """Run equal-shaped training micro-batches in one layer invocation. - 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``. + Keeping the micro-batch loop inside the decoder layer lets outer FSDP and checkpointing + materialize the layer only once, while each attention call keeps its own + ``SequenceContext``. + + Args: + hidden_states (torch.Tensor | list[torch.Tensor]): Input hidden states, one tensor per + micro-batch. + position_embeddings (tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]]): + Rotary position embeddings ``(cos, sin)``, aligned with ``hidden_states``. + seq_ctx (SequenceContext | list[SequenceContext]): Sequence context, aligned with + ``hidden_states``. + Returns: + DenseDecoderLayerOutput | DenseDecoderLayerMicroBatchOutput: Output hidden states. A + single tensor for a single ``hidden_states`` tensor, a per-micro-batch list for a list + of them. """ - if len(hidden_states) == 1: + if not isinstance(hidden_states, list): assert isinstance(position_embeddings, tuple) and len(position_embeddings) == 2 assert isinstance(seq_ctx, SequenceContext) - return self._forward( - hidden_states=hidden_states[0], - position_embeddings=position_embeddings, - seq_ctx=seq_ctx, - ) + return { + "hidden_states": self._forward( + hidden_states=hidden_states, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + ) + } assert isinstance(position_embeddings, list) and len(position_embeddings) == len(hidden_states) assert isinstance(seq_ctx, list) and len(seq_ctx) == len(hidden_states) assert all(hidden.shape == hidden_states[0].shape for hidden in hidden_states) - return tuple( - self._forward( - hidden_states=hidden, - position_embeddings=position_embedding, - seq_ctx=context, - ) - for hidden, position_embedding, context in zip(hidden_states, position_embeddings, seq_ctx) - ) + return { + "hidden_states": [ + self._forward( + hidden_states=hidden, + position_embeddings=position_embedding, + seq_ctx=context, + ) + for hidden, position_embedding, context in zip(hidden_states, position_embeddings, seq_ctx) + ] + } def _forward( self, diff --git a/xtuner/v1/module/decoder_layer/moe_decoder_layer.py b/xtuner/v1/module/decoder_layer/moe_decoder_layer.py index 862d09c030..64dbb2f37a 100644 --- a/xtuner/v1/module/decoder_layer/moe_decoder_layer.py +++ b/xtuner/v1/module/decoder_layer/moe_decoder_layer.py @@ -1,5 +1,5 @@ from functools import partial -from typing import Literal, Protocol, TypeAlias, cast +from typing import Callable, Literal, Protocol, TypeAlias, TypedDict, cast import torch import torch.nn as nn @@ -47,6 +47,28 @@ HiddenStates: TypeAlias = torch.Tensor +class MoEDecoderLayerOutput(TypedDict): + """Per-micro-batch outputs of one :class:`MoEDecoderLayer` forward.""" + + hidden_states: HiddenStates + router_logits: RouterLogits + router_weights: RouterWeights + router_topk_ids: RouterTopKIds + + +class MoEDecoderLayerMicroBatchOutput(TypedDict): + """Outputs of one :class:`MoEDecoderLayer` forward over several micro- + batches (domino EP). + + Each field holds one entry per micro-batch, in input order. + """ + + hidden_states: list[HiddenStates] + router_logits: list[RouterLogits] + router_weights: list[RouterWeights] + router_topk_ids: list[RouterTopKIds] + + class MoEActFnProtocol(Protocol): def __call__(self, fused_x: torch.Tensor, split_dim: int = -1) -> torch.Tensor: ... @@ -309,23 +331,31 @@ def __init__( def forward( self, - *hidden_states: torch.Tensor, + hidden_states: torch.Tensor | list[torch.Tensor], + *, seq_ctx: SequenceContext | list[SequenceContext], - position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]] | None = None, - ) -> tuple[HiddenStates, RouterLogits, RouterWeights, RouterTopKIds] | tuple[torch.Tensor, ...]: + position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], + ) -> MoEDecoderLayerOutput | MoEDecoderLayerMicroBatchOutput: """Forward pass of the MoE decoder layer. + Passing lists runs several equal-shaped micro-batches in one layer invocation (domino EP), + so that the expert dispatch/combine communication of one micro-batch overlaps the expert + compute of another. + Args: - hidden_states (torch.Tensor): Input hidden states. - seq_ctx (SequenceContext): Sequence context. - position_embeddings (tuple[torch.Tensor, torch.Tensor]): Position embeddings. - past_key_values (list[list[torch.Tensor]], optional): Past key values for pre-filling or decoding. + hidden_states (torch.Tensor | list[torch.Tensor]): Input hidden states, one tensor per + micro-batch. + seq_ctx (SequenceContext | list[SequenceContext]): Sequence context, aligned with + ``hidden_states``. + position_embeddings (tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]]): + Rotary position embeddings ``(cos, sin)``, aligned with ``hidden_states``. Returns: - tuple: Output hidden states, router logits, router weights, and the - expert IDs selected by the router. + MoEDecoderLayerOutput | MoEDecoderLayerMicroBatchOutput: Hidden states and router + results. Scalar fields for a single ``hidden_states`` tensor, per-micro-batch lists for + a list of them. """ - if len(hidden_states) == 1: + if not isinstance(hidden_states, list): assert isinstance(seq_ctx, SequenceContext), ( f"seq_ctx should be a SequenceContext instance but got {seq_ctx}" ) @@ -333,23 +363,23 @@ def forward( "position_embeddings should be a tuple of two tensors (position_ids, position_embeds)" ) return self._forward( - hidden_states=hidden_states[0], + hidden_states=hidden_states, seq_ctx=seq_ctx, position_embeddings=position_embeddings, ) - else: - assert isinstance(seq_ctx, list) and len(seq_ctx) == len(hidden_states), ( - "seq_ctx should be a list of SequenceContext instances with the same length as hidden_states" - ) - assert isinstance(position_embeddings, list) and len(position_embeddings) == len(hidden_states), ( - "position_embeddings should be a list of tuples with the same length as hidden_states" - ) - return self._micro_batch_forward( - hidden_states_list=list(hidden_states), - seq_ctx_list=seq_ctx, - position_embeddings_list=position_embeddings, - ) + assert isinstance(seq_ctx, list) and len(seq_ctx) == len(hidden_states), ( + "seq_ctx should be a list of SequenceContext instances with the same length as hidden_states" + ) + assert isinstance(position_embeddings, list) and len(position_embeddings) == len(hidden_states), ( + "position_embeddings should be a list of tuples with the same length as hidden_states" + ) + + return self._micro_batch_forward( + hidden_states_list=hidden_states, + seq_ctx_list=seq_ctx, + position_embeddings_list=position_embeddings, + ) def _hf_expert_forward_for_debug(self, hidden_states: torch.Tensor, router_results: RouterResults, origin_shape): # xtuner: num_experts * 2 * expert_dim, hidden_size @@ -394,12 +424,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( + attention_kwargs: dict[str, object] | None = None, + ) -> MoEDecoderLayerOutput: + residual, hidden_states, router_results, attn_outputs = self._pre_moe_forward( hidden_states=hidden_states, seq_ctx=seq_ctx, position_embeddings=position_embeddings, state=ForwardState.TRAINING, + attention_kwargs=attention_kwargs, ) origin_shape = hidden_states.shape @@ -429,6 +461,11 @@ def _forward( # post_dispatched.get("row_ids_map"), # type: ignore[arg-type] # dispatched["topk_weights"], # ) + if self.ep_mesh is not None: + # MoEBlock is fullgraph-compiled and shared by all decoder layers. Only the routed-token + # dimension varies, so make it dynamic before entering the compile boundary to keep one + # AOT Autograd save plan across the original forward and checkpoint replay. + torch._dynamo.mark_dynamic(post_dispatched["hidden_states"], 0) experts_out = self.experts( post_dispatched["hidden_states"], post_dispatched["tokens_per_expert"], @@ -480,19 +517,34 @@ def _forward( residual=residual, shared_experts_out=shared_experts_out, ) - return ( - hidden_states, - router_results["logits"], - router_results["router_weights"], - router_results["topk_ids"], + return self._build_output( + hidden_states=hidden_states, + router_results=router_results, + attn_outputs=attn_outputs, ) + def _build_output( + self, + *, + hidden_states: torch.Tensor, + router_results: RouterResults, + attn_outputs: AttnOutputs, + ) -> MoEDecoderLayerOutput: + """Build the public output; model-specific decoders may extend it.""" + return { + "hidden_states": hidden_states, + "router_logits": router_results["logits"], + "router_weights": router_results["router_weights"], + "router_topk_ids": router_results["topk_ids"], + } + def _micro_batch_forward( self, hidden_states_list: list[torch.Tensor], seq_ctx_list: list[SequenceContext], position_embeddings_list: list[tuple[torch.Tensor, torch.Tensor]], - ) -> tuple[torch.Tensor, ...]: + attention_kwargs_list: list[dict[str, object]] | None = None, + ) -> MoEDecoderLayerMicroBatchOutput: origin_shape = hidden_states_list[0].shape assert all(hidden_states.shape == origin_shape for hidden_states in hidden_states_list), ( "All hidden states should have the same shape" @@ -500,6 +552,10 @@ def _micro_batch_forward( intra_layer_micro_batch = len(hidden_states_list) residual_list: list[torch.Tensor] = [] router_results_list: list[RouterResults] = [] + attn_outputs_list: list[AttnOutputs] = [] + if attention_kwargs_list is None: + attention_kwargs_list = [{} for _ in hidden_states_list] + assert len(attention_kwargs_list) == intra_layer_micro_batch pre_dispatched_list: list[PreDispatchResult] = [] dispatched_list: list[DispatchResult] = [] @@ -508,18 +564,21 @@ def _micro_batch_forward( # Attention + gate + pre-dispatch for ( hidden_states, + attention_kwargs, seq_ctx, position_embeddings, ) in zip( hidden_states_list, + attention_kwargs_list, seq_ctx_list, position_embeddings_list, ): - residual, hidden_states, router_results = self._pre_moe_forward( + residual, hidden_states, router_results, attn_outputs = self._pre_moe_forward( hidden_states=hidden_states, seq_ctx=seq_ctx, position_embeddings=position_embeddings, state=ForwardState.TRAINING, + attention_kwargs=attention_kwargs, ) pre_moe_forward_out_list.append(hidden_states) hidden_states = hidden_states.view(-1, hidden_states.shape[-1]) @@ -532,6 +591,7 @@ def _micro_batch_forward( pre_dispatched_list.append(pre_dispatched) residual_list.append(residual) router_results_list.append(router_results) + attn_outputs_list.append(attn_outputs) post_dispatched_list: list[PostDispatchResult] = [] experts_out_list: list[torch.Tensor] = [] @@ -554,6 +614,9 @@ def _micro_batch_forward( dispatched=dispatched, async_op=True, ) + if self.ep_mesh is not None: + # Preserve the same dynamic-token compile contract for every in-layer micro-batch. + torch._dynamo.mark_dynamic(post_dispatched["hidden_states"], 0) experts_out = self.experts( post_dispatched["hidden_states"], post_dispatched["tokens_per_expert"], @@ -618,10 +681,27 @@ def _micro_batch_forward( ) hidden_states_out_list.append(hidden_states) - 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) + return self._build_micro_batch_output( + hidden_states_list=hidden_states_out_list, + router_results_list=router_results_list, + attn_outputs_list=attn_outputs_list, + ) + + def _build_micro_batch_output( + self, + *, + hidden_states_list: list[torch.Tensor], + router_results_list: list[RouterResults], + attn_outputs_list: list[AttnOutputs], + ) -> MoEDecoderLayerMicroBatchOutput: + """Build the public micro-batch output; model-specific decoders may + extend it.""" + return { + "hidden_states": hidden_states_list, + "router_logits": [router_results["logits"] for router_results in router_results_list], + "router_weights": [router_results["router_weights"] for router_results in router_results_list], + "router_topk_ids": [router_results["topk_ids"] for router_results in router_results_list], + } def _pre_moe_forward( self, @@ -630,7 +710,8 @@ def _pre_moe_forward( position_embeddings: tuple[torch.Tensor, torch.Tensor], state: ForwardState, past_key_values: list[list[torch.Tensor]] | None = None, - ) -> tuple[torch.Tensor, torch.Tensor, RouterResults]: + attention_kwargs: dict[str, object] | None = None, + ) -> tuple[torch.Tensor, torch.Tensor, RouterResults, AttnOutputs]: # NOTE: In order to allow `torch.compile` to compile the ops before and after attention as much as possible, # attention, post-layernorm and gate are implemented in one function residual = hidden_states @@ -638,13 +719,16 @@ def _pre_moe_forward( # Self Attention if state == ForwardState.TRAINING: - attn_outputs: AttnOutputs = self.self_attn( + attention_forward = cast(Callable[..., AttnOutputs], self.self_attn) + attn_outputs = attention_forward( hidden_states=hidden_states, position_embeddings=position_embeddings, seq_ctx=seq_ctx, + **(attention_kwargs or {}), ) hidden_states = attn_outputs["projected_output"] elif state == ForwardState.PREFILLING: + attn_outputs = {} assert past_key_values is not None, "past_key_values should be provided in pre-filling state" hidden_states = self.self_attn.prefilling( # type: ignore hidden_states=hidden_states, @@ -653,6 +737,7 @@ def _pre_moe_forward( past_key_values=past_key_values, ) elif state == ForwardState.DECODING: + attn_outputs = {} assert past_key_values is not None, "past_key_values should be provided in decoding state" hidden_states = self.self_attn.decoding( # type: ignore hidden_states=hidden_states, @@ -676,7 +761,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, attn_outputs def _shared_experts_forward( self, diff --git a/xtuner/v1/module/mtp/mtp_block.py b/xtuner/v1/module/mtp/mtp_block.py index 9d43f685e0..f13f1c2235 100644 --- a/xtuner/v1/module/mtp/mtp_block.py +++ b/xtuner/v1/module/mtp/mtp_block.py @@ -1,18 +1,26 @@ """Multi-Token Prediction (MTP) Block implementation.""" -from typing import Callable +from typing import Callable, cast import torch import torch.nn as nn from xtuner.v1.data_proto import SequenceContext +from xtuner.v1.module.decoder_layer.moe_decoder_layer import ( + MoEDecoderLayerMicroBatchOutput, + MoEDecoderLayerOutput, +) from .config import MTPConfig from .mtp_layer import MTPLayer from .utils import roll_sequence_context -MTPDepthOutput = tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] +MTPDepthOutput = MoEDecoderLayerOutput +"""One MTP depth produces the same keyed outputs as the decoder layer it +wraps.""" + +MTPInternalOutput = MoEDecoderLayerOutput | MoEDecoderLayerMicroBatchOutput class MTPBlock(nn.Module): @@ -62,12 +70,12 @@ class MTPBlock(nn.Module): >>> >>> # Multi-microbatch (domino EP) forward >>> outputs_per_mb = mtp_block( - ... h0, h1, + ... [h0, h1], ... embed_tokens_fn=embed_fn, ... position_embeddings=[pos_emb_0, pos_emb_1], ... seq_ctx=[ctx_0, ctx_1], ... ) - >>> # outputs_per_mb[mb_idx][depth_idx] -> (hidden, router_logits, router_weights, router_topk_ids) + >>> # outputs_per_mb[mb_idx][depth_idx] -> MTPDepthOutput """ def __init__(self, *, mtp_config: MTPConfig, mtp_layers: list[MTPLayer]): @@ -84,7 +92,8 @@ def __init__(self, *, mtp_config: MTPConfig, mtp_layers: list[MTPLayer]): def forward( self, - *hidden_states: torch.Tensor, + hidden_states: torch.Tensor | list[torch.Tensor], + *, embed_tokens_fn: Callable[[torch.Tensor], torch.Tensor], position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], seq_ctx: SequenceContext | list[SequenceContext], @@ -97,9 +106,8 @@ def forward( the wrapped decoder layer can be overlapped across micro-batches (domino EP). Args: - hidden_states (torch.Tensor): One or more hidden state tensors from the main - model, shape ``[batch, seq_len, hidden_size]`` each. Single tensor → single- - microbatch path; multiple tensors → multi-microbatch (domino EP) path. + hidden_states (torch.Tensor | list[torch.Tensor]): Hidden states from the main model, + shape ``[batch, seq_len, hidden_size]`` each, one tensor per micro-batch. embed_tokens_fn (Callable): Function to embed tokens. Takes token IDs and returns embeddings. Should have signature ``embed_tokens_fn(token_ids: Tensor) -> Tensor``. position_embeddings (tuple | list[tuple]): Rotary position embeddings (cos, sin), @@ -107,14 +115,11 @@ def forward( seq_ctx (SequenceContext | list[SequenceContext]): Sequence context per micro-batch. Returns: - list: For single-microbatch input, - ``list[(hidden, router_logits, router_weights, router_topk_ids)]`` - of length ``D``, where ``outputs[k]`` is the prediction for token ``i+k+1``. - For ``N`` micro-batches, - ``list[list[(hidden, router_logits, router_weights, router_topk_ids)]]`` - with outer length ``N`` and inner length ``D``: ``outputs[mb_idx][depth_idx]``. + list[MTPDepthOutput] | list[list[MTPDepthOutput]]: For a single ``hidden_states`` + tensor, one entry per MTP depth ``D``, where ``outputs[k]`` is the prediction for + token ``i+k+1``. For ``N`` micro-batches, ``outputs[mb_idx][depth_idx]``. """ - if len(hidden_states) == 1: + if not isinstance(hidden_states, list): assert isinstance(seq_ctx, SequenceContext), ( "seq_ctx should be a SequenceContext instance in single-microbatch mode" ) @@ -122,7 +127,7 @@ def forward( "position_embeddings should be a (cos, sin) tuple in single-microbatch mode" ) return self._forward( - hidden_states=hidden_states[0], + hidden_states=hidden_states, embed_tokens_fn=embed_tokens_fn, position_embeddings=position_embeddings, seq_ctx=seq_ctx, @@ -136,7 +141,7 @@ def forward( "position_embeddings should be a list aligned with hidden_states in multi-microbatch mode" ) return self._micro_batch_forward( - hidden_states_list=list(hidden_states), + hidden_states_list=hidden_states, embed_tokens_fn=embed_tokens_fn, position_embeddings_list=position_embeddings, seq_ctx_list=seq_ctx, @@ -164,9 +169,7 @@ def _forward( attention mask, etc. Returns: - list[MTPDepthOutput]: List of 4-tuples - (hidden_states, router_logits, router_weights, router_topk_ids) - for each MTP depth. + list[MTPDepthOutput]: One entry per MTP depth. Length equals num_layers. - outputs[0]: Outputs for predicting token at position (i+1) - outputs[k]: Outputs for predicting token at position (i+k+1) @@ -174,10 +177,11 @@ def _forward( mtp_outputs: list[MTPDepthOutput] = [] current_hidden_states = hidden_states.detach() if self.mtp_config.detach_mtp_inputs else hidden_states current_seq_ctx = seq_ctx + previous_layer_results: MTPInternalOutput | None = None num_steps = self.mtp_config.num_layers for step in range(num_steps): - layer = self.layers[0] if self.mtp_config.share_weights else self.layers[step] + layer = cast(MTPLayer, self.layers[0] if self.mtp_config.share_weights else self.layers[step]) # Roll each packed sequence independently so we get the (i+k)-th token while # respecting per-sequence boundaries inside the packed batch. current_seq_ctx = roll_sequence_context(current_seq_ctx, shifts=-1) @@ -186,16 +190,52 @@ 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, - future_embeddings=future_embeddings, - position_embeddings=position_embeddings, - seq_ctx=current_seq_ctx, + layer_results = cast( + MoEDecoderLayerOutput, + self._call_decoder_layer( + layer, + current_hidden_states, + previous_layer_results=previous_layer_results, + future_embeddings=future_embeddings, + position_embeddings=position_embeddings, + seq_ctx=current_seq_ctx, + ), + ) + previous_layer_results = layer_results + current_hidden_states = layer_results["hidden_states"] + mtp_outputs.append( + { + "hidden_states": current_hidden_states, + "router_logits": layer_results["router_logits"], + "router_weights": layer_results["router_weights"], + "router_topk_ids": layer_results["router_topk_ids"], + } ) - mtp_outputs.append((current_hidden_states, router_logits, router_weights, router_topk_ids)) return mtp_outputs + def _call_decoder_layer( + self, + layer: MTPLayer, + hidden_states: torch.Tensor | list[torch.Tensor], + *, + previous_layer_results: MTPInternalOutput | None, + future_embeddings: torch.Tensor | list[torch.Tensor], + position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], + seq_ctx: SequenceContext | list[SequenceContext], + ) -> MTPInternalOutput: + """Call one MTP decoder layer. + + Subclasses can consume model-specific fields from ``previous_layer_results`` while the + generic MTP input and public output stay model-agnostic. + """ + return layer( + hidden_states, + future_embeddings=future_embeddings, + position_embeddings=position_embeddings, + seq_ctx=seq_ctx, + ) + def _micro_batch_forward( self, *, @@ -210,35 +250,39 @@ def _micro_batch_forward( outputs_per_mb: list[list[MTPDepthOutput]] = [[] for _ in range(n)] current_hidden_states_list = list(hidden_states_list) current_seq_ctx_list = list(seq_ctx_list) + previous_layer_results: MTPInternalOutput | None = None num_steps = self.mtp_config.num_layers for step in range(num_steps): - layer = self.layers[0] if self.mtp_config.share_weights else self.layers[step] + layer = cast(MTPLayer, self.layers[0] if self.mtp_config.share_weights else self.layers[step]) current_seq_ctx_list = [roll_sequence_context(ctx, shifts=-1) for ctx in current_seq_ctx_list] future_embeddings_list = [self._embed_future(ctx, embed_tokens_fn) for ctx in current_seq_ctx_list] - layer_results = layer( - *current_hidden_states_list, - 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}, " - f"got {len(layer_results) if isinstance(layer_results, tuple) else type(layer_results)}" + layer_results = cast( + MoEDecoderLayerMicroBatchOutput, + self._call_decoder_layer( + layer, + current_hidden_states_list, + previous_layer_results=previous_layer_results, + future_embeddings=future_embeddings_list, + position_embeddings=position_embeddings_list, + seq_ctx=current_seq_ctx_list, + ), ) - 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 :]) + previous_layer_results = layer_results for mb_idx in range(n): outputs_per_mb[mb_idx].append( - (new_hidden[mb_idx], router_logits[mb_idx], router_weights[mb_idx], router_topk_ids[mb_idx]) + { + "hidden_states": layer_results["hidden_states"][mb_idx], + "router_logits": layer_results["router_logits"][mb_idx], + "router_weights": layer_results["router_weights"][mb_idx], + "router_topk_ids": layer_results["router_topk_ids"][mb_idx], + } ) - current_hidden_states_list = new_hidden + current_hidden_states_list = layer_results["hidden_states"] return outputs_per_mb diff --git a/xtuner/v1/module/mtp/mtp_layer.py b/xtuner/v1/module/mtp/mtp_layer.py index 711c7f4b29..4106912756 100644 --- a/xtuner/v1/module/mtp/mtp_layer.py +++ b/xtuner/v1/module/mtp/mtp_layer.py @@ -7,6 +7,10 @@ from xtuner.v1.data_proto import SequenceContext from xtuner.v1.module import RMSNorm +from xtuner.v1.module.decoder_layer.moe_decoder_layer import ( + MoEDecoderLayerMicroBatchOutput, + MoEDecoderLayerOutput, +) from xtuner.v1.module.linear import build_linear @@ -81,40 +85,33 @@ def __init__( def forward( self, - *hidden_states: torch.Tensor, + hidden_states: torch.Tensor | list[torch.Tensor], + *, future_embeddings: torch.Tensor | list[torch.Tensor], position_embeddings: tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]], seq_ctx: SequenceContext | list[SequenceContext], - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] | tuple[torch.Tensor, ...]: + ) -> MoEDecoderLayerOutput | MoEDecoderLayerMicroBatchOutput: """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`: passing lists runs ``N`` micro-batches together + (intra-layer micro-batching / domino EP). 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. Args: - hidden_states (torch.Tensor): One or more hidden state tensors. A single tensor - triggers the single-microbatch path; multiple tensors trigger the - multi-microbatch path. - future_embeddings (torch.Tensor | list[torch.Tensor]): Embeddings of the future - tokens, aligned per-microbatch with ``hidden_states``. - position_embeddings (tuple | list[tuple]): Rotary position embeddings (cos, sin), - aligned per-microbatch with ``hidden_states``. - seq_ctx (SequenceContext | list[SequenceContext]): Sequence context per micro-batch. - + hidden_states (torch.Tensor | list[torch.Tensor]): Hidden states, one tensor per + micro-batch. + future_embeddings (torch.Tensor | list[torch.Tensor]): Embeddings of the future tokens, + aligned with ``hidden_states``. + position_embeddings (tuple[torch.Tensor, torch.Tensor] | list[tuple[torch.Tensor, torch.Tensor]]): + Rotary position embeddings ``(cos, sin)``, aligned with ``hidden_states``. + seq_ctx (SequenceContext | list[SequenceContext]): Sequence context, aligned with + ``hidden_states``. 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})``. + MoEDecoderLayerOutput | MoEDecoderLayerMicroBatchOutput: The wrapped decoder layer's + outputs with the MTP final layernorm applied to the hidden states. """ - if len(hidden_states) == 1: + if not isinstance(hidden_states, list): assert isinstance(future_embeddings, torch.Tensor), ( "future_embeddings should be a Tensor in single-microbatch mode" ) @@ -125,7 +122,7 @@ def forward( "position_embeddings should be a (cos, sin) tuple in single-microbatch mode" ) return self._forward( - hidden_states=hidden_states[0], + hidden_states=hidden_states, future_embeddings=future_embeddings, position_embeddings=position_embeddings, seq_ctx=seq_ctx, @@ -141,7 +138,7 @@ def forward( "position_embeddings should be a list aligned with hidden_states in multi-microbatch mode" ) return self._micro_batch_forward( - hidden_states_list=list(hidden_states), + hidden_states_list=hidden_states, future_embeddings_list=future_embeddings, position_embeddings_list=position_embeddings, seq_ctx_list=seq_ctx, @@ -153,17 +150,20 @@ def _forward( future_embeddings: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], seq_ctx: SequenceContext, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + ) -> MoEDecoderLayerOutput: projected = self._preprocess(hidden_states=hidden_states, future_embeddings=future_embeddings) - hidden_states, router_results, router_weights, router_topk_ids = self.decoder_layer( + layer_results: MoEDecoderLayerOutput = self.decoder_layer( projected, position_embeddings=position_embeddings, seq_ctx=seq_ctx, ) - - hidden_states = self.final_layernorm(hidden_states) - return hidden_states, router_results, router_weights, router_topk_ids + return { + "hidden_states": self.final_layernorm(layer_results["hidden_states"]), + "router_logits": layer_results["router_logits"], + "router_weights": layer_results["router_weights"], + "router_topk_ids": layer_results["router_topk_ids"], + } def _micro_batch_forward( self, @@ -172,7 +172,7 @@ def _micro_batch_forward( future_embeddings_list: list[torch.Tensor], position_embeddings_list: list[tuple[torch.Tensor, torch.Tensor]], seq_ctx_list: list[SequenceContext], - ) -> tuple[torch.Tensor, ...]: + ) -> MoEDecoderLayerMicroBatchOutput: n = len(hidden_states_list) assert len(future_embeddings_list) == n and len(position_embeddings_list) == n and len(seq_ctx_list) == n, ( "All per-microbatch inputs must share the same length" @@ -184,23 +184,17 @@ def _micro_batch_forward( self._preprocess(hidden_states=h, future_embeddings=e) for h, e in zip(hidden_states_list, future_embeddings_list) ] - - layer_results = self.decoder_layer( - *projected_list, + layer_results: MoEDecoderLayerMicroBatchOutput = self.decoder_layer( + projected_list, position_embeddings=position_embeddings_list, seq_ctx=seq_ctx_list, ) - assert isinstance(layer_results, tuple) and len(layer_results) == 4 * 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}; " - 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) + return { + "hidden_states": [self.final_layernorm(hidden) for hidden in layer_results["hidden_states"]], + "router_logits": layer_results["router_logits"], + "router_weights": layer_results["router_weights"], + "router_topk_ids": layer_results["router_topk_ids"], + } def _preprocess( self, diff --git a/xtuner/v1/profiler/prober.py b/xtuner/v1/profiler/prober.py index e3555c6642..c393e78b13 100644 --- a/xtuner/v1/profiler/prober.py +++ b/xtuner/v1/profiler/prober.py @@ -312,9 +312,11 @@ def wrapped_forward(self, *args, **kwargs): hidden_states = kwargs["hidden_states"] ProberList.before_layer(name, hidden_states) outputs = forward(*args, **kwargs) - if isinstance(outputs, tuple): # for MoEDecoderLayer + if isinstance(outputs, dict): + hidden_states = outputs["hidden_states"] + elif isinstance(outputs, tuple): # for legacy decoder layers hidden_states = outputs[0] - else: # for DenseDecoderLayer + else: hidden_states = outputs ProberList.after_layer(name, hidden_states) return outputs