From 8ca566f781b01f77cf7c5d9d5f5efee4c0ce0cef Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Mon, 10 Aug 2026 20:16:18 +1000 Subject: [PATCH 01/11] relax cp chunkwise intergration --- relax/backends/megatron/arguments.py | 37 +++ relax/backends/megatron/model.py | 158 ++--------- relax/backends/megatron/model_provider.py | 1 + .../megatron/test_gdn_chunkwise_cp_layout.py | 256 ++++++++++++++++++ .../megatron/test_gdn_cp_mode_stage2.py | 229 ++++++++++++++++ 5 files changed, 551 insertions(+), 130 deletions(-) create mode 100644 tests/backends/megatron/test_gdn_chunkwise_cp_layout.py create mode 100644 tests/backends/megatron/test_gdn_cp_mode_stage2.py diff --git a/relax/backends/megatron/arguments.py b/relax/backends/megatron/arguments.py index ca8d21e60..96a66258e 100644 --- a/relax/backends/megatron/arguments.py +++ b/relax/backends/megatron/arguments.py @@ -128,6 +128,41 @@ def _validate_dynamic_context_parallel(args): args.max_seqlen_per_dp_cp_rank = args.max_tokens_per_gpu +def _validate_linear_cp_mode(args) -> None: + """Fail fast on `--linear-cp-mode` / flag combinations that are invalid for + every model, without needing the HF config. + + Geometry-dependent rejections (e.g. explicit `headwise` on heads not + divisible by `tp*max_cp`) can only be checked once the real GDN head counts + are known, which happens in MCore's `TransformerConfig.__post_init__` gate + -- not here. + """ + mode = getattr(args, "linear_cp_mode", "headwise") + allowed_modes = {"headwise", "chunkwise", "all_gather"} + if mode not in allowed_modes: + raise ValueError( + f"--linear-cp-mode must be one of {sorted(allowed_modes)!r}; got {mode!r}. v1 does not support 'auto'." + ) + + if mode == "chunkwise" and getattr(args, "allgather_cp", False): + raise ValueError( + "--linear-cp-mode=chunkwise is incompatible with --allgather-cp: chunkwise CP requires " + "Megatron's zig-zag THD packing, while --allgather-cp switches the data path to a single " + "contiguous per-rank chunk. Note --allgather-cp is a data/attention packing flag, unrelated " + "to the GDN `all_gather` CP mode." + ) + + cp_may_exceed_one = ( + getattr(args, "dynamic_context_parallel", False) or getattr(args, "context_parallel_size", 1) > 1 + ) + if mode == "chunkwise" and cp_may_exceed_one and getattr(args, "deterministic_mode", False): + raise ValueError( + "--linear-cp-mode=chunkwise does not support --deterministic-mode while CP>1 may occur: " + "the deterministic torch reference path only accepts cp_context=None. Use " + "--linear-cp-mode=headwise or =all_gather for deterministic CP>1 runs." + ) + + def validate_args(args): """Run megatron's own validate_args plus slime-specific megatron validations.""" @@ -174,6 +209,8 @@ class _DeviceProperty: assert args.calculate_per_token_loss, ( "--calculate-per-token-loss must be set when context_parallel_size > 1 or dynamic_context_parallel is enabled (required by Megatron-Bridge)." ) + + _validate_linear_cp_mode(args) return args diff --git a/relax/backends/megatron/model.py b/relax/backends/megatron/model.py index 43ca8c2c0..eb9083712 100644 --- a/relax/backends/megatron/model.py +++ b/relax/backends/megatron/model.py @@ -387,14 +387,6 @@ def setup_model_and_optimizer( assert not args.moe_use_upcycling assert args.load is not None or args.pretrained_checkpoint is not None - # Relax the Megatron GDN head-vs-(tp*cp) config gate down to (tp) BEFORE the model - # provider finalizes the TransformerConfig (get_model_provider_func below triggers - # __post_init__), so high-CP GDN configs (e.g. TP2/CP16) validate. The matching - # forward all-gather path is installed by _patch_gdn_for_dynamic_cp after the model - # is built; see both functions for why % tp suffices (GDN weights are TP-only). - if getattr(args, "dynamic_context_parallel", False) or getattr(args, "context_parallel_size", 1) > 1: - _relax_gdn_cp_config_assert() - model = get_model( wrap_model_provider_with_freeze(get_model_provider_func(args, role), args), ModelType.encoder_or_decoder, @@ -421,6 +413,15 @@ def setup_model_and_optimizer( # (dynamic CP, or static context_parallel_size > 1), incl. weight-only # roles that still run forward. _patch_gdn_for_dynamic_cp() + model_config = get_model_config(model[0]) + if getattr(model_config, "experimental_attention_variant", None) == "gated_delta_net" and ( + not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0 + ): + logger.info( + f"[GDN CP] role={role} linear_cp_mode={model_config.linear_cp_mode} " + f"TP={model_config.tensor_model_parallel_size} max_CP={model_config.context_parallel_size} " + f"key_heads={model_config.linear_num_key_heads} value_heads={model_config.linear_num_value_heads}" + ) if args.only_load_weight: return model, None, None @@ -495,84 +496,13 @@ def _gdn_cp_gather_full(qkvzba, cu_seqlens_cpu, cp_size, cp_group): return gdn_cp_gather_full(qkvzba, cu_seqlens_cpu, cp_size, cp_group) -def _relax_gdn_cp_config_assert() -> None: - """Relax Megatron's GDN config gate ``linear_num_{key,value}_heads % - (tp*cp) == 0`` down to ``% tp`` so high-CP GDN configs (e.g. TP2/CP16) - finalize. - - Megatron's ``TransformerConfig.__post_init__`` enforces the *native* cp2hp - (split-sequence -> split-head) divisibility ``heads % (tp * cp)``. - ``_patch_gdn_for_dynamic_cp`` replaces that forward with an all-gather + duplicated - scan whose weights stay **TP-only** (``qk_dim_local_tp = qk_dim // tp``, etc.), so - only ``heads % tp`` is actually required. Without relaxing this config gate, TP2/CP16 - (16 % 32 != 0) aborts at config finalize (``get_model_provider_func`` -> ``finalize`` - -> ``__post_init__``) *before* the forward patch is installed. - - Only intervenes when the native check would reject but the relaxed ``% tp`` check - passes: it temporarily scales the two GDN head counts by ``cp`` (which preserves - ``value % key`` and makes ``heads % (tp*cp)`` hold), runs the original - ``__post_init__``, then restores them. Those head counts are validation-only in - ``__post_init__`` (no stored value is derived from them -- verified against Megatron - core), and ``GatedDeltaNet.__init__`` reads the restored config later, so nothing - downstream sees the temporary values. Idempotent; monkey-patch only (no upstream - edit), matching ``_patch_gdn_for_dynamic_cp``. - """ - try: - from megatron.core.transformer.transformer_config import TransformerConfig - except ImportError: - return - - if getattr(TransformerConfig, "_gdn_cp_relaxed", False): - return - - _orig_post_init = TransformerConfig.__post_init__ - - def _relaxed_post_init(self, *post_init_args, **post_init_kwargs): - if getattr(self, "experimental_attention_variant", None) == "gated_delta_net": - tp = self.tensor_model_parallel_size - cp = self.context_parallel_size - key = self.linear_num_key_heads or 0 - val = self.linear_num_value_heads or 0 - native_bad = cp > 1 and ((key % (tp * cp)) != 0 or (val % (tp * cp)) != 0) - relaxed_ok = tp > 0 and (key % tp) == 0 and (val % tp) == 0 - if native_bad and relaxed_ok: - # key%tp==0 => (key*cp)%(tp*cp)==0, and (val*cp)%(key*cp)==(val%key) so the - # value%key assert is preserved. Restored in `finally` before anything else - # (incl. GatedDeltaNet.__init__) reads the config. - self.linear_num_key_heads = key * cp - self.linear_num_value_heads = val * cp - try: - _orig_post_init(self, *post_init_args, **post_init_kwargs) - finally: - self.linear_num_key_heads = key - self.linear_num_value_heads = val - return - _orig_post_init(self, *post_init_args, **post_init_kwargs) - - TransformerConfig.__post_init__ = _relaxed_post_init - TransformerConfig._gdn_cp_relaxed = True - - def _patch_gdn_for_dynamic_cp() -> None: - """Monkey-patch GatedDeltaNet.forward for CP via all-gather + duplicated - scan. - - Megatron's native GDN forward implements CP by converting "split sequence" - into "split head" (``cp2hp`` all-to-all, ``num_value_heads // tp // cp``), - which forces ``num_heads % (tp * cp) == 0`` and breaks at high CP for - head-light models (e.g. Qwen3.5). This patch keeps that efficient native path - whenever the heads still divide ``tp * cp`` (``native_ok``), and only when - native would break does it fall back to all-gathering the full sequence across - CP, running the recurrent scan duplicated on each rank while keeping relax's - **TP** head-split intact, then re-slicing this rank's shard. The effective - constraint drops to ``num_heads % tp == 0`` (CP16 works), and weight - conversion / DCS sync / checkpoint (all TP-only) are untouched. - - Dynamic CP: size/group are read per micro-batch from ``packed_seq_params`` - (set in get_batch), falling back to the static CP group. The ``cp == 1``, - non-thd, and ``native_ok`` cases keep upstream behavior (swap the dynamic CP - group, call the original forward). Idempotent; avoids editing upstream - Megatron source. + """Patch GDN forward for dynamic CP and Relax's all-gather mode. + + CP=1 and MCore-native headwise/chunkwise modes call the patched MCore + forward directly. Only static ``linear_cp_mode='all_gather'`` with CP>1 + executes Relax's existing fallback. No shared module/config state is + modified. """ try: from megatron.core.ssm.gated_delta_net import GatedDeltaNet @@ -584,23 +514,6 @@ def _patch_gdn_for_dynamic_cp() -> None: _orig_forward = GatedDeltaNet.forward - def _call_orig_with_dynamic_cp( - self, cp_size, cp_group, hidden_states, attention_mask, inference_context, packed_seq_params, *args, **kwargs - ): - # cp == 1 or non-thd: preserve upstream behavior; just point the module at - # the (possibly dynamic) CP group for the original forward. - _orig_cp_size = self.cp_size - _orig_cp_group = self.pg_collection.cp - self.cp_size = cp_size - self.pg_collection.cp = cp_group - try: - return _orig_forward( - self, hidden_states, attention_mask, inference_context, packed_seq_params, *args, **kwargs - ) - finally: - self.cp_size = _orig_cp_size - self.pg_collection.cp = _orig_cp_group - def _dcp_gdn_forward( self, hidden_states, attention_mask, inference_context=None, packed_seq_params=None, *args, **kwargs ): @@ -610,26 +523,16 @@ def _dcp_gdn_forward( from .cp_utils import gdn_cp_slice cp_size, cp_group, cp_rank = _resolve_gdn_cp(self, packed_seq_params) - is_thd = packed_seq_params is not None and getattr(packed_seq_params, "qkv_format", None) == "thd" - # Native cp2hp (head-split) is exact and cheaper (no duplicated scan, GDN - # activation sharded by CP) whenever the heads divide tp*cp. Only fall back - # to the all-gather path when native would break the head split — i.e. when - # num_key_heads is not divisible by tp*cp (covers tp*cp > num_key_heads). - # num_value_heads is a multiple of num_key_heads, so this one check suffices. - native_ok = self.num_key_heads % (self.tp_size * cp_size) == 0 - if cp_size == 1 or not is_thd or native_ok: - return _call_orig_with_dynamic_cp( - self, - cp_size, - cp_group, - hidden_states, - attention_mask, - inference_context, - packed_seq_params, - *args, - **kwargs, + if cp_size == 1 or self.config.linear_cp_mode != "all_gather": + return _orig_forward( + self, hidden_states, attention_mask, inference_context, packed_seq_params, *args, **kwargs ) + is_thd = packed_seq_params is not None and getattr(packed_seq_params, "qkv_format", None) == "thd" + assert is_thd, ( + "GDN linear_cp_mode='all_gather' with cp_size>1 only supports packed (thd) sequences; " + "use linear_cp_mode='headwise' or 'chunkwise' for SBHD/static-batch inputs." + ) assert inference_context is None, "GDN all-gather CP path does not support inference." # Packed (thd) + deterministic is unsupported: a single conv/scan over the # concatenated samples would bleed state across cu_seqlens boundaries, and @@ -689,19 +592,14 @@ def _dcp_gdn_forward( ) # Reuse the module's own prep (split/l2norm/GQA-expand) with CP disabled so - # its internal `// self.cp_size` becomes a no-op. Wrap in the dynamo-disable + # its internal `// cp_size_headwise` becomes a no-op. Wrap in the dynamo-disable # guard added by docker/patch/megatron/20260506-85bced0ae.patch (Qwen3.6 GDN # torch.compile failure); calling _prepare_qkv_for_gated_delta_rule directly # would re-trigger that compile failure. - _saved_cp = self.cp_size - self.cp_size = 1 - try: - with torch._dynamo.config.patch(disable=True): - query, key, value, gate, beta, alpha = self._prepare_qkv_for_gated_delta_rule( - qkv, gate, beta, alpha, batch, seq_len - ) - finally: - self.cp_size = _saved_cp + with torch._dynamo.config.patch(disable=True): + query, key, value, gate, beta, alpha = self._prepare_qkv_for_gated_delta_rule( + qkv, gate, beta, alpha, batch, seq_len, cp_size_headwise=1 + ) # g/beta from the full (un-CP-sliced) A_log / dt_bias. g, beta = self._compute_g_and_beta(self.A_log, self.dt_bias, alpha, beta) diff --git a/relax/backends/megatron/model_provider.py b/relax/backends/megatron/model_provider.py index 8e5594253..32b80f0e7 100644 --- a/relax/backends/megatron/model_provider.py +++ b/relax/backends/megatron/model_provider.py @@ -289,6 +289,7 @@ def wrapped_model_provider( "pipeline_model_parallel_size", "virtual_pipeline_model_parallel_size", "context_parallel_size", + "linear_cp_mode", "expert_model_parallel_size", "expert_tensor_parallel_size", "variable_seq_lengths", diff --git a/tests/backends/megatron/test_gdn_chunkwise_cp_layout.py b/tests/backends/megatron/test_gdn_chunkwise_cp_layout.py new file mode 100644 index 000000000..a651c8c2a --- /dev/null +++ b/tests/backends/megatron/test_gdn_chunkwise_cp_layout.py @@ -0,0 +1,256 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. +"""Unit tests for the GDN chunkwise-CP layout backport (Task 32, phase 1). + +Covers the pure-tensor half of the backported MCore capability: the two THD CP +partitions (``zigzag`` / ``contiguous``), their agreement with Relax's existing +zigzag sharding, dynamic group resolution, and the construction-time +``linear_cp_mode`` gate. + +Everything here runs on CPU with no process group. It validates partition +definitions only; the actual all-to-all round trip is exercised with NCCL in +``test_gdn_chunkwise_cp_gpu.py``. + +The real-kernel / real-collective half lives in +``test_gdn_chunkwise_cp_gpu.py``. +""" + +from __future__ import annotations + +import inspect + +import pytest +import torch + + +cpl = pytest.importorskip("megatron.core.context_parallel_layout", reason="requires the patched Megatron-LM") + +from megatron.core.packed_seq_params import PackedSeqParams, resolve_cp_group # noqa: E402 + +from relax.backends.megatron.cp_utils import gdn_cp_slice, slice_with_cp # noqa: E402 + + +def _cu(lengths: list[int]) -> torch.Tensor: + cu = [0] + for n in lengths: + cu.append(cu[-1] + n) + return torch.tensor(cu, dtype=torch.int64) + + +def _tagged_tokens(total: int, width: int = 3) -> torch.Tensor: + """[total, width] where row t is (t, t+1e6, t+2e6): token identity is + unambiguous.""" + base = torch.arange(total, dtype=torch.float64).unsqueeze(1) + return base + torch.arange(width, dtype=torch.float64).unsqueeze(0) * 1e6 + + +# --------------------------------------------------------------------------- +# Partition definitions +# --------------------------------------------------------------------------- +@pytest.mark.parametrize("cp_size", [1, 2, 4, 8]) +@pytest.mark.parametrize("layout", ["zigzag", "contiguous"]) +@pytest.mark.parametrize("lengths_factor", [[1], [1, 2, 3], [3, 1, 1, 2]]) +def test_thd_rank_indices_partition_all_tokens_exactly_once(cp_size, layout, lengths_factor): + lengths = [2 * cp_size * f for f in lengths_factor] + cu = _cu(lengths) + owned = torch.cat([cpl.get_thd_context_parallel_rank_indices(cu, cp_size, r, layout) for r in range(cp_size)]) + assert owned.numel() == int(cu[-1]) + assert torch.equal(torch.sort(owned).values, torch.arange(int(cu[-1]))) + + +@pytest.mark.parametrize("cp_size", [2, 4, 8]) +def test_zigzag_rank_indices_match_relax_data_sharding(cp_size): + """MCore's zigzag partition must be token-for-token what Relax's data path + produces. + + If these ever disagree, chunkwise CP would silently permute tokens relative + to the all-gather fallback and the attention layers. + """ + lengths = [2 * cp_size * f for f in (1, 3, 2)] + cu = _cu(lengths) + full = _tagged_tokens(int(cu[-1])).reshape(-1, 1, 3) # [s, b=1, C] + + for rank in range(cp_size): + mcore_idx = cpl.get_thd_context_parallel_rank_indices(cu, cp_size, rank, "zigzag") + mcore_shard = full[mcore_idx] + + # Relax data.py: per-sample slice_with_cp then concat. + relax_shard = torch.cat( + [ + slice_with_cp( + full[cu[i] : cu[i + 1]], + pad_value=0.0, + qkv_format="thd", + dynamic_cp_size=cp_size, + dynamic_cp_rank=rank, + ) + for i in range(len(lengths)) + ], + dim=0, + ) + assert torch.equal(mcore_shard, relax_shard) + + # Relax model.py (all-gather fallback) re-slices with gdn_cp_slice. + assert torch.equal(mcore_shard, gdn_cp_slice(full, cu, cp_size, rank)) + + +@pytest.mark.parametrize("cp_size", [2, 4, 8]) +@pytest.mark.parametrize("lengths_factor", [[1], [1, 2, 3], [3, 1, 1, 2]]) +def test_both_layouts_are_permutations_of_each_other(cp_size, lengths_factor): + """The two partitions must describe the same token set with the same per- + rank size. + + That is the precondition for the all-to-all between them to be a pure + permutation -- no token invented, dropped, or duplicated. The real collective + round trip is asserted in ``test_gdn_chunkwise_cp_gpu.py``. + """ + lengths = [2 * cp_size * f for f in lengths_factor] + cu = _cu(lengths) + total = int(cu[-1]) + zig_by_rank = [] + con_by_rank = [] + for rank in range(cp_size): + zig = cpl.get_thd_context_parallel_rank_indices(cu, cp_size, rank, "zigzag") + con = cpl.get_thd_context_parallel_rank_indices(cu, cp_size, rank, "contiguous") + zig_by_rank.append(zig) + con_by_rank.append(con) + assert zig.numel() == con.numel() == total // cp_size + # contiguous is exactly this rank's span of the flattened buffer + assert torch.equal(con, torch.arange(rank * (total // cp_size), (rank + 1) * (total // cp_size))) + + # Across the whole CP group, both layouts are permutations of exactly the + # same global token rows. + assert torch.equal( + torch.cat(zig_by_rank).sort().values, + torch.cat(con_by_rank).sort().values, + ) + + +@pytest.mark.parametrize("cp_size", [2, 4]) +def test_rank_indices_reject_lengths_not_divisible_by_two_cp(cp_size): + bad = _cu([2 * cp_size, 2 * cp_size + 1]) + with pytest.raises(ValueError, match="divisible by"): + cpl.get_thd_context_parallel_rank_indices(bad, cp_size, 0, "zigzag") + + +def test_gdn_rejects_packed_lengths_not_divisible_by_cp(): + from megatron.core.ssm.gated_delta_net import GatedDeltaNet + + cu = _cu([8, 6]) + with pytest.raises(ValueError, match="divisible by cp_size=4"): + GatedDeltaNet._resolve_cu_seqlens(None, None, cu, int(cu[-1]), "cu_seqlens_q", cp_size=4) + + +def test_rank_indices_reject_unknown_layout(): + with pytest.raises(ValueError, match="Unsupported context-parallel layout"): + cpl.get_thd_context_parallel_rank_indices(_cu([16, 16]), 2, 0, "contiguous_ish") + + +@pytest.mark.parametrize("layout", ["zigzag", "contiguous"]) +def test_rank_indices_ignore_duplicate_boundaries(layout): + compact = torch.tensor([0, 16, 40], dtype=torch.int64) + padded = torch.tensor([0, 16, 40, 40, 40], dtype=torch.int64) + for rank in range(2): + assert torch.equal( + cpl.get_thd_context_parallel_rank_indices(compact, 2, rank, layout), + cpl.get_thd_context_parallel_rank_indices(padded, 2, rank, layout), + ) + + +@pytest.mark.parametrize("layout", ["zigzag", "contiguous"]) +def test_rank_indices_reject_decreasing_boundaries(layout): + with pytest.raises(ValueError, match="nondecreasing"): + cpl.get_thd_context_parallel_rank_indices(torch.tensor([0, 16, 8]), 2, 0, layout) + + +# --------------------------------------------------------------------------- +# Dynamic CP group resolution +# --------------------------------------------------------------------------- +def test_resolve_cp_group_prefers_packed_seq_params(): + static = object() + dynamic = object() + assert resolve_cp_group(static, None) is static + assert resolve_cp_group(static, PackedSeqParams(qkv_format="thd")) is static + assert resolve_cp_group(static, PackedSeqParams(qkv_format="thd", cp_group=dynamic)) is dynamic + + +# --------------------------------------------------------------------------- +# Construction-time capability gate +# --------------------------------------------------------------------------- +def _gdn_config(**overrides): + import torch.nn.functional as F + from megatron.core.transformer.transformer_config import TransformerConfig + + kwargs = dict( + hidden_size=2048, + num_layers=1, + num_attention_heads=16, + num_query_groups=2, + normalization="RMSNorm", + use_cpu_initialization=True, + activation_func=F.silu, + bf16=True, + experimental_attention_variant="gated_delta_net", + linear_attention_freq=[1], + linear_conv_kernel_dim=4, + linear_key_head_dim=128, + linear_value_head_dim=128, + linear_num_key_heads=16, + linear_num_value_heads=32, + ) + kwargs.update(overrides) + return TransformerConfig(**kwargs) + + +def test_config_default_mode_is_headwise(): + """Upgrading the image must not silently reroute an existing recipe.""" + assert _gdn_config().linear_cp_mode == "headwise" + + +def test_headwise_config_requires_heads_divisible_by_tp_times_cp(): + # 16 key heads, tp=2, cp=4 -> 16 % 8 == 0: fine. + _gdn_config(tensor_model_parallel_size=2, context_parallel_size=4) + # tp=2, cp=16 -> 16 % 32 != 0: the geometry headwise cannot express. + with pytest.raises(AssertionError, match="linear_num_key_heads"): + _gdn_config(tensor_model_parallel_size=2, context_parallel_size=16) + + +def test_chunkwise_config_only_requires_heads_divisible_by_tp(): + """This is what replaces Relax's temporary head-count rewrite hack.""" + cfg = _gdn_config(tensor_model_parallel_size=2, context_parallel_size=16, linear_cp_mode="chunkwise") + assert cfg.linear_num_key_heads == 16 and cfg.linear_num_value_heads == 32 + # ... but TP divisibility is still enforced: GDN weights stay TP-sharded. + with pytest.raises(AssertionError, match="linear_num_key_heads"): + _gdn_config( + tensor_model_parallel_size=8, + context_parallel_size=2, + linear_cp_mode="chunkwise", + num_query_groups=8, + linear_num_key_heads=4, + linear_num_value_heads=8, + ) + + +def test_all_gather_config_uses_the_tp_only_head_rule(): + """`--linear-cp-mode=all_gather` must be constructible on a non-divisible + geometry. + + Relax's all-gather fallback keeps GDN weights TP-only, so declaring it + should relax the head check exactly as chunkwise does. + """ + cfg = _gdn_config(tensor_model_parallel_size=2, context_parallel_size=16, linear_cp_mode="all_gather") + assert cfg.linear_num_key_heads == 16 and cfg.linear_num_value_heads == 32 + + +def test_config_rejects_unresolved_and_unknown_linear_cp_mode(): + """MCore only accepts the three concrete execution modes.""" + for bad in ("auto", "allgather", "chunk", ""): + with pytest.raises(AssertionError, match="linear_cp_mode"): + _gdn_config(context_parallel_size=2, linear_cp_mode=bad) + with pytest.raises(AssertionError, match="linear_cp_mode"): + _gdn_config(context_parallel_size=4, tensor_model_parallel_size=2, linear_cp_mode=bad) + + +def test_gdn_forward_has_no_per_call_mode_override(): + from megatron.core.ssm.gated_delta_net import GatedDeltaNet + + assert "linear_cp_mode" not in inspect.signature(GatedDeltaNet.forward).parameters diff --git a/tests/backends/megatron/test_gdn_cp_mode_stage2.py b/tests/backends/megatron/test_gdn_cp_mode_stage2.py new file mode 100644 index 000000000..96034560a --- /dev/null +++ b/tests/backends/megatron/test_gdn_cp_mode_stage2.py @@ -0,0 +1,229 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. +"""CPU-only tests for the Task 32 Stage 2 Relax-side GDN CP routing. + +Covers the pieces added on top of the Stage 1 FLA/MCore backport +(``test_gdn_chunkwise_cp_layout.py``): the ``--linear-cp-mode`` CLI, invalid +Chunkwise combinations, and the thin runtime dispatcher installed on +``GatedDeltaNet.forward``. + +The dispatcher tests drive ``GatedDeltaNet.forward`` through duck-typed fakes +and hook/counter spies instead of a real distributed process group or FLA +kernel call, per task32-stage2-handoff.md §6.2 ("use hooks/counters, don't +infer routing from numerics"). Real-kernel / real-collective coverage stays in +``test_gdn_chunkwise_cp_gpu.py``. +""" + +from __future__ import annotations + +import argparse +from types import SimpleNamespace + +import pytest + + +pytest.importorskip("megatron.core.context_parallel_layout", reason="requires the patched Megatron-LM") + +from megatron.core.packed_seq_params import PackedSeqParams # noqa: E402 +from megatron.core.ssm.gated_delta_net import GatedDeltaNet # noqa: E402 + +from relax.backends.megatron import arguments as megatron_arguments # noqa: E402 +from relax.backends.megatron import model as gdn_model # noqa: E402 +from relax.backends.megatron.arguments import _validate_linear_cp_mode # noqa: E402 + + +# --------------------------------------------------------------------------- +# Step 1: CLI flag +# --------------------------------------------------------------------------- +def _parse_megatron_args(monkeypatch, *argv): + pytest.importorskip("sglang.srt.server_args") + from relax.utils.arguments import get_slime_extra_args_provider + + monkeypatch.setattr("sys.argv", ["test-linear-cp-mode", *argv]) + return megatron_arguments._megatron_parse_args( + extra_args_provider=get_slime_extra_args_provider(), + ignore_unknown_args=False, + ) + + +def test_linear_cp_mode_flag_defaults_to_headwise(monkeypatch): + args = _parse_megatron_args(monkeypatch) + assert args.linear_cp_mode == "headwise" + + +@pytest.mark.parametrize("mode", ["headwise", "chunkwise", "all_gather"]) +def test_linear_cp_mode_flag_accepts_all_concrete_modes(monkeypatch, mode): + args = _parse_megatron_args(monkeypatch, "--linear-cp-mode", mode) + assert args.linear_cp_mode == mode + + +# --------------------------------------------------------------------------- +# Step 2: argument validation +# --------------------------------------------------------------------------- +def _args(**overrides): + base = dict( + linear_cp_mode="headwise", + allgather_cp=False, + deterministic_mode=False, + dynamic_context_parallel=False, + context_parallel_size=1, + ) + base.update(overrides) + return argparse.Namespace(**base) + + +@pytest.mark.parametrize("bad", ["auto", "allgather"]) +def test_validate_linear_cp_mode_rejects_unsupported_value(bad): + with pytest.raises(ValueError, match="does not support 'auto'|must be one of"): + _validate_linear_cp_mode(_args(linear_cp_mode=bad)) + + +def test_validate_linear_cp_mode_rejects_chunkwise_with_allgather_cp(): + with pytest.raises(ValueError, match="allgather-cp"): + _validate_linear_cp_mode(_args(linear_cp_mode="chunkwise", allgather_cp=True)) + + +@pytest.mark.parametrize("cp_kwargs", [{"context_parallel_size": 2}, {"dynamic_context_parallel": True}]) +def test_validate_linear_cp_mode_rejects_chunkwise_deterministic_when_cp_may_exceed_one(cp_kwargs): + with pytest.raises(ValueError, match="deterministic"): + _validate_linear_cp_mode(_args(linear_cp_mode="chunkwise", deterministic_mode=True, **cp_kwargs)) + + +# --------------------------------------------------------------------------- +# Steps 4-6: runtime dispatcher +# --------------------------------------------------------------------------- +class _FakeGroup: + """Minimal process-group stand-in exposing only .size()/.rank(): the + dispatcher and resolve_cp_group() never issue a real collective.""" + + def __init__(self, size, rank=0): + self._size = size + self._rank = rank + + def size(self): + return self._size + + def rank(self): + return self._rank + + +def _fake_packed_seq_params(qkv_format="thd", cp_group=None, local_cp_size=None): + params = PackedSeqParams(qkv_format=qkv_format) + if cp_group is not None: + params.cp_group = cp_group + if local_cp_size is not None: + params.local_cp_size = local_cp_size + return params + + +def _fake_gdn_module(*, linear_cp_mode, static_cp_size=1, deterministic_mode=False): + """Duck-typed GatedDeltaNet 'self': just enough attributes for the + dispatcher guards to read -- no real nn.Module/CUDA/FLA state.""" + return SimpleNamespace( + pg_collection=SimpleNamespace(cp=_FakeGroup(static_cp_size)), + cp_size=static_cp_size, + config=SimpleNamespace(linear_cp_mode=linear_cp_mode, deterministic_mode=deterministic_mode), + ) + + +@pytest.fixture(autouse=True) +def _isolate_gdn_forward_patch(): + """`_patch_gdn_for_dynamic_cp` idempotently monkey-patches the *shared* + GatedDeltaNet class attribute; save/restore it around every test so it + cannot leak into test_gdn_chunkwise_cp_gpu.py.""" + orig_forward = GatedDeltaNet.forward + orig_patched_flag = getattr(GatedDeltaNet, "_dcp_patched", False) + yield + GatedDeltaNet.forward = orig_forward + GatedDeltaNet._dcp_patched = orig_patched_flag + + +def _install_dispatcher_with_spies(): + """Install the dispatcher over a spy for the original MCore forward.""" + calls = {"orig": 0} + + def spy_orig(self, hidden_states, attention_mask, inference_context=None, packed_seq_params=None, *a, **kw): + calls["orig"] += 1 + return "orig", hidden_states + + GatedDeltaNet.forward = spy_orig + GatedDeltaNet._dcp_patched = False + gdn_model._patch_gdn_for_dynamic_cp() + return calls + + +def test_dispatcher_cp1_goes_to_original_forward_regardless_of_mode(): + for mode in ("headwise", "chunkwise", "all_gather"): + calls = _install_dispatcher_with_spies() + m = _fake_gdn_module(linear_cp_mode=mode, static_cp_size=1) + out = GatedDeltaNet.forward(m, "hs", None, None, None) + assert calls == {"orig": 1}, mode + assert out == ("orig", "hs") + + +@pytest.mark.parametrize("mode", ["headwise", "chunkwise"]) +def test_dispatcher_headwise_and_chunkwise_cp_gt_1_go_to_original_forward(mode): + calls = _install_dispatcher_with_spies() + m = _fake_gdn_module(linear_cp_mode=mode, static_cp_size=4) + psp = _fake_packed_seq_params(cp_group=_FakeGroup(4, rank=2), local_cp_size=4) + GatedDeltaNet.forward(m, "hs", None, None, psp) + assert calls == {"orig": 1} + + +def test_dispatcher_all_gather_cp_gt_1_goes_to_relax_fallback(): + calls = _install_dispatcher_with_spies() + m = _fake_gdn_module(linear_cp_mode="all_gather", static_cp_size=4) + psp = _fake_packed_seq_params(qkv_format="sbhd", cp_group=_FakeGroup(4, rank=3), local_cp_size=4) + with pytest.raises(AssertionError, match=r"packed \(thd\) sequences"): + GatedDeltaNet.forward(m, "hs", None, None, psp) + assert calls == {"orig": 0} + + +def test_dispatcher_prefers_dynamic_group_over_static_group(): + """Runtime CP (from packed_seq_params) must win over the module's static + max-CP group -- e.g. a static CP=8 model running a CP=1 micro-batch must + not take the all_gather branch just because the static group has size 8.""" + calls = _install_dispatcher_with_spies() + m = _fake_gdn_module(linear_cp_mode="all_gather", static_cp_size=8) + psp = _fake_packed_seq_params(cp_group=_FakeGroup(1, rank=0), local_cp_size=1) + GatedDeltaNet.forward(m, "hs", None, None, psp) + assert calls == {"orig": 1} + + +@pytest.mark.parametrize( + ("mode", "runtime_cp_size"), + [("headwise", 4), ("chunkwise", 4), ("all_gather", 1)], +) +def test_dispatcher_never_mutates_shared_module_or_config_state(mode, runtime_cp_size): + """Covers handoff §6.3: self.cp_size / self.pg_collection.cp / + self.config.linear_cp_mode must be bit-identical before and after, across + every mode and every runtime CP size (static max CP fixed at 8, so a + smaller runtime CP can only come from the dynamic packed_seq_params).""" + _install_dispatcher_with_spies() + m = _fake_gdn_module(linear_cp_mode=mode, static_cp_size=8) + psp = _fake_packed_seq_params(cp_group=_FakeGroup(runtime_cp_size, rank=0), local_cp_size=runtime_cp_size) + before_cp_size = m.cp_size + before_pg_cp = m.pg_collection.cp + before_mode = m.config.linear_cp_mode + GatedDeltaNet.forward(m, "hs", None, None, psp) + assert m.cp_size == before_cp_size + assert m.pg_collection.cp is before_pg_cp + assert m.config.linear_cp_mode == before_mode + + +# --------------------------------------------------------------------------- +# All-gather guard clauses (checked before any real tensor operation). +# --------------------------------------------------------------------------- +def test_all_gather_fallback_rejects_inference(): + _install_dispatcher_with_spies() + m = _fake_gdn_module(linear_cp_mode="all_gather", static_cp_size=4) + psp = _fake_packed_seq_params(cp_group=_FakeGroup(4), local_cp_size=4) + with pytest.raises(AssertionError, match="inference"): + GatedDeltaNet.forward(m, "hs", None, object(), psp) + + +def test_all_gather_fallback_rejects_deterministic_mode(): + _install_dispatcher_with_spies() + m = _fake_gdn_module(linear_cp_mode="all_gather", static_cp_size=4, deterministic_mode=True) + psp = _fake_packed_seq_params(cp_group=_FakeGroup(4), local_cp_size=4) + with pytest.raises(AssertionError, match="deterministic mode"): + GatedDeltaNet.forward(m, "hs", None, None, psp) From 4e8976a7586e4d60767e252a66e54be01f305cce Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Thu, 13 Aug 2026 23:18:11 +1000 Subject: [PATCH 02/11] add cp chunkwise route --- .../patch/megatron/20260728-0e6ac576f.patch | 883 +++++++++++++++++- .../megatron/test_gdn_chunkwise_cp_route.py | 231 +++++ 2 files changed, 1103 insertions(+), 11 deletions(-) create mode 100644 tests/backends/megatron/test_gdn_chunkwise_cp_route.py diff --git a/docker/patch/megatron/20260728-0e6ac576f.patch b/docker/patch/megatron/20260728-0e6ac576f.patch index 9a3404d19..4962d2aaa 100644 --- a/docker/patch/megatron/20260728-0e6ac576f.patch +++ b/docker/patch/megatron/20260728-0e6ac576f.patch @@ -1115,10 +1115,135 @@ index 465e83f..232caef 100644 tensor_recv_prev = None tensor_recv_next = None diff --git a/megatron/core/ssm/gated_delta_net.py b/megatron/core/ssm/gated_delta_net.py -index 7e28691..3c19562 100644 +index 7e28691..f1410f8 100644 --- a/megatron/core/ssm/gated_delta_net.py +++ b/megatron/core/ssm/gated_delta_net.py -@@ -584,16 +584,17 @@ class GatedDeltaNet(MegatronModule): +@@ -18,6 +18,7 @@ from torch import Tensor + from megatron.core import tensor_parallel + from megatron.core.context_parallel_layout import ( + contiguous_to_zigzag_chunks, ++ get_thd_cp_partition_route, + zigzag_to_contiguous_chunks, + ) + from megatron.core.fp8_utils import get_fp8_align_size +@@ -342,6 +343,21 @@ class GatedDeltaNet(MegatronModule): + + inference_context = deprecate_inference_params(inference_context, inference_params) + ++ # Validate the final post-repack metadata without synchronizing device tensors. ++ if packed_seq_params is not None: ++ dynamic_cp_group = packed_seq_params.cp_group ++ local_cp_size = packed_seq_params.local_cp_size ++ if (dynamic_cp_group is None) != (local_cp_size is None): ++ raise ValueError( ++ "PackedSeqParams.cp_group and local_cp_size must either both be set " ++ "or both be None." ++ ) ++ if dynamic_cp_group is not None and local_cp_size != dynamic_cp_group.size(): ++ raise ValueError( ++ f"PackedSeqParams.local_cp_size ({local_cp_size}) does not match " ++ f"cp_group.size() ({dynamic_cp_group.size()})." ++ ) ++ + # Route the CP group to either the headwise (Ulysses-style) path or the + # chunkwise CP path according to config.linear_cp_mode. The two paths + # are mutually exclusive — whichever one is active owns the full CP +@@ -350,15 +366,20 @@ class GatedDeltaNet(MegatronModule): + # CUDA-graph-unsafe `torch.distributed.new_group` on every forward. + base_cp_group = pg_collection.cp if pg_collection is not None else self.pg_collection.cp + cp_group = resolve_cp_group(base_cp_group, packed_seq_params) ++ cp_size = cp_group.size() if cp_group is not None else 1 + if self.config.linear_cp_mode == "chunkwise": + cp_group_chunkwise = cp_group + cp_group_headwise = None + elif self.config.linear_cp_mode == "headwise": + cp_group_chunkwise = None + cp_group_headwise = cp_group +- elif cp_group.size() == 1: ++ elif cp_size == 1: + cp_group_chunkwise = None + cp_group_headwise = None ++ elif self.config.linear_cp_mode == "all_gather": ++ raise RuntimeError( ++ "linear_cp_mode='all_gather' with CP>1 requires the Relax GatedDeltaNet wrapper." ++ ) + else: + raise ValueError( + f"Unsupported linear_cp_mode {self.config.linear_cp_mode!r}; " +@@ -385,31 +406,9 @@ class GatedDeltaNet(MegatronModule): + not self.config.deterministic_mode + ), "Packed sequence does not support deterministic mode." + +- # Resolve cu_seqlens with alignment padding handling. +- # cu_seqlens in packed_seq_params is the global (pre-CP-split) cu_seqlens, so we +- # validate against the global sequence length. +- cu_seqlens_q = self._resolve_cu_seqlens( +- packed_seq_params.cu_seqlens_q_padded, +- packed_seq_params.cu_seqlens_q, +- seq_len_global, +- "cu_seqlens_q", +- cp_size=self.cp_size, +- ) +- cu_seqlens_kv = self._resolve_cu_seqlens( +- packed_seq_params.cu_seqlens_kv_padded, +- packed_seq_params.cu_seqlens_kv, +- seq_len_global, +- "cu_seqlens_kv", +- cp_size=self.cp_size, +- ) +- assert torch.equal(cu_seqlens_q, cu_seqlens_kv), ( +- "Currently only support cu_seqlens_q equals to cu_seqlens_kv, " +- f"but got {cu_seqlens_q=} and {cu_seqlens_kv=}" +- ) +- num_packed_seqs = cu_seqlens_q.shape[0] - 1 +- assert num_packed_seqs > 0, ( +- "Number of packed sequences must be greater than 0, " +- f"but got {cu_seqlens_q=} and {cu_seqlens_kv=}" ++ # Cache validation on the final PackedSeqParams shared by layers/recompute. ++ cu_seqlens_q, cu_seqlens_kv = self._resolve_thd_cu_seqlens( ++ packed_seq_params, seq_len_global, cp_size + ) + else: + cu_seqlens_q = None +@@ -422,7 +421,7 @@ class GatedDeltaNet(MegatronModule): + # tensor and the resulting chunkwise CP context so we don't + # reallocate them on every forward — those reallocations break + # CUDA graph capture. +- cache_key = (seq_len_global, batch) ++ cache_key = (seq_len_global, batch, cp_group_chunkwise) + cached = self._chunkwise_cp_context_cache.get(cache_key) + if cached is None: + cached_cu_seqlens = ( +@@ -503,6 +502,18 @@ class GatedDeltaNet(MegatronModule): + Returns: + Tuple of (output, output_bias). + """ ++ route_to_contiguous = None ++ route_to_zigzag = None ++ if cp_size_chunkwise > 1 and packed_seq_params is not None and packed_seq_params.qkv_format == 'thd': ++ route_to_contiguous = get_thd_cp_partition_route( ++ packed_seq_params, cu_seqlens_q, cp_size_chunkwise, ++ cp_group_chunkwise.rank(), "zigzag", "contiguous", device=hidden_states.device, ++ ) ++ route_to_zigzag = get_thd_cp_partition_route( ++ packed_seq_params, cu_seqlens_q, cp_size_chunkwise, ++ cp_group_chunkwise.rank(), "contiguous", "zigzag", device=hidden_states.device, ++ ) ++ + # Input projection + nvtx_range_push(suffix="in_proj") + qkvzba, _ = self.in_proj(hidden_states) +@@ -519,7 +530,8 @@ class GatedDeltaNet(MegatronModule): + nvtx_range_push(suffix="zigzag_to_contiguous") + if packed_seq_params is not None and packed_seq_params.qkv_format == 'thd': + qkvzba = zigzag_to_contiguous_chunks( +- qkvzba, cp_group_chunkwise, seq_dim=0, cu_seqlens=cu_seqlens_q ++ qkvzba, cp_group_chunkwise, seq_dim=0, cu_seqlens=cu_seqlens_q, ++ thd_cp_partition_route=route_to_contiguous, + ) + else: + qkvzba = zigzag_to_contiguous_chunks(qkvzba, cp_group_chunkwise, seq_dim=0) +@@ -584,16 +596,17 @@ class GatedDeltaNet(MegatronModule): "gdn_conv_pad_alignment is incompatible with GDN chunkwise CP. Padding " "chunk-local causal-conv inputs can change later chunk numerics." ) @@ -1144,17 +1269,93 @@ index 7e28691..3c19562 100644 + packed_seq_params=packed_seq_params, + ) nvtx_range_pop(suffix="pre_gated_delta_rule") - + nvtx_range_push(suffix="gated_delta_rule") -@@ -1211,7 +1212,7 @@ def get_parameter_local_cp_headwise( +@@ -630,7 +643,8 @@ class GatedDeltaNet(MegatronModule): + nvtx_range_push(suffix="contiguous_to_zigzag") + if packed_seq_params is not None and packed_seq_params.qkv_format == 'thd': + norm_out = contiguous_to_zigzag_chunks( +- norm_out, cp_group=cp_group_chunkwise, seq_dim=0, cu_seqlens=cu_seqlens_q ++ norm_out, cp_group=cp_group_chunkwise, seq_dim=0, cu_seqlens=cu_seqlens_q, ++ thd_cp_partition_route=route_to_zigzag, + ) + else: + norm_out = contiguous_to_zigzag_chunks( +@@ -984,6 +998,65 @@ class GatedDeltaNet(MegatronModule): + beta = beta.sigmoid() + return g, beta + ++ def _resolve_thd_cu_seqlens(self, packed_seq_params, seq_len_global, cp_size): ++ """Resolve and validate this micro-batch's global THD boundaries, once. ++ ++ Every GDN layer is handed the same ``PackedSeqParams``, so the checks below -- ++ each of which reads a device tensor from Python and therefore stalls on the ++ GPU -- only need to run for the first layer that asks. The result is cached on ++ the ``PackedSeqParams`` instance and reused while it still describes the very ++ same boundary tensors, mirroring how the CP layout routes are cached. ++ ++ A cached result is reused only when all four source tensors are still the same ++ objects *and* none of them has been written in place. Identity alone would not ++ notice a caller refilling a preallocated boundary buffer, so autograd's version ++ counter -- which every in-place op bumps, and which costs a plain Python ++ attribute read -- is checked as well. Comparing values instead would reintroduce ++ the device-host sync this cache exists to remove. ++ """ ++ sources = ( ++ packed_seq_params.cu_seqlens_q_padded, ++ packed_seq_params.cu_seqlens_q, ++ packed_seq_params.cu_seqlens_kv_padded, ++ packed_seq_params.cu_seqlens_kv, ++ ) ++ cacheable = all(t is None or not t.is_inference() for t in sources) ++ versions = tuple(None if t is None or t.is_inference() else t._version for t in sources) ++ cached = getattr(packed_seq_params, "_gdn_resolved_cu_seqlens", None) ++ if cacheable and cached is not None: ++ cached_meta, cached_sources, cached_versions, cached_value = cached ++ if ( ++ cached_meta == (seq_len_global, cp_size) ++ and all(a is b for a, b in zip(cached_sources, sources)) ++ and cached_versions == versions ++ ): ++ return cached_value ++ ++ cu_seqlens_q = self._resolve_cu_seqlens( ++ sources[0], sources[1], seq_len_global, "cu_seqlens_q", cp_size=cp_size ++ ) ++ cu_seqlens_kv = self._resolve_cu_seqlens( ++ sources[2], sources[3], seq_len_global, "cu_seqlens_kv", cp_size=cp_size ++ ) ++ assert torch.equal(cu_seqlens_q, cu_seqlens_kv), ( ++ "Currently only support cu_seqlens_q equals to cu_seqlens_kv, " ++ f"but got {cu_seqlens_q=} and {cu_seqlens_kv=}" ++ ) ++ num_packed_seqs = cu_seqlens_q.shape[0] - 1 ++ assert num_packed_seqs > 0, ( ++ "Number of packed sequences must be greater than 0, " ++ f"but got {cu_seqlens_q=} and {cu_seqlens_kv=}" ++ ) ++ ++ resolved = (cu_seqlens_q, cu_seqlens_kv) ++ packed_seq_params._gdn_resolved_cu_seqlens = ( ++ (seq_len_global, cp_size), ++ sources, ++ versions, ++ resolved, ++ ) ++ return resolved ++ + def _resolve_cu_seqlens( + self, cu_seqlens_padded, cu_seqlens_actual, total_seq_len, name, cp_size: int = 1 + ) -> torch.Tensor: +@@ -1211,7 +1284,7 @@ def get_parameter_local_cp_headwise( slices = [slice(None)] * param.dim() dim_size = param.size(dim=dim) slices[dim] = slice(cp_rank * dim_size // cp_size, (cp_rank + 1) * dim_size // cp_size) - param = param[slices] + param = param[tuple(slices)] return param - - + + diff --git a/megatron/core/tensor_parallel/random.py b/megatron/core/tensor_parallel/random.py index dbecd73..1c52478 100644 --- a/megatron/core/tensor_parallel/random.py @@ -1351,13 +1552,13 @@ diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core for iteration in range(self.config.mtp_num_layers): layer_idx = 0 if self.mtp_use_repeated_layer else iteration diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py -index 43e45b7..8166f94 100644 +index 43e45b7..2cb8ad6 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py -@@ -224,6 +224,10 @@ class TransformerConfig(ModelParallelConfig): +@@ -222,6 +222,10 @@ class TransformerConfig(ModelParallelConfig): """Clamp the output of the linear_fc1 in the activation function. Only used when activation_func is quick_gelu or weighted SwiGLU (MoE only).""" - + + activation_func_clamp_shared_expert: bool = True + """If False, skip activation_func_clamp_value inside SharedExpertMLP so only routed MoE + experts get the clamp.""" @@ -1368,13 +1569,47 @@ index 43e45b7..8166f94 100644 @@ -283,6 +287,9 @@ class TransformerConfig(ModelParallelConfig): num_query_groups >= tp, and under fp8/fp4 a per-partition linear_qkv_out_dim aligned to 16/32.""" - + + post_self_attn_layernorm: bool = False + post_mlp_layernorm: bool = False + test_mode: bool = False """Whether to run real-time tests.""" - + +@@ -1045,7 +1052,7 @@ class TransformerConfig(ModelParallelConfig): + linear_cp_mode: Optional[str] = "chunkwise" + """Context-parallel execution mode for linear-attention layers + (e.g. Gated Delta Net). Independent of `cp_comm_type`, which only controls standard attention. +- Can be "chunkwise" or "headwise": ++ Can be "chunkwise", "headwise", or "all_gather" (Relax wrapper only): + "chunkwise": Keep sequence chunks sharded across CP ranks and use CP-aware linear kernels + (e.g. chunk_gated_delta_rule + causal_conv1d with a CP context). This follows the chunkwise + DeltaNet idea of storing state at chunk boundaries and doing chunk-local matrix work, avoiding +@@ -1053,6 +1060,8 @@ class TransformerConfig(ModelParallelConfig): + See https://sustcsonglin.github.io/blog/2024/deltanet-2/#a-chunkwise-algorithm-for-deltanet. + "headwise": Scatter heads across the CP group with all-to-all (Ulysses-style); each rank runs + the linear-attention kernel on the full sequence for a shard of heads. Correct but memory-heavy. ++ "all_gather": Relax-only mode that gathers the full sequence and keeps TP-local heads. ++ CP>1 execution requires Relax's GDN wrapper; native Megatron does not implement this mode. + """ + + ################## +@@ -1590,11 +1599,11 @@ class TransformerConfig(ModelParallelConfig): + f"got {self.gdn_conv_pad_alignment}." + ) + ++ assert self.linear_cp_mode in ("headwise", "chunkwise", "all_gather"), ( ++ "linear_cp_mode must be one of 'headwise', 'chunkwise', or 'all_gather'; " ++ f"got {self.linear_cp_mode!r}. 'auto' must be resolved before construction." ++ ) + if self.context_parallel_size > 1: +- assert self.linear_cp_mode in ("headwise", "chunkwise"), ( +- f"linear_cp_mode must be one of 'headwise' or 'chunkwise', " +- f"got {self.linear_cp_mode!r}." +- ) + if self.gdn_conv_pad_alignment is not None: + assert self.linear_cp_mode != "chunkwise", ( + "gdn_conv_pad_alignment is incompatible with " diff --git a/megatron/core/transformer/transformer_layer.py b/megatron/core/transformer/transformer_layer.py index 34a456b..8182c7c 100644 --- a/megatron/core/transformer/transformer_layer.py @@ -1864,3 +2099,629 @@ diff --git a/megatron/bridge/models/conversion/model_bridge.py b/megatron/bridge if isinstance(task.param_weight, DTensor): from megatron.core.distributed.fsdp.src.megatron_fsdp.uneven_dtensor import ( uneven_dtensor_to_full_tensor, +diff --git a/megatron/core/context_parallel_layout.py b/megatron/core/context_parallel_layout.py +index 4401458..a68dc2c 100644 +--- a/megatron/core/context_parallel_layout.py ++++ b/megatron/core/context_parallel_layout.py +@@ -2,12 +2,30 @@ + + """Context parallel tensor layout helpers.""" + +-from typing import List, Optional, Tuple ++from contextlib import contextmanager ++from typing import Any, List, NamedTuple, Optional, Tuple + + import torch + + from megatron.core.tensor_parallel import all_to_all + ++_THD_CP_ROUTE_ATTRS = { ++ ("zigzag", "contiguous"): "cp_partition_route_zigzag_to_contiguous", ++ ("contiguous", "zigzag"): "cp_partition_route_contiguous_to_zigzag", ++} ++ ++ ++@contextmanager ++def _cp_layout_nvtx_range(message: str): ++ active = torch.cuda.is_available() ++ if active: ++ torch.cuda.nvtx.range_push(message) ++ try: ++ yield ++ finally: ++ if active: ++ torch.cuda.nvtx.range_pop() ++ + + def get_thd_context_parallel_rank_indices( + cu_seqlens: torch.Tensor, cp_size: int, cp_rank: int, layout: str +@@ -82,11 +100,391 @@ def get_thd_context_parallel_rank_indices( + return rank_positions[torch.argsort(rank_local_pos)] + + ++class ThdCpPartitionRoute(NamedTuple): ++ """Everything one THD zigzag<->contiguous all-to-all needs, precomputed. ++ ++ ``send_rows`` / ``recv_rows`` are ``None`` when the permutation on that side is ++ the identity, which lets the swap skip the gather/scatter entirely. ++ The first six fields are exactly what #5664's ``decode_thd_cp_partition_route()`` ++ returns; see :func:`build_thd_cp_partition_route` for why the encoded tensor form is ++ not used here. ``cu_seqlens``, the layout pair and ``cp_size`` / ``cp_rank`` are ++ Relax-only additions carrying no computation -- they exist so a route cached on a ++ ``PackedSeqParams`` can prove it still describes the conversion being asked for. ++ """ ++ ++ local_source_length: int ++ local_target_length: int ++ send_rows: Optional[torch.Tensor] ++ recv_rows: Optional[torch.Tensor] ++ input_split_sizes: List[int] ++ output_split_sizes: List[int] ++ cu_seqlens: torch.Tensor ++ cu_seqlens_version: Optional[int] ++ cp_size: int ++ cp_rank: int ++ source_layout: str ++ target_layout: str ++ ++ ++_ThdLayoutSegment = Tuple[int, int, int] ++ ++ ++def _compact_thd_cu_seqlens_to_list(cu_seqlens: torch.Tensor) -> List[int]: ++ if cu_seqlens.dim() != 1: ++ raise ValueError(f"cu_seqlens must be 1-D, got shape {tuple(cu_seqlens.shape)}.") ++ ++ cu = cu_seqlens.detach().to(device="cpu", dtype=torch.long).tolist() ++ if not cu or cu[0] != 0: ++ raise ValueError(f"cu_seqlens must start at 0, got {cu_seqlens}.") ++ ++ compact_cu: List[int] = [cu[0]] ++ prev = cu[0] ++ for value in cu[1:]: ++ if value < prev: ++ raise ValueError(f"cu_seqlens must be nondecreasing, got {cu_seqlens}.") ++ if value != prev: ++ compact_cu.append(value) ++ prev = value ++ return compact_cu ++ ++ ++def _validate_thd_route_partitioning(cu: List[int], cp_size: int) -> None: ++ total_tokens = cu[-1] ++ if total_tokens % cp_size != 0: ++ raise ValueError( ++ f"Contiguous CP partitioning requires total_tokens={total_tokens} " ++ f"to be divisible by cp_size={cp_size}." ++ ) ++ ++ chunk_divisor = 2 * cp_size ++ bad_seq_lens = [ ++ seq_end - seq_start ++ for seq_start, seq_end in zip(cu[:-1], cu[1:]) ++ if (seq_end - seq_start) % chunk_divisor != 0 ++ ] ++ if bad_seq_lens: ++ raise ValueError( ++ "All packed sequence lengths must be divisible by " ++ f"2 * cp_size ({chunk_divisor}) for zigzag/contiguous CP layout conversion, " ++ f"got {bad_seq_lens}." ++ ) ++ ++ ++def _build_thd_layout_segments( ++ cu: List[int], cp_size: int, cp_rank: int, layout: str ++) -> Tuple[List[_ThdLayoutSegment], int]: ++ """Describe a rank's THD partition as (global_start, length, local_start) runs. ++ ++ Both layouts are unions of contiguous global spans, so the whole route can be ++ derived by intersecting spans instead of materialising per-token index tensors. ++ """ ++ total_tokens = cu[-1] ++ if layout == "contiguous": ++ part_len = total_tokens // cp_size ++ if part_len == 0: ++ return [], 0 ++ return [(cp_rank * part_len, part_len, 0)], part_len ++ ++ if layout != "zigzag": ++ raise ValueError( ++ f"Unsupported context-parallel layout {layout!r} for THD layout segments " ++ f"with cp_size={cp_size}, rank={cp_rank}." ++ ) ++ ++ segments: List[_ThdLayoutSegment] = [] ++ local_start = 0 ++ for seq_start, seq_end in zip(cu[:-1], cu[1:]): ++ seq_len = seq_end - seq_start ++ chunk_len = seq_len // (2 * cp_size) ++ first_chunk = cp_rank ++ second_chunk = 2 * cp_size - cp_rank - 1 ++ segments.append((seq_start + first_chunk * chunk_len, chunk_len, local_start)) ++ segments.append((seq_start + second_chunk * chunk_len, chunk_len, local_start + chunk_len)) ++ local_start += 2 * chunk_len ++ ++ return segments, local_start ++ ++ ++def _intersect_thd_layout_segments( ++ source_segments: List[_ThdLayoutSegment], target_segments: List[_ThdLayoutSegment] ++) -> List[Tuple[int, int, int]]: ++ """Overlap two sorted segment lists into (source_row, target_row, length) runs.""" ++ intersections: List[Tuple[int, int, int]] = [] ++ source_index = 0 ++ target_index = 0 ++ while source_index < len(source_segments) and target_index < len(target_segments): ++ source_global_start, source_len, source_local_start = source_segments[source_index] ++ target_global_start, target_len, target_local_start = target_segments[target_index] ++ source_global_end = source_global_start + source_len ++ target_global_end = target_global_start + target_len ++ ++ overlap_start = max(source_global_start, target_global_start) ++ overlap_end = min(source_global_end, target_global_end) ++ if overlap_start < overlap_end: ++ intersections.append( ++ ( ++ source_local_start + overlap_start - source_global_start, ++ target_local_start + overlap_start - target_global_start, ++ overlap_end - overlap_start, ++ ) ++ ) ++ ++ if source_global_end <= target_global_end: ++ source_index += 1 ++ else: ++ target_index += 1 ++ ++ return intersections ++ ++ ++def _append_range(rows: List[int], start: int, length: int) -> None: ++ rows.extend(range(start, start + length)) ++ ++ ++def _row_list_is_identity(rows: List[int]) -> bool: ++ return all(row == index for index, row in enumerate(rows)) ++ ++ ++def _thd_cp_partition_route_attr_name(source_layout: str, target_layout: str) -> str: ++ try: ++ return _THD_CP_ROUTE_ATTRS[(source_layout, target_layout)] ++ except KeyError as exc: ++ raise ValueError( ++ f"Unsupported CP layout conversion {source_layout!r} -> {target_layout!r} " ++ "for THD route." ++ ) from exc ++ ++ ++def build_thd_cp_partition_route( ++ cu_seqlens: torch.Tensor, ++ cp_size: int, ++ cp_rank: int, ++ source_layout: str, ++ target_layout: str, ++ *, ++ device: Optional[torch.device] = None, ++) -> ThdCpPartitionRoute: ++ """Precompute one THD CP layout conversion route. ++ ++ The route depends only on packed sequence metadata, CP rank/size and the ++ source/target layouts, so it can be reused by every tensor with the same THD ++ sequence axis in the same microbatch. The whole derivation runs on CPU ints ++ after a single ``cu_seqlens`` transfer, which is what keeps the conversion ++ itself free of device-host synchronisation. ++ ++ Deviation from #5664: upstream serialises the result into a single flat ++ ``torch.Tensor`` and decodes it again inside every conversion, so that the route ++ can be a CUDA-graph capture input. Decoding costs three device-to-host copies, ++ which on this path (~200 conversions per microbatch under full recompute) puts ++ back the synchronisation this route exists to remove, and Megatron's full-iteration ++ CUDA graph is rejected for THD layout conversion anyway. We therefore return the ++ decoded form directly -- the fields below are exactly upstream's ++ ``decode_thd_cp_partition_route()`` tuple. Revisit if GDN ever needs to run under ++ graph capture: the route tensors would then have to be stable capture inputs again. ++ """ ++ _thd_cp_partition_route_attr_name(source_layout, target_layout) ++ if cp_size < 1: ++ raise ValueError(f"cp_size must be >= 1, got {cp_size}.") ++ if not 0 <= cp_rank < cp_size: ++ raise ValueError(f"cp_rank must be in [0, {cp_size}), got {cp_rank}.") ++ if device is None: ++ device = cu_seqlens.device ++ ++ with _cp_layout_nvtx_range(f"cp_layout/thd/route/{source_layout}_to_{target_layout}"): ++ cu = _compact_thd_cu_seqlens_to_list(cu_seqlens) ++ _validate_thd_route_partitioning(cu, cp_size) ++ ++ source_segments_by_rank: List[List[_ThdLayoutSegment]] = [] ++ source_lengths: List[int] = [] ++ target_segments_by_rank: List[List[_ThdLayoutSegment]] = [] ++ target_lengths: List[int] = [] ++ for rank in range(cp_size): ++ source_segments, source_length = _build_thd_layout_segments( ++ cu, cp_size, rank, source_layout ++ ) ++ target_segments, target_length = _build_thd_layout_segments( ++ cu, cp_size, rank, target_layout ++ ) ++ source_segments_by_rank.append(source_segments) ++ source_lengths.append(source_length) ++ target_segments_by_rank.append(target_segments) ++ target_lengths.append(target_length) ++ ++ local_source_segments = source_segments_by_rank[cp_rank] ++ local_target_segments = target_segments_by_rank[cp_rank] ++ ++ send_rows_list: List[int] = [] ++ input_split_sizes: List[int] = [] ++ for dst_rank in range(cp_size): ++ intersections = _intersect_thd_layout_segments( ++ local_source_segments, target_segments_by_rank[dst_rank] ++ ) ++ intersections.sort(key=lambda item: item[1]) ++ input_split_size = 0 ++ for source_row, _, length in intersections: ++ _append_range(send_rows_list, source_row, length) ++ input_split_size += length ++ input_split_sizes.append(input_split_size) ++ ++ recv_rows_list: List[int] = [] ++ output_split_sizes: List[int] = [] ++ for src_rank in range(cp_size): ++ intersections = _intersect_thd_layout_segments( ++ source_segments_by_rank[src_rank], local_target_segments ++ ) ++ intersections.sort(key=lambda item: item[1]) ++ output_split_size = 0 ++ for _, target_row, length in intersections: ++ _append_range(recv_rows_list, target_row, length) ++ output_split_size += length ++ output_split_sizes.append(output_split_size) ++ ++ assert len(send_rows_list) == source_lengths[cp_rank] ++ assert len(recv_rows_list) == target_lengths[cp_rank] ++ ++ send_rows = ( ++ None ++ if _row_list_is_identity(send_rows_list) ++ else torch.tensor(send_rows_list, device=device, dtype=torch.long) ++ ) ++ recv_rows = ( ++ None ++ if _row_list_is_identity(recv_rows_list) ++ else torch.tensor(recv_rows_list, device=device, dtype=torch.long) ++ ) ++ return ThdCpPartitionRoute( ++ local_source_length=source_lengths[cp_rank], ++ local_target_length=target_lengths[cp_rank], ++ send_rows=send_rows, ++ recv_rows=recv_rows, ++ input_split_sizes=input_split_sizes, ++ output_split_sizes=output_split_sizes, ++ cu_seqlens=cu_seqlens, ++ cu_seqlens_version=None if cu_seqlens.is_inference() else cu_seqlens._version, ++ cp_size=cp_size, ++ cp_rank=cp_rank, ++ source_layout=source_layout, ++ target_layout=target_layout, ++ ) ++ ++ ++def _thd_cp_partition_route_is_reusable( ++ route: Optional[ThdCpPartitionRoute], ++ cu_seqlens: torch.Tensor, ++ cp_size: int, ++ cp_rank: int, ++ source_layout: str, ++ target_layout: str, ++ device: torch.device, ++) -> bool: ++ """A cached route is only valid for the exact packed boundaries it was built from. ++ ++ ``cu_seqlens`` is matched on object identity *plus* autograd's version counter, ++ never on value: a value comparison would itself be a device-host synchronisation, ++ which is precisely what the route exists to avoid. Identity alone would not survive ++ a caller that refills a preallocated boundary buffer in place -- a pattern that ++ already exists elsewhere in Megatron -- so the version counter, which every in-place ++ op bumps (views share the counter with their base) and which costs a plain Python ++ attribute read, closes that hole. The one case neither check sees is a write that ++ bypasses the dispatcher entirely, e.g. through ``data_ptr()``. ++ """ ++ if route is None: ++ return False ++ if route.cu_seqlens is not cu_seqlens: ++ return False ++ if cu_seqlens.is_inference() or route.cu_seqlens_version != cu_seqlens._version: ++ return False ++ if route.cp_size != cp_size or route.cp_rank != cp_rank: ++ return False ++ if route.source_layout != source_layout or route.target_layout != target_layout: ++ return False ++ rows = route.send_rows if route.send_rows is not None else route.recv_rows ++ return rows is None or rows.device == device ++ ++ ++def get_thd_cp_partition_route( ++ packed_seq_params: Optional[Any], ++ cu_seqlens: torch.Tensor, ++ cp_size: int, ++ cp_rank: int, ++ source_layout: str, ++ target_layout: str, ++ *, ++ device: Optional[torch.device] = None, ++) -> ThdCpPartitionRoute: ++ """Return this microbatch's route, building and caching it on first use. ++ ++ Packed boundaries change every microbatch, so the cache is attached to the ++ ``PackedSeqParams`` instance the forward was handed and is only reused while it ++ still describes the very same ``cu_seqlens`` tensor. Every GDN layer in a ++ microbatch, plus its recompute replay, shares one build. ++ ++ Deviation from #5664: upstream expects the routes to have been prebuilt by the data ++ pipeline, so its lookup is a bare ``getattr`` and the build-on-miss path emits a ++ ``FutureWarning``. Here build-on-miss is the intended path -- the object a GDN ++ forward receives is not necessarily the one the data pipeline created, because the ++ Bridge/VLM path can repack ``PackedSeqParams`` after embedding -- so there is ++ nothing to warn about, and the cache instead has to defend itself against reuse ++ across microbatches and against dynamic CP changing ``cp_size``/``cp_rank`` between ++ microbatches on the same module. Callers that do own the final object can still ++ prebuild eagerly via :func:`prebuild_thd_cp_partition_routes`. ++ """ ++ if device is None: ++ device = cu_seqlens.device ++ attr_name = _thd_cp_partition_route_attr_name(source_layout, target_layout) ++ cached = getattr(packed_seq_params, attr_name, None) if packed_seq_params is not None else None ++ if _thd_cp_partition_route_is_reusable( ++ cached, cu_seqlens, cp_size, cp_rank, source_layout, target_layout, device ++ ): ++ return cached ++ ++ route = build_thd_cp_partition_route( ++ cu_seqlens, cp_size, cp_rank, source_layout, target_layout, device=device ++ ) ++ if packed_seq_params is not None: ++ setattr(packed_seq_params, attr_name, route) ++ return route ++ ++ ++def prebuild_thd_cp_partition_routes( ++ packed_seq_params: Optional[Any], ++ cp_group: Optional[torch.distributed.ProcessGroup] = None, ++ cu_seqlens: Optional[torch.Tensor] = None, ++ *, ++ device: Optional[torch.device] = None, ++) -> None: ++ """Eagerly populate both THD CP layout routes for a packed microbatch.""" ++ if packed_seq_params is None or getattr(packed_seq_params, "qkv_format", None) != "thd": ++ return ++ if cp_group is None: ++ cp_group = getattr(packed_seq_params, "cp_group", None) ++ if cp_group is None or cp_group.size() <= 1: ++ return ++ if cu_seqlens is None: ++ cu_seqlens = getattr(packed_seq_params, "cu_seqlens_q_padded", None) ++ if cu_seqlens is None: ++ cu_seqlens = getattr(packed_seq_params, "cu_seqlens_q", None) ++ if cu_seqlens is None: ++ return ++ ++ for source_layout, target_layout in _THD_CP_ROUTE_ATTRS: ++ get_thd_cp_partition_route( ++ packed_seq_params, ++ cu_seqlens, ++ cp_group.size(), ++ cp_group.rank(), ++ source_layout, ++ target_layout, ++ device=device, ++ ) ++ ++ + def zigzag_to_contiguous_chunks( + x: torch.Tensor, + cp_group: torch.distributed.ProcessGroup, + seq_dim: int = 0, + cu_seqlens: Optional[torch.Tensor] = None, ++ thd_cp_partition_route: Optional[ThdCpPartitionRoute] = None, + ) -> torch.Tensor: + """Permute CP chunks from Megatron zigzag layout to contiguous-time layout. + +@@ -96,7 +494,13 @@ def zigzag_to_contiguous_chunks( + """ + if cu_seqlens is not None: + return _zigzag_contiguous_thd_swap( +- x, cp_group, seq_dim, cu_seqlens, source_layout="zigzag", target_layout="contiguous" ++ x, ++ cp_group, ++ seq_dim, ++ cu_seqlens, ++ source_layout="zigzag", ++ target_layout="contiguous", ++ thd_cp_partition_route=thd_cp_partition_route, + ) + return _zigzag_contiguous_chunk_swap(x, cp_group, seq_dim, to_contiguous=True) + +@@ -106,15 +510,43 @@ def contiguous_to_zigzag_chunks( + cp_group: torch.distributed.ProcessGroup, + seq_dim: int = 0, + cu_seqlens: Optional[torch.Tensor] = None, ++ thd_cp_partition_route: Optional[ThdCpPartitionRoute] = None, + ) -> torch.Tensor: + """Inverse of :func:`zigzag_to_contiguous_chunks`.""" + if cu_seqlens is not None: + return _zigzag_contiguous_thd_swap( +- x, cp_group, seq_dim, cu_seqlens, source_layout="contiguous", target_layout="zigzag" ++ x, ++ cp_group, ++ seq_dim, ++ cu_seqlens, ++ source_layout="contiguous", ++ target_layout="zigzag", ++ thd_cp_partition_route=thd_cp_partition_route, + ) + return _zigzag_contiguous_chunk_swap(x, cp_group, seq_dim, to_contiguous=False) + + ++def _pack_thd_cp_route_send_buffer( ++ x: torch.Tensor, local_source_length: int, send_rows: Optional[torch.Tensor] ++) -> torch.Tensor: ++ if local_source_length == 0: ++ return x.narrow(0, 0, 0) ++ if send_rows is None: ++ return x ++ return x.index_select(0, send_rows) ++ ++ ++def _scatter_thd_cp_route_recv_buffer( ++ recv_buf: torch.Tensor, recv_rows: Optional[torch.Tensor], out_shape: Tuple[int, ...] ++) -> torch.Tensor: ++ if recv_rows is None: ++ return recv_buf ++ out = recv_buf.new_empty(out_shape) ++ if recv_rows.numel() > 0: ++ out.index_copy_(0, recv_rows, recv_buf) ++ return out ++ ++ + def _zigzag_contiguous_thd_swap( + x: torch.Tensor, + cp_group: Optional[torch.distributed.ProcessGroup], +@@ -122,95 +554,59 @@ def _zigzag_contiguous_thd_swap( + cu_seqlens: torch.Tensor, + source_layout: str, + target_layout: str, ++ thd_cp_partition_route: Optional[ThdCpPartitionRoute] = None, + ) -> torch.Tensor: + """Single-all-to-all THD permutation between zigzag and contiguous layouts. + + The packed THD tensor stays packed: we first group local tokens by their + target CP rank, exchange those groups once, then scatter received tokens +- back into the target rank-local order. ++ back into the target rank-local order. Which rows go where is described by a ++ :class:`ThdCpPartitionRoute` the caller should have precomputed once for the ++ microbatch; without one this rebuilds it, which is correct but pays the build ++ on every conversion. + """ + cp_size = cp_group.size() if cp_group is not None else 1 + if cp_size == 1: + return x + cp_rank = cp_group.rank() + +- if seq_dim != 0: +- x = x.movedim(seq_dim, 0) +- x = x.contiguous() +- +- cu = cu_seqlens.to(device=x.device, dtype=torch.long) +- # TODO: Let a future CP layout scheduler precompute this routing once per +- # microbatch from immutable cu_seqlens and pass it through both THD swaps. +- # Do not cache it across microbatches because packed sequence boundaries change. +- source_by_rank = [ +- get_thd_context_parallel_rank_indices(cu, cp_size, rank, source_layout) +- for rank in range(cp_size) +- ] +- target_by_rank = [ +- get_thd_context_parallel_rank_indices(cu, cp_size, rank, target_layout) +- for rank in range(cp_size) +- ] +- +- local_source_indices = source_by_rank[cp_rank] +- local_target_indices = target_by_rank[cp_rank] +- if x.size(0) != local_source_indices.numel(): +- raise ValueError( +- f"Local THD tensor length ({x.size(0)}) does not match {source_layout} " +- f"rank-{cp_rank} partition length ({local_source_indices.numel()})." +- ) +- +- total_tokens = int(cu[-1].item()) +- target_owner = torch.empty(total_tokens, device=x.device, dtype=torch.long) +- target_local_pos = torch.empty(total_tokens, device=x.device, dtype=torch.long) +- for rank, indices in enumerate(target_by_rank): +- target_owner[indices] = rank +- target_local_pos[indices] = torch.arange(indices.numel(), device=x.device) +- +- local_target_owner = target_owner[local_source_indices] +- local_target_pos = target_local_pos[local_source_indices] +- +- send_parts: List[torch.Tensor] = [] +- input_split_sizes: List[int] = [] +- for dst_rank in range(cp_size): +- dst_mask = local_target_owner == dst_rank +- dst_rows = dst_mask.nonzero(as_tuple=False).flatten() +- if dst_rows.numel() > 0: +- dst_rows = dst_rows[torch.argsort(local_target_pos[dst_rows])] +- send_part = x.index_select(0, dst_rows) +- else: +- send_part = x.narrow(0, 0, 0) +- send_parts.append(send_part) +- input_split_sizes.append(send_part.size(0)) +- send_buf = torch.cat(send_parts, dim=0).contiguous() +- +- output_split_sizes: List[int] = [] +- recv_target_positions: List[torch.Tensor] = [] +- for src_rank in range(cp_size): +- src_indices = source_by_rank[src_rank] +- src_to_this_rank = target_owner[src_indices] == cp_rank +- recv_global_indices = src_indices[src_to_this_rank] +- if recv_global_indices.numel() > 0: +- recv_positions = target_local_pos[recv_global_indices] +- recv_positions = recv_positions[torch.argsort(recv_positions)] +- else: +- recv_positions = local_target_indices.narrow(0, 0, 0) +- recv_target_positions.append(recv_positions) +- output_split_sizes.append(recv_positions.numel()) +- +- recv_buf = all_to_all(cp_group, send_buf, output_split_sizes, input_split_sizes) +- +- out_shape = (local_target_indices.numel(),) + tuple(x.shape[1:]) +- out = x.new_empty(out_shape) +- offset = 0 +- for recv_positions in recv_target_positions: +- recv_len = recv_positions.numel() +- if recv_len > 0: +- out[recv_positions] = recv_buf[offset : offset + recv_len] +- offset += recv_len +- +- if seq_dim != 0: +- out = out.movedim(0, seq_dim) +- return out.contiguous() ++ conversion_name = f"{source_layout}_to_{target_layout}" ++ with _cp_layout_nvtx_range(f"cp_layout/thd/swap/{conversion_name}"): ++ if seq_dim != 0: ++ x = x.movedim(seq_dim, 0) ++ x = x.contiguous() ++ ++ route = thd_cp_partition_route ++ if not _thd_cp_partition_route_is_reusable( ++ route, cu_seqlens, cp_size, cp_rank, source_layout, target_layout, x.device ++ ): ++ route = build_thd_cp_partition_route( ++ cu_seqlens, cp_size, cp_rank, source_layout, target_layout, device=x.device ++ ) ++ ++ if x.size(0) != route.local_source_length: ++ raise ValueError( ++ f"Local THD tensor length ({x.size(0)}) does not match {source_layout} " ++ f"rank-{cp_rank} partition length ({route.local_source_length})." ++ ) ++ ++ with _cp_layout_nvtx_range(f"cp_layout/thd/pack/{conversion_name}"): ++ send_buf = _pack_thd_cp_route_send_buffer(x, route.local_source_length, route.send_rows) ++ if not send_buf.is_contiguous(): ++ send_buf = send_buf.contiguous() ++ ++ with _cp_layout_nvtx_range(f"cp_layout/thd/all_to_all/{conversion_name}"): ++ recv_buf = all_to_all( ++ cp_group, send_buf, route.output_split_sizes, route.input_split_sizes ++ ) ++ ++ with _cp_layout_nvtx_range(f"cp_layout/thd/scatter/{conversion_name}"): ++ out_shape = (route.local_target_length,) + tuple(x.shape[1:]) ++ out = _scatter_thd_cp_route_recv_buffer(recv_buf, route.recv_rows, out_shape) ++ ++ if seq_dim != 0: ++ out = out.movedim(0, seq_dim) ++ return out.contiguous() + + + def _zigzag_contiguous_chunk_swap( diff --git a/tests/backends/megatron/test_gdn_chunkwise_cp_route.py b/tests/backends/megatron/test_gdn_chunkwise_cp_route.py new file mode 100644 index 000000000..6d426c117 --- /dev/null +++ b/tests/backends/megatron/test_gdn_chunkwise_cp_route.py @@ -0,0 +1,231 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. +"""Unit tests for the prebuilt THD CP layout route (Task 32, phase 3). + +Phase 1 derived the zigzag<->contiguous all-to-all plan inside every conversion, +from device tensors, which cost a device-host synchronisation per CP rank per +call. Phase 3 backports NVIDIA/Megatron-LM#5664's idea instead: derive the plan +once per micro-batch on CPU and hand the same route to every GDN layer. + +Two things therefore need proving on CPU, with no process group: + +1. the segment-based route describes *exactly* the permutation the phase-1 + index-based partition described -- otherwise chunkwise CP silently reorders + tokens; +2. a cached route is only ever reused for the micro-batch, CP geometry and + direction it was built for. + +The real all-to-all round trip over NCCL stays in +``test_gdn_chunkwise_cp_gpu.py``. +""" + +from __future__ import annotations + +import pytest +import torch + + +cpl = pytest.importorskip("megatron.core.context_parallel_layout", reason="requires the patched Megatron-LM") + +from megatron.core.packed_seq_params import PackedSeqParams # noqa: E402 + + +DIRECTIONS = [("zigzag", "contiguous"), ("contiguous", "zigzag")] + +# Packed boundary shapes worth covering: single sequence, uneven multi-sequence, +# and a duplicated boundary (an empty padding slot), which the compaction step +# has to drop before the segments line up. +LENGTH_CASES = [ + [1], + [3, 1, 2], + [2, 0, 1, 3], +] + + +def _cu(lengths: list[int], unit: int) -> torch.Tensor: + cu = [0] + for n in lengths: + cu.append(cu[-1] + n * unit) + return torch.tensor(cu, dtype=torch.int64) + + +def _packed_seq_params(cu: torch.Tensor) -> PackedSeqParams: + return PackedSeqParams( + qkv_format="thd", + cu_seqlens_q=cu, + cu_seqlens_kv=cu, + max_seqlen_q=int(cu[-1]), + max_seqlen_kv=int(cu[-1]), + ) + + +def _apply_route_across_ranks( + x: torch.Tensor, cu: torch.Tensor, cp_size: int, source: str, target: str +) -> list[torch.Tensor]: + """Run the route-driven swap for every rank, emulating the all-to-all locally.""" + source_by_rank = [ + cpl.get_thd_context_parallel_rank_indices(cu, cp_size, r, source) for r in range(cp_size) + ] + routes = [ + cpl.build_thd_cp_partition_route(cu, cp_size, r, source, target) for r in range(cp_size) + ] + + send_bufs = [] + for rank, route in enumerate(routes): + local = x[source_by_rank[rank]] + assert local.size(0) == route.local_source_length + send_bufs.append(local if route.send_rows is None else local.index_select(0, route.send_rows)) + + outputs = [] + for dst, route in enumerate(routes): + parts = [] + for src in range(cp_size): + offset = sum(routes[src].input_split_sizes[:dst]) + length = routes[src].input_split_sizes[dst] + assert length == route.output_split_sizes[src], "split sizes disagree between peers" + parts.append(send_bufs[src][offset : offset + length]) + recv = torch.cat(parts, dim=0) + if route.recv_rows is None: + outputs.append(recv) + else: + out = recv.new_empty((route.local_target_length,) + tuple(x.shape[1:])) + out.index_copy_(0, route.recv_rows, recv) + outputs.append(out) + return outputs + + +# --------------------------------------------------------------------------- +# The route is the same permutation phase 1 computed +# --------------------------------------------------------------------------- +@pytest.mark.parametrize("cp_size", [1, 2, 4, 8]) +@pytest.mark.parametrize("source,target", DIRECTIONS) +@pytest.mark.parametrize("lengths", LENGTH_CASES) +def test_route_reproduces_index_based_partition(cp_size, source, target, lengths): + cu = _cu(lengths, unit=2 * cp_size) + total = int(cu[-1]) + x = torch.arange(total * 3, dtype=torch.float64).reshape(total, 3) + + got = _apply_route_across_ranks(x, cu, cp_size, source, target) + for rank in range(cp_size): + want = x[cpl.get_thd_context_parallel_rank_indices(cu, cp_size, rank, target)] + assert torch.equal(got[rank], want), f"cp_size={cp_size} rank={rank} {source}->{target}" + + +# --------------------------------------------------------------------------- +# Fail-fast parity with the index-based builder +# --------------------------------------------------------------------------- +def test_route_rejects_lengths_not_divisible_by_two_cp(): + cu = torch.tensor([0, 12], dtype=torch.int64) # 12 % (2 * 4) != 0 + with pytest.raises(ValueError, match="divisible by"): + cpl.get_thd_context_parallel_rank_indices(cu, 4, 0, "zigzag") + with pytest.raises(ValueError, match="divisible by"): + cpl.build_thd_cp_partition_route(cu, 4, 0, "zigzag", "contiguous") + + +def test_route_rejects_malformed_cu_seqlens(): + with pytest.raises(ValueError, match="must start at 0"): + cpl.build_thd_cp_partition_route( + torch.tensor([8, 16], dtype=torch.int64), 2, 0, "zigzag", "contiguous" + ) + with pytest.raises(ValueError, match="nondecreasing"): + cpl.build_thd_cp_partition_route( + torch.tensor([0, 16, 8], dtype=torch.int64), 2, 0, "zigzag", "contiguous" + ) + + +def test_route_rejects_unknown_layout(): + cu = _cu([1], unit=4) + with pytest.raises(ValueError, match="Unsupported CP layout conversion"): + cpl.build_thd_cp_partition_route(cu, 2, 0, "zigzag", "interleaved") + + +# --------------------------------------------------------------------------- +# Caching: reuse only within the micro-batch it was built for +# --------------------------------------------------------------------------- +def test_route_is_cached_per_packed_seq_params(): + cu = _cu([3, 1], unit=8) + psp = _packed_seq_params(cu) + + first = cpl.get_thd_cp_partition_route(psp, cu, 4, 1, "zigzag", "contiguous") + second = cpl.get_thd_cp_partition_route(psp, cu, 4, 1, "zigzag", "contiguous") + assert second is first, "a second layer of the same micro-batch must reuse the route" + assert psp.cp_partition_route_zigzag_to_contiguous is first + + +def test_both_directions_are_cached_separately(): + cu = _cu([3, 1], unit=8) + psp = _packed_seq_params(cu) + + to_contiguous = cpl.get_thd_cp_partition_route(psp, cu, 4, 1, "zigzag", "contiguous") + to_zigzag = cpl.get_thd_cp_partition_route(psp, cu, 4, 1, "contiguous", "zigzag") + assert to_contiguous is not to_zigzag + assert psp.cp_partition_route_zigzag_to_contiguous is to_contiguous + assert psp.cp_partition_route_contiguous_to_zigzag is to_zigzag + + +def test_route_is_rebuilt_for_new_packed_boundaries(): + """Packed boundaries move every micro-batch; a stale route would corrupt tokens.""" + cu = _cu([3, 1], unit=8) + psp = _packed_seq_params(cu) + first = cpl.get_thd_cp_partition_route(psp, cu, 4, 1, "zigzag", "contiguous") + + next_cu = _cu([2, 2], unit=8) + psp.cu_seqlens_q = next_cu + rebuilt = cpl.get_thd_cp_partition_route(psp, next_cu, 4, 1, "zigzag", "contiguous") + assert rebuilt is not first + assert rebuilt.cu_seqlens is next_cu + + want = cpl.build_thd_cp_partition_route(next_cu, 4, 1, "zigzag", "contiguous") + assert rebuilt.input_split_sizes == want.input_split_sizes + assert rebuilt.output_split_sizes == want.output_split_sizes + for field in ("send_rows", "recv_rows"): + got_rows, want_rows = getattr(rebuilt, field), getattr(want, field) + assert (got_rows is None) == (want_rows is None) + if want_rows is not None: + assert torch.equal(got_rows, want_rows) + + +def test_route_is_rebuilt_when_the_dynamic_cp_geometry_changes(): + """Dynamic CP varies cp_size/cp_rank across micro-batches on one module.""" + cu = _cu([3, 1], unit=8) + psp = _packed_seq_params(cu) + cp4 = cpl.get_thd_cp_partition_route(psp, cu, 4, 1, "zigzag", "contiguous") + + cp2 = cpl.get_thd_cp_partition_route(psp, cu, 2, 1, "zigzag", "contiguous") + assert cp2 is not cp4 + assert (cp2.cp_size, cp2.cp_rank) == (2, 1) + + other_rank = cpl.get_thd_cp_partition_route(psp, cu, 2, 0, "zigzag", "contiguous") + assert other_rank is not cp2 + assert other_rank.cp_rank == 0 + + +def test_prebuild_populates_both_directions(): + cu = _cu([3, 1], unit=8) + psp = _packed_seq_params(cu) + + class _FakeGroup: + def size(self): + return 4 + + def rank(self): + return 2 + + psp.cp_group = _FakeGroup() + psp.local_cp_size = 4 + cpl.prebuild_thd_cp_partition_routes(psp) + + for attr in ("cp_partition_route_zigzag_to_contiguous", "cp_partition_route_contiguous_to_zigzag"): + route = getattr(psp, attr) + assert route is not None + assert (route.cp_size, route.cp_rank) == (4, 2) + + +def test_prebuild_is_a_noop_without_context_parallelism(): + cu = _cu([3, 1], unit=8) + psp = _packed_seq_params(cu) + cpl.prebuild_thd_cp_partition_routes(psp) + assert getattr(psp, "cp_partition_route_zigzag_to_contiguous", None) is None + + non_thd = PackedSeqParams(qkv_format="sbhd") + cpl.prebuild_thd_cp_partition_routes(non_thd) + assert getattr(non_thd, "cp_partition_route_zigzag_to_contiguous", None) is None From 1122107c3e57c89a7f8b018f8142451c67cc0c54 Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Thu, 13 Aug 2026 23:26:54 +1000 Subject: [PATCH 03/11] style: apply ruff-format and docformatter to the CP route tests Pre-commit was not run before pushing the previous commit, so CI's ruff-format and docformatter hooks failed on the new test file. Formatting only -- no test logic changed. Co-authored-by: Cursor --- .../megatron/test_gdn_chunkwise_cp_route.py | 22 +++++++------------ 1 file changed, 8 insertions(+), 14 deletions(-) diff --git a/tests/backends/megatron/test_gdn_chunkwise_cp_route.py b/tests/backends/megatron/test_gdn_chunkwise_cp_route.py index 6d426c117..886a88a07 100644 --- a/tests/backends/megatron/test_gdn_chunkwise_cp_route.py +++ b/tests/backends/megatron/test_gdn_chunkwise_cp_route.py @@ -61,13 +61,10 @@ def _packed_seq_params(cu: torch.Tensor) -> PackedSeqParams: def _apply_route_across_ranks( x: torch.Tensor, cu: torch.Tensor, cp_size: int, source: str, target: str ) -> list[torch.Tensor]: - """Run the route-driven swap for every rank, emulating the all-to-all locally.""" - source_by_rank = [ - cpl.get_thd_context_parallel_rank_indices(cu, cp_size, r, source) for r in range(cp_size) - ] - routes = [ - cpl.build_thd_cp_partition_route(cu, cp_size, r, source, target) for r in range(cp_size) - ] + """Run the route-driven swap for every rank, emulating the all-to-all + locally.""" + source_by_rank = [cpl.get_thd_context_parallel_rank_indices(cu, cp_size, r, source) for r in range(cp_size)] + routes = [cpl.build_thd_cp_partition_route(cu, cp_size, r, source, target) for r in range(cp_size)] send_bufs = [] for rank, route in enumerate(routes): @@ -123,13 +120,9 @@ def test_route_rejects_lengths_not_divisible_by_two_cp(): def test_route_rejects_malformed_cu_seqlens(): with pytest.raises(ValueError, match="must start at 0"): - cpl.build_thd_cp_partition_route( - torch.tensor([8, 16], dtype=torch.int64), 2, 0, "zigzag", "contiguous" - ) + cpl.build_thd_cp_partition_route(torch.tensor([8, 16], dtype=torch.int64), 2, 0, "zigzag", "contiguous") with pytest.raises(ValueError, match="nondecreasing"): - cpl.build_thd_cp_partition_route( - torch.tensor([0, 16, 8], dtype=torch.int64), 2, 0, "zigzag", "contiguous" - ) + cpl.build_thd_cp_partition_route(torch.tensor([0, 16, 8], dtype=torch.int64), 2, 0, "zigzag", "contiguous") def test_route_rejects_unknown_layout(): @@ -163,7 +156,8 @@ def test_both_directions_are_cached_separately(): def test_route_is_rebuilt_for_new_packed_boundaries(): - """Packed boundaries move every micro-batch; a stale route would corrupt tokens.""" + """Packed boundaries move every micro-batch; a stale route would corrupt + tokens.""" cu = _cu([3, 1], unit=8) psp = _packed_seq_params(cu) first = cpl.get_thd_cp_partition_route(psp, cu, 4, 1, "zigzag", "contiguous") From 62e627dae52b24094c322ab696e727f24685a87a Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Fri, 14 Aug 2026 14:18:03 +1000 Subject: [PATCH 04/11] guard cp route cache against in-place cu_seqlens writes --- .../megatron/test_gdn_chunkwise_cp_route.py | 33 +++++++++++++++++++ 1 file changed, 33 insertions(+) diff --git a/tests/backends/megatron/test_gdn_chunkwise_cp_route.py b/tests/backends/megatron/test_gdn_chunkwise_cp_route.py index 886a88a07..6276150c6 100644 --- a/tests/backends/megatron/test_gdn_chunkwise_cp_route.py +++ b/tests/backends/megatron/test_gdn_chunkwise_cp_route.py @@ -178,6 +178,39 @@ def test_route_is_rebuilt_for_new_packed_boundaries(): assert torch.equal(got_rows, want_rows) +def test_route_is_rebuilt_when_cu_seqlens_is_mutated_in_place(): + """Identity alone would miss a caller refilling a preallocated boundary + buffer. + + Megatron already has that pattern elsewhere (persistent ``_cu_seqlens_buffer`` + written with ``buf[0] = 0``), so the cache also fingerprints autograd's version + counter, which every in-place write bumps. + """ + cu = _cu([3, 1], unit=8) + psp = _packed_seq_params(cu) + first = cpl.get_thd_cp_partition_route(psp, cu, 4, 1, "zigzag", "contiguous") + + # Same tensor object, refilled with different boundaries. + cu.copy_(_cu([2, 2], unit=8)) + rebuilt = cpl.get_thd_cp_partition_route(psp, cu, 4, 1, "zigzag", "contiguous") + assert rebuilt is not first, "an in-place refill must invalidate the cached route" + + want = cpl.build_thd_cp_partition_route(cu, 4, 1, "zigzag", "contiguous") + assert rebuilt.input_split_sizes == want.input_split_sizes + assert rebuilt.output_split_sizes == want.output_split_sizes + + +def test_route_is_rebuilt_when_a_view_of_cu_seqlens_is_mutated(): + """Views share the version counter with their base, so writes through one + count.""" + cu = _cu([3, 1], unit=8) + psp = _packed_seq_params(cu) + first = cpl.get_thd_cp_partition_route(psp, cu, 4, 1, "zigzag", "contiguous") + + cu[1:] = _cu([2, 2], unit=8)[1:] + assert cpl.get_thd_cp_partition_route(psp, cu, 4, 1, "zigzag", "contiguous") is not first + + def test_route_is_rebuilt_when_the_dynamic_cp_geometry_changes(): """Dynamic CP varies cp_size/cp_rank across micro-batches on one module.""" cu = _cu([3, 1], unit=8) From 73f57e3e76ee643fe62234a807cb2525d44c2bc1 Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Mon, 21 Sep 2026 21:20:21 +1000 Subject: [PATCH 05/11] perf(megatron): adapt GDN CP to current MCore MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit # ⚡ Performance - Adapt the rebased Task 32 work to current main using reviewer PR #348 (b2cd009970100b7d4a86c1d963eb796b4895b12d). - Preserve the pinned dependencies, chunkwise default and native selective GDN recompute while keeping all_gather as an explicit fallback. - Keep route and boundary caches tied to final packed metadata and validate tensor versions, inference tensors and runtime CP geometry. # 🐛 Bug Fix - Move pure GDN mode validation into a standard-library-only module so model providers do not import the optimizer and training argument stack. - Retain provider-level packing validation without breaking the existing VPP dependency-isolation tests. # ✅ Tests - Port the current CPU and CUDA/NCCL regression suites from #348. - Pass 159 CPU regressions against exact pinned MCore sources, including all 17 VPP cases that fail on the reference branch. - Pass all pre-commit checks and verify the complete cumulative patch applies to the pinned Bridge/MCore assembly using Dockerfile's patch command. References: redai-studio/Relax#213, redai-studio/Relax#273, redai-studio/Relax#348. --- relax/backends/megatron/arguments.py | 37 +- relax/backends/megatron/gdn_cp_config.py | 43 + relax/backends/megatron/model.py | 49 +- relax/backends/megatron/model_provider.py | 6 + .../megatron/test_gdn_chunkwise_cp_gpu.py | 1008 +++++++++++++++++ .../megatron/test_gdn_chunkwise_cp_layout.py | 8 +- .../megatron/test_gdn_chunkwise_cp_route.py | 96 ++ .../megatron/test_gdn_cp_mode_stage2.py | 99 +- 8 files changed, 1275 insertions(+), 71 deletions(-) create mode 100644 relax/backends/megatron/gdn_cp_config.py create mode 100644 tests/backends/megatron/test_gdn_chunkwise_cp_gpu.py diff --git a/relax/backends/megatron/arguments.py b/relax/backends/megatron/arguments.py index 96a66258e..d84a6b38a 100644 --- a/relax/backends/megatron/arguments.py +++ b/relax/backends/megatron/arguments.py @@ -20,6 +20,8 @@ from relax.utils.logging_utils import get_logger from relax.utils.model_source import ModelSource +from .gdn_cp_config import _validate_linear_cp_mode + __all__ = ["validate_args", "megatron_parse_args", "set_default_megatron_args"] @@ -128,41 +130,6 @@ def _validate_dynamic_context_parallel(args): args.max_seqlen_per_dp_cp_rank = args.max_tokens_per_gpu -def _validate_linear_cp_mode(args) -> None: - """Fail fast on `--linear-cp-mode` / flag combinations that are invalid for - every model, without needing the HF config. - - Geometry-dependent rejections (e.g. explicit `headwise` on heads not - divisible by `tp*max_cp`) can only be checked once the real GDN head counts - are known, which happens in MCore's `TransformerConfig.__post_init__` gate - -- not here. - """ - mode = getattr(args, "linear_cp_mode", "headwise") - allowed_modes = {"headwise", "chunkwise", "all_gather"} - if mode not in allowed_modes: - raise ValueError( - f"--linear-cp-mode must be one of {sorted(allowed_modes)!r}; got {mode!r}. v1 does not support 'auto'." - ) - - if mode == "chunkwise" and getattr(args, "allgather_cp", False): - raise ValueError( - "--linear-cp-mode=chunkwise is incompatible with --allgather-cp: chunkwise CP requires " - "Megatron's zig-zag THD packing, while --allgather-cp switches the data path to a single " - "contiguous per-rank chunk. Note --allgather-cp is a data/attention packing flag, unrelated " - "to the GDN `all_gather` CP mode." - ) - - cp_may_exceed_one = ( - getattr(args, "dynamic_context_parallel", False) or getattr(args, "context_parallel_size", 1) > 1 - ) - if mode == "chunkwise" and cp_may_exceed_one and getattr(args, "deterministic_mode", False): - raise ValueError( - "--linear-cp-mode=chunkwise does not support --deterministic-mode while CP>1 may occur: " - "the deterministic torch reference path only accepts cp_context=None. Use " - "--linear-cp-mode=headwise or =all_gather for deterministic CP>1 runs." - ) - - def validate_args(args): """Run megatron's own validate_args plus slime-specific megatron validations.""" diff --git a/relax/backends/megatron/gdn_cp_config.py b/relax/backends/megatron/gdn_cp_config.py new file mode 100644 index 000000000..a74e42637 --- /dev/null +++ b/relax/backends/megatron/gdn_cp_config.py @@ -0,0 +1,43 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. + +from argparse import Namespace +from typing import Optional + + +def _validate_linear_cp_mode(args: Namespace, config: Optional[object] = None) -> None: + """Validate the mode name, then GDN-specific flags once the model is known. + + Geometry-dependent rejections (e.g. explicit `headwise` on heads not + divisible by `tp*max_cp`) can only be checked once the real GDN head counts + are known, which happens in MCore's `TransformerConfig.__post_init__` gate + -- not here. Bridge calls this again with the provider's actual config. + """ + model_config = config if config is not None else args + mode = getattr(model_config, "linear_cp_mode", "chunkwise") + allowed_modes = {"headwise", "chunkwise", "all_gather"} + if mode not in allowed_modes: + raise ValueError( + f"--linear-cp-mode must be one of {sorted(allowed_modes)!r}; got {mode!r}. Resolve 'auto' before construction." + ) + + # Bridge determines the attention variant from the HF checkpoint. The default + # linear_cp_mode on an ordinary-attention model does not make it a GDN model. + if getattr(model_config, "experimental_attention_variant", None) != "gated_delta_net": + return + + cp_may_exceed_one = ( + getattr(args, "dynamic_context_parallel", False) or getattr(model_config, "context_parallel_size", 1) > 1 + ) + if cp_may_exceed_one and getattr(args, "allgather_cp", False): + raise ValueError( + "GDN CP requires zig-zag THD packing in every linear_cp_mode; --allgather-cp uses " + "contiguous per-rank packing and is incompatible with GDN CP>1. --allgather-cp is a " + "data/attention packing flag, separate from --linear-cp-mode=all_gather." + ) + + if mode == "chunkwise" and cp_may_exceed_one and getattr(model_config, "deterministic_mode", False): + raise ValueError( + "--linear-cp-mode=chunkwise does not support --deterministic-mode while CP>1 may occur: " + "the deterministic torch reference path only accepts cp_context=None. " + "Packed GDN inputs also do not support deterministic mode in the other CP modes." + ) diff --git a/relax/backends/megatron/model.py b/relax/backends/megatron/model.py index eb9083712..ee4bb6441 100644 --- a/relax/backends/megatron/model.py +++ b/relax/backends/megatron/model.py @@ -415,7 +415,8 @@ def setup_model_and_optimizer( _patch_gdn_for_dynamic_cp() model_config = get_model_config(model[0]) if getattr(model_config, "experimental_attention_variant", None) == "gated_delta_net" and ( - not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0 + not torch.distributed.is_initialized() + or torch.distributed.get_rank(group=torch.distributed.group.WORLD) == 0 ): logger.info( f"[GDN CP] role={role} linear_cp_mode={model_config.linear_cp_mode} " @@ -441,19 +442,24 @@ def setup_model_and_optimizer( return model, optimizer, opt_param_scheduler -def _resolve_gdn_cp(self, packed_seq_params): +def _resolve_gdn_cp(self, packed_seq_params, pg_collection=None): """Resolve (cp_size, cp_group, cp_rank) for a GDN forward. Prefers the per-micro-batch dynamic CP group carried on ``packed_seq_params`` (set in ``data.py``); falls back to the module's static CP group. """ - if packed_seq_params is not None and getattr(packed_seq_params, "local_cp_size", None) is not None: - cp_group = packed_seq_params.cp_group - cp_size = packed_seq_params.local_cp_size - else: - cp_group = self.pg_collection.cp - cp_size = cp_group.size() + cp_group = pg_collection.cp if pg_collection is not None else self.pg_collection.cp + if packed_seq_params is not None: + dynamic_group = getattr(packed_seq_params, "cp_group", None) + local_cp_size = getattr(packed_seq_params, "local_cp_size", None) + if (dynamic_group is None) != (local_cp_size is None): + raise ValueError("PackedSeqParams.cp_group and local_cp_size must both be set or both be None.") + if dynamic_group is not None: + if local_cp_size != dynamic_group.size(): + raise ValueError("PackedSeqParams.local_cp_size does not match cp_group.size().") + cp_group = dynamic_group + cp_size = cp_group.size() if cp_group is not None else 1 cp_rank = cp_group.rank() if cp_size > 1 else 0 return cp_size, cp_group, cp_rank @@ -463,9 +469,9 @@ def _assert_gdn_full_recompute() -> None: The all-gather path below runs the recurrent scan on the *full* sequence duplicated on every CP rank, so the GDN activation scales with the full - context length. Only ``--recompute-granularity full`` (whole-layer - checkpointing) keeps that a per-layer transient; ``selective`` does not - cover GDN (its module list has no gdn/mamba entry) and silently OOMs. + context length. ``--recompute-granularity full`` (whole-layer checkpointing) + keeps that a per-layer transient. The fallback bypasses native GDN forward, + including its selective recompute wrapper. Only relevant to training forwards that build a graph (and thus retain activations): skipped when grad is disabled (weight-only / inference roles @@ -478,11 +484,11 @@ def _assert_gdn_full_recompute() -> None: args = get_args() if getattr(args, "recompute_granularity", None) != "full": raise ValueError( - "GatedDeltaNet context-parallel (cp>1) requires whole-layer activation recompute: " + "GatedDeltaNet all_gather context-parallel (cp>1) requires whole-layer activation recompute: " "pass `--recompute-granularity full --recompute-method uniform --recompute-num-layers 1`. " f"Got recompute_granularity={getattr(args, 'recompute_granularity', None)!r}. " - "`selective` recompute does not cover GDN and will OOM (its full-sequence duplicated scan " - "activation stays resident)." + "The Relax all_gather fallback bypasses native GDN selective recompute; " + "its full-sequence duplicated scan activation would stay resident." ) _assert_gdn_full_recompute._checked = True @@ -522,7 +528,7 @@ def _dcp_gdn_forward( from .cp_utils import gdn_cp_slice - cp_size, cp_group, cp_rank = _resolve_gdn_cp(self, packed_seq_params) + cp_size, cp_group, cp_rank = _resolve_gdn_cp(self, packed_seq_params, kwargs.get("pg_collection")) if cp_size == 1 or self.config.linear_cp_mode != "all_gather": return _orig_forward( self, hidden_states, attention_mask, inference_context, packed_seq_params, *args, **kwargs @@ -545,14 +551,19 @@ def _dcp_gdn_forward( ) _assert_gdn_full_recompute() - cu_seqlens = packed_seq_params.cu_seqlens_q + cu_seqlens, _ = self._resolve_thd_cu_seqlens( + packed_seq_params, hidden_states.shape[0] * self.sp_size * cp_size, cp_size + ) # Precompute the host-side boundary list once per micro-batch (cached on the # shared packed_seq_params object) so the gather/slice below don't force a # per-GDN-layer .tolist() device sync — repeated under full recompute. - cu_seqlens_cpu = getattr(packed_seq_params, "_gdn_cu_seqlens_cpu", None) - if cu_seqlens_cpu is None: + cached = getattr(packed_seq_params, "_gdn_cu_seqlens_cpu", None) + version = None if cu_seqlens.is_inference() else cu_seqlens._version + if cached is not None and version is not None and cached[0] is cu_seqlens and cached[1] == version: + cu_seqlens_cpu = cached[2] + else: cu_seqlens_cpu = cu_seqlens.tolist() - packed_seq_params._gdn_cu_seqlens_cpu = cu_seqlens_cpu + packed_seq_params._gdn_cu_seqlens_cpu = (cu_seqlens, version, cu_seqlens_cpu) _, batch, _ = hidden_states.shape # Input projection on the CP-sharded (and SP-sharded) sequence. diff --git a/relax/backends/megatron/model_provider.py b/relax/backends/megatron/model_provider.py index 32b80f0e7..257191770 100644 --- a/relax/backends/megatron/model_provider.py +++ b/relax/backends/megatron/model_provider.py @@ -49,6 +49,7 @@ ) from .conditional_branch_sync import install_conditional_branch_sync +from .gdn_cp_config import _validate_linear_cp_mode logger = get_logger(__name__) @@ -263,6 +264,9 @@ def wrapped_model_provider( model = custom_model_provider(pre_process=pre_process, post_process=post_process, vp_stage=vp_stage) else: model = custom_model_provider(pre_process=pre_process, post_process=post_process) + for module in model.modules(): + if getattr(module, "config", None) is not None: + _validate_linear_cp_mode(args, module.config) configure_mtp_detach_paths(args, model) # Apply critic output layer if needed install_critic_value_head_in_provider(model, role, post_process) @@ -409,6 +413,7 @@ def wrapped_model_provider( provider.bf16 = True provider.params_dtype = torch.bfloat16 + _validate_linear_cp_mode(args, provider) provider.finalize() # Pickle provider for offline inspection / reproducibility (only on rank 0) @@ -457,6 +462,7 @@ def model_provider(pre_process: bool = True, post_process: bool = True, vp_stage # Experimental loading arguments from yaml config: TransformerConfig = core_transformer_config_from_args(args) + _validate_linear_cp_mode(args, config) if args.spec is not None: transformer_layer_spec = import_module(args.spec) diff --git a/tests/backends/megatron/test_gdn_chunkwise_cp_gpu.py b/tests/backends/megatron/test_gdn_chunkwise_cp_gpu.py new file mode 100644 index 000000000..ea1a934d4 --- /dev/null +++ b/tests/backends/megatron/test_gdn_chunkwise_cp_gpu.py @@ -0,0 +1,1008 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. +"""Real-kernel / real-collective tests for the GDN chunkwise-CP backport. + +RFC Task 32 phase-1 acceptance items 3 and 4: + +* a minimal CP=2 chunkwise case must drive the *actual* FLA kernels and match a + CP=1 reference in forward and backward within tolerance; +* the GDN parameter keys and shard dimensions in ``state_dict`` / + ``sharded_state_dict`` must not move, i.e. GDN weights stay TP-only and + checkpoints are unaffected by the CP mode. + +Three layers are covered: + +1. the FLA kernels directly (``causal_conv1d`` / ``chunk_gated_delta_rule`` with + a ``cp_context``); +2. the whole MCore ``GatedDeltaNet`` module in fp32 -- the *algebraic* check. In + fp32 the only difference between CP=1 and CP=2 is float summation order, so + the tolerances can be tight enough to catch a genuinely wrong permutation or + a dropped boundary term; +3. the same module in bf16 -- the *production* check, at the dtype training + actually uses, where the achievable agreement is bounded by the storage + format rather than by the algorithm. + +A headwise CP=2 run is included at every layer as a control: if headwise and +chunkwise both drift the same way, the cause is shared plumbing, not the new +code. + +Most tests need 2 visible GPUs; the TP2/CP2 matrix test needs 4: + pytest tests/backends/megatron/test_gdn_chunkwise_cp_gpu.py +""" + +from __future__ import annotations + +import os + +import pytest +import torch +import torch.multiprocessing as mp + + +WORLD_SIZE = 2 + +# --- tolerances ------------------------------------------------------------ +# bf16 module output / input gradient: RFC section 5, "MCore GDN, CP>1 vs CP=1". +ATOL_BF16 = 2e-3 +RTOL_BF16 = 1e-2 +MIN_COSINE = 0.9999 +# FLA kernels: RFC section 5, normalised RMS error thresholds. +CONV_RMS_RATIO = 1e-3 +GDN_RMS_RATIO = 2e-3 +# fp32 module run: both CP algorithms must land far below any bf16 threshold. +# 1e-3 is 2x below the bf16 element-wise atol of the RFC gate, i.e. "fp32 must be +# comfortably better than the dtype we actually ship". +RMS_RATIO_FP32 = 1e-3 +# fp32 kernel-level: measured ~1e-7, so 1e-5 is a real gate, not a rubber stamp. +KERNEL_RMS_RATIO_FP32 = 1e-5 +# chunkwise vs headwise. headwise is the already-shipped CP algorithm, so whatever +# CP-vs-no-CP disagreement it shows is the floor this environment imposes (reduced +# precision inside the Triton dots, changed summation order, bf16 storage) rather +# than anything about the algorithm. Requiring chunkwise to be no worse than that +# floor is the assertion that actually means something; a fixed atol on a bf16 +# token-sum gradient mostly measures rounding luck. +CHUNKWISE_VS_HEADWISE_RMS_FACTOR = 4.0 +# ...with a floor, so a headwise value that happens to land at or near zero on a given +# run cannot turn into an impossible budget. 1e-6 is still ~1000x tighter than the fp32 +# absolute gate, so the comparison keeps its teeth. +RMS_FLOOR_FP32 = 1e-6 +# ...applied only where headwise is not bit-exact. headwise hands each rank the +# whole sequence and 1/cp of the heads, so none of its reductions are +# repartitioned and it can land exactly on the CP=1 result; "4x zero" would be a +# budget no correct implementation could meet. Those tensors are covered by the +# absolute gates instead. + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.device_count() < WORLD_SIZE, + reason=f"requires {WORLD_SIZE} CUDA devices", +) + + +def _has_backport() -> bool: + try: + import megatron.core.context_parallel_layout # noqa: F401 + from fla.ops.cp import build_cp_context # noqa: F401 + except ImportError: + return False + return True + + +needs_backport = pytest.mark.skipif(not _has_backport(), reason="requires patched Megatron-LM + FLA >= 0.4.2") + + +# --------------------------------------------------------------------------- +# comparison helpers (run inside the workers) +# --------------------------------------------------------------------------- +def _prep(name, got, want): + got32 = got.detach().float().flatten() + want32 = want.detach().float().flatten() + assert got32.shape == want32.shape, f"{name}: shape {got32.shape} vs {want32.shape}" + assert torch.isfinite(got32).all(), f"{name}: non-finite values in candidate" + assert torch.isfinite(want32).all(), f"{name}: non-finite values in reference" + return got32, want32 + + +def _stats(got32, want32): + diff = (got32 - want32).abs() + rms = (diff.square().mean().sqrt() / (want32.square().mean().sqrt() + 1e-12)).item() + cos = torch.nn.functional.cosine_similarity(got32, want32, dim=0).item() + return diff, rms, cos + + +def _report_elementwise(name, got, want, atol, rtol): + """Per-token tensors: every element within atol + rtol * |ref|.""" + got32, want32 = _prep(name, got, want) + diff, rms, cos = _stats(got32, want32) + worst = (diff - (atol + rtol * want32.abs())).max().item() + assert worst <= 0, ( + f"{name}: max |diff| {diff.max().item():.3e} exceeds atol({atol:.0e})+rtol({rtol:.0e})*|ref| " + f"by {worst:.3e} (rms {rms:.3e}, cosine {cos:.8f})" + ) + assert cos >= MIN_COSINE, f"{name}: cosine {cos:.8f} < {MIN_COSINE}" + + +def _report_rms(name, got, want, ratio): + """Whole-tensor normalised RMS error -- the metric FLA's own CP tests + use.""" + got32, want32 = _prep(name, got, want) + diff, rms, cos = _stats(got32, want32) + assert rms < ratio, ( + f"{name}: normalised RMS error {rms:.3e} >= {ratio:.1e} (max |diff| {diff.max().item():.3e}, cosine {cos:.8f})" + ) + assert cos >= MIN_COSINE, f"{name}: cosine {cos:.8f} < {MIN_COSINE} (rms {rms:.3e})" + + +def _zigzag_shard(full: torch.Tensor, cu, cp_size: int, cp_rank: int) -> torch.Tensor: + from relax.backends.megatron.cp_utils import gdn_cp_slice + + return gdn_cp_slice(full, cu, cp_size, cp_rank) + + +def _init_dist(rank, world_size): + os.environ.setdefault("MASTER_ADDR", "127.0.0.1") + os.environ.setdefault("MASTER_PORT", "29531") + torch.cuda.set_device(rank) + import torch.distributed as dist + + dist.init_process_group("nccl", rank=rank, world_size=world_size) + + +# --------------------------------------------------------------------------- +# worker: FLA kernel level +# --------------------------------------------------------------------------- +def _worker_fla_kernels(rank, world_size, dtype_name, _unused): + seq_lens = [256, 128] + dtype = {"fp32": torch.float32, "bf16": torch.bfloat16}[dtype_name] + _init_dist(rank, world_size) + import torch.distributed as dist + from fla.modules.convolution import causal_conv1d + from fla.modules.l2norm import l2norm + from fla.ops.cp import build_cp_context + from fla.ops.gated_delta_rule import chunk_gated_delta_rule + + device = torch.device("cuda", rank) + cp_group = dist.new_group(list(range(world_size))) + + H, DK, DV, W = 2, 64, 64, 4 + total = sum(seq_lens) + cu = torch.tensor([0] + torch.tensor(seq_lens).cumsum(0).tolist(), device=device, dtype=torch.int32) + part = total // world_size + lo, hi = rank * part, (rank + 1) * part + + gen = torch.Generator(device="cpu").manual_seed(1234) + + def _mk(*shape, dtype=dtype): + return torch.randn(*shape, generator=gen, dtype=torch.float32).to(device=device, dtype=dtype) + + conv_ratio = KERNEL_RMS_RATIO_FP32 if dtype is torch.float32 else CONV_RMS_RATIO + gdn_ratio = KERNEL_RMS_RATIO_FP32 if dtype is torch.float32 else GDN_RMS_RATIO + tag = f"[rank{rank}/{dtype_name}]" + + # ---- causal conv ---- + # weight/bias are leaves here on purpose: their gradients are sums over every + # token, which is exactly the quantity chunkwise CP repartitions. Checking only + # dx would leave that untested. + x = _mk(1, total, H * DK) + w0 = _mk(H * DK, W) + b0 = _mk(H * DK) + conv_grad = _mk(1, total, H * DK) + + x_ref = x.clone().requires_grad_(True) + w_ref = w0.clone().requires_grad_(True) + b_ref = b0.clone().requires_grad_(True) + out_ref, _ = causal_conv1d(x=x_ref, weight=w_ref, bias=b_ref, activation="silu", cu_seqlens=cu) + (out_ref.float() * conv_grad.float()).sum().backward() + + x_cp = x[:, lo:hi].clone().requires_grad_(True) + w_cp = w0.clone().requires_grad_(True) + b_cp = b0.clone().requires_grad_(True) + ctx = build_cp_context(cu_seqlens=cu, group=cp_group, conv1d_kernel_size=W) + out_cp, _ = causal_conv1d(x=x_cp, weight=w_cp, bias=b_cp, activation="silu", cu_seqlens=cu, cp_context=ctx) + _report_rms(f"{tag} conv fwd", out_cp, out_ref[:, lo:hi], conv_ratio) + (out_cp.float() * conv_grad[:, lo:hi].float()).sum().backward() + _report_rms(f"{tag} conv dx", x_cp.grad, x_ref.grad[:, lo:hi], conv_ratio) + + for pname, cp_leaf, ref_leaf in (("dweight", w_cp, w_ref), ("dbias", b_cp, b_ref)): + summed = cp_leaf.grad.detach().float().clone() + dist.all_reduce(summed, group=cp_group) + if dtype is torch.float32: + _report_rms(f"{tag} conv {pname}", summed, ref_leaf.grad, conv_ratio) + else: + # These are 384-token sums landing in bf16. The fp32 parametrisation of + # this very test pins the algebra at ~1e-7; in bf16 the achievable + # agreement is set by the storage format, so assert direction and report + # the size rather than pretend a sub-ULP threshold is meaningful. + got32, want32 = _prep(f"{tag} conv {pname}", summed, ref_leaf.grad) + _, rms, cos = _stats(got32, want32) + assert cos >= MIN_COSINE, f"{tag} conv {pname}: cosine {cos:.8f} (rms {rms:.3e})" + if rank == 0: + print(f" {tag} conv {pname}: rms {rms:.3e} cosine {cos:.10f}") + + # ---- gated delta rule ---- + # Inputs must look like what GatedDeltaNet actually feeds the kernel: + # * q/k are L2-normalised (the module sets use_qk_l2norm=True). Un-normalised + # q/k make the recurrent state diverge over hundreds of steps and the + # reference itself goes to NaN -- that would test nothing. + # * g is a log-domain decay built as -A.exp() * softplus(...), hence <= 0. + q = l2norm(_mk(1, total, H, DK).contiguous()) + k = l2norm(_mk(1, total, H, DK).contiguous()) + v = _mk(1, total, H, DV) + g0 = -_mk(1, total, H, dtype=torch.float32).abs() * 0.1 + beta0 = _mk(1, total, H, dtype=torch.float32).sigmoid() + + leaves_ref = [t.detach().clone().requires_grad_(True) for t in (q, k, v)] + g_ref = g0.detach().clone().requires_grad_(True) + beta_ref = beta0.detach().clone().requires_grad_(True) + o_ref, _ = chunk_gated_delta_rule( + *leaves_ref, + g=g_ref, + beta=beta_ref, + initial_state=None, + output_final_state=False, + use_qk_l2norm_in_kernel=False, + cu_seqlens=cu, + ) + o_grad = _mk(1, total, H, DV) + (o_ref.float() * o_grad.float()).sum().backward() + + leaves_cp = [t.detach()[:, lo:hi].clone().requires_grad_(True) for t in (q, k, v)] + g_cp = g0.detach()[:, lo:hi].clone().requires_grad_(True) + beta_cp = beta0.detach()[:, lo:hi].clone().requires_grad_(True) + ctx2 = build_cp_context(cu_seqlens=cu, group=cp_group, conv1d_kernel_size=W) + o_cp, _ = chunk_gated_delta_rule( + *leaves_cp, + g=g_cp, + beta=beta_cp, + initial_state=None, + output_final_state=False, + use_qk_l2norm_in_kernel=False, + cu_seqlens=cu, + cp_context=ctx2, + ) + _report_rms(f"{tag} gdr fwd", o_cp, o_ref[:, lo:hi], gdn_ratio) + (o_cp.float() * o_grad[:, lo:hi].float()).sum().backward() + for name, a, b in zip("qkv", leaves_cp, leaves_ref): + _report_rms(f"{tag} gdr d{name}", a.grad, b.grad[:, lo:hi], gdn_ratio) + _report_rms(f"{tag} gdr dg", g_cp.grad, g_ref.grad[:, lo:hi], gdn_ratio) + _report_rms(f"{tag} gdr dbeta", beta_cp.grad, beta_ref.grad[:, lo:hi], gdn_ratio) + + dist.barrier() + dist.destroy_process_group() + + +# --------------------------------------------------------------------------- +# worker: full MCore GatedDeltaNet +# --------------------------------------------------------------------------- +def _build_gdn( + cp_size, + linear_cp_mode, + dtype, + num_key_heads=4, + num_value_heads=8, + tp_size=1, + deterministic_mode=False, +): + import torch.nn.functional as F + from megatron.core import parallel_state + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_experimental_attention_variant_module_spec, + ) + from megatron.core.process_groups_config import ProcessGroupCollection + from megatron.core.ssm.gated_delta_net import GatedDeltaNet + from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed + from megatron.core.transformer.transformer_config import TransformerConfig + + model_parallel_cuda_manual_seed(123) + config = TransformerConfig( + hidden_size=512, + num_layers=1, + num_attention_heads=8, + num_query_groups=2, + normalization="RMSNorm", + use_cpu_initialization=True, + layernorm_zero_centered_gamma=True, + activation_func=F.silu, + bf16=dtype is torch.bfloat16, + tensor_model_parallel_size=tp_size, + context_parallel_size=cp_size, + deterministic_mode=deterministic_mode, + experimental_attention_variant="gated_delta_net", + linear_attention_freq=[1], + linear_conv_kernel_dim=4, + linear_key_head_dim=64, + linear_value_head_dim=64, + linear_num_key_heads=num_key_heads, + linear_num_value_heads=num_value_heads, + linear_cp_mode=linear_cp_mode, + transformer_impl="transformer_engine", + ) + pg_collection = ProcessGroupCollection( + tp=parallel_state.get_tensor_model_parallel_group(), + cp=parallel_state.get_context_parallel_group(), + ) + gdn = GatedDeltaNet( + config, + submodules=get_experimental_attention_variant_module_spec(config=config).submodules, + layer_number=1, + bias=False, + conv_bias=False, + conv_init=1.0, + use_qk_l2norm=True, + A_init_range=(1, 16), + pg_collection=pg_collection, + ) + return gdn.cuda().to(dtype), config + + +def _run_gdn_once(gdn, hidden, psp, grad_out, *, recompute=False, **forward_kwargs): + """One forward+backward; returns (out, d_hidden, {param: grad}).""" + gdn.zero_grad(set_to_none=True) + h = hidden.clone().requires_grad_(True) + if recompute: + from torch.utils.checkpoint import checkpoint + + out = checkpoint( + lambda x: gdn(x, None, packed_seq_params=psp, **forward_kwargs)[0], + h, + use_reentrant=False, + ) + else: + out, _ = gdn(h, None, packed_seq_params=psp, **forward_kwargs) + (out.float() * grad_out).sum().backward() + grads = {n: p.grad.detach().float().clone() for n, p in gdn.named_parameters()} + return out.detach().clone(), h.grad.detach().clone(), grads + + +def _worker_gdn_module(rank, world_size, dtype_name, _unused): + """CP=1 reference vs CP=N, for BOTH CP algorithms, in one process. + + Running headwise and chunkwise side by side is the point: it turns "is + chunkwise close enough to CP=1" (which needs an absolute threshold, and in + bf16 lands on the noise floor of the storage format) into "is chunkwise as + close to CP=1 as the algorithm we already ship" -- a comparison with no + free parameters to tune. + """ + dtype = {"fp32": torch.float32, "bf16": torch.bfloat16}[dtype_name] + + _init_dist(rank, world_size) + import torch.distributed as dist + from megatron.core import parallel_state + from megatron.core.packed_seq_params import PackedSeqParams + + parallel_state.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + context_parallel_size=world_size, + ) + device = torch.device("cuda", rank) + # The CP algorithm is static config now, so each mode needs its own module. They + # share weights, so the comparison is still like-for-like. + modules = {} + gdn, config = _build_gdn(world_size, "headwise", dtype) + for p in gdn.parameters(): + # Same weights on every rank so the CP=1 reference is rank-independent. + dist.broadcast(p.data, src=0) + modules["headwise"] = gdn + modules["chunkwise"], _ = _build_gdn(world_size, "chunkwise", dtype) + modules["chunkwise"].load_state_dict(gdn.state_dict()) + + cp_group = parallel_state.get_context_parallel_group() + # A per-rank size-1 group gives us the CP=1 reference *inside* the same + # process, driving the very same weights through the very same forward. + solo = [dist.new_group([r]) for r in range(world_size)][rank] + + seq_lens = [256, 128] + total = sum(seq_lens) + cu = torch.tensor([0, seq_lens[0], total], device=device, dtype=torch.int32) + + gen = torch.Generator(device="cpu").manual_seed(7) + hidden_full = torch.randn(total, 1, config.hidden_size, generator=gen, dtype=torch.float32).to( + device=device, dtype=dtype + ) + grad_seed = torch.randn(total, 1, config.hidden_size, generator=gen, dtype=torch.float32).to(device) + + def _psp(group, local_cp_size): + return PackedSeqParams( + qkv_format="thd", + cu_seqlens_q=cu, + cu_seqlens_kv=cu, + cu_seqlens_q_padded=cu, + cu_seqlens_kv_padded=cu, + max_seqlen_q=max(seq_lens), + max_seqlen_kv=max(seq_lens), + cp_group=group, + local_cp_size=local_cp_size, + ) + + # CP=1 reference. A size-1 group short-circuits before the mode is read, so either + # module gives the same reference; use the headwise one. + out_ref, in_grad_ref, param_grads_ref = _run_gdn_once(gdn, hidden_full, _psp(solo, 1), grad_seed) + ref = { + "out": _zigzag_shard(out_ref, cu, world_size, rank), + "d_hidden": _zigzag_shard(in_grad_ref, cu, world_size, rank), + } + ref.update({f"grad {n}": g for n, g in param_grads_ref.items()}) + + shard = _zigzag_shard(hidden_full, cu, world_size, rank) + grad_shard = _zigzag_shard(grad_seed, cu, world_size, rank) + + metrics = {} + for mode in ("headwise", "chunkwise"): + out_cp, in_grad_cp, param_grads_cp = _run_gdn_once( + modules[mode], shard, _psp(cp_group, world_size), grad_shard + ) + got = {"out": out_cp, "d_hidden": in_grad_cp} + # Each CP rank holds a partial parameter gradient; the total is the CP sum. + for name, g in param_grads_cp.items(): + summed = g.clone() + dist.all_reduce(summed, group=cp_group) + got[f"grad {name}"] = summed + metrics[mode] = {} + for key, value in got.items(): + got32, want32 = _prep(f"[rank{rank}][{mode}/{dtype_name}] {key}", value, ref[key]) + diff, rms, cos = _stats(got32, want32) + metrics[mode][key] = (rms, cos, diff.max().item()) + assert cos >= MIN_COSINE, ( + f"[rank{rank}][{mode}/{dtype_name}] {key}: cosine {cos:.8f} < {MIN_COSINE} (rms {rms:.3e})" + ) + + # RFC section 5 absolute gate, on the per-token tensors, in the dtype the + # RFC specifies it for. Applied to headwise too, so a drift in the shared + # plumbing cannot hide behind the comparative check below. + if dtype is torch.bfloat16: + for key in ("out", "d_hidden"): + _report_elementwise( + f"[rank{rank}][{mode}/{dtype_name}] {key}", got[key], ref[key], ATOL_BF16, RTOL_BF16 + ) + # Parameter gradients are token-sum reductions stored in bf16. Bound them + # by the RFC's own relative tolerance for this comparison row (rtol=1e-2) + # applied to the whole tensor, plus the RFC's cosine floor. What actually + # pins the algebra is the fp32 parametrisation of this same test. + for key, (rms, cos, _) in metrics[mode].items(): + if key in ("out", "d_hidden"): + continue + assert rms < RTOL_BF16, ( + f"[rank{rank}][{mode}/bf16] {key}: relative RMS {rms:.3e} >= {RTOL_BF16:.0e} (cosine {cos:.8f})" + ) + else: + for key, (rms, _, _) in metrics[mode].items(): + assert rms < RMS_RATIO_FP32, f"[rank{rank}][{mode}/fp32] {key}: rms {rms:.3e} >= {RMS_RATIO_FP32:.0e}" + + # The comparative assertion -- fp32 only, on purpose. + # + # Its premise is "headwise's disagreement with CP=1 is the floor this environment + # imposes". That holds only while both algorithms perform the *same* reductions. + # They do not: headwise hands each rank the whole sequence and 1/cp of the heads, so + # a gradient like conv1d.weight / dt_bias / A_log (a sum over every token) is summed + # in one go exactly as at CP=1 and can come out bit-exact. Chunkwise splits the + # tokens, so that same sum really is partitioned and re-added. In fp32 the mantissa + # absorbs it and the two are directly comparable (observed 1.00x-1.10x). In bf16 the + # repartitioned sum sits on the format's ULP floor while headwise sits near zero, so + # their *ratio* measures the dtype, not the algorithm -- bf16 is covered by the + # absolute gates above instead. + worst = [] + for key, (rms_c, cos_c, max_c) in metrics["chunkwise"].items(): + rms_h = metrics["headwise"][key][0] + worst.append((rms_c / max(rms_h, 1e-12), key, rms_c, rms_h)) + if dtype is not torch.float32: + continue + assert rms_c <= CHUNKWISE_VS_HEADWISE_RMS_FACTOR * max(rms_h, RMS_FLOOR_FP32), ( + f"[rank{rank}][{dtype_name}] {key}: chunkwise rms {rms_c:.3e} exceeds " + f"{CHUNKWISE_VS_HEADWISE_RMS_FACTOR}x the shipped headwise rms {rms_h:.3e} " + f"(cosine {cos_c:.8f}, max |diff| {max_c:.3e})" + ) + worst.sort(reverse=True) + if rank == 0: + print(f"\n[{dtype_name}] chunkwise vs headwise, worst 6 by rms ratio:") + for ratio, key, rms_c, rms_h in worst[:6]: + shown = f"{ratio:6.2f}x" if rms_h > 0 else " n/a" + print(f" {key:38s} {shown} chunkwise {rms_c:.3e} headwise {rms_h:.3e}") + + dist.barrier() + parallel_state.destroy_model_parallel() + dist.destroy_process_group() + + +def _worker_deterministic_reference(rank, world_size, _spec, _unused): + """The torch-native deterministic rule must accept cp_context=None and + preserve headwise CP correctness.""" + _init_dist(rank, world_size) + import torch.distributed as dist + from megatron.core import parallel_state + from megatron.core.process_groups_config import ProcessGroupCollection + + parallel_state.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + context_parallel_size=world_size, + ) + device = torch.device("cuda", rank) + cp_group = parallel_state.get_context_parallel_group() + tp_group = parallel_state.get_tensor_model_parallel_group() + solo = [dist.new_group([r]) for r in range(world_size)][rank] + + gdn, config = _build_gdn( + world_size, + "headwise", + torch.float32, + deterministic_mode=True, + ) + assert gdn.gated_delta_rule.__name__ == "torch_chunk_gated_delta_rule" + for p in gdn.parameters(): + dist.broadcast(p.data, src=0, group=cp_group) + + total = 64 + cu = torch.tensor([0, total], device=device, dtype=torch.int32) + gen = torch.Generator(device="cpu").manual_seed(17) + hidden_full = torch.randn(total, 1, config.hidden_size, generator=gen).to(device) + grad_full = torch.randn(total, 1, config.hidden_size, generator=gen).to(device) + + solo_pg = ProcessGroupCollection(tp=tp_group, cp=solo) + out_ref, in_grad_ref, param_grads_ref = _run_gdn_once( + gdn, + hidden_full, + None, + grad_full, + pg_collection=solo_pg, + ) + + hidden_shard = _zigzag_shard(hidden_full, cu, world_size, rank) + grad_shard = _zigzag_shard(grad_full, cu, world_size, rank) + out_cp, in_grad_cp, param_grads_cp = _run_gdn_once(gdn, hidden_shard, None, grad_shard) + + _report_rms( + f"[rank{rank}] deterministic out", + out_cp, + _zigzag_shard(out_ref, cu, world_size, rank), + RMS_RATIO_FP32, + ) + _report_rms( + f"[rank{rank}] deterministic d_hidden", + in_grad_cp, + _zigzag_shard(in_grad_ref, cu, world_size, rank), + RMS_RATIO_FP32, + ) + for name, grad in param_grads_cp.items(): + summed = grad.clone() + dist.all_reduce(summed, group=cp_group) + _report_rms( + f"[rank{rank}] deterministic grad {name}", + summed, + param_grads_ref[name], + RMS_RATIO_FP32, + ) + + dist.barrier() + parallel_state.destroy_model_parallel() + dist.destroy_process_group() + + +def _worker_recompute_parity(rank, world_size, recompute_kind, _unused): + """External activation checkpointing must replay chunkwise collectives + without changing outputs or gradients.""" + _init_dist(rank, world_size) + import torch.distributed as dist + from megatron.core import parallel_state + from megatron.core.packed_seq_params import PackedSeqParams + + parallel_state.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + context_parallel_size=world_size, + ) + device = torch.device("cuda", rank) + cp_group = parallel_state.get_context_parallel_group() + + gdn, config = _build_gdn(world_size, "chunkwise", torch.float32) + for p in gdn.parameters(): + dist.broadcast(p.data, src=0, group=cp_group) + + seq_lens = [128, 64] + total = sum(seq_lens) + cu = torch.tensor([0, seq_lens[0], total], device=device, dtype=torch.int32) + psp = PackedSeqParams( + qkv_format="thd", + cu_seqlens_q=cu, + cu_seqlens_kv=cu, + cu_seqlens_q_padded=cu, + cu_seqlens_kv_padded=cu, + max_seqlen_q=max(seq_lens), + max_seqlen_kv=max(seq_lens), + cp_group=cp_group, + local_cp_size=world_size, + ) + + gen = torch.Generator(device="cpu").manual_seed(29) + hidden_full = torch.randn(total, 1, config.hidden_size, generator=gen).to(device) + grad_full = torch.randn(total, 1, config.hidden_size, generator=gen).to(device) + hidden = _zigzag_shard(hidden_full, cu, world_size, rank) + grad = _zigzag_shard(grad_full, cu, world_size, rank) + + out_eager, in_grad_eager, param_grads_eager = _run_gdn_once(gdn, hidden, psp, grad) + if recompute_kind == "selective": + # Use the upgraded MCore's own GDN checkpoint path, not the external wrapper. + gdn.recompute_gdn = True + out_recompute, in_grad_recompute, param_grads_recompute = _run_gdn_once( + gdn, + hidden, + psp, + grad, + recompute=recompute_kind == "full", + ) + + _report_rms(f"[rank{rank}] recompute out", out_recompute, out_eager, KERNEL_RMS_RATIO_FP32) + _report_rms( + f"[rank{rank}] recompute d_hidden", + in_grad_recompute, + in_grad_eager, + KERNEL_RMS_RATIO_FP32, + ) + assert set(param_grads_recompute) == set(param_grads_eager) + for name in param_grads_eager: + _report_rms( + f"[rank{rank}] recompute grad {name}", + param_grads_recompute[name], + param_grads_eager[name], + KERNEL_RMS_RATIO_FP32, + ) + + dist.barrier() + parallel_state.destroy_model_parallel() + dist.destroy_process_group() + + +def _worker_tp2_cp2(rank, world_size, _spec, _unused): + """Exercise TP head sharding and CP routing together.""" + assert world_size == 4 + _init_dist(rank, world_size) + import torch.distributed as dist + from megatron.core import parallel_state + from megatron.core.packed_seq_params import PackedSeqParams + from megatron.core.process_groups_config import ProcessGroupCollection + + parallel_state.initialize_model_parallel( + tensor_model_parallel_size=2, + pipeline_model_parallel_size=1, + context_parallel_size=2, + ) + device = torch.device("cuda", rank) + cp_group = parallel_state.get_context_parallel_group() + cp_rank = cp_group.rank() + tp_group = parallel_state.get_tensor_model_parallel_group() + cp_source = dist.get_process_group_ranks(cp_group)[0] + solo = [dist.new_group([r]) for r in range(world_size)][rank] + + modules = {} + headwise, config = _build_gdn(2, "headwise", torch.float32, tp_size=2) + for p in headwise.parameters(): + dist.broadcast(p.data, src=cp_source, group=cp_group) + modules["headwise"] = headwise + modules["chunkwise"], _ = _build_gdn(2, "chunkwise", torch.float32, tp_size=2) + modules["chunkwise"].load_state_dict(headwise.state_dict()) + sharded_signatures = {} + for mode, module in modules.items(): + sharded_signatures[mode] = { + key: ( + tuple(getattr(value, "global_shape", ())), + tuple(getattr(value, "local_shape", ())), + getattr(value, "axis_fragmentations", None), + ) + for key, value in sorted(module.sharded_state_dict(prefix="mixer.").items()) + } + assert sharded_signatures["headwise"] == sharded_signatures["chunkwise"] + + seq_lens = [128, 64] + total = sum(seq_lens) + cu = torch.tensor([0, seq_lens[0], total], device=device, dtype=torch.int32) + gen = torch.Generator(device="cpu").manual_seed(41) + hidden_full = torch.randn(total, 1, config.hidden_size, generator=gen).to(device) + grad_full = torch.randn(total, 1, config.hidden_size, generator=gen).to(device) + + def _psp(group, cp_size): + return PackedSeqParams( + qkv_format="thd", + cu_seqlens_q=cu, + cu_seqlens_kv=cu, + cu_seqlens_q_padded=cu, + cu_seqlens_kv_padded=cu, + max_seqlen_q=max(seq_lens), + max_seqlen_kv=max(seq_lens), + cp_group=group, + local_cp_size=cp_size, + ) + + solo_pg = ProcessGroupCollection(tp=tp_group, cp=solo) + out_ref, in_grad_ref, param_grads_ref = _run_gdn_once( + headwise, + hidden_full, + _psp(solo, 1), + grad_full, + pg_collection=solo_pg, + ) + hidden = _zigzag_shard(hidden_full, cu, 2, cp_rank) + grad = _zigzag_shard(grad_full, cu, 2, cp_rank) + out_want = _zigzag_shard(out_ref, cu, 2, cp_rank) + in_grad_want = _zigzag_shard(in_grad_ref, cu, 2, cp_rank) + + for mode, module in modules.items(): + out, in_grad, param_grads = _run_gdn_once(module, hidden, _psp(cp_group, 2), grad) + _report_rms(f"[rank{rank}][{mode}] TP2/CP2 out", out, out_want, RMS_RATIO_FP32) + _report_rms( + f"[rank{rank}][{mode}] TP2/CP2 d_hidden", + in_grad, + in_grad_want, + RMS_RATIO_FP32, + ) + assert set(param_grads) == set(param_grads_ref) + for name, param_grad in param_grads.items(): + summed = param_grad.clone() + dist.all_reduce(summed, group=cp_group) + _report_rms( + f"[rank{rank}][{mode}] TP2/CP2 grad {name}", + summed, + param_grads_ref[name], + RMS_RATIO_FP32, + ) + + dist.barrier() + parallel_state.destroy_model_parallel() + dist.destroy_process_group() + + +def _worker_layout_round_trip(rank, world_size, _spec, _unused): + """zigzag -> contiguous -> zigzag over a real CP group must be token-exact. + + This is the collective-level version of RFC 5.1: it drives the actual + ``all_to_all`` in ``context_parallel_layout``, for packed THD (several + unequal-length samples) and for SBHD. + """ + _init_dist(rank, world_size) + import torch.distributed as dist + from megatron.core.context_parallel_layout import ( + contiguous_to_zigzag_chunks, + get_thd_context_parallel_rank_indices, + zigzag_to_contiguous_chunks, + ) + + device = torch.device("cuda", rank) + cp_group = dist.new_group(list(range(world_size))) + + # --- packed THD, three samples of different lengths --- + lengths = [2 * world_size * f for f in (5, 1, 3)] + cu = torch.tensor([0] + torch.tensor(lengths).cumsum(0).tolist(), device=device, dtype=torch.int32) + total = int(cu[-1]) + # Row t is (t, t+1e6, t+2e6): a permuted token is impossible to miss. + full = ( + torch.arange(total, dtype=torch.float64, device=device).unsqueeze(1) + + torch.arange(3, dtype=torch.float64, device=device).unsqueeze(0) * 1e6 + ) + + zig_idx = get_thd_context_parallel_rank_indices(cu, world_size, rank, "zigzag") + con_idx = get_thd_context_parallel_rank_indices(cu, world_size, rank, "contiguous") + local_zig = full[zig_idx] + + got_con = zigzag_to_contiguous_chunks(local_zig, cp_group, seq_dim=0, cu_seqlens=cu) + assert torch.equal(got_con, full[con_idx]), f"rank {rank}: THD zigzag->contiguous is wrong" + got_zig = contiguous_to_zigzag_chunks(got_con, cp_group=cp_group, seq_dim=0, cu_seqlens=cu) + assert torch.equal(got_zig, local_zig), f"rank {rank}: THD round trip is not identity" + + # --- SBHD (chunk-level swap, no cu_seqlens) --- + seq_local = 2 * world_size * 4 + sbhd = torch.arange(seq_local * 2 * 3, dtype=torch.float64, device=device).reshape(seq_local, 2, 3) + rank * 1e9 + swapped = zigzag_to_contiguous_chunks(sbhd, cp_group, seq_dim=0) + back = contiguous_to_zigzag_chunks(swapped, cp_group=cp_group, seq_dim=0) + assert torch.equal(back, sbhd), f"rank {rank}: SBHD round trip is not identity" + + dist.barrier() + dist.destroy_process_group() + + +def _worker_illegal_modes_fail_fast(rank, world_size, _spec, _unused): + """Illegal / unresolved GDN CP modes must raise, not silently pick an + algorithm.""" + _init_dist(rank, world_size) + import torch.distributed as dist + from megatron.core import parallel_state + from megatron.core.packed_seq_params import PackedSeqParams + + parallel_state.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + context_parallel_size=world_size, + ) + device = torch.device("cuda", rank) + cp_group = parallel_state.get_context_parallel_group() + solo = [dist.new_group([r]) for r in range(world_size)][rank] + + seq_lens = [256, 128] + total = sum(seq_lens) + cu = torch.tensor([0, seq_lens[0], total], device=device, dtype=torch.int32) + + def _psp(group, local_cp_size): + return PackedSeqParams( + qkv_format="thd", + cu_seqlens_q=cu, + cu_seqlens_kv=cu, + cu_seqlens_q_padded=cu, + cu_seqlens_kv_padded=cu, + max_seqlen_q=max(seq_lens), + max_seqlen_kv=max(seq_lens), + cp_group=group, + local_cp_size=local_cp_size, + ) + + gdn, config = _build_gdn(world_size, "all_gather", torch.bfloat16) + hidden = torch.randn(total // world_size, 1, config.hidden_size, device=device, dtype=torch.bfloat16) + + # 1. all_gather is a Relax-wrapper mode: reaching MCore's forward with cp>1 means the + # wrapper is missing, and that must be loud. + with pytest.raises(RuntimeError, match="requires the Relax"): + gdn(hidden, None, packed_seq_params=_psp(cp_group, world_size)) + + # 2. ...but a CP=1 micro-batch is legal under any declared mode: it needs no CP + # communication at all, so the mode is never consulted. + full = torch.randn(total, 1, config.hidden_size, device=device, dtype=torch.bfloat16) + solo_psp = _psp(solo, 1) + gdn(full, None, packed_seq_params=solo_psp) + + # 3. The final PackedSeqParams must describe the same runtime CP geometry + # through both fields. + gdn_hw, _ = _build_gdn(world_size, "headwise", torch.bfloat16) + with pytest.raises(ValueError, match=r"does not match.*cp_group.size"): + gdn_hw(hidden, None, packed_seq_params=_psp(cp_group, world_size + 1)) + + # 4. deterministic mode has no CP-context scan. + gdn_cw, _ = _build_gdn(world_size, "chunkwise", torch.bfloat16) + gdn_cw.config.deterministic_mode = True + try: + with pytest.raises((ValueError, AssertionError)): + gdn_cw(hidden, None, packed_seq_params=_psp(cp_group, world_size)) + finally: + gdn_cw.config.deterministic_mode = False + + # 5. inference is not supported. + class _Ctx: + def is_static_batching(self): + return True + + with pytest.raises(NotImplementedError): + gdn_cw(hidden, None, inference_context=_Ctx(), packed_seq_params=_psp(cp_group, world_size)) + + dist.barrier() + parallel_state.destroy_model_parallel() + dist.destroy_process_group() + + +def _worker_state_dict_invariance(rank, world_size, _spec, _unused): + """GDN weights must stay TP-only: same keys and shard dims in every CP + mode.""" + _init_dist(rank, world_size) + import io + + import torch.distributed as dist + from megatron.core import parallel_state + + parallel_state.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + context_parallel_size=world_size, + ) + device = torch.device("cuda", rank) + + signatures = {} + for cp_size, mode in ((1, "headwise"), (world_size, "headwise"), (world_size, "chunkwise")): + gdn, _ = _build_gdn(cp_size, mode, torch.bfloat16) + sd = gdn.state_dict() + sharded = gdn.sharded_state_dict(prefix="mixer.") + signatures[(cp_size, mode)] = ( + {k: tuple(v.shape) for k, v in sd.items() if torch.is_tensor(v)}, + { + k: ( + tuple(getattr(v, "global_shape", ())), + tuple(getattr(v, "local_shape", ())), + getattr(v, "axis_fragmentations", None), + ) + for k, v in sorted(sharded.items()) + }, + ) + + # Exercise real serialization and loading, not just key comparison. + checkpoint = io.BytesIO() + torch.save(sd, checkpoint) + checkpoint.seek(0) + loaded = torch.load(checkpoint, map_location=device, weights_only=True) + restored, _ = _build_gdn(cp_size, mode, torch.bfloat16) + restored.load_state_dict(loaded, strict=True) + for name, tensor in sd.items(): + if torch.is_tensor(tensor): + assert torch.equal(restored.state_dict()[name], tensor), ( + f"state_dict round trip changed {name} for {(cp_size, mode)}" + ) + del restored + del gdn + + baseline = signatures[(1, "headwise")] + assert baseline[0], "state_dict is empty; the invariance check would be vacuous" + assert baseline[1], "sharded_state_dict is empty; the invariance check would be vacuous" + for key, sig in signatures.items(): + assert sig[0] == baseline[0], f"state_dict shapes changed for {key}" + assert set(sig[1]) == set(baseline[1]), f"sharded_state_dict keys changed for {key}" + for k in baseline[1]: + assert sig[1][k] == baseline[1][k], f"sharded shard dims changed for {key} at {k}" + + parallel_state.destroy_model_parallel() + dist.destroy_process_group() + + +# --------------------------------------------------------------------------- +# pytest entry points +# --------------------------------------------------------------------------- +def _spawn(fn, spec, port, extra=None): + _spawn_world(fn, WORLD_SIZE, spec, port, extra=extra) + + +def _spawn_world(fn, world_size, spec, port, extra=None): + os.environ["MASTER_PORT"] = str(port) + mp.spawn( + fn, + args=(world_size, extra if extra is not None else spec, None), + nprocs=world_size, + join=True, + ) + + +@needs_backport +@pytest.mark.parametrize("dtype_name,port", [("fp32", 29540), ("bf16", 29541)]) +def test_fla_cp_kernels_match_single_rank(dtype_name, port): + """FLA causal_conv1d / chunk_gated_delta_rule under cp_context vs no CP.""" + _spawn(_worker_fla_kernels, dtype_name, port) + + +@needs_backport +@pytest.mark.parametrize("dtype_name,port", [("fp32", 29542), ("bf16", 29543)]) +def test_gdn_cp_matches_cp1(dtype_name, port): + """RFC 3.4 item 3: CP=2 vs CP=1, forward and backward, both CP + algorithms.""" + _spawn(_worker_gdn_module, dtype_name, port) + + +@needs_backport +def test_deterministic_headwise_cp_matches_cp1(): + """The torch-native rule accepts cp_context=None and stays correct under + headwise CP.""" + _spawn(_worker_deterministic_reference, "n/a", 29547) + + +@needs_backport +@pytest.mark.parametrize("recompute_kind", ["full", "selective"]) +def test_chunkwise_recompute_matches_eager(recompute_kind): + """Replaying chunkwise forward during backward preserves every tested + gradient.""" + _spawn(_worker_recompute_parity, recompute_kind, 29548) + + +@needs_backport +@pytest.mark.skipif(torch.cuda.device_count() < 4, reason="requires 4 CUDA devices for TP2/CP2") +def test_gdn_tp2_cp2_matches_cp1(): + """TP2/CP2 exercises TP head shards and both CP algorithms together.""" + _spawn_world(_worker_tp2_cp2, 4, "n/a", 29549) + + +@needs_backport +def test_layout_round_trip_over_real_cp_group(): + """RFC 5.1 at the collective level: the layout swap is a pure + permutation.""" + _spawn(_worker_layout_round_trip, "n/a", 29544) + + +@needs_backport +def test_illegal_gdn_cp_modes_fail_fast(): + """RFC emphasis 1: illegal combinations must fail fast, never pick + silently.""" + _spawn(_worker_illegal_modes_fail_fast, "n/a", 29545) + + +@needs_backport +def test_gdn_state_dict_invariant_across_cp_modes(): + """RFC 3.4 item 4: checkpoint keys and shard dims do not depend on the CP + mode.""" + _spawn(_worker_state_dict_invariance, "n/a", 29546) diff --git a/tests/backends/megatron/test_gdn_chunkwise_cp_layout.py b/tests/backends/megatron/test_gdn_chunkwise_cp_layout.py index a651c8c2a..05ad47c03 100644 --- a/tests/backends/megatron/test_gdn_chunkwise_cp_layout.py +++ b/tests/backends/megatron/test_gdn_chunkwise_cp_layout.py @@ -201,17 +201,17 @@ def _gdn_config(**overrides): return TransformerConfig(**kwargs) -def test_config_default_mode_is_headwise(): +def test_config_default_mode_is_chunkwise(): """Upgrading the image must not silently reroute an existing recipe.""" - assert _gdn_config().linear_cp_mode == "headwise" + assert _gdn_config().linear_cp_mode == "chunkwise" def test_headwise_config_requires_heads_divisible_by_tp_times_cp(): # 16 key heads, tp=2, cp=4 -> 16 % 8 == 0: fine. - _gdn_config(tensor_model_parallel_size=2, context_parallel_size=4) + _gdn_config(tensor_model_parallel_size=2, context_parallel_size=4, linear_cp_mode="headwise") # tp=2, cp=16 -> 16 % 32 != 0: the geometry headwise cannot express. with pytest.raises(AssertionError, match="linear_num_key_heads"): - _gdn_config(tensor_model_parallel_size=2, context_parallel_size=16) + _gdn_config(tensor_model_parallel_size=2, context_parallel_size=16, linear_cp_mode="headwise") def test_chunkwise_config_only_requires_heads_divisible_by_tp(): diff --git a/tests/backends/megatron/test_gdn_chunkwise_cp_route.py b/tests/backends/megatron/test_gdn_chunkwise_cp_route.py index 6276150c6..03ea96637 100644 --- a/tests/backends/megatron/test_gdn_chunkwise_cp_route.py +++ b/tests/backends/megatron/test_gdn_chunkwise_cp_route.py @@ -256,3 +256,99 @@ def test_prebuild_is_a_noop_without_context_parallelism(): non_thd = PackedSeqParams(qkv_format="sbhd") cpl.prebuild_thd_cp_partition_routes(non_thd) assert getattr(non_thd, "cp_partition_route_zigzag_to_contiguous", None) is None + + +def test_route_does_not_cache_unversioned_inference_boundaries(): + with torch.inference_mode(): + cu = _cu([3, 1], unit=8) + psp = _packed_seq_params(cu) + first = cpl.get_thd_cp_partition_route(psp, cu, 4, 1, "zigzag", "contiguous") + cu.copy_(_cu([2, 2], unit=8)) + second = cpl.get_thd_cp_partition_route(psp, cu, 4, 1, "zigzag", "contiguous") + assert second is not first + expected = cpl.build_thd_cp_partition_route(cu, 4, 1, "zigzag", "contiguous") + assert second.input_split_sizes == expected.input_split_sizes + + +def _boundary_resolver(): + from types import SimpleNamespace + + from megatron.core.ssm.gated_delta_net import GatedDeltaNet + + calls = [] + module = SimpleNamespace(cp_size=8) + + def resolve(*args, **kwargs): + calls.append(kwargs["cp_size"]) + return GatedDeltaNet._resolve_cu_seqlens(module, *args, **kwargs) + + module._resolve_cu_seqlens = resolve + return module, calls, GatedDeltaNet._resolve_thd_cu_seqlens + + +def test_boundary_validation_is_shared_across_layers_and_recompute(): + module, calls, resolve = _boundary_resolver() + # Static max CP=8 would reject a length of 12; runtime CP=2 is legal. + cu = _cu([3, 1], unit=4) + psp = _packed_seq_params(cu) + first = resolve(module, psp, 16, 2) + assert calls == [2, 2] + assert resolve(module, psp, 16, 2) is first + assert calls == [2, 2] + + other_module, other_calls, _ = _boundary_resolver() + assert resolve(other_module, psp, 16, 2) is first + assert other_calls == [] + + +@pytest.mark.parametrize("change", ["replace", "inplace", "view", "new_pack", "runtime_cp", "global_length", "padded"]) +def test_boundary_cache_is_invalidated_when_inputs_change(change): + module, calls, resolve = _boundary_resolver() + cu = _cu([3, 1], unit=8) + psp = _packed_seq_params(cu) + resolve(module, psp, 32, 2) + cp_size, total = 2, 32 + if change == "replace": + psp.cu_seqlens_q = cu.clone() + psp.cu_seqlens_kv = psp.cu_seqlens_q + elif change == "inplace": + cu.copy_(_cu([2, 2], unit=8)) + elif change == "view": + cu[1:2].fill_(16) + elif change == "new_pack": + psp = _packed_seq_params(cu) + elif change == "runtime_cp": + cp_size = 4 + elif change == "global_length": + total = 64 + elif change == "padded": + psp.cu_seqlens_q_padded = _cu([2, 2], unit=8) + psp.cu_seqlens_kv_padded = psp.cu_seqlens_q_padded + if change == "global_length": + with pytest.raises(ValueError, match="total_sequence_length"): + resolve(module, psp, total, cp_size) + assert len(calls) == 3 + else: + resolve(module, psp, total, cp_size) + assert len(calls) == 4 + + +def test_boundary_validation_rechecks_inference_tensors(): + module, calls, resolve = _boundary_resolver() + with torch.inference_mode(): + cu = _cu([3, 1], unit=8) + psp = _packed_seq_params(cu) + resolve(module, psp, 32, 2) + cu[1] = 16 + resolve(module, psp, 32, 2) + assert len(calls) == 4 + + +def test_boundary_cache_does_not_hide_q_kv_mismatch(): + module, calls, resolve = _boundary_resolver() + cu = _cu([3, 1], unit=8) + psp = _packed_seq_params(cu) + resolve(module, psp, 32, 2) + psp.cu_seqlens_kv = _cu([2, 2], unit=8) + with pytest.raises(AssertionError, match="cu_seqlens_q equals"): + resolve(module, psp, 32, 2) diff --git a/tests/backends/megatron/test_gdn_cp_mode_stage2.py b/tests/backends/megatron/test_gdn_cp_mode_stage2.py index 96034560a..83f776b8f 100644 --- a/tests/backends/megatron/test_gdn_cp_mode_stage2.py +++ b/tests/backends/megatron/test_gdn_cp_mode_stage2.py @@ -16,9 +16,12 @@ from __future__ import annotations import argparse -from types import SimpleNamespace +import ast +from pathlib import Path +from types import ModuleType, SimpleNamespace import pytest +import torch pytest.importorskip("megatron.core.context_parallel_layout", reason="requires the patched Megatron-LM") @@ -26,28 +29,47 @@ from megatron.core.packed_seq_params import PackedSeqParams # noqa: E402 from megatron.core.ssm.gated_delta_net import GatedDeltaNet # noqa: E402 -from relax.backends.megatron import arguments as megatron_arguments # noqa: E402 -from relax.backends.megatron import model as gdn_model # noqa: E402 -from relax.backends.megatron.arguments import _validate_linear_cp_mode # noqa: E402 +from relax.backends.megatron.gdn_cp_config import _validate_linear_cp_mode # noqa: E402 + + +def _load_relax_functions(filename, names): + """Execute the real functions without importing the Ray/optimizer stack.""" + path = Path(__file__).resolve().parents[3] / "relax/backends/megatron" / filename + tree = ast.parse(path.read_text()) + selected = [node for node in tree.body if isinstance(node, ast.FunctionDef) and node.name in names] + assert {node.name for node in selected} == set(names) + module = ModuleType("_gdn_cp_test_" + path.stem) + module.__package__ = "relax.backends.megatron" + module.torch = torch + tree = ast.Module( + body=[ast.ImportFrom(module="__future__", names=[ast.alias(name="annotations")], level=0), *selected], + type_ignores=[], + ) + exec(compile(ast.fix_missing_locations(tree), str(path), "exec"), vars(module)) + return module + + +gdn_model = _load_relax_functions( + "model.py", ["_patch_gdn_for_dynamic_cp", "_resolve_gdn_cp", "_assert_gdn_full_recompute"] +) # --------------------------------------------------------------------------- # Step 1: CLI flag # --------------------------------------------------------------------------- def _parse_megatron_args(monkeypatch, *argv): - pytest.importorskip("sglang.srt.server_args") - from relax.utils.arguments import get_slime_extra_args_provider + pytest.importorskip("triton", reason="the full Megatron training CLI imports Triton kernels") + training_arguments = pytest.importorskip("megatron.training.arguments") monkeypatch.setattr("sys.argv", ["test-linear-cp-mode", *argv]) - return megatron_arguments._megatron_parse_args( - extra_args_provider=get_slime_extra_args_provider(), + return training_arguments.parse_args( ignore_unknown_args=False, ) -def test_linear_cp_mode_flag_defaults_to_headwise(monkeypatch): +def test_linear_cp_mode_flag_defaults_to_chunkwise(monkeypatch): args = _parse_megatron_args(monkeypatch) - assert args.linear_cp_mode == "headwise" + assert args.linear_cp_mode == "chunkwise" @pytest.mark.parametrize("mode", ["headwise", "chunkwise", "all_gather"]) @@ -61,7 +83,8 @@ def test_linear_cp_mode_flag_accepts_all_concrete_modes(monkeypatch, mode): # --------------------------------------------------------------------------- def _args(**overrides): base = dict( - linear_cp_mode="headwise", + linear_cp_mode="chunkwise", + experimental_attention_variant="gated_delta_net", allgather_cp=False, deterministic_mode=False, dynamic_context_parallel=False, @@ -73,13 +96,13 @@ def _args(**overrides): @pytest.mark.parametrize("bad", ["auto", "allgather"]) def test_validate_linear_cp_mode_rejects_unsupported_value(bad): - with pytest.raises(ValueError, match="does not support 'auto'|must be one of"): + with pytest.raises(ValueError, match="must be one of"): _validate_linear_cp_mode(_args(linear_cp_mode=bad)) def test_validate_linear_cp_mode_rejects_chunkwise_with_allgather_cp(): with pytest.raises(ValueError, match="allgather-cp"): - _validate_linear_cp_mode(_args(linear_cp_mode="chunkwise", allgather_cp=True)) + _validate_linear_cp_mode(_args(linear_cp_mode="chunkwise", allgather_cp=True, context_parallel_size=2)) @pytest.mark.parametrize("cp_kwargs", [{"context_parallel_size": 2}, {"dynamic_context_parallel": True}]) @@ -227,3 +250,53 @@ def test_all_gather_fallback_rejects_deterministic_mode(): psp = _fake_packed_seq_params(cp_group=_FakeGroup(4), local_cp_size=4) with pytest.raises(AssertionError, match="deterministic mode"): GatedDeltaNet.forward(m, "hs", None, None, psp) + + +@pytest.mark.parametrize("mode", ["headwise", "chunkwise", "all_gather"]) +def test_gdn_modes_reject_contiguous_attention_packing(mode): + with pytest.raises(ValueError, match="allgather-cp"): + _validate_linear_cp_mode(_args(linear_cp_mode=mode, context_parallel_size=2, allgather_cp=True)) + + +def test_default_chunkwise_does_not_restrict_non_gdn_models(): + _validate_linear_cp_mode( + _args( + experimental_attention_variant="dsa", context_parallel_size=4, allgather_cp=True, deterministic_mode=True + ) + ) + + +def test_bridge_validation_uses_provider_attention_variant(): + args = _args(experimental_attention_variant=None, allgather_cp=True, context_parallel_size=2) + _validate_linear_cp_mode(args) + provider = SimpleNamespace( + experimental_attention_variant="gated_delta_net", linear_cp_mode="chunkwise", context_parallel_size=2 + ) + with pytest.raises(ValueError, match="allgather-cp"): + _validate_linear_cp_mode(args, provider) + + +@pytest.mark.parametrize("group,size", [(None, 2), (_FakeGroup(2), None), (_FakeGroup(2), 4)]) +def test_dispatcher_rejects_inconsistent_dynamic_metadata(group, size): + _install_dispatcher_with_spies() + m = _fake_gdn_module(linear_cp_mode="chunkwise", static_cp_size=8) + with pytest.raises(ValueError, match="PackedSeqParams"): + GatedDeltaNet.forward( + m, "hs", None, packed_seq_params=_fake_packed_seq_params(cp_group=group, local_cp_size=size) + ) + + +def test_dispatcher_respects_explicit_process_group_collection(): + calls = _install_dispatcher_with_spies() + m = _fake_gdn_module(linear_cp_mode="all_gather", static_cp_size=8) + GatedDeltaNet.forward(m, "hs", None, pg_collection=SimpleNamespace(cp=_FakeGroup(1))) + assert calls == {"orig": 1} + + +def test_all_gather_still_requires_full_recompute(monkeypatch): + monkeypatch.setattr(gdn_model, "get_args", lambda: _args(recompute_granularity="selective"), raising=False) + monkeypatch.setattr(gdn_model._assert_gdn_full_recompute, "_checked", False, raising=False) + with torch.enable_grad(), pytest.raises(ValueError, match="all_gather.*whole-layer"): + gdn_model._assert_gdn_full_recompute() + with torch.no_grad(): + gdn_model._assert_gdn_full_recompute() From 7e7e73c4018ec2c5013af15acca4ed8da57310de Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Tue, 22 Sep 2026 16:57:48 +1000 Subject: [PATCH 06/11] test(megatron): isolate VPP import dependencies MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit # ✅ Tests - Preserve reviewer PR #348 production code, including the GDN validator in arguments.py and the original model-provider import. - Complete fake Megatron optimizer, parser, and tokenizer dependencies for isolated VPP tests without replacing the real GDN validator. - Restore arguments/model-provider module caches and package attributes after each test, including when those modules were initially absent. - Restore the reference stage2 test import strategy. - Pass 159 CPU regressions, 17 standalone VPP tests, and 46 combined VPP/FP16 tests with zero skips, plus real production imports and all pre-commit hooks. References: redai-studio/Relax#273, redai-studio/Relax#348. --- relax/backends/megatron/arguments.py | 42 +++++++++++++++++- relax/backends/megatron/gdn_cp_config.py | 43 ------------------- relax/backends/megatron/model_provider.py | 2 +- .../megatron/test_gdn_cp_mode_stage2.py | 4 +- .../megatron/test_model_provider_vpp.py | 23 +++++++++- 5 files changed, 65 insertions(+), 49 deletions(-) delete mode 100644 relax/backends/megatron/gdn_cp_config.py diff --git a/relax/backends/megatron/arguments.py b/relax/backends/megatron/arguments.py index d84a6b38a..ae17f8dc3 100644 --- a/relax/backends/megatron/arguments.py +++ b/relax/backends/megatron/arguments.py @@ -2,6 +2,7 @@ import ast import math +from argparse import Namespace from typing import Optional from megatron.core.optimizer import OptimizerConfig @@ -20,8 +21,6 @@ from relax.utils.logging_utils import get_logger from relax.utils.model_source import ModelSource -from .gdn_cp_config import _validate_linear_cp_mode - __all__ = ["validate_args", "megatron_parse_args", "set_default_megatron_args"] @@ -130,6 +129,45 @@ def _validate_dynamic_context_parallel(args): args.max_seqlen_per_dp_cp_rank = args.max_tokens_per_gpu +def _validate_linear_cp_mode(args: Namespace, config: Optional[object] = None) -> None: + """Validate the mode name, then GDN-specific flags once the model is known. + + Geometry-dependent rejections (e.g. explicit `headwise` on heads not + divisible by `tp*max_cp`) can only be checked once the real GDN head counts + are known, which happens in MCore's `TransformerConfig.__post_init__` gate + -- not here. Bridge calls this again with the provider's actual config. + """ + model_config = config if config is not None else args + mode = getattr(model_config, "linear_cp_mode", "chunkwise") + allowed_modes = {"headwise", "chunkwise", "all_gather"} + if mode not in allowed_modes: + raise ValueError( + f"--linear-cp-mode must be one of {sorted(allowed_modes)!r}; got {mode!r}. Resolve 'auto' before construction." + ) + + # Bridge determines the attention variant from the HF checkpoint. The default + # linear_cp_mode on an ordinary-attention model does not make it a GDN model. + if getattr(model_config, "experimental_attention_variant", None) != "gated_delta_net": + return + + cp_may_exceed_one = ( + getattr(args, "dynamic_context_parallel", False) or getattr(model_config, "context_parallel_size", 1) > 1 + ) + if cp_may_exceed_one and getattr(args, "allgather_cp", False): + raise ValueError( + "GDN CP requires zig-zag THD packing in every linear_cp_mode; --allgather-cp uses " + "contiguous per-rank packing and is incompatible with GDN CP>1. --allgather-cp is a " + "data/attention packing flag, separate from --linear-cp-mode=all_gather." + ) + + if mode == "chunkwise" and cp_may_exceed_one and getattr(model_config, "deterministic_mode", False): + raise ValueError( + "--linear-cp-mode=chunkwise does not support --deterministic-mode while CP>1 may occur: " + "the deterministic torch reference path only accepts cp_context=None. " + "Packed GDN inputs also do not support deterministic mode in the other CP modes." + ) + + def validate_args(args): """Run megatron's own validate_args plus slime-specific megatron validations.""" diff --git a/relax/backends/megatron/gdn_cp_config.py b/relax/backends/megatron/gdn_cp_config.py deleted file mode 100644 index a74e42637..000000000 --- a/relax/backends/megatron/gdn_cp_config.py +++ /dev/null @@ -1,43 +0,0 @@ -# Copyright (c) 2026 Relax Authors. All Rights Reserved. - -from argparse import Namespace -from typing import Optional - - -def _validate_linear_cp_mode(args: Namespace, config: Optional[object] = None) -> None: - """Validate the mode name, then GDN-specific flags once the model is known. - - Geometry-dependent rejections (e.g. explicit `headwise` on heads not - divisible by `tp*max_cp`) can only be checked once the real GDN head counts - are known, which happens in MCore's `TransformerConfig.__post_init__` gate - -- not here. Bridge calls this again with the provider's actual config. - """ - model_config = config if config is not None else args - mode = getattr(model_config, "linear_cp_mode", "chunkwise") - allowed_modes = {"headwise", "chunkwise", "all_gather"} - if mode not in allowed_modes: - raise ValueError( - f"--linear-cp-mode must be one of {sorted(allowed_modes)!r}; got {mode!r}. Resolve 'auto' before construction." - ) - - # Bridge determines the attention variant from the HF checkpoint. The default - # linear_cp_mode on an ordinary-attention model does not make it a GDN model. - if getattr(model_config, "experimental_attention_variant", None) != "gated_delta_net": - return - - cp_may_exceed_one = ( - getattr(args, "dynamic_context_parallel", False) or getattr(model_config, "context_parallel_size", 1) > 1 - ) - if cp_may_exceed_one and getattr(args, "allgather_cp", False): - raise ValueError( - "GDN CP requires zig-zag THD packing in every linear_cp_mode; --allgather-cp uses " - "contiguous per-rank packing and is incompatible with GDN CP>1. --allgather-cp is a " - "data/attention packing flag, separate from --linear-cp-mode=all_gather." - ) - - if mode == "chunkwise" and cp_may_exceed_one and getattr(model_config, "deterministic_mode", False): - raise ValueError( - "--linear-cp-mode=chunkwise does not support --deterministic-mode while CP>1 may occur: " - "the deterministic torch reference path only accepts cp_context=None. " - "Packed GDN inputs also do not support deterministic mode in the other CP modes." - ) diff --git a/relax/backends/megatron/model_provider.py b/relax/backends/megatron/model_provider.py index 257191770..f55eb3385 100644 --- a/relax/backends/megatron/model_provider.py +++ b/relax/backends/megatron/model_provider.py @@ -48,8 +48,8 @@ install_sequence_classification_head_in_provider, ) +from .arguments import _validate_linear_cp_mode from .conditional_branch_sync import install_conditional_branch_sync -from .gdn_cp_config import _validate_linear_cp_mode logger = get_logger(__name__) diff --git a/tests/backends/megatron/test_gdn_cp_mode_stage2.py b/tests/backends/megatron/test_gdn_cp_mode_stage2.py index 83f776b8f..6c4679f15 100644 --- a/tests/backends/megatron/test_gdn_cp_mode_stage2.py +++ b/tests/backends/megatron/test_gdn_cp_mode_stage2.py @@ -29,8 +29,6 @@ from megatron.core.packed_seq_params import PackedSeqParams # noqa: E402 from megatron.core.ssm.gated_delta_net import GatedDeltaNet # noqa: E402 -from relax.backends.megatron.gdn_cp_config import _validate_linear_cp_mode # noqa: E402 - def _load_relax_functions(filename, names): """Execute the real functions without importing the Ray/optimizer stack.""" @@ -49,9 +47,11 @@ def _load_relax_functions(filename, names): return module +megatron_arguments = _load_relax_functions("arguments.py", ["_validate_linear_cp_mode"]) gdn_model = _load_relax_functions( "model.py", ["_patch_gdn_for_dynamic_cp", "_resolve_gdn_cp", "_assert_gdn_full_recompute"] ) +_validate_linear_cp_mode = megatron_arguments._validate_linear_cp_mode # --------------------------------------------------------------------------- diff --git a/tests/backends/megatron/test_model_provider_vpp.py b/tests/backends/megatron/test_model_provider_vpp.py index ae9f7e545..e928aafd1 100644 --- a/tests/backends/megatron/test_model_provider_vpp.py +++ b/tests/backends/megatron/test_model_provider_vpp.py @@ -49,6 +49,7 @@ def _install_fake_megatron(monkeypatch, provider=None): megatron = types.ModuleType("megatron") core = types.ModuleType("megatron.core") + optimizer = types.ModuleType("megatron.core.optimizer") mpu = types.ModuleType("megatron.core.mpu") tensor_parallel = types.ModuleType("megatron.core.tensor_parallel") models = types.ModuleType("megatron.core.models") @@ -59,6 +60,8 @@ def _install_fake_megatron(monkeypatch, provider=None): transformer_config = types.ModuleType("megatron.core.transformer.transformer_config") training = types.ModuleType("megatron.training") arguments = types.ModuleType("megatron.training.arguments") + tokenizer_package = types.ModuleType("megatron.training.tokenizer") + tokenizer = types.ModuleType("megatron.training.tokenizer.tokenizer") bridge = types.ModuleType("megatron.bridge") misc = types.ModuleType("relax.utils.misc") @@ -76,6 +79,9 @@ def from_hf_pretrained(cls, *args, **kwargs): def to_megatron_provider(self, load_weights=False): return provider + def _unexpected_training_setup(*args, **kwargs): + raise AssertionError("VPP provider tests must not invoke Megatron training setup") + mpu.get_virtual_pipeline_model_parallel_world_size = lambda: 2 mpu.get_virtual_pipeline_model_parallel_rank = lambda: 1 mpu.get_context_parallel_world_size = lambda: 1 @@ -83,6 +89,7 @@ def to_megatron_provider(self, load_weights=False): mpu.get_tensor_model_parallel_rank = lambda: 0 core.mpu = mpu core.tensor_parallel = tensor_parallel + optimizer.OptimizerConfig = SimpleNamespace gpt.GPTModel = _FakeGPTModel gpt_layer_specs.get_gpt_decoder_block_spec = lambda *args, **kwargs: object() gpt_layer_specs.get_gpt_layer_local_spec = lambda *args, **kwargs: object() @@ -90,12 +97,16 @@ def to_megatron_provider(self, load_weights=False): spec_utils.import_module = lambda path: object() transformer_config.TransformerConfig = _FakeTransformerConfig arguments.core_transformer_config_from_args = lambda args: _FakeTransformerConfig() + arguments.parse_args = _unexpected_training_setup + arguments.validate_args = _unexpected_training_setup + tokenizer._vocab_size_with_padding = _unexpected_training_setup bridge.AutoBridge = _FakeAutoBridge misc.load_function = lambda path: None modules = { "megatron": megatron, "megatron.core": core, + "megatron.core.optimizer": optimizer, "megatron.core.mpu": mpu, "megatron.core.tensor_parallel": tensor_parallel, "megatron.core.models": models, @@ -106,6 +117,8 @@ def to_megatron_provider(self, load_weights=False): "megatron.core.transformer.transformer_config": transformer_config, "megatron.training": training, "megatron.training.arguments": arguments, + "megatron.training.tokenizer": tokenizer_package, + "megatron.training.tokenizer.tokenizer": tokenizer, "megatron.bridge": bridge, "relax.utils.misc": misc, } @@ -117,7 +130,15 @@ def to_megatron_provider(self, load_weights=False): def _load_model_provider(monkeypatch, provider=None): provider = _install_fake_megatron(monkeypatch, provider=provider) - sys.modules.pop("relax.backends.megatron.model_provider", None) + package = importlib.import_module("relax.backends.megatron") + # Both real modules bind fake dependencies. Restore their cache entries and + # package attributes after the test, including when initially absent. + for name in ("arguments", "model_provider"): + fullname = f"{package.__name__}.{name}" + monkeypatch.setitem(sys.modules, fullname, None) + monkeypatch.delitem(sys.modules, fullname) + monkeypatch.setattr(package, name, None, raising=False) + monkeypatch.delattr(package, name) module = importlib.import_module("relax.backends.megatron.model_provider") monkeypatch.setattr(module.dist, "is_initialized", lambda: True) monkeypatch.setattr(module.dist, "get_rank", lambda: 1) From 6d9ef05221c9cbeeb1086c6a6621b75fb35328e2 Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Tue, 22 Sep 2026 17:19:53 +1000 Subject: [PATCH 07/11] test(megatron): keep only fake import additions MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit # ✅ Tests - Restore the model-provider test loader to reviewer PR #348 unchanged. - Limit the reference diff to 13 lines completing fake Megatron import dependencies in _install_fake_megatron. - Preserve all production code and existing test assertions from #348. - Pass 159 CPU regressions, 17 isolated VPP tests, 46 combined VPP/FP16 tests, and all pre-commit hooks with zero skipped tests. References: redai-studio/Relax#273, redai-studio/Relax#348. --- tests/backends/megatron/test_model_provider_vpp.py | 10 +--------- 1 file changed, 1 insertion(+), 9 deletions(-) diff --git a/tests/backends/megatron/test_model_provider_vpp.py b/tests/backends/megatron/test_model_provider_vpp.py index e928aafd1..de81fcf6d 100644 --- a/tests/backends/megatron/test_model_provider_vpp.py +++ b/tests/backends/megatron/test_model_provider_vpp.py @@ -130,15 +130,7 @@ def _unexpected_training_setup(*args, **kwargs): def _load_model_provider(monkeypatch, provider=None): provider = _install_fake_megatron(monkeypatch, provider=provider) - package = importlib.import_module("relax.backends.megatron") - # Both real modules bind fake dependencies. Restore their cache entries and - # package attributes after the test, including when initially absent. - for name in ("arguments", "model_provider"): - fullname = f"{package.__name__}.{name}" - monkeypatch.setitem(sys.modules, fullname, None) - monkeypatch.delitem(sys.modules, fullname) - monkeypatch.setattr(package, name, None, raising=False) - monkeypatch.delattr(package, name) + sys.modules.pop("relax.backends.megatron.model_provider", None) module = importlib.import_module("relax.backends.megatron.model_provider") monkeypatch.setattr(module.dist, "is_initialized", lambda: True) monkeypatch.setattr(module.dist, "get_rank", lambda: 1) From 6c190fedffbc4f027eacee03b028e6b47d9b643b Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Tue, 22 Sep 2026 17:29:35 +1000 Subject: [PATCH 08/11] test(megatron): match existing fake callables MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit # ✅ Tests - Replace the special raising helper with the existing lambda placeholder style for fake parser, validator, and tokenizer imports. - Keep the reference diff limited to ten fake-dependency additions. - Pass 17 isolated VPP tests, 46 VPP/FP16 tests, and all pre-commit hooks. --- tests/backends/megatron/test_model_provider_vpp.py | 9 +++------ 1 file changed, 3 insertions(+), 6 deletions(-) diff --git a/tests/backends/megatron/test_model_provider_vpp.py b/tests/backends/megatron/test_model_provider_vpp.py index de81fcf6d..edf89b7cd 100644 --- a/tests/backends/megatron/test_model_provider_vpp.py +++ b/tests/backends/megatron/test_model_provider_vpp.py @@ -79,9 +79,6 @@ def from_hf_pretrained(cls, *args, **kwargs): def to_megatron_provider(self, load_weights=False): return provider - def _unexpected_training_setup(*args, **kwargs): - raise AssertionError("VPP provider tests must not invoke Megatron training setup") - mpu.get_virtual_pipeline_model_parallel_world_size = lambda: 2 mpu.get_virtual_pipeline_model_parallel_rank = lambda: 1 mpu.get_context_parallel_world_size = lambda: 1 @@ -97,9 +94,9 @@ def _unexpected_training_setup(*args, **kwargs): spec_utils.import_module = lambda path: object() transformer_config.TransformerConfig = _FakeTransformerConfig arguments.core_transformer_config_from_args = lambda args: _FakeTransformerConfig() - arguments.parse_args = _unexpected_training_setup - arguments.validate_args = _unexpected_training_setup - tokenizer._vocab_size_with_padding = _unexpected_training_setup + arguments.parse_args = lambda *args, **kwargs: None + arguments.validate_args = lambda *args, **kwargs: None + tokenizer._vocab_size_with_padding = lambda *args, **kwargs: None bridge.AutoBridge = _FakeAutoBridge misc.load_function = lambda path: None From 41d5296f3982d4a57c85d2e65f3cc7e3e0c18acf Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Tue, 22 Sep 2026 18:23:29 +1000 Subject: [PATCH 09/11] fix(megatron): tolerate legacy GDN log config MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit # 🐛 Bug Fix - Read linear_cp_mode with getattr when logging GDN CP initialization so legacy NPU model configurations without this field do not raise. - Preserve existing forward dispatch and report None for an absent field. # ✅ Tests - Check the actual log expression with a missing field and all three existing modes on CPU; existing mode output is unchanged. - Verify the repository Qwen3.5 NPU CP4 provider uses its own Bridge GDN class. - Run all pre-commit hooks. References: redai-studio/Relax#273. --- relax/backends/megatron/model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/relax/backends/megatron/model.py b/relax/backends/megatron/model.py index ee4bb6441..a3a5f4ae8 100644 --- a/relax/backends/megatron/model.py +++ b/relax/backends/megatron/model.py @@ -419,7 +419,7 @@ def setup_model_and_optimizer( or torch.distributed.get_rank(group=torch.distributed.group.WORLD) == 0 ): logger.info( - f"[GDN CP] role={role} linear_cp_mode={model_config.linear_cp_mode} " + f"[GDN CP] role={role} linear_cp_mode={getattr(model_config, 'linear_cp_mode', None)} " f"TP={model_config.tensor_model_parallel_size} max_CP={model_config.context_parallel_size} " f"key_heads={model_config.linear_num_key_heads} value_heads={model_config.linear_num_value_heads}" ) From 785c56533911e9358c095f444e52a5680fec3ef9 Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Tue, 22 Sep 2026 22:06:46 +1000 Subject: [PATCH 10/11] test(megatron): focus GDN CP regression tests MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit # ✅ Tests - Keep one real NCCL THD/SBHD layout round-trip test and remove the ten expensive full GPU validation cases and their dedicated helpers. - Merge CPU partition, route, and cache coverage into test_gdn_cp_layout. - Group configuration and dispatch coverage in test_gdn_cp_mode, removing historical stage names while preserving all 126 core CPU cases. - Pass 159 related CPU regressions, verify the GPU file collects one test, and run all pre-commit hooks. No new GPU execution is claimed. --- .../megatron/test_gdn_chunkwise_cp_gpu.py | 1008 ----------------- .../megatron/test_gdn_chunkwise_cp_layout.py | 256 ----- ...wise_cp_route.py => test_gdn_cp_layout.py} | 150 ++- .../megatron/test_gdn_cp_layout_gpu.py | 114 ++ ..._cp_mode_stage2.py => test_gdn_cp_mode.py} | 123 +- 5 files changed, 353 insertions(+), 1298 deletions(-) delete mode 100644 tests/backends/megatron/test_gdn_chunkwise_cp_gpu.py delete mode 100644 tests/backends/megatron/test_gdn_chunkwise_cp_layout.py rename tests/backends/megatron/{test_gdn_chunkwise_cp_route.py => test_gdn_cp_layout.py} (69%) create mode 100644 tests/backends/megatron/test_gdn_cp_layout_gpu.py rename tests/backends/megatron/{test_gdn_cp_mode_stage2.py => test_gdn_cp_mode.py} (72%) diff --git a/tests/backends/megatron/test_gdn_chunkwise_cp_gpu.py b/tests/backends/megatron/test_gdn_chunkwise_cp_gpu.py deleted file mode 100644 index ea1a934d4..000000000 --- a/tests/backends/megatron/test_gdn_chunkwise_cp_gpu.py +++ /dev/null @@ -1,1008 +0,0 @@ -# Copyright (c) 2026 Relax Authors. All Rights Reserved. -"""Real-kernel / real-collective tests for the GDN chunkwise-CP backport. - -RFC Task 32 phase-1 acceptance items 3 and 4: - -* a minimal CP=2 chunkwise case must drive the *actual* FLA kernels and match a - CP=1 reference in forward and backward within tolerance; -* the GDN parameter keys and shard dimensions in ``state_dict`` / - ``sharded_state_dict`` must not move, i.e. GDN weights stay TP-only and - checkpoints are unaffected by the CP mode. - -Three layers are covered: - -1. the FLA kernels directly (``causal_conv1d`` / ``chunk_gated_delta_rule`` with - a ``cp_context``); -2. the whole MCore ``GatedDeltaNet`` module in fp32 -- the *algebraic* check. In - fp32 the only difference between CP=1 and CP=2 is float summation order, so - the tolerances can be tight enough to catch a genuinely wrong permutation or - a dropped boundary term; -3. the same module in bf16 -- the *production* check, at the dtype training - actually uses, where the achievable agreement is bounded by the storage - format rather than by the algorithm. - -A headwise CP=2 run is included at every layer as a control: if headwise and -chunkwise both drift the same way, the cause is shared plumbing, not the new -code. - -Most tests need 2 visible GPUs; the TP2/CP2 matrix test needs 4: - pytest tests/backends/megatron/test_gdn_chunkwise_cp_gpu.py -""" - -from __future__ import annotations - -import os - -import pytest -import torch -import torch.multiprocessing as mp - - -WORLD_SIZE = 2 - -# --- tolerances ------------------------------------------------------------ -# bf16 module output / input gradient: RFC section 5, "MCore GDN, CP>1 vs CP=1". -ATOL_BF16 = 2e-3 -RTOL_BF16 = 1e-2 -MIN_COSINE = 0.9999 -# FLA kernels: RFC section 5, normalised RMS error thresholds. -CONV_RMS_RATIO = 1e-3 -GDN_RMS_RATIO = 2e-3 -# fp32 module run: both CP algorithms must land far below any bf16 threshold. -# 1e-3 is 2x below the bf16 element-wise atol of the RFC gate, i.e. "fp32 must be -# comfortably better than the dtype we actually ship". -RMS_RATIO_FP32 = 1e-3 -# fp32 kernel-level: measured ~1e-7, so 1e-5 is a real gate, not a rubber stamp. -KERNEL_RMS_RATIO_FP32 = 1e-5 -# chunkwise vs headwise. headwise is the already-shipped CP algorithm, so whatever -# CP-vs-no-CP disagreement it shows is the floor this environment imposes (reduced -# precision inside the Triton dots, changed summation order, bf16 storage) rather -# than anything about the algorithm. Requiring chunkwise to be no worse than that -# floor is the assertion that actually means something; a fixed atol on a bf16 -# token-sum gradient mostly measures rounding luck. -CHUNKWISE_VS_HEADWISE_RMS_FACTOR = 4.0 -# ...with a floor, so a headwise value that happens to land at or near zero on a given -# run cannot turn into an impossible budget. 1e-6 is still ~1000x tighter than the fp32 -# absolute gate, so the comparison keeps its teeth. -RMS_FLOOR_FP32 = 1e-6 -# ...applied only where headwise is not bit-exact. headwise hands each rank the -# whole sequence and 1/cp of the heads, so none of its reductions are -# repartitioned and it can land exactly on the CP=1 result; "4x zero" would be a -# budget no correct implementation could meet. Those tensors are covered by the -# absolute gates instead. - -pytestmark = pytest.mark.skipif( - not torch.cuda.is_available() or torch.cuda.device_count() < WORLD_SIZE, - reason=f"requires {WORLD_SIZE} CUDA devices", -) - - -def _has_backport() -> bool: - try: - import megatron.core.context_parallel_layout # noqa: F401 - from fla.ops.cp import build_cp_context # noqa: F401 - except ImportError: - return False - return True - - -needs_backport = pytest.mark.skipif(not _has_backport(), reason="requires patched Megatron-LM + FLA >= 0.4.2") - - -# --------------------------------------------------------------------------- -# comparison helpers (run inside the workers) -# --------------------------------------------------------------------------- -def _prep(name, got, want): - got32 = got.detach().float().flatten() - want32 = want.detach().float().flatten() - assert got32.shape == want32.shape, f"{name}: shape {got32.shape} vs {want32.shape}" - assert torch.isfinite(got32).all(), f"{name}: non-finite values in candidate" - assert torch.isfinite(want32).all(), f"{name}: non-finite values in reference" - return got32, want32 - - -def _stats(got32, want32): - diff = (got32 - want32).abs() - rms = (diff.square().mean().sqrt() / (want32.square().mean().sqrt() + 1e-12)).item() - cos = torch.nn.functional.cosine_similarity(got32, want32, dim=0).item() - return diff, rms, cos - - -def _report_elementwise(name, got, want, atol, rtol): - """Per-token tensors: every element within atol + rtol * |ref|.""" - got32, want32 = _prep(name, got, want) - diff, rms, cos = _stats(got32, want32) - worst = (diff - (atol + rtol * want32.abs())).max().item() - assert worst <= 0, ( - f"{name}: max |diff| {diff.max().item():.3e} exceeds atol({atol:.0e})+rtol({rtol:.0e})*|ref| " - f"by {worst:.3e} (rms {rms:.3e}, cosine {cos:.8f})" - ) - assert cos >= MIN_COSINE, f"{name}: cosine {cos:.8f} < {MIN_COSINE}" - - -def _report_rms(name, got, want, ratio): - """Whole-tensor normalised RMS error -- the metric FLA's own CP tests - use.""" - got32, want32 = _prep(name, got, want) - diff, rms, cos = _stats(got32, want32) - assert rms < ratio, ( - f"{name}: normalised RMS error {rms:.3e} >= {ratio:.1e} (max |diff| {diff.max().item():.3e}, cosine {cos:.8f})" - ) - assert cos >= MIN_COSINE, f"{name}: cosine {cos:.8f} < {MIN_COSINE} (rms {rms:.3e})" - - -def _zigzag_shard(full: torch.Tensor, cu, cp_size: int, cp_rank: int) -> torch.Tensor: - from relax.backends.megatron.cp_utils import gdn_cp_slice - - return gdn_cp_slice(full, cu, cp_size, cp_rank) - - -def _init_dist(rank, world_size): - os.environ.setdefault("MASTER_ADDR", "127.0.0.1") - os.environ.setdefault("MASTER_PORT", "29531") - torch.cuda.set_device(rank) - import torch.distributed as dist - - dist.init_process_group("nccl", rank=rank, world_size=world_size) - - -# --------------------------------------------------------------------------- -# worker: FLA kernel level -# --------------------------------------------------------------------------- -def _worker_fla_kernels(rank, world_size, dtype_name, _unused): - seq_lens = [256, 128] - dtype = {"fp32": torch.float32, "bf16": torch.bfloat16}[dtype_name] - _init_dist(rank, world_size) - import torch.distributed as dist - from fla.modules.convolution import causal_conv1d - from fla.modules.l2norm import l2norm - from fla.ops.cp import build_cp_context - from fla.ops.gated_delta_rule import chunk_gated_delta_rule - - device = torch.device("cuda", rank) - cp_group = dist.new_group(list(range(world_size))) - - H, DK, DV, W = 2, 64, 64, 4 - total = sum(seq_lens) - cu = torch.tensor([0] + torch.tensor(seq_lens).cumsum(0).tolist(), device=device, dtype=torch.int32) - part = total // world_size - lo, hi = rank * part, (rank + 1) * part - - gen = torch.Generator(device="cpu").manual_seed(1234) - - def _mk(*shape, dtype=dtype): - return torch.randn(*shape, generator=gen, dtype=torch.float32).to(device=device, dtype=dtype) - - conv_ratio = KERNEL_RMS_RATIO_FP32 if dtype is torch.float32 else CONV_RMS_RATIO - gdn_ratio = KERNEL_RMS_RATIO_FP32 if dtype is torch.float32 else GDN_RMS_RATIO - tag = f"[rank{rank}/{dtype_name}]" - - # ---- causal conv ---- - # weight/bias are leaves here on purpose: their gradients are sums over every - # token, which is exactly the quantity chunkwise CP repartitions. Checking only - # dx would leave that untested. - x = _mk(1, total, H * DK) - w0 = _mk(H * DK, W) - b0 = _mk(H * DK) - conv_grad = _mk(1, total, H * DK) - - x_ref = x.clone().requires_grad_(True) - w_ref = w0.clone().requires_grad_(True) - b_ref = b0.clone().requires_grad_(True) - out_ref, _ = causal_conv1d(x=x_ref, weight=w_ref, bias=b_ref, activation="silu", cu_seqlens=cu) - (out_ref.float() * conv_grad.float()).sum().backward() - - x_cp = x[:, lo:hi].clone().requires_grad_(True) - w_cp = w0.clone().requires_grad_(True) - b_cp = b0.clone().requires_grad_(True) - ctx = build_cp_context(cu_seqlens=cu, group=cp_group, conv1d_kernel_size=W) - out_cp, _ = causal_conv1d(x=x_cp, weight=w_cp, bias=b_cp, activation="silu", cu_seqlens=cu, cp_context=ctx) - _report_rms(f"{tag} conv fwd", out_cp, out_ref[:, lo:hi], conv_ratio) - (out_cp.float() * conv_grad[:, lo:hi].float()).sum().backward() - _report_rms(f"{tag} conv dx", x_cp.grad, x_ref.grad[:, lo:hi], conv_ratio) - - for pname, cp_leaf, ref_leaf in (("dweight", w_cp, w_ref), ("dbias", b_cp, b_ref)): - summed = cp_leaf.grad.detach().float().clone() - dist.all_reduce(summed, group=cp_group) - if dtype is torch.float32: - _report_rms(f"{tag} conv {pname}", summed, ref_leaf.grad, conv_ratio) - else: - # These are 384-token sums landing in bf16. The fp32 parametrisation of - # this very test pins the algebra at ~1e-7; in bf16 the achievable - # agreement is set by the storage format, so assert direction and report - # the size rather than pretend a sub-ULP threshold is meaningful. - got32, want32 = _prep(f"{tag} conv {pname}", summed, ref_leaf.grad) - _, rms, cos = _stats(got32, want32) - assert cos >= MIN_COSINE, f"{tag} conv {pname}: cosine {cos:.8f} (rms {rms:.3e})" - if rank == 0: - print(f" {tag} conv {pname}: rms {rms:.3e} cosine {cos:.10f}") - - # ---- gated delta rule ---- - # Inputs must look like what GatedDeltaNet actually feeds the kernel: - # * q/k are L2-normalised (the module sets use_qk_l2norm=True). Un-normalised - # q/k make the recurrent state diverge over hundreds of steps and the - # reference itself goes to NaN -- that would test nothing. - # * g is a log-domain decay built as -A.exp() * softplus(...), hence <= 0. - q = l2norm(_mk(1, total, H, DK).contiguous()) - k = l2norm(_mk(1, total, H, DK).contiguous()) - v = _mk(1, total, H, DV) - g0 = -_mk(1, total, H, dtype=torch.float32).abs() * 0.1 - beta0 = _mk(1, total, H, dtype=torch.float32).sigmoid() - - leaves_ref = [t.detach().clone().requires_grad_(True) for t in (q, k, v)] - g_ref = g0.detach().clone().requires_grad_(True) - beta_ref = beta0.detach().clone().requires_grad_(True) - o_ref, _ = chunk_gated_delta_rule( - *leaves_ref, - g=g_ref, - beta=beta_ref, - initial_state=None, - output_final_state=False, - use_qk_l2norm_in_kernel=False, - cu_seqlens=cu, - ) - o_grad = _mk(1, total, H, DV) - (o_ref.float() * o_grad.float()).sum().backward() - - leaves_cp = [t.detach()[:, lo:hi].clone().requires_grad_(True) for t in (q, k, v)] - g_cp = g0.detach()[:, lo:hi].clone().requires_grad_(True) - beta_cp = beta0.detach()[:, lo:hi].clone().requires_grad_(True) - ctx2 = build_cp_context(cu_seqlens=cu, group=cp_group, conv1d_kernel_size=W) - o_cp, _ = chunk_gated_delta_rule( - *leaves_cp, - g=g_cp, - beta=beta_cp, - initial_state=None, - output_final_state=False, - use_qk_l2norm_in_kernel=False, - cu_seqlens=cu, - cp_context=ctx2, - ) - _report_rms(f"{tag} gdr fwd", o_cp, o_ref[:, lo:hi], gdn_ratio) - (o_cp.float() * o_grad[:, lo:hi].float()).sum().backward() - for name, a, b in zip("qkv", leaves_cp, leaves_ref): - _report_rms(f"{tag} gdr d{name}", a.grad, b.grad[:, lo:hi], gdn_ratio) - _report_rms(f"{tag} gdr dg", g_cp.grad, g_ref.grad[:, lo:hi], gdn_ratio) - _report_rms(f"{tag} gdr dbeta", beta_cp.grad, beta_ref.grad[:, lo:hi], gdn_ratio) - - dist.barrier() - dist.destroy_process_group() - - -# --------------------------------------------------------------------------- -# worker: full MCore GatedDeltaNet -# --------------------------------------------------------------------------- -def _build_gdn( - cp_size, - linear_cp_mode, - dtype, - num_key_heads=4, - num_value_heads=8, - tp_size=1, - deterministic_mode=False, -): - import torch.nn.functional as F - from megatron.core import parallel_state - from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( - get_experimental_attention_variant_module_spec, - ) - from megatron.core.process_groups_config import ProcessGroupCollection - from megatron.core.ssm.gated_delta_net import GatedDeltaNet - from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed - from megatron.core.transformer.transformer_config import TransformerConfig - - model_parallel_cuda_manual_seed(123) - config = TransformerConfig( - hidden_size=512, - num_layers=1, - num_attention_heads=8, - num_query_groups=2, - normalization="RMSNorm", - use_cpu_initialization=True, - layernorm_zero_centered_gamma=True, - activation_func=F.silu, - bf16=dtype is torch.bfloat16, - tensor_model_parallel_size=tp_size, - context_parallel_size=cp_size, - deterministic_mode=deterministic_mode, - experimental_attention_variant="gated_delta_net", - linear_attention_freq=[1], - linear_conv_kernel_dim=4, - linear_key_head_dim=64, - linear_value_head_dim=64, - linear_num_key_heads=num_key_heads, - linear_num_value_heads=num_value_heads, - linear_cp_mode=linear_cp_mode, - transformer_impl="transformer_engine", - ) - pg_collection = ProcessGroupCollection( - tp=parallel_state.get_tensor_model_parallel_group(), - cp=parallel_state.get_context_parallel_group(), - ) - gdn = GatedDeltaNet( - config, - submodules=get_experimental_attention_variant_module_spec(config=config).submodules, - layer_number=1, - bias=False, - conv_bias=False, - conv_init=1.0, - use_qk_l2norm=True, - A_init_range=(1, 16), - pg_collection=pg_collection, - ) - return gdn.cuda().to(dtype), config - - -def _run_gdn_once(gdn, hidden, psp, grad_out, *, recompute=False, **forward_kwargs): - """One forward+backward; returns (out, d_hidden, {param: grad}).""" - gdn.zero_grad(set_to_none=True) - h = hidden.clone().requires_grad_(True) - if recompute: - from torch.utils.checkpoint import checkpoint - - out = checkpoint( - lambda x: gdn(x, None, packed_seq_params=psp, **forward_kwargs)[0], - h, - use_reentrant=False, - ) - else: - out, _ = gdn(h, None, packed_seq_params=psp, **forward_kwargs) - (out.float() * grad_out).sum().backward() - grads = {n: p.grad.detach().float().clone() for n, p in gdn.named_parameters()} - return out.detach().clone(), h.grad.detach().clone(), grads - - -def _worker_gdn_module(rank, world_size, dtype_name, _unused): - """CP=1 reference vs CP=N, for BOTH CP algorithms, in one process. - - Running headwise and chunkwise side by side is the point: it turns "is - chunkwise close enough to CP=1" (which needs an absolute threshold, and in - bf16 lands on the noise floor of the storage format) into "is chunkwise as - close to CP=1 as the algorithm we already ship" -- a comparison with no - free parameters to tune. - """ - dtype = {"fp32": torch.float32, "bf16": torch.bfloat16}[dtype_name] - - _init_dist(rank, world_size) - import torch.distributed as dist - from megatron.core import parallel_state - from megatron.core.packed_seq_params import PackedSeqParams - - parallel_state.initialize_model_parallel( - tensor_model_parallel_size=1, - pipeline_model_parallel_size=1, - context_parallel_size=world_size, - ) - device = torch.device("cuda", rank) - # The CP algorithm is static config now, so each mode needs its own module. They - # share weights, so the comparison is still like-for-like. - modules = {} - gdn, config = _build_gdn(world_size, "headwise", dtype) - for p in gdn.parameters(): - # Same weights on every rank so the CP=1 reference is rank-independent. - dist.broadcast(p.data, src=0) - modules["headwise"] = gdn - modules["chunkwise"], _ = _build_gdn(world_size, "chunkwise", dtype) - modules["chunkwise"].load_state_dict(gdn.state_dict()) - - cp_group = parallel_state.get_context_parallel_group() - # A per-rank size-1 group gives us the CP=1 reference *inside* the same - # process, driving the very same weights through the very same forward. - solo = [dist.new_group([r]) for r in range(world_size)][rank] - - seq_lens = [256, 128] - total = sum(seq_lens) - cu = torch.tensor([0, seq_lens[0], total], device=device, dtype=torch.int32) - - gen = torch.Generator(device="cpu").manual_seed(7) - hidden_full = torch.randn(total, 1, config.hidden_size, generator=gen, dtype=torch.float32).to( - device=device, dtype=dtype - ) - grad_seed = torch.randn(total, 1, config.hidden_size, generator=gen, dtype=torch.float32).to(device) - - def _psp(group, local_cp_size): - return PackedSeqParams( - qkv_format="thd", - cu_seqlens_q=cu, - cu_seqlens_kv=cu, - cu_seqlens_q_padded=cu, - cu_seqlens_kv_padded=cu, - max_seqlen_q=max(seq_lens), - max_seqlen_kv=max(seq_lens), - cp_group=group, - local_cp_size=local_cp_size, - ) - - # CP=1 reference. A size-1 group short-circuits before the mode is read, so either - # module gives the same reference; use the headwise one. - out_ref, in_grad_ref, param_grads_ref = _run_gdn_once(gdn, hidden_full, _psp(solo, 1), grad_seed) - ref = { - "out": _zigzag_shard(out_ref, cu, world_size, rank), - "d_hidden": _zigzag_shard(in_grad_ref, cu, world_size, rank), - } - ref.update({f"grad {n}": g for n, g in param_grads_ref.items()}) - - shard = _zigzag_shard(hidden_full, cu, world_size, rank) - grad_shard = _zigzag_shard(grad_seed, cu, world_size, rank) - - metrics = {} - for mode in ("headwise", "chunkwise"): - out_cp, in_grad_cp, param_grads_cp = _run_gdn_once( - modules[mode], shard, _psp(cp_group, world_size), grad_shard - ) - got = {"out": out_cp, "d_hidden": in_grad_cp} - # Each CP rank holds a partial parameter gradient; the total is the CP sum. - for name, g in param_grads_cp.items(): - summed = g.clone() - dist.all_reduce(summed, group=cp_group) - got[f"grad {name}"] = summed - metrics[mode] = {} - for key, value in got.items(): - got32, want32 = _prep(f"[rank{rank}][{mode}/{dtype_name}] {key}", value, ref[key]) - diff, rms, cos = _stats(got32, want32) - metrics[mode][key] = (rms, cos, diff.max().item()) - assert cos >= MIN_COSINE, ( - f"[rank{rank}][{mode}/{dtype_name}] {key}: cosine {cos:.8f} < {MIN_COSINE} (rms {rms:.3e})" - ) - - # RFC section 5 absolute gate, on the per-token tensors, in the dtype the - # RFC specifies it for. Applied to headwise too, so a drift in the shared - # plumbing cannot hide behind the comparative check below. - if dtype is torch.bfloat16: - for key in ("out", "d_hidden"): - _report_elementwise( - f"[rank{rank}][{mode}/{dtype_name}] {key}", got[key], ref[key], ATOL_BF16, RTOL_BF16 - ) - # Parameter gradients are token-sum reductions stored in bf16. Bound them - # by the RFC's own relative tolerance for this comparison row (rtol=1e-2) - # applied to the whole tensor, plus the RFC's cosine floor. What actually - # pins the algebra is the fp32 parametrisation of this same test. - for key, (rms, cos, _) in metrics[mode].items(): - if key in ("out", "d_hidden"): - continue - assert rms < RTOL_BF16, ( - f"[rank{rank}][{mode}/bf16] {key}: relative RMS {rms:.3e} >= {RTOL_BF16:.0e} (cosine {cos:.8f})" - ) - else: - for key, (rms, _, _) in metrics[mode].items(): - assert rms < RMS_RATIO_FP32, f"[rank{rank}][{mode}/fp32] {key}: rms {rms:.3e} >= {RMS_RATIO_FP32:.0e}" - - # The comparative assertion -- fp32 only, on purpose. - # - # Its premise is "headwise's disagreement with CP=1 is the floor this environment - # imposes". That holds only while both algorithms perform the *same* reductions. - # They do not: headwise hands each rank the whole sequence and 1/cp of the heads, so - # a gradient like conv1d.weight / dt_bias / A_log (a sum over every token) is summed - # in one go exactly as at CP=1 and can come out bit-exact. Chunkwise splits the - # tokens, so that same sum really is partitioned and re-added. In fp32 the mantissa - # absorbs it and the two are directly comparable (observed 1.00x-1.10x). In bf16 the - # repartitioned sum sits on the format's ULP floor while headwise sits near zero, so - # their *ratio* measures the dtype, not the algorithm -- bf16 is covered by the - # absolute gates above instead. - worst = [] - for key, (rms_c, cos_c, max_c) in metrics["chunkwise"].items(): - rms_h = metrics["headwise"][key][0] - worst.append((rms_c / max(rms_h, 1e-12), key, rms_c, rms_h)) - if dtype is not torch.float32: - continue - assert rms_c <= CHUNKWISE_VS_HEADWISE_RMS_FACTOR * max(rms_h, RMS_FLOOR_FP32), ( - f"[rank{rank}][{dtype_name}] {key}: chunkwise rms {rms_c:.3e} exceeds " - f"{CHUNKWISE_VS_HEADWISE_RMS_FACTOR}x the shipped headwise rms {rms_h:.3e} " - f"(cosine {cos_c:.8f}, max |diff| {max_c:.3e})" - ) - worst.sort(reverse=True) - if rank == 0: - print(f"\n[{dtype_name}] chunkwise vs headwise, worst 6 by rms ratio:") - for ratio, key, rms_c, rms_h in worst[:6]: - shown = f"{ratio:6.2f}x" if rms_h > 0 else " n/a" - print(f" {key:38s} {shown} chunkwise {rms_c:.3e} headwise {rms_h:.3e}") - - dist.barrier() - parallel_state.destroy_model_parallel() - dist.destroy_process_group() - - -def _worker_deterministic_reference(rank, world_size, _spec, _unused): - """The torch-native deterministic rule must accept cp_context=None and - preserve headwise CP correctness.""" - _init_dist(rank, world_size) - import torch.distributed as dist - from megatron.core import parallel_state - from megatron.core.process_groups_config import ProcessGroupCollection - - parallel_state.initialize_model_parallel( - tensor_model_parallel_size=1, - pipeline_model_parallel_size=1, - context_parallel_size=world_size, - ) - device = torch.device("cuda", rank) - cp_group = parallel_state.get_context_parallel_group() - tp_group = parallel_state.get_tensor_model_parallel_group() - solo = [dist.new_group([r]) for r in range(world_size)][rank] - - gdn, config = _build_gdn( - world_size, - "headwise", - torch.float32, - deterministic_mode=True, - ) - assert gdn.gated_delta_rule.__name__ == "torch_chunk_gated_delta_rule" - for p in gdn.parameters(): - dist.broadcast(p.data, src=0, group=cp_group) - - total = 64 - cu = torch.tensor([0, total], device=device, dtype=torch.int32) - gen = torch.Generator(device="cpu").manual_seed(17) - hidden_full = torch.randn(total, 1, config.hidden_size, generator=gen).to(device) - grad_full = torch.randn(total, 1, config.hidden_size, generator=gen).to(device) - - solo_pg = ProcessGroupCollection(tp=tp_group, cp=solo) - out_ref, in_grad_ref, param_grads_ref = _run_gdn_once( - gdn, - hidden_full, - None, - grad_full, - pg_collection=solo_pg, - ) - - hidden_shard = _zigzag_shard(hidden_full, cu, world_size, rank) - grad_shard = _zigzag_shard(grad_full, cu, world_size, rank) - out_cp, in_grad_cp, param_grads_cp = _run_gdn_once(gdn, hidden_shard, None, grad_shard) - - _report_rms( - f"[rank{rank}] deterministic out", - out_cp, - _zigzag_shard(out_ref, cu, world_size, rank), - RMS_RATIO_FP32, - ) - _report_rms( - f"[rank{rank}] deterministic d_hidden", - in_grad_cp, - _zigzag_shard(in_grad_ref, cu, world_size, rank), - RMS_RATIO_FP32, - ) - for name, grad in param_grads_cp.items(): - summed = grad.clone() - dist.all_reduce(summed, group=cp_group) - _report_rms( - f"[rank{rank}] deterministic grad {name}", - summed, - param_grads_ref[name], - RMS_RATIO_FP32, - ) - - dist.barrier() - parallel_state.destroy_model_parallel() - dist.destroy_process_group() - - -def _worker_recompute_parity(rank, world_size, recompute_kind, _unused): - """External activation checkpointing must replay chunkwise collectives - without changing outputs or gradients.""" - _init_dist(rank, world_size) - import torch.distributed as dist - from megatron.core import parallel_state - from megatron.core.packed_seq_params import PackedSeqParams - - parallel_state.initialize_model_parallel( - tensor_model_parallel_size=1, - pipeline_model_parallel_size=1, - context_parallel_size=world_size, - ) - device = torch.device("cuda", rank) - cp_group = parallel_state.get_context_parallel_group() - - gdn, config = _build_gdn(world_size, "chunkwise", torch.float32) - for p in gdn.parameters(): - dist.broadcast(p.data, src=0, group=cp_group) - - seq_lens = [128, 64] - total = sum(seq_lens) - cu = torch.tensor([0, seq_lens[0], total], device=device, dtype=torch.int32) - psp = PackedSeqParams( - qkv_format="thd", - cu_seqlens_q=cu, - cu_seqlens_kv=cu, - cu_seqlens_q_padded=cu, - cu_seqlens_kv_padded=cu, - max_seqlen_q=max(seq_lens), - max_seqlen_kv=max(seq_lens), - cp_group=cp_group, - local_cp_size=world_size, - ) - - gen = torch.Generator(device="cpu").manual_seed(29) - hidden_full = torch.randn(total, 1, config.hidden_size, generator=gen).to(device) - grad_full = torch.randn(total, 1, config.hidden_size, generator=gen).to(device) - hidden = _zigzag_shard(hidden_full, cu, world_size, rank) - grad = _zigzag_shard(grad_full, cu, world_size, rank) - - out_eager, in_grad_eager, param_grads_eager = _run_gdn_once(gdn, hidden, psp, grad) - if recompute_kind == "selective": - # Use the upgraded MCore's own GDN checkpoint path, not the external wrapper. - gdn.recompute_gdn = True - out_recompute, in_grad_recompute, param_grads_recompute = _run_gdn_once( - gdn, - hidden, - psp, - grad, - recompute=recompute_kind == "full", - ) - - _report_rms(f"[rank{rank}] recompute out", out_recompute, out_eager, KERNEL_RMS_RATIO_FP32) - _report_rms( - f"[rank{rank}] recompute d_hidden", - in_grad_recompute, - in_grad_eager, - KERNEL_RMS_RATIO_FP32, - ) - assert set(param_grads_recompute) == set(param_grads_eager) - for name in param_grads_eager: - _report_rms( - f"[rank{rank}] recompute grad {name}", - param_grads_recompute[name], - param_grads_eager[name], - KERNEL_RMS_RATIO_FP32, - ) - - dist.barrier() - parallel_state.destroy_model_parallel() - dist.destroy_process_group() - - -def _worker_tp2_cp2(rank, world_size, _spec, _unused): - """Exercise TP head sharding and CP routing together.""" - assert world_size == 4 - _init_dist(rank, world_size) - import torch.distributed as dist - from megatron.core import parallel_state - from megatron.core.packed_seq_params import PackedSeqParams - from megatron.core.process_groups_config import ProcessGroupCollection - - parallel_state.initialize_model_parallel( - tensor_model_parallel_size=2, - pipeline_model_parallel_size=1, - context_parallel_size=2, - ) - device = torch.device("cuda", rank) - cp_group = parallel_state.get_context_parallel_group() - cp_rank = cp_group.rank() - tp_group = parallel_state.get_tensor_model_parallel_group() - cp_source = dist.get_process_group_ranks(cp_group)[0] - solo = [dist.new_group([r]) for r in range(world_size)][rank] - - modules = {} - headwise, config = _build_gdn(2, "headwise", torch.float32, tp_size=2) - for p in headwise.parameters(): - dist.broadcast(p.data, src=cp_source, group=cp_group) - modules["headwise"] = headwise - modules["chunkwise"], _ = _build_gdn(2, "chunkwise", torch.float32, tp_size=2) - modules["chunkwise"].load_state_dict(headwise.state_dict()) - sharded_signatures = {} - for mode, module in modules.items(): - sharded_signatures[mode] = { - key: ( - tuple(getattr(value, "global_shape", ())), - tuple(getattr(value, "local_shape", ())), - getattr(value, "axis_fragmentations", None), - ) - for key, value in sorted(module.sharded_state_dict(prefix="mixer.").items()) - } - assert sharded_signatures["headwise"] == sharded_signatures["chunkwise"] - - seq_lens = [128, 64] - total = sum(seq_lens) - cu = torch.tensor([0, seq_lens[0], total], device=device, dtype=torch.int32) - gen = torch.Generator(device="cpu").manual_seed(41) - hidden_full = torch.randn(total, 1, config.hidden_size, generator=gen).to(device) - grad_full = torch.randn(total, 1, config.hidden_size, generator=gen).to(device) - - def _psp(group, cp_size): - return PackedSeqParams( - qkv_format="thd", - cu_seqlens_q=cu, - cu_seqlens_kv=cu, - cu_seqlens_q_padded=cu, - cu_seqlens_kv_padded=cu, - max_seqlen_q=max(seq_lens), - max_seqlen_kv=max(seq_lens), - cp_group=group, - local_cp_size=cp_size, - ) - - solo_pg = ProcessGroupCollection(tp=tp_group, cp=solo) - out_ref, in_grad_ref, param_grads_ref = _run_gdn_once( - headwise, - hidden_full, - _psp(solo, 1), - grad_full, - pg_collection=solo_pg, - ) - hidden = _zigzag_shard(hidden_full, cu, 2, cp_rank) - grad = _zigzag_shard(grad_full, cu, 2, cp_rank) - out_want = _zigzag_shard(out_ref, cu, 2, cp_rank) - in_grad_want = _zigzag_shard(in_grad_ref, cu, 2, cp_rank) - - for mode, module in modules.items(): - out, in_grad, param_grads = _run_gdn_once(module, hidden, _psp(cp_group, 2), grad) - _report_rms(f"[rank{rank}][{mode}] TP2/CP2 out", out, out_want, RMS_RATIO_FP32) - _report_rms( - f"[rank{rank}][{mode}] TP2/CP2 d_hidden", - in_grad, - in_grad_want, - RMS_RATIO_FP32, - ) - assert set(param_grads) == set(param_grads_ref) - for name, param_grad in param_grads.items(): - summed = param_grad.clone() - dist.all_reduce(summed, group=cp_group) - _report_rms( - f"[rank{rank}][{mode}] TP2/CP2 grad {name}", - summed, - param_grads_ref[name], - RMS_RATIO_FP32, - ) - - dist.barrier() - parallel_state.destroy_model_parallel() - dist.destroy_process_group() - - -def _worker_layout_round_trip(rank, world_size, _spec, _unused): - """zigzag -> contiguous -> zigzag over a real CP group must be token-exact. - - This is the collective-level version of RFC 5.1: it drives the actual - ``all_to_all`` in ``context_parallel_layout``, for packed THD (several - unequal-length samples) and for SBHD. - """ - _init_dist(rank, world_size) - import torch.distributed as dist - from megatron.core.context_parallel_layout import ( - contiguous_to_zigzag_chunks, - get_thd_context_parallel_rank_indices, - zigzag_to_contiguous_chunks, - ) - - device = torch.device("cuda", rank) - cp_group = dist.new_group(list(range(world_size))) - - # --- packed THD, three samples of different lengths --- - lengths = [2 * world_size * f for f in (5, 1, 3)] - cu = torch.tensor([0] + torch.tensor(lengths).cumsum(0).tolist(), device=device, dtype=torch.int32) - total = int(cu[-1]) - # Row t is (t, t+1e6, t+2e6): a permuted token is impossible to miss. - full = ( - torch.arange(total, dtype=torch.float64, device=device).unsqueeze(1) - + torch.arange(3, dtype=torch.float64, device=device).unsqueeze(0) * 1e6 - ) - - zig_idx = get_thd_context_parallel_rank_indices(cu, world_size, rank, "zigzag") - con_idx = get_thd_context_parallel_rank_indices(cu, world_size, rank, "contiguous") - local_zig = full[zig_idx] - - got_con = zigzag_to_contiguous_chunks(local_zig, cp_group, seq_dim=0, cu_seqlens=cu) - assert torch.equal(got_con, full[con_idx]), f"rank {rank}: THD zigzag->contiguous is wrong" - got_zig = contiguous_to_zigzag_chunks(got_con, cp_group=cp_group, seq_dim=0, cu_seqlens=cu) - assert torch.equal(got_zig, local_zig), f"rank {rank}: THD round trip is not identity" - - # --- SBHD (chunk-level swap, no cu_seqlens) --- - seq_local = 2 * world_size * 4 - sbhd = torch.arange(seq_local * 2 * 3, dtype=torch.float64, device=device).reshape(seq_local, 2, 3) + rank * 1e9 - swapped = zigzag_to_contiguous_chunks(sbhd, cp_group, seq_dim=0) - back = contiguous_to_zigzag_chunks(swapped, cp_group=cp_group, seq_dim=0) - assert torch.equal(back, sbhd), f"rank {rank}: SBHD round trip is not identity" - - dist.barrier() - dist.destroy_process_group() - - -def _worker_illegal_modes_fail_fast(rank, world_size, _spec, _unused): - """Illegal / unresolved GDN CP modes must raise, not silently pick an - algorithm.""" - _init_dist(rank, world_size) - import torch.distributed as dist - from megatron.core import parallel_state - from megatron.core.packed_seq_params import PackedSeqParams - - parallel_state.initialize_model_parallel( - tensor_model_parallel_size=1, - pipeline_model_parallel_size=1, - context_parallel_size=world_size, - ) - device = torch.device("cuda", rank) - cp_group = parallel_state.get_context_parallel_group() - solo = [dist.new_group([r]) for r in range(world_size)][rank] - - seq_lens = [256, 128] - total = sum(seq_lens) - cu = torch.tensor([0, seq_lens[0], total], device=device, dtype=torch.int32) - - def _psp(group, local_cp_size): - return PackedSeqParams( - qkv_format="thd", - cu_seqlens_q=cu, - cu_seqlens_kv=cu, - cu_seqlens_q_padded=cu, - cu_seqlens_kv_padded=cu, - max_seqlen_q=max(seq_lens), - max_seqlen_kv=max(seq_lens), - cp_group=group, - local_cp_size=local_cp_size, - ) - - gdn, config = _build_gdn(world_size, "all_gather", torch.bfloat16) - hidden = torch.randn(total // world_size, 1, config.hidden_size, device=device, dtype=torch.bfloat16) - - # 1. all_gather is a Relax-wrapper mode: reaching MCore's forward with cp>1 means the - # wrapper is missing, and that must be loud. - with pytest.raises(RuntimeError, match="requires the Relax"): - gdn(hidden, None, packed_seq_params=_psp(cp_group, world_size)) - - # 2. ...but a CP=1 micro-batch is legal under any declared mode: it needs no CP - # communication at all, so the mode is never consulted. - full = torch.randn(total, 1, config.hidden_size, device=device, dtype=torch.bfloat16) - solo_psp = _psp(solo, 1) - gdn(full, None, packed_seq_params=solo_psp) - - # 3. The final PackedSeqParams must describe the same runtime CP geometry - # through both fields. - gdn_hw, _ = _build_gdn(world_size, "headwise", torch.bfloat16) - with pytest.raises(ValueError, match=r"does not match.*cp_group.size"): - gdn_hw(hidden, None, packed_seq_params=_psp(cp_group, world_size + 1)) - - # 4. deterministic mode has no CP-context scan. - gdn_cw, _ = _build_gdn(world_size, "chunkwise", torch.bfloat16) - gdn_cw.config.deterministic_mode = True - try: - with pytest.raises((ValueError, AssertionError)): - gdn_cw(hidden, None, packed_seq_params=_psp(cp_group, world_size)) - finally: - gdn_cw.config.deterministic_mode = False - - # 5. inference is not supported. - class _Ctx: - def is_static_batching(self): - return True - - with pytest.raises(NotImplementedError): - gdn_cw(hidden, None, inference_context=_Ctx(), packed_seq_params=_psp(cp_group, world_size)) - - dist.barrier() - parallel_state.destroy_model_parallel() - dist.destroy_process_group() - - -def _worker_state_dict_invariance(rank, world_size, _spec, _unused): - """GDN weights must stay TP-only: same keys and shard dims in every CP - mode.""" - _init_dist(rank, world_size) - import io - - import torch.distributed as dist - from megatron.core import parallel_state - - parallel_state.initialize_model_parallel( - tensor_model_parallel_size=1, - pipeline_model_parallel_size=1, - context_parallel_size=world_size, - ) - device = torch.device("cuda", rank) - - signatures = {} - for cp_size, mode in ((1, "headwise"), (world_size, "headwise"), (world_size, "chunkwise")): - gdn, _ = _build_gdn(cp_size, mode, torch.bfloat16) - sd = gdn.state_dict() - sharded = gdn.sharded_state_dict(prefix="mixer.") - signatures[(cp_size, mode)] = ( - {k: tuple(v.shape) for k, v in sd.items() if torch.is_tensor(v)}, - { - k: ( - tuple(getattr(v, "global_shape", ())), - tuple(getattr(v, "local_shape", ())), - getattr(v, "axis_fragmentations", None), - ) - for k, v in sorted(sharded.items()) - }, - ) - - # Exercise real serialization and loading, not just key comparison. - checkpoint = io.BytesIO() - torch.save(sd, checkpoint) - checkpoint.seek(0) - loaded = torch.load(checkpoint, map_location=device, weights_only=True) - restored, _ = _build_gdn(cp_size, mode, torch.bfloat16) - restored.load_state_dict(loaded, strict=True) - for name, tensor in sd.items(): - if torch.is_tensor(tensor): - assert torch.equal(restored.state_dict()[name], tensor), ( - f"state_dict round trip changed {name} for {(cp_size, mode)}" - ) - del restored - del gdn - - baseline = signatures[(1, "headwise")] - assert baseline[0], "state_dict is empty; the invariance check would be vacuous" - assert baseline[1], "sharded_state_dict is empty; the invariance check would be vacuous" - for key, sig in signatures.items(): - assert sig[0] == baseline[0], f"state_dict shapes changed for {key}" - assert set(sig[1]) == set(baseline[1]), f"sharded_state_dict keys changed for {key}" - for k in baseline[1]: - assert sig[1][k] == baseline[1][k], f"sharded shard dims changed for {key} at {k}" - - parallel_state.destroy_model_parallel() - dist.destroy_process_group() - - -# --------------------------------------------------------------------------- -# pytest entry points -# --------------------------------------------------------------------------- -def _spawn(fn, spec, port, extra=None): - _spawn_world(fn, WORLD_SIZE, spec, port, extra=extra) - - -def _spawn_world(fn, world_size, spec, port, extra=None): - os.environ["MASTER_PORT"] = str(port) - mp.spawn( - fn, - args=(world_size, extra if extra is not None else spec, None), - nprocs=world_size, - join=True, - ) - - -@needs_backport -@pytest.mark.parametrize("dtype_name,port", [("fp32", 29540), ("bf16", 29541)]) -def test_fla_cp_kernels_match_single_rank(dtype_name, port): - """FLA causal_conv1d / chunk_gated_delta_rule under cp_context vs no CP.""" - _spawn(_worker_fla_kernels, dtype_name, port) - - -@needs_backport -@pytest.mark.parametrize("dtype_name,port", [("fp32", 29542), ("bf16", 29543)]) -def test_gdn_cp_matches_cp1(dtype_name, port): - """RFC 3.4 item 3: CP=2 vs CP=1, forward and backward, both CP - algorithms.""" - _spawn(_worker_gdn_module, dtype_name, port) - - -@needs_backport -def test_deterministic_headwise_cp_matches_cp1(): - """The torch-native rule accepts cp_context=None and stays correct under - headwise CP.""" - _spawn(_worker_deterministic_reference, "n/a", 29547) - - -@needs_backport -@pytest.mark.parametrize("recompute_kind", ["full", "selective"]) -def test_chunkwise_recompute_matches_eager(recompute_kind): - """Replaying chunkwise forward during backward preserves every tested - gradient.""" - _spawn(_worker_recompute_parity, recompute_kind, 29548) - - -@needs_backport -@pytest.mark.skipif(torch.cuda.device_count() < 4, reason="requires 4 CUDA devices for TP2/CP2") -def test_gdn_tp2_cp2_matches_cp1(): - """TP2/CP2 exercises TP head shards and both CP algorithms together.""" - _spawn_world(_worker_tp2_cp2, 4, "n/a", 29549) - - -@needs_backport -def test_layout_round_trip_over_real_cp_group(): - """RFC 5.1 at the collective level: the layout swap is a pure - permutation.""" - _spawn(_worker_layout_round_trip, "n/a", 29544) - - -@needs_backport -def test_illegal_gdn_cp_modes_fail_fast(): - """RFC emphasis 1: illegal combinations must fail fast, never pick - silently.""" - _spawn(_worker_illegal_modes_fail_fast, "n/a", 29545) - - -@needs_backport -def test_gdn_state_dict_invariant_across_cp_modes(): - """RFC 3.4 item 4: checkpoint keys and shard dims do not depend on the CP - mode.""" - _spawn(_worker_state_dict_invariance, "n/a", 29546) diff --git a/tests/backends/megatron/test_gdn_chunkwise_cp_layout.py b/tests/backends/megatron/test_gdn_chunkwise_cp_layout.py deleted file mode 100644 index 05ad47c03..000000000 --- a/tests/backends/megatron/test_gdn_chunkwise_cp_layout.py +++ /dev/null @@ -1,256 +0,0 @@ -# Copyright (c) 2026 Relax Authors. All Rights Reserved. -"""Unit tests for the GDN chunkwise-CP layout backport (Task 32, phase 1). - -Covers the pure-tensor half of the backported MCore capability: the two THD CP -partitions (``zigzag`` / ``contiguous``), their agreement with Relax's existing -zigzag sharding, dynamic group resolution, and the construction-time -``linear_cp_mode`` gate. - -Everything here runs on CPU with no process group. It validates partition -definitions only; the actual all-to-all round trip is exercised with NCCL in -``test_gdn_chunkwise_cp_gpu.py``. - -The real-kernel / real-collective half lives in -``test_gdn_chunkwise_cp_gpu.py``. -""" - -from __future__ import annotations - -import inspect - -import pytest -import torch - - -cpl = pytest.importorskip("megatron.core.context_parallel_layout", reason="requires the patched Megatron-LM") - -from megatron.core.packed_seq_params import PackedSeqParams, resolve_cp_group # noqa: E402 - -from relax.backends.megatron.cp_utils import gdn_cp_slice, slice_with_cp # noqa: E402 - - -def _cu(lengths: list[int]) -> torch.Tensor: - cu = [0] - for n in lengths: - cu.append(cu[-1] + n) - return torch.tensor(cu, dtype=torch.int64) - - -def _tagged_tokens(total: int, width: int = 3) -> torch.Tensor: - """[total, width] where row t is (t, t+1e6, t+2e6): token identity is - unambiguous.""" - base = torch.arange(total, dtype=torch.float64).unsqueeze(1) - return base + torch.arange(width, dtype=torch.float64).unsqueeze(0) * 1e6 - - -# --------------------------------------------------------------------------- -# Partition definitions -# --------------------------------------------------------------------------- -@pytest.mark.parametrize("cp_size", [1, 2, 4, 8]) -@pytest.mark.parametrize("layout", ["zigzag", "contiguous"]) -@pytest.mark.parametrize("lengths_factor", [[1], [1, 2, 3], [3, 1, 1, 2]]) -def test_thd_rank_indices_partition_all_tokens_exactly_once(cp_size, layout, lengths_factor): - lengths = [2 * cp_size * f for f in lengths_factor] - cu = _cu(lengths) - owned = torch.cat([cpl.get_thd_context_parallel_rank_indices(cu, cp_size, r, layout) for r in range(cp_size)]) - assert owned.numel() == int(cu[-1]) - assert torch.equal(torch.sort(owned).values, torch.arange(int(cu[-1]))) - - -@pytest.mark.parametrize("cp_size", [2, 4, 8]) -def test_zigzag_rank_indices_match_relax_data_sharding(cp_size): - """MCore's zigzag partition must be token-for-token what Relax's data path - produces. - - If these ever disagree, chunkwise CP would silently permute tokens relative - to the all-gather fallback and the attention layers. - """ - lengths = [2 * cp_size * f for f in (1, 3, 2)] - cu = _cu(lengths) - full = _tagged_tokens(int(cu[-1])).reshape(-1, 1, 3) # [s, b=1, C] - - for rank in range(cp_size): - mcore_idx = cpl.get_thd_context_parallel_rank_indices(cu, cp_size, rank, "zigzag") - mcore_shard = full[mcore_idx] - - # Relax data.py: per-sample slice_with_cp then concat. - relax_shard = torch.cat( - [ - slice_with_cp( - full[cu[i] : cu[i + 1]], - pad_value=0.0, - qkv_format="thd", - dynamic_cp_size=cp_size, - dynamic_cp_rank=rank, - ) - for i in range(len(lengths)) - ], - dim=0, - ) - assert torch.equal(mcore_shard, relax_shard) - - # Relax model.py (all-gather fallback) re-slices with gdn_cp_slice. - assert torch.equal(mcore_shard, gdn_cp_slice(full, cu, cp_size, rank)) - - -@pytest.mark.parametrize("cp_size", [2, 4, 8]) -@pytest.mark.parametrize("lengths_factor", [[1], [1, 2, 3], [3, 1, 1, 2]]) -def test_both_layouts_are_permutations_of_each_other(cp_size, lengths_factor): - """The two partitions must describe the same token set with the same per- - rank size. - - That is the precondition for the all-to-all between them to be a pure - permutation -- no token invented, dropped, or duplicated. The real collective - round trip is asserted in ``test_gdn_chunkwise_cp_gpu.py``. - """ - lengths = [2 * cp_size * f for f in lengths_factor] - cu = _cu(lengths) - total = int(cu[-1]) - zig_by_rank = [] - con_by_rank = [] - for rank in range(cp_size): - zig = cpl.get_thd_context_parallel_rank_indices(cu, cp_size, rank, "zigzag") - con = cpl.get_thd_context_parallel_rank_indices(cu, cp_size, rank, "contiguous") - zig_by_rank.append(zig) - con_by_rank.append(con) - assert zig.numel() == con.numel() == total // cp_size - # contiguous is exactly this rank's span of the flattened buffer - assert torch.equal(con, torch.arange(rank * (total // cp_size), (rank + 1) * (total // cp_size))) - - # Across the whole CP group, both layouts are permutations of exactly the - # same global token rows. - assert torch.equal( - torch.cat(zig_by_rank).sort().values, - torch.cat(con_by_rank).sort().values, - ) - - -@pytest.mark.parametrize("cp_size", [2, 4]) -def test_rank_indices_reject_lengths_not_divisible_by_two_cp(cp_size): - bad = _cu([2 * cp_size, 2 * cp_size + 1]) - with pytest.raises(ValueError, match="divisible by"): - cpl.get_thd_context_parallel_rank_indices(bad, cp_size, 0, "zigzag") - - -def test_gdn_rejects_packed_lengths_not_divisible_by_cp(): - from megatron.core.ssm.gated_delta_net import GatedDeltaNet - - cu = _cu([8, 6]) - with pytest.raises(ValueError, match="divisible by cp_size=4"): - GatedDeltaNet._resolve_cu_seqlens(None, None, cu, int(cu[-1]), "cu_seqlens_q", cp_size=4) - - -def test_rank_indices_reject_unknown_layout(): - with pytest.raises(ValueError, match="Unsupported context-parallel layout"): - cpl.get_thd_context_parallel_rank_indices(_cu([16, 16]), 2, 0, "contiguous_ish") - - -@pytest.mark.parametrize("layout", ["zigzag", "contiguous"]) -def test_rank_indices_ignore_duplicate_boundaries(layout): - compact = torch.tensor([0, 16, 40], dtype=torch.int64) - padded = torch.tensor([0, 16, 40, 40, 40], dtype=torch.int64) - for rank in range(2): - assert torch.equal( - cpl.get_thd_context_parallel_rank_indices(compact, 2, rank, layout), - cpl.get_thd_context_parallel_rank_indices(padded, 2, rank, layout), - ) - - -@pytest.mark.parametrize("layout", ["zigzag", "contiguous"]) -def test_rank_indices_reject_decreasing_boundaries(layout): - with pytest.raises(ValueError, match="nondecreasing"): - cpl.get_thd_context_parallel_rank_indices(torch.tensor([0, 16, 8]), 2, 0, layout) - - -# --------------------------------------------------------------------------- -# Dynamic CP group resolution -# --------------------------------------------------------------------------- -def test_resolve_cp_group_prefers_packed_seq_params(): - static = object() - dynamic = object() - assert resolve_cp_group(static, None) is static - assert resolve_cp_group(static, PackedSeqParams(qkv_format="thd")) is static - assert resolve_cp_group(static, PackedSeqParams(qkv_format="thd", cp_group=dynamic)) is dynamic - - -# --------------------------------------------------------------------------- -# Construction-time capability gate -# --------------------------------------------------------------------------- -def _gdn_config(**overrides): - import torch.nn.functional as F - from megatron.core.transformer.transformer_config import TransformerConfig - - kwargs = dict( - hidden_size=2048, - num_layers=1, - num_attention_heads=16, - num_query_groups=2, - normalization="RMSNorm", - use_cpu_initialization=True, - activation_func=F.silu, - bf16=True, - experimental_attention_variant="gated_delta_net", - linear_attention_freq=[1], - linear_conv_kernel_dim=4, - linear_key_head_dim=128, - linear_value_head_dim=128, - linear_num_key_heads=16, - linear_num_value_heads=32, - ) - kwargs.update(overrides) - return TransformerConfig(**kwargs) - - -def test_config_default_mode_is_chunkwise(): - """Upgrading the image must not silently reroute an existing recipe.""" - assert _gdn_config().linear_cp_mode == "chunkwise" - - -def test_headwise_config_requires_heads_divisible_by_tp_times_cp(): - # 16 key heads, tp=2, cp=4 -> 16 % 8 == 0: fine. - _gdn_config(tensor_model_parallel_size=2, context_parallel_size=4, linear_cp_mode="headwise") - # tp=2, cp=16 -> 16 % 32 != 0: the geometry headwise cannot express. - with pytest.raises(AssertionError, match="linear_num_key_heads"): - _gdn_config(tensor_model_parallel_size=2, context_parallel_size=16, linear_cp_mode="headwise") - - -def test_chunkwise_config_only_requires_heads_divisible_by_tp(): - """This is what replaces Relax's temporary head-count rewrite hack.""" - cfg = _gdn_config(tensor_model_parallel_size=2, context_parallel_size=16, linear_cp_mode="chunkwise") - assert cfg.linear_num_key_heads == 16 and cfg.linear_num_value_heads == 32 - # ... but TP divisibility is still enforced: GDN weights stay TP-sharded. - with pytest.raises(AssertionError, match="linear_num_key_heads"): - _gdn_config( - tensor_model_parallel_size=8, - context_parallel_size=2, - linear_cp_mode="chunkwise", - num_query_groups=8, - linear_num_key_heads=4, - linear_num_value_heads=8, - ) - - -def test_all_gather_config_uses_the_tp_only_head_rule(): - """`--linear-cp-mode=all_gather` must be constructible on a non-divisible - geometry. - - Relax's all-gather fallback keeps GDN weights TP-only, so declaring it - should relax the head check exactly as chunkwise does. - """ - cfg = _gdn_config(tensor_model_parallel_size=2, context_parallel_size=16, linear_cp_mode="all_gather") - assert cfg.linear_num_key_heads == 16 and cfg.linear_num_value_heads == 32 - - -def test_config_rejects_unresolved_and_unknown_linear_cp_mode(): - """MCore only accepts the three concrete execution modes.""" - for bad in ("auto", "allgather", "chunk", ""): - with pytest.raises(AssertionError, match="linear_cp_mode"): - _gdn_config(context_parallel_size=2, linear_cp_mode=bad) - with pytest.raises(AssertionError, match="linear_cp_mode"): - _gdn_config(context_parallel_size=4, tensor_model_parallel_size=2, linear_cp_mode=bad) - - -def test_gdn_forward_has_no_per_call_mode_override(): - from megatron.core.ssm.gated_delta_net import GatedDeltaNet - - assert "linear_cp_mode" not in inspect.signature(GatedDeltaNet.forward).parameters diff --git a/tests/backends/megatron/test_gdn_chunkwise_cp_route.py b/tests/backends/megatron/test_gdn_cp_layout.py similarity index 69% rename from tests/backends/megatron/test_gdn_chunkwise_cp_route.py rename to tests/backends/megatron/test_gdn_cp_layout.py index 03ea96637..2a43e1879 100644 --- a/tests/backends/megatron/test_gdn_chunkwise_cp_route.py +++ b/tests/backends/megatron/test_gdn_cp_layout.py @@ -1,21 +1,9 @@ # Copyright (c) 2026 Relax Authors. All Rights Reserved. -"""Unit tests for the prebuilt THD CP layout route (Task 32, phase 3). +"""CPU tests for GDN CP partitions, layout routes, and boundary caches. -Phase 1 derived the zigzag<->contiguous all-to-all plan inside every conversion, -from device tensors, which cost a device-host synchronisation per CP rank per -call. Phase 3 backports NVIDIA/Megatron-LM#5664's idea instead: derive the plan -once per micro-batch on CPU and hand the same route to every GDN layer. - -Two things therefore need proving on CPU, with no process group: - -1. the segment-based route describes *exactly* the permutation the phase-1 - index-based partition described -- otherwise chunkwise CP silently reorders - tokens; -2. a cached route is only ever reused for the micro-batch, CP geometry and - direction it was built for. - -The real all-to-all round trip over NCCL stays in -``test_gdn_chunkwise_cp_gpu.py``. +Checks token ownership, agreement with Relax's sharding, route equivalence, and +cache reuse/invalidation. Real NCCL layout communication is covered in +``test_gdn_cp_layout_gpu.py``. """ from __future__ import annotations @@ -28,6 +16,8 @@ from megatron.core.packed_seq_params import PackedSeqParams # noqa: E402 +from relax.backends.megatron.cp_utils import gdn_cp_slice, slice_with_cp # noqa: E402 + DIRECTIONS = [("zigzag", "contiguous"), ("contiguous", "zigzag")] @@ -41,13 +31,139 @@ ] -def _cu(lengths: list[int], unit: int) -> torch.Tensor: +def _cu(lengths: list[int], unit: int = 1) -> torch.Tensor: cu = [0] for n in lengths: cu.append(cu[-1] + n * unit) return torch.tensor(cu, dtype=torch.int64) +def _tagged_tokens(total: int, width: int = 3) -> torch.Tensor: + """[total, width] where row t is (t, t+1e6, t+2e6): token identity is + unambiguous.""" + base = torch.arange(total, dtype=torch.float64).unsqueeze(1) + return base + torch.arange(width, dtype=torch.float64).unsqueeze(0) * 1e6 + + +# --------------------------------------------------------------------------- +# Partition definitions +# --------------------------------------------------------------------------- +@pytest.mark.parametrize("cp_size", [1, 2, 4, 8]) +@pytest.mark.parametrize("layout", ["zigzag", "contiguous"]) +@pytest.mark.parametrize("lengths_factor", [[1], [1, 2, 3], [3, 1, 1, 2]]) +def test_thd_rank_indices_partition_all_tokens_exactly_once(cp_size, layout, lengths_factor): + lengths = [2 * cp_size * f for f in lengths_factor] + cu = _cu(lengths) + owned = torch.cat([cpl.get_thd_context_parallel_rank_indices(cu, cp_size, r, layout) for r in range(cp_size)]) + assert owned.numel() == int(cu[-1]) + assert torch.equal(torch.sort(owned).values, torch.arange(int(cu[-1]))) + + +@pytest.mark.parametrize("cp_size", [2, 4, 8]) +def test_zigzag_rank_indices_match_relax_data_sharding(cp_size): + """MCore's zigzag partition must be token-for-token what Relax's data path + produces. + + If these ever disagree, chunkwise CP would silently permute tokens relative + to the all-gather fallback and the attention layers. + """ + lengths = [2 * cp_size * f for f in (1, 3, 2)] + cu = _cu(lengths) + full = _tagged_tokens(int(cu[-1])).reshape(-1, 1, 3) # [s, b=1, C] + + for rank in range(cp_size): + mcore_idx = cpl.get_thd_context_parallel_rank_indices(cu, cp_size, rank, "zigzag") + mcore_shard = full[mcore_idx] + + # Relax data.py: per-sample slice_with_cp then concat. + relax_shard = torch.cat( + [ + slice_with_cp( + full[cu[i] : cu[i + 1]], + pad_value=0.0, + qkv_format="thd", + dynamic_cp_size=cp_size, + dynamic_cp_rank=rank, + ) + for i in range(len(lengths)) + ], + dim=0, + ) + assert torch.equal(mcore_shard, relax_shard) + + # Relax model.py (all-gather fallback) re-slices with gdn_cp_slice. + assert torch.equal(mcore_shard, gdn_cp_slice(full, cu, cp_size, rank)) + + +@pytest.mark.parametrize("cp_size", [2, 4, 8]) +@pytest.mark.parametrize("lengths_factor", [[1], [1, 2, 3], [3, 1, 1, 2]]) +def test_both_layouts_are_permutations_of_each_other(cp_size, lengths_factor): + """The two partitions must describe the same token set with the same per- + rank size. + + That is the precondition for the all-to-all between them to be a pure + permutation -- no token invented, dropped, or duplicated. The real collective + round trip is asserted in ``test_gdn_cp_layout_gpu.py``. + """ + lengths = [2 * cp_size * f for f in lengths_factor] + cu = _cu(lengths) + total = int(cu[-1]) + zig_by_rank = [] + con_by_rank = [] + for rank in range(cp_size): + zig = cpl.get_thd_context_parallel_rank_indices(cu, cp_size, rank, "zigzag") + con = cpl.get_thd_context_parallel_rank_indices(cu, cp_size, rank, "contiguous") + zig_by_rank.append(zig) + con_by_rank.append(con) + assert zig.numel() == con.numel() == total // cp_size + # contiguous is exactly this rank's span of the flattened buffer + assert torch.equal(con, torch.arange(rank * (total // cp_size), (rank + 1) * (total // cp_size))) + + # Across the whole CP group, both layouts are permutations of exactly the + # same global token rows. + assert torch.equal( + torch.cat(zig_by_rank).sort().values, + torch.cat(con_by_rank).sort().values, + ) + + +@pytest.mark.parametrize("cp_size", [2, 4]) +def test_rank_indices_reject_lengths_not_divisible_by_two_cp(cp_size): + bad = _cu([2 * cp_size, 2 * cp_size + 1]) + with pytest.raises(ValueError, match="divisible by"): + cpl.get_thd_context_parallel_rank_indices(bad, cp_size, 0, "zigzag") + + +def test_gdn_rejects_packed_lengths_not_divisible_by_cp(): + from megatron.core.ssm.gated_delta_net import GatedDeltaNet + + cu = _cu([8, 6]) + with pytest.raises(ValueError, match="divisible by cp_size=4"): + GatedDeltaNet._resolve_cu_seqlens(None, None, cu, int(cu[-1]), "cu_seqlens_q", cp_size=4) + + +def test_rank_indices_reject_unknown_layout(): + with pytest.raises(ValueError, match="Unsupported context-parallel layout"): + cpl.get_thd_context_parallel_rank_indices(_cu([16, 16]), 2, 0, "contiguous_ish") + + +@pytest.mark.parametrize("layout", ["zigzag", "contiguous"]) +def test_rank_indices_ignore_duplicate_boundaries(layout): + compact = torch.tensor([0, 16, 40], dtype=torch.int64) + padded = torch.tensor([0, 16, 40, 40, 40], dtype=torch.int64) + for rank in range(2): + assert torch.equal( + cpl.get_thd_context_parallel_rank_indices(compact, 2, rank, layout), + cpl.get_thd_context_parallel_rank_indices(padded, 2, rank, layout), + ) + + +@pytest.mark.parametrize("layout", ["zigzag", "contiguous"]) +def test_rank_indices_reject_decreasing_boundaries(layout): + with pytest.raises(ValueError, match="nondecreasing"): + cpl.get_thd_context_parallel_rank_indices(torch.tensor([0, 16, 8]), 2, 0, layout) + + def _packed_seq_params(cu: torch.Tensor) -> PackedSeqParams: return PackedSeqParams( qkv_format="thd", diff --git a/tests/backends/megatron/test_gdn_cp_layout_gpu.py b/tests/backends/megatron/test_gdn_cp_layout_gpu.py new file mode 100644 index 000000000..5e97629ae --- /dev/null +++ b/tests/backends/megatron/test_gdn_cp_layout_gpu.py @@ -0,0 +1,114 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. +"""Real NCCL layout round-trip regression for GDN context parallelism. + +Checks token-exact zigzag/contiguous conversion for packed THD inputs and an +SBHD round trip. Requires two CUDA devices and the patched Megatron-LM. + +Run with: + pytest tests/backends/megatron/test_gdn_cp_layout_gpu.py +""" + +from __future__ import annotations + +import os + +import pytest +import torch +import torch.multiprocessing as mp + + +WORLD_SIZE = 2 + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.device_count() < WORLD_SIZE, + reason=f"requires {WORLD_SIZE} CUDA devices", +) + + +def _has_backport() -> bool: + try: + import megatron.core.context_parallel_layout # noqa: F401 + except ImportError: + return False + return True + + +needs_backport = pytest.mark.skipif(not _has_backport(), reason="requires patched Megatron-LM") + + +def _init_dist(rank, world_size): + os.environ.setdefault("MASTER_ADDR", "127.0.0.1") + os.environ.setdefault("MASTER_PORT", "29531") + torch.cuda.set_device(rank) + import torch.distributed as dist + + dist.init_process_group("nccl", rank=rank, world_size=world_size) + + +def _worker_layout_round_trip(rank, world_size, _spec, _unused): + """zigzag -> contiguous -> zigzag over a real CP group must be token-exact. + + This is the collective-level version of RFC 5.1: it drives the actual + ``all_to_all`` in ``context_parallel_layout``, for packed THD (several + unequal-length samples) and for SBHD. + """ + _init_dist(rank, world_size) + import torch.distributed as dist + from megatron.core.context_parallel_layout import ( + contiguous_to_zigzag_chunks, + get_thd_context_parallel_rank_indices, + zigzag_to_contiguous_chunks, + ) + + device = torch.device("cuda", rank) + cp_group = dist.new_group(list(range(world_size))) + + # --- packed THD, three samples of different lengths --- + lengths = [2 * world_size * f for f in (5, 1, 3)] + cu = torch.tensor([0] + torch.tensor(lengths).cumsum(0).tolist(), device=device, dtype=torch.int32) + total = int(cu[-1]) + # Row t is (t, t+1e6, t+2e6): a permuted token is impossible to miss. + full = ( + torch.arange(total, dtype=torch.float64, device=device).unsqueeze(1) + + torch.arange(3, dtype=torch.float64, device=device).unsqueeze(0) * 1e6 + ) + + zig_idx = get_thd_context_parallel_rank_indices(cu, world_size, rank, "zigzag") + con_idx = get_thd_context_parallel_rank_indices(cu, world_size, rank, "contiguous") + local_zig = full[zig_idx] + + got_con = zigzag_to_contiguous_chunks(local_zig, cp_group, seq_dim=0, cu_seqlens=cu) + assert torch.equal(got_con, full[con_idx]), f"rank {rank}: THD zigzag->contiguous is wrong" + got_zig = contiguous_to_zigzag_chunks(got_con, cp_group=cp_group, seq_dim=0, cu_seqlens=cu) + assert torch.equal(got_zig, local_zig), f"rank {rank}: THD round trip is not identity" + + # --- SBHD (chunk-level swap, no cu_seqlens) --- + seq_local = 2 * world_size * 4 + sbhd = torch.arange(seq_local * 2 * 3, dtype=torch.float64, device=device).reshape(seq_local, 2, 3) + rank * 1e9 + swapped = zigzag_to_contiguous_chunks(sbhd, cp_group, seq_dim=0) + back = contiguous_to_zigzag_chunks(swapped, cp_group=cp_group, seq_dim=0) + assert torch.equal(back, sbhd), f"rank {rank}: SBHD round trip is not identity" + + dist.barrier() + dist.destroy_process_group() + + +def _spawn(fn, spec, port, extra=None): + _spawn_world(fn, WORLD_SIZE, spec, port, extra=extra) + + +def _spawn_world(fn, world_size, spec, port, extra=None): + os.environ["MASTER_PORT"] = str(port) + mp.spawn( + fn, + args=(world_size, extra if extra is not None else spec, None), + nprocs=world_size, + join=True, + ) + + +@needs_backport +def test_layout_round_trip_over_real_cp_group(): + """RFC 5.1 at the collective level: the layout swap is a pure + permutation.""" + _spawn(_worker_layout_round_trip, "n/a", 29544) diff --git a/tests/backends/megatron/test_gdn_cp_mode_stage2.py b/tests/backends/megatron/test_gdn_cp_mode.py similarity index 72% rename from tests/backends/megatron/test_gdn_cp_mode_stage2.py rename to tests/backends/megatron/test_gdn_cp_mode.py index 6c4679f15..7c098b589 100644 --- a/tests/backends/megatron/test_gdn_cp_mode_stage2.py +++ b/tests/backends/megatron/test_gdn_cp_mode.py @@ -1,22 +1,17 @@ # Copyright (c) 2026 Relax Authors. All Rights Reserved. -"""CPU-only tests for the Task 32 Stage 2 Relax-side GDN CP routing. - -Covers the pieces added on top of the Stage 1 FLA/MCore backport -(``test_gdn_chunkwise_cp_layout.py``): the ``--linear-cp-mode`` CLI, invalid -Chunkwise combinations, and the thin runtime dispatcher installed on -``GatedDeltaNet.forward``. - -The dispatcher tests drive ``GatedDeltaNet.forward`` through duck-typed fakes -and hook/counter spies instead of a real distributed process group or FLA -kernel call, per task32-stage2-handoff.md §6.2 ("use hooks/counters, don't -infer routing from numerics"). Real-kernel / real-collective coverage stays in -``test_gdn_chunkwise_cp_gpu.py``. +"""CPU tests for GDN CP configuration and runtime dispatch. + +Covers the CLI, construction-time mode constraints, packing validation, dynamic +CP group selection, and Relax's all-gather fallback guards. Uses fakes and +spies for dispatch; real layout communication is covered in +``test_gdn_cp_layout_gpu.py``. """ from __future__ import annotations import argparse import ast +import inspect from pathlib import Path from types import ModuleType, SimpleNamespace @@ -26,7 +21,7 @@ pytest.importorskip("megatron.core.context_parallel_layout", reason="requires the patched Megatron-LM") -from megatron.core.packed_seq_params import PackedSeqParams # noqa: E402 +from megatron.core.packed_seq_params import PackedSeqParams, resolve_cp_group # noqa: E402 from megatron.core.ssm.gated_delta_net import GatedDeltaNet # noqa: E402 @@ -55,7 +50,7 @@ def _load_relax_functions(filename, names): # --------------------------------------------------------------------------- -# Step 1: CLI flag +# CLI flag # --------------------------------------------------------------------------- def _parse_megatron_args(monkeypatch, *argv): pytest.importorskip("triton", reason="the full Megatron training CLI imports Triton kernels") @@ -79,7 +74,7 @@ def test_linear_cp_mode_flag_accepts_all_concrete_modes(monkeypatch, mode): # --------------------------------------------------------------------------- -# Step 2: argument validation +# Argument validation # --------------------------------------------------------------------------- def _args(**overrides): base = dict( @@ -112,7 +107,7 @@ def test_validate_linear_cp_mode_rejects_chunkwise_deterministic_when_cp_may_exc # --------------------------------------------------------------------------- -# Steps 4-6: runtime dispatcher +# Runtime dispatcher # --------------------------------------------------------------------------- class _FakeGroup: """Minimal process-group stand-in exposing only .size()/.rank(): the @@ -152,7 +147,7 @@ def _fake_gdn_module(*, linear_cp_mode, static_cp_size=1, deterministic_mode=Fal def _isolate_gdn_forward_patch(): """`_patch_gdn_for_dynamic_cp` idempotently monkey-patches the *shared* GatedDeltaNet class attribute; save/restore it around every test so it - cannot leak into test_gdn_chunkwise_cp_gpu.py.""" + cannot leak into test_gdn_cp_layout_gpu.py.""" orig_forward = GatedDeltaNet.forward orig_patched_flag = getattr(GatedDeltaNet, "_dcp_patched", False) yield @@ -300,3 +295,97 @@ def test_all_gather_still_requires_full_recompute(monkeypatch): gdn_model._assert_gdn_full_recompute() with torch.no_grad(): gdn_model._assert_gdn_full_recompute() + + +# --------------------------------------------------------------------------- +# Dynamic CP group resolution +# --------------------------------------------------------------------------- +def test_resolve_cp_group_prefers_packed_seq_params(): + static = object() + dynamic = object() + assert resolve_cp_group(static, None) is static + assert resolve_cp_group(static, PackedSeqParams(qkv_format="thd")) is static + assert resolve_cp_group(static, PackedSeqParams(qkv_format="thd", cp_group=dynamic)) is dynamic + + +# --------------------------------------------------------------------------- +# Construction-time capability gate +# --------------------------------------------------------------------------- +def _gdn_config(**overrides): + import torch.nn.functional as F + from megatron.core.transformer.transformer_config import TransformerConfig + + kwargs = dict( + hidden_size=2048, + num_layers=1, + num_attention_heads=16, + num_query_groups=2, + normalization="RMSNorm", + use_cpu_initialization=True, + activation_func=F.silu, + bf16=True, + experimental_attention_variant="gated_delta_net", + linear_attention_freq=[1], + linear_conv_kernel_dim=4, + linear_key_head_dim=128, + linear_value_head_dim=128, + linear_num_key_heads=16, + linear_num_value_heads=32, + ) + kwargs.update(overrides) + return TransformerConfig(**kwargs) + + +def test_config_default_mode_is_chunkwise(): + """Upgrading the image must not silently reroute an existing recipe.""" + assert _gdn_config().linear_cp_mode == "chunkwise" + + +def test_headwise_config_requires_heads_divisible_by_tp_times_cp(): + # 16 key heads, tp=2, cp=4 -> 16 % 8 == 0: fine. + _gdn_config(tensor_model_parallel_size=2, context_parallel_size=4, linear_cp_mode="headwise") + # tp=2, cp=16 -> 16 % 32 != 0: the geometry headwise cannot express. + with pytest.raises(AssertionError, match="linear_num_key_heads"): + _gdn_config(tensor_model_parallel_size=2, context_parallel_size=16, linear_cp_mode="headwise") + + +def test_chunkwise_config_only_requires_heads_divisible_by_tp(): + """This is what replaces Relax's temporary head-count rewrite hack.""" + cfg = _gdn_config(tensor_model_parallel_size=2, context_parallel_size=16, linear_cp_mode="chunkwise") + assert cfg.linear_num_key_heads == 16 and cfg.linear_num_value_heads == 32 + # ... but TP divisibility is still enforced: GDN weights stay TP-sharded. + with pytest.raises(AssertionError, match="linear_num_key_heads"): + _gdn_config( + tensor_model_parallel_size=8, + context_parallel_size=2, + linear_cp_mode="chunkwise", + num_query_groups=8, + linear_num_key_heads=4, + linear_num_value_heads=8, + ) + + +def test_all_gather_config_uses_the_tp_only_head_rule(): + """`--linear-cp-mode=all_gather` must be constructible on a non-divisible + geometry. + + Relax's all-gather fallback keeps GDN weights TP-only, so declaring it + should relax the head check exactly as chunkwise does. + """ + cfg = _gdn_config(tensor_model_parallel_size=2, context_parallel_size=16, linear_cp_mode="all_gather") + assert cfg.linear_num_key_heads == 16 and cfg.linear_num_value_heads == 32 + + +def test_config_rejects_unresolved_and_unknown_linear_cp_mode(): + """MCore only accepts the three concrete execution modes.""" + for bad in ("auto", "allgather", "chunk", ""): + with pytest.raises(AssertionError, match="linear_cp_mode"): + _gdn_config(context_parallel_size=2, linear_cp_mode=bad) + with pytest.raises(AssertionError, match="linear_cp_mode"): + _gdn_config(context_parallel_size=4, tensor_model_parallel_size=2, linear_cp_mode=bad) + + +def test_gdn_forward_has_no_per_call_mode_override(): + from megatron.core.ssm.gated_delta_net import GatedDeltaNet + + assert "linear_cp_mode" not in inspect.signature(GatedDeltaNet.forward).parameters From e09e8f5f1f481680d1d50177b4c7f3e945add59c Mon Sep 17 00:00:00 2001 From: jambow0320 <1193785411@qq.com> Date: Wed, 23 Sep 2026 18:36:21 +1000 Subject: [PATCH 11/11] test(megatron): cover GDN CP forward and backward MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit # ✅ Tests - Compare real BF16 GDN CP=1 with CP=2 chunkwise and all_gather outputs, input gradients, and all parameter gradients using 128-wide heads and unequal packed sequences. - Run the existing exact THD/SBHD layout checks in the same two-GPU spawn. - Select one FLA autotune candidate only inside test workers, retaining real JIT kernels and NCCL while avoiding repeated candidate benchmarks. - Update references to the consolidated GPU test filename. ## Validation - Two H200s, empty compilation caches: 1 passed, 0 skipped; 60.224 seconds for the combined regression (61.740 seconds including pytest overhead). - Python 3.11 CPU-only collection skips the GPU case as intended. - pre-commit run --all-files --show-diff-on-failure passed. --- tests/backends/megatron/test_gdn_cp_gpu.py | 244 ++++++++++++++++++ tests/backends/megatron/test_gdn_cp_layout.py | 4 +- .../megatron/test_gdn_cp_layout_gpu.py | 114 -------- tests/backends/megatron/test_gdn_cp_mode.py | 4 +- 4 files changed, 248 insertions(+), 118 deletions(-) create mode 100644 tests/backends/megatron/test_gdn_cp_gpu.py delete mode 100644 tests/backends/megatron/test_gdn_cp_layout_gpu.py diff --git a/tests/backends/megatron/test_gdn_cp_gpu.py b/tests/backends/megatron/test_gdn_cp_gpu.py new file mode 100644 index 000000000..d2d952f3a --- /dev/null +++ b/tests/backends/megatron/test_gdn_cp_gpu.py @@ -0,0 +1,244 @@ +# Copyright (c) 2026 Relax Authors. All Rights Reserved. +"""Real layout and GDN forward/backward regression for context parallelism. + +Checks exact THD/SBHD layout round trips and BF16 CP=1/CP=2 GDN parity in one +two-GPU spawn. FLA kernels use a fixed candidate configuration to avoid +autotune benchmarks. Run with: pytest +tests/backends/megatron/test_gdn_cp_gpu.py +""" + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +import torch +import torch.multiprocessing as mp + + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.device_count() < 2, + reason="requires two CUDA devices and the patched Megatron-LM image", +) + + +def _check_layout_round_trip(cp_group): + from megatron.core.context_parallel_layout import ( + contiguous_to_zigzag_chunks, + get_thd_context_parallel_rank_indices, + zigzag_to_contiguous_chunks, + ) + + rank, world_size = cp_group.rank(), cp_group.size() + device = torch.device("cuda", rank) + + # Packed THD: short, unequal samples exercise the original route checks. + lengths = [2 * world_size * f for f in (5, 1, 3)] + cu = torch.tensor([0] + torch.tensor(lengths).cumsum(0).tolist(), device=device, dtype=torch.int32) + total = int(cu[-1]) + # Row t is (t, t+1e6, t+2e6): a permuted token is impossible to miss. + full = ( + torch.arange(total, dtype=torch.float64, device=device).unsqueeze(1) + + torch.arange(3, dtype=torch.float64, device=device).unsqueeze(0) * 1e6 + ) + + zig_idx = get_thd_context_parallel_rank_indices(cu, world_size, rank, "zigzag") + con_idx = get_thd_context_parallel_rank_indices(cu, world_size, rank, "contiguous") + local_zig = full[zig_idx] + + got_con = zigzag_to_contiguous_chunks(local_zig, cp_group, seq_dim=0, cu_seqlens=cu) + assert torch.equal(got_con, full[con_idx]), f"rank {rank}: THD zigzag->contiguous is wrong" + got_zig = contiguous_to_zigzag_chunks(got_con, cp_group=cp_group, seq_dim=0, cu_seqlens=cu) + assert torch.equal(got_zig, local_zig), f"rank {rank}: THD round trip is not identity" + + # SBHD: chunk-level swap without packed sequence metadata. + seq_local = 2 * world_size * 4 + sbhd = torch.arange(seq_local * 2 * 3, dtype=torch.float64, device=device).reshape(seq_local, 2, 3) + rank * 1e9 + swapped = zigzag_to_contiguous_chunks(sbhd, cp_group, seq_dim=0) + back = contiguous_to_zigzag_chunks(swapped, cp_group=cp_group, seq_dim=0) + assert torch.equal(back, sbhd), f"rank {rank}: SBHD round trip is not identity" + + +def _build_gdn(mode): + from megatron.core import parallel_state + from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( + get_experimental_attention_variant_module_spec, + ) + from megatron.core.process_groups_config import ProcessGroupCollection + from megatron.core.ssm.gated_delta_net import GatedDeltaNet + from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed + from megatron.core.transformer.transformer_config import TransformerConfig + + torch.manual_seed(123) + model_parallel_cuda_manual_seed(123) + config = TransformerConfig( + hidden_size=512, + num_layers=1, + num_attention_heads=8, + normalization="RMSNorm", + use_cpu_initialization=True, + layernorm_zero_centered_gamma=True, + activation_func=torch.nn.functional.silu, + bf16=True, + context_parallel_size=2, + experimental_attention_variant="gated_delta_net", + linear_attention_freq=[1], + linear_conv_kernel_dim=4, + # Qwen3.5 uses 128-wide heads with a 1:2 key/value head ratio. + linear_key_head_dim=128, + linear_value_head_dim=128, + linear_num_key_heads=4, + linear_num_value_heads=8, + linear_cp_mode=mode, + transformer_impl="transformer_engine", + ) + return ( + GatedDeltaNet( + config, + submodules=get_experimental_attention_variant_module_spec(config=config).submodules, + layer_number=1, + bias=False, + conv_bias=False, + conv_init=1.0, + use_qk_l2norm=True, + A_init_range=(1, 16), + pg_collection=ProcessGroupCollection( + tp=parallel_state.get_tensor_model_parallel_group(), + cp=parallel_state.get_context_parallel_group(), + ), + ) + .cuda() + .to(torch.bfloat16) + ) + + +def _forward_backward(gdn, hidden, packed, grad_out, *, recompute=False): + from torch.utils.checkpoint import checkpoint + + gdn.zero_grad(set_to_none=True) + hidden = hidden.detach().clone().requires_grad_(True) + + def forward(x): + return gdn(x, None, packed_seq_params=packed)[0] + + out = checkpoint(forward, hidden, use_reentrant=False) if recompute else forward(hidden) + (out.float() * grad_out).sum().backward() + grads = {name: param.grad.detach().float().clone() for name, param in gdn.named_parameters()} + return out.detach(), hidden.grad.detach(), grads + + +def _assert_grad_close(got, expected, name): + # BF16 parameter gradients sum over different token partitions. Bound both + # magnitude and direction, without imposing sub-ULP elementwise agreement. + got, expected = got.flatten().float(), expected.flatten().float() + assert torch.isfinite(got).all() and torch.isfinite(expected).all(), name + relative_rms = (got - expected).norm() / expected.norm().clamp_min(1e-12) + cosine = torch.nn.functional.cosine_similarity(got, expected, dim=0) + assert relative_rms < 1e-2, f"{name}: relative RMS error {relative_rms.item():.3e}" + assert cosine >= 0.9999, f"{name}: cosine {cosine.item():.8f}" + + +def _use_one_fla_config(): + """Skip FLA tuning in this spawned worker while retaining real JIT + kernels.""" + import triton + + autotune = triton.autotune + + def single_config(configs, *args, **kwargs): + def decorate(fn): + selected = configs[:1] if fn.__module__.startswith("fla.") else configs + return autotune(selected, *args, **kwargs)(fn) + + return decorate + + # The override covers lazy imports and ends with this isolated worker. + triton.autotune = single_config + + +def _worker_gdn_cp(rank, init_method): + torch.cuda.set_device(rank) + _use_one_fla_config() # Install before Megatron/FLA imports create kernels. + import torch.distributed as dist + from megatron.core import parallel_state + from megatron.core.packed_seq_params import PackedSeqParams + from megatron.training.global_vars import set_args + + from relax.backends.megatron.model import _patch_gdn_for_dynamic_cp + + dist.init_process_group("nccl", init_method=init_method, rank=rank, world_size=2) + parallel_state.initialize_model_parallel(context_parallel_size=2) + try: + set_args(SimpleNamespace(recompute_granularity="full")) + _patch_gdn_for_dynamic_cp() + cp_group = parallel_state.get_context_parallel_group() + _check_layout_round_trip(cp_group) + solo = [dist.new_group([r]) for r in range(2)][rank] + models = {mode: _build_gdn(mode) for mode in ("chunkwise", "all_gather")} + for param in models["chunkwise"].parameters(): + dist.broadcast(param.data, src=0) + models["all_gather"].load_state_dict(models["chunkwise"].state_dict()) + + # Unequal packed samples; the contiguous CP boundary at 512 splits the + # first sample, exercising recurrent/convolution state across ranks. + lengths = [768, 256] + total = sum(lengths) + hidden_size = models["chunkwise"].config.hidden_size + cu = torch.tensor([0, lengths[0], total], device="cuda", dtype=torch.int32) + generator = torch.Generator().manual_seed(7) + hidden = torch.randn(total, 1, hidden_size, generator=generator).cuda().to(torch.bfloat16) + grad_out = torch.randn(total, 1, hidden_size, generator=generator).cuda() + # Build expected zigzag ownership independently of the production route. + indices = torch.cat( + [ + part.reshape(4, -1)[[rank, 3 - rank]].flatten() + for part in torch.arange(total, device="cuda").split(lengths) + ] + ) + + def packed(group, size): + return PackedSeqParams( + qkv_format="thd", + cu_seqlens_q=cu, + cu_seqlens_kv=cu, + cu_seqlens_q_padded=cu, + cu_seqlens_kv_padded=cu, + max_seqlen_q=max(lengths), + max_seqlen_kv=max(lengths), + cp_group=group, + local_cp_size=size, + ) + + ref_out, ref_dx, ref_grads = _forward_backward(models["chunkwise"], hidden, packed(solo, 1), grad_out) + shard = hidden[indices] + grad_shard = grad_out[indices] + for mode, gdn in models.items(): + out, dx, grads = _forward_backward( + gdn, shard, packed(cp_group, 2), grad_shard, recompute=mode == "all_gather" + ) + for name, got, expected in (("output", out, ref_out), ("input gradient", dx, ref_dx)): + expected = expected[indices] + torch.testing.assert_close( + got, + expected, + atol=2e-3, + rtol=1e-2, + msg=lambda message: f"{mode} {name}: {message}", + ) + cosine = torch.nn.functional.cosine_similarity( + got.flatten().float(), expected.flatten().float(), dim=0 + ) + assert cosine >= 0.9999, f"{mode} {name}: cosine {cosine.item():.8f}" + for name, grad in grads.items(): + # CP owns partial token sums; DDP would sum these replicas. + dist.all_reduce(grad, group=cp_group) + _assert_grad_close(grad, ref_grads[name], f"{mode} {name}") + finally: + parallel_state.destroy_model_parallel() + dist.destroy_process_group() + + +def test_gdn_cp_layout_and_gradients(tmp_path): + """Check exact layouts and CP=2 output/input/parameter gradients against + CP=1.""" + mp.spawn(_worker_gdn_cp, args=(f"file://{tmp_path / 'gdn_cp_init'}",), nprocs=2, join=True) diff --git a/tests/backends/megatron/test_gdn_cp_layout.py b/tests/backends/megatron/test_gdn_cp_layout.py index 2a43e1879..7c85863d6 100644 --- a/tests/backends/megatron/test_gdn_cp_layout.py +++ b/tests/backends/megatron/test_gdn_cp_layout.py @@ -3,7 +3,7 @@ Checks token ownership, agreement with Relax's sharding, route equivalence, and cache reuse/invalidation. Real NCCL layout communication is covered in -``test_gdn_cp_layout_gpu.py``. +``test_gdn_cp_gpu.py``. """ from __future__ import annotations @@ -103,7 +103,7 @@ def test_both_layouts_are_permutations_of_each_other(cp_size, lengths_factor): That is the precondition for the all-to-all between them to be a pure permutation -- no token invented, dropped, or duplicated. The real collective - round trip is asserted in ``test_gdn_cp_layout_gpu.py``. + round trip is asserted in ``test_gdn_cp_gpu.py``. """ lengths = [2 * cp_size * f for f in lengths_factor] cu = _cu(lengths) diff --git a/tests/backends/megatron/test_gdn_cp_layout_gpu.py b/tests/backends/megatron/test_gdn_cp_layout_gpu.py deleted file mode 100644 index 5e97629ae..000000000 --- a/tests/backends/megatron/test_gdn_cp_layout_gpu.py +++ /dev/null @@ -1,114 +0,0 @@ -# Copyright (c) 2026 Relax Authors. All Rights Reserved. -"""Real NCCL layout round-trip regression for GDN context parallelism. - -Checks token-exact zigzag/contiguous conversion for packed THD inputs and an -SBHD round trip. Requires two CUDA devices and the patched Megatron-LM. - -Run with: - pytest tests/backends/megatron/test_gdn_cp_layout_gpu.py -""" - -from __future__ import annotations - -import os - -import pytest -import torch -import torch.multiprocessing as mp - - -WORLD_SIZE = 2 - -pytestmark = pytest.mark.skipif( - not torch.cuda.is_available() or torch.cuda.device_count() < WORLD_SIZE, - reason=f"requires {WORLD_SIZE} CUDA devices", -) - - -def _has_backport() -> bool: - try: - import megatron.core.context_parallel_layout # noqa: F401 - except ImportError: - return False - return True - - -needs_backport = pytest.mark.skipif(not _has_backport(), reason="requires patched Megatron-LM") - - -def _init_dist(rank, world_size): - os.environ.setdefault("MASTER_ADDR", "127.0.0.1") - os.environ.setdefault("MASTER_PORT", "29531") - torch.cuda.set_device(rank) - import torch.distributed as dist - - dist.init_process_group("nccl", rank=rank, world_size=world_size) - - -def _worker_layout_round_trip(rank, world_size, _spec, _unused): - """zigzag -> contiguous -> zigzag over a real CP group must be token-exact. - - This is the collective-level version of RFC 5.1: it drives the actual - ``all_to_all`` in ``context_parallel_layout``, for packed THD (several - unequal-length samples) and for SBHD. - """ - _init_dist(rank, world_size) - import torch.distributed as dist - from megatron.core.context_parallel_layout import ( - contiguous_to_zigzag_chunks, - get_thd_context_parallel_rank_indices, - zigzag_to_contiguous_chunks, - ) - - device = torch.device("cuda", rank) - cp_group = dist.new_group(list(range(world_size))) - - # --- packed THD, three samples of different lengths --- - lengths = [2 * world_size * f for f in (5, 1, 3)] - cu = torch.tensor([0] + torch.tensor(lengths).cumsum(0).tolist(), device=device, dtype=torch.int32) - total = int(cu[-1]) - # Row t is (t, t+1e6, t+2e6): a permuted token is impossible to miss. - full = ( - torch.arange(total, dtype=torch.float64, device=device).unsqueeze(1) - + torch.arange(3, dtype=torch.float64, device=device).unsqueeze(0) * 1e6 - ) - - zig_idx = get_thd_context_parallel_rank_indices(cu, world_size, rank, "zigzag") - con_idx = get_thd_context_parallel_rank_indices(cu, world_size, rank, "contiguous") - local_zig = full[zig_idx] - - got_con = zigzag_to_contiguous_chunks(local_zig, cp_group, seq_dim=0, cu_seqlens=cu) - assert torch.equal(got_con, full[con_idx]), f"rank {rank}: THD zigzag->contiguous is wrong" - got_zig = contiguous_to_zigzag_chunks(got_con, cp_group=cp_group, seq_dim=0, cu_seqlens=cu) - assert torch.equal(got_zig, local_zig), f"rank {rank}: THD round trip is not identity" - - # --- SBHD (chunk-level swap, no cu_seqlens) --- - seq_local = 2 * world_size * 4 - sbhd = torch.arange(seq_local * 2 * 3, dtype=torch.float64, device=device).reshape(seq_local, 2, 3) + rank * 1e9 - swapped = zigzag_to_contiguous_chunks(sbhd, cp_group, seq_dim=0) - back = contiguous_to_zigzag_chunks(swapped, cp_group=cp_group, seq_dim=0) - assert torch.equal(back, sbhd), f"rank {rank}: SBHD round trip is not identity" - - dist.barrier() - dist.destroy_process_group() - - -def _spawn(fn, spec, port, extra=None): - _spawn_world(fn, WORLD_SIZE, spec, port, extra=extra) - - -def _spawn_world(fn, world_size, spec, port, extra=None): - os.environ["MASTER_PORT"] = str(port) - mp.spawn( - fn, - args=(world_size, extra if extra is not None else spec, None), - nprocs=world_size, - join=True, - ) - - -@needs_backport -def test_layout_round_trip_over_real_cp_group(): - """RFC 5.1 at the collective level: the layout swap is a pure - permutation.""" - _spawn(_worker_layout_round_trip, "n/a", 29544) diff --git a/tests/backends/megatron/test_gdn_cp_mode.py b/tests/backends/megatron/test_gdn_cp_mode.py index 7c098b589..582c91cc0 100644 --- a/tests/backends/megatron/test_gdn_cp_mode.py +++ b/tests/backends/megatron/test_gdn_cp_mode.py @@ -4,7 +4,7 @@ Covers the CLI, construction-time mode constraints, packing validation, dynamic CP group selection, and Relax's all-gather fallback guards. Uses fakes and spies for dispatch; real layout communication is covered in -``test_gdn_cp_layout_gpu.py``. +``test_gdn_cp_gpu.py``. """ from __future__ import annotations @@ -147,7 +147,7 @@ def _fake_gdn_module(*, linear_cp_mode, static_cp_size=1, deterministic_mode=Fal def _isolate_gdn_forward_patch(): """`_patch_gdn_for_dynamic_cp` idempotently monkey-patches the *shared* GatedDeltaNet class attribute; save/restore it around every test so it - cannot leak into test_gdn_cp_layout_gpu.py.""" + cannot leak into test_gdn_cp_gpu.py.""" orig_forward = GatedDeltaNet.forward orig_patched_flag = getattr(GatedDeltaNet, "_dcp_patched", False) yield