From 695e93b91f6ce8590c3cdaaf8fbac64201ae8546 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Fri, 27 Feb 2026 08:12:18 +0800 Subject: [PATCH] Make reorderer support packed format with CP in dump comparator (#19462) --- .../comparator/aligner/reorderer/executor.py | 72 +++++- .../comparator/aligner/reorderer/planner.py | 40 ++- .../comparator/aligner/reorderer/types.py | 16 +- .../aligner/reorderer/test_executor.py | 239 ++++++++++++++++++ .../aligner/reorderer/test_planner.py | 49 +++- 5 files changed, 398 insertions(+), 18 deletions(-) diff --git a/python/sglang/srt/debug_utils/comparator/aligner/reorderer/executor.py b/python/sglang/srt/debug_utils/comparator/aligner/reorderer/executor.py index 6355f4926..a44f5edce 100644 --- a/python/sglang/srt/debug_utils/comparator/aligner/reorderer/executor.py +++ b/python/sglang/srt/debug_utils/comparator/aligner/reorderer/executor.py @@ -1,6 +1,12 @@ +from typing import Optional + import torch -from sglang.srt.debug_utils.comparator.aligner.reorderer.types import ReordererPlan +from sglang.srt.debug_utils.comparator.aligner.reorderer.types import ( + ReordererPlan, + ZigzagToNaturalParams, + ZigzagToNaturalThdParams, +) from sglang.srt.debug_utils.comparator.dims import ( resolve_dim_by_name, strip_dim_names, @@ -11,12 +17,66 @@ def execute_reorderer_plan( plan: ReordererPlan, tensors: list[torch.Tensor], ) -> list[torch.Tensor]: - dim: int = resolve_dim_by_name(tensors[0], plan.params.dim_name) - return [ - _reorder_zigzag_to_natural(tensor, dim=dim, cp_size=plan.params.cp_size) - for tensor in tensors + if isinstance(plan.params, ZigzagToNaturalThdParams): + thd_dim: int = resolve_dim_by_name(tensors[0], plan.params.dim_name) + return [ + _reorder_zigzag_to_natural_thd( + tensor, + dim=thd_dim, + cp_size=plan.params.cp_size, + seq_lens=plan.params.seq_lens, + ) + for tensor in tensors + ] + + if isinstance(plan.params, ZigzagToNaturalParams): + dim: int = resolve_dim_by_name(tensors[0], plan.params.dim_name) + return [ + _reorder_zigzag_to_natural(tensor, dim=dim, cp_size=plan.params.cp_size) + for tensor in tensors + ] + + raise ValueError(f"Unsupported reorderer params type: {type(plan.params).__name__}") + + +def _reorder_zigzag_to_natural_thd( + tensor: torch.Tensor, *, dim: int, cp_size: int, seq_lens: list[int] +) -> torch.Tensor: + """Undo CP zigzag interleaving for THD (packed-seq) format. + + Each seq in seq_lens is independently reordered from zigzag to natural order + along the given dim. + """ + stripped: torch.Tensor = strip_dim_names(tensor) + names: tuple[Optional[str], ...] = tensor.names + + split_sizes: list[int] = list(seq_lens) + remainder: int = stripped.shape[dim] - sum(split_sizes) + if remainder < 0: + raise ValueError( + f"sum(seq_lens)={sum(split_sizes)} exceeds tensor dim size " + f"{stripped.shape[dim]} along dim={dim}" + ) + if remainder > 0: + split_sizes.append(remainder) + + segments: list[torch.Tensor] = list(stripped.split(split_sizes, dim=dim)) + + reordered_segments: list[torch.Tensor] = [ + _reorder_zigzag_to_natural(seg, dim=dim, cp_size=cp_size) + for seg in segments[: len(seq_lens)] ] + # Tail padding — pass through unchanged + if remainder > 0: + reordered_segments.append(segments[-1]) + + result: torch.Tensor = torch.cat(reordered_segments, dim=dim) + + if names[0] is not None: + result = result.refine_names(*names) + return result + def _reorder_zigzag_to_natural( tensor: torch.Tensor, *, dim: int, cp_size: int @@ -27,7 +87,7 @@ def _reorder_zigzag_to_natural( (megatron/core/ssm/mamba_context_parallel.py:360-373). """ stripped: torch.Tensor = strip_dim_names(tensor) - names: tuple = tensor.names + names: tuple[Optional[str], ...] = tensor.names num_chunks: int = cp_size * 2 chunks: tuple[torch.Tensor, ...] = stripped.chunk(num_chunks, dim=dim) diff --git a/python/sglang/srt/debug_utils/comparator/aligner/reorderer/planner.py b/python/sglang/srt/debug_utils/comparator/aligner/reorderer/planner.py index ee867e596..11c81cd0e 100644 --- a/python/sglang/srt/debug_utils/comparator/aligner/reorderer/planner.py +++ b/python/sglang/srt/debug_utils/comparator/aligner/reorderer/planner.py @@ -1,21 +1,27 @@ +from typing import Optional + from sglang.srt.debug_utils.comparator.aligner.reorderer.types import ( ReordererPlan, ZigzagToNaturalParams, + ZigzagToNaturalThdParams, ) from sglang.srt.debug_utils.comparator.aligner.unsharder.types import AxisInfo from sglang.srt.debug_utils.comparator.dims import ( SEQ_DIM_NAME, + TOKEN_DIM_NAME, DimSpec, Ordering, ParallelAxis, ) -_ALLOWED_ZIGZAG_DIM_NAMES: set[str] = {SEQ_DIM_NAME} +_ALLOWED_ZIGZAG_DIM_NAMES: set[str] = {SEQ_DIM_NAME, TOKEN_DIM_NAME} def compute_reorderer_plans( dim_specs: list[DimSpec], parallel_infos: list[dict[ParallelAxis, AxisInfo]], + *, + thd_global_seq_lens: Optional[list[int]] = None, ) -> list[ReordererPlan]: plans: list[ReordererPlan] = [] @@ -28,17 +34,35 @@ def compute_reorderer_plans( if spec.name not in _ALLOWED_ZIGZAG_DIM_NAMES: raise ValueError( f"Zigzag ordering is only supported on sequence dims " - f"(bshd/sbhd format, dim name must be one of " + f"(dim name must be one of " f"{sorted(_ALLOWED_ZIGZAG_DIM_NAMES)}), " f"but got dim name {spec.name!r} in {spec}" ) - assert spec.ordering == Ordering.ZIGZAG - axis_size: int = parallel_infos[0][spec.parallel].axis_size - plans.append( - ReordererPlan( - params=ZigzagToNaturalParams(dim_name=spec.name, cp_size=axis_size), + if spec.ordering != Ordering.ZIGZAG: + raise ValueError( + f"Unsupported ordering {spec.ordering!r} for dim {spec.name!r}" ) - ) + axis_size: int = parallel_infos[0][spec.parallel].axis_size + + if spec.name == TOKEN_DIM_NAME: + if thd_global_seq_lens is None: + raise ValueError( + "thd_global_seq_lens is required for zigzag reorder on 't' dimension" + ) + params = ZigzagToNaturalThdParams( + dim_name=spec.name, + cp_size=axis_size, + seq_lens=thd_global_seq_lens, + ) + elif spec.name == SEQ_DIM_NAME: + params = ZigzagToNaturalParams(dim_name=spec.name, cp_size=axis_size) + else: + raise ValueError( + f"Unsupported zigzag dim name {spec.name!r}, " + f"expected one of {sorted(_ALLOWED_ZIGZAG_DIM_NAMES)}" + ) + + plans.append(ReordererPlan(params=params)) return plans diff --git a/python/sglang/srt/debug_utils/comparator/aligner/reorderer/types.py b/python/sglang/srt/debug_utils/comparator/aligner/reorderer/types.py index c247dce62..614a28b9d 100644 --- a/python/sglang/srt/debug_utils/comparator/aligner/reorderer/types.py +++ b/python/sglang/srt/debug_utils/comparator/aligner/reorderer/types.py @@ -1,4 +1,6 @@ -from typing import Literal +from typing import Annotated, Literal, Union + +from pydantic import Field from sglang.srt.debug_utils.comparator.utils import _FrozenBase @@ -9,7 +11,17 @@ class ZigzagToNaturalParams(_FrozenBase): cp_size: int -ReordererParams = ZigzagToNaturalParams +class ZigzagToNaturalThdParams(_FrozenBase): + op: Literal["zigzag_to_natural_thd"] = "zigzag_to_natural_thd" + dim_name: str + cp_size: int + seq_lens: list[int] # unshard-ed per-seq token counts, e.g. [100, 64, 92] + + +ReordererParams = Annotated[ + Union[ZigzagToNaturalParams, ZigzagToNaturalThdParams], + Field(discriminator="op"), +] class ReordererPlan(_FrozenBase): diff --git a/test/registered/debug_utils/comparator/aligner/reorderer/test_executor.py b/test/registered/debug_utils/comparator/aligner/reorderer/test_executor.py index 1029e5475..da0764e9b 100644 --- a/test/registered/debug_utils/comparator/aligner/reorderer/test_executor.py +++ b/test/registered/debug_utils/comparator/aligner/reorderer/test_executor.py @@ -5,12 +5,49 @@ import torch from sglang.srt.debug_utils.comparator.aligner.reorderer.executor import ( _reorder_zigzag_to_natural, + _reorder_zigzag_to_natural_thd, + execute_reorderer_plan, ) +from sglang.srt.debug_utils.comparator.aligner.reorderer.types import ( + ReordererPlan, + ZigzagToNaturalThdParams, +) +from sglang.srt.debug_utils.comparator.aligner.unsharder.executor import ( + execute_unsharder_plan, +) +from sglang.srt.debug_utils.comparator.aligner.unsharder.types import ( + CpThdConcatParams, + UnsharderPlan, +) +from sglang.srt.debug_utils.comparator.dims import ParallelAxis +from sglang.srt.debug_utils.comparator.warning_sink import warning_sink from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=10, suite="default", nightly=True) +def _zigzag_order(cp_size: int) -> list[int]: + """Build zigzag interleaving order for 2*cp_size chunks.""" + order: list[int] = [] + num_chunks: int = cp_size * 2 + for i in range(cp_size): + order.append(i) + order.append(num_chunks - 1 - i) + return order + + +def _zigzag_split_seq(seq_natural: torch.Tensor, *, cp_size: int) -> list[torch.Tensor]: + """Split a natural-order seq into per-rank zigzag segments. + + Returns: list of per-rank tensors, where rank_i holds chunks assigned by zigzag. + """ + num_chunks: int = cp_size * 2 + chunks: list[torch.Tensor] = list(seq_natural.chunk(num_chunks, dim=0)) + order: list[int] = _zigzag_order(cp_size) + zigzagged: torch.Tensor = torch.cat([chunks[i] for i in order], dim=0) + return list(zigzagged.chunk(cp_size, dim=0)) + + class TestZigzagToNatural: def test_zigzag_to_natural_cp2(self) -> None: """cp_size=2: zigzag order [0,3,1,2] -> natural [0,1,2,3].""" @@ -46,5 +83,207 @@ class TestZigzagToNatural: assert torch.equal(result, natural) +class TestZigzagToNaturalThd: + def test_single_seq(self) -> None: + """Single seq THD reorder: equivalent to whole-tensor reorder.""" + natural = torch.arange(100) + zigzag_ranks: list[torch.Tensor] = _zigzag_split_seq(natural, cp_size=2) + zigzagged: torch.Tensor = torch.cat(zigzag_ranks, dim=0) + + result = _reorder_zigzag_to_natural_thd( + zigzagged, dim=0, cp_size=2, seq_lens=[100] + ) + assert torch.equal(result, natural) + + def test_multi_seq(self) -> None: + """Two seqs of different lengths, each independently reordered.""" + seq_a_natural = torch.arange(100) + seq_b_natural = torch.arange(100, 164) + + seq_a_zigzag: torch.Tensor = torch.cat( + _zigzag_split_seq(seq_a_natural, cp_size=2), dim=0 + ) + seq_b_zigzag: torch.Tensor = torch.cat( + _zigzag_split_seq(seq_b_natural, cp_size=2), dim=0 + ) + + combined_zigzag: torch.Tensor = torch.cat([seq_a_zigzag, seq_b_zigzag], dim=0) + result = _reorder_zigzag_to_natural_thd( + combined_zigzag, dim=0, cp_size=2, seq_lens=[100, 64] + ) + + expected: torch.Tensor = torch.cat([seq_a_natural, seq_b_natural], dim=0) + assert torch.equal(result, expected) + + def test_with_tail_pad(self) -> None: + """THD reorder with trailing global padding preserved unchanged.""" + seq_natural = torch.arange(100) + pad: torch.Tensor = torch.full((56,), fill_value=-1) + + seq_zigzag: torch.Tensor = torch.cat( + _zigzag_split_seq(seq_natural, cp_size=2), dim=0 + ) + combined: torch.Tensor = torch.cat([seq_zigzag, pad], dim=0) + + result = _reorder_zigzag_to_natural_thd( + combined, dim=0, cp_size=2, seq_lens=[100] + ) + + assert torch.equal(result[:100], seq_natural) + assert torch.equal(result[100:], pad) + + def test_with_hidden_dim(self) -> None: + """THD reorder with trailing hidden dimension (shape [T, H]).""" + torch.manual_seed(42) + hidden: int = 8 + seq_natural = torch.randn(100, hidden) + + seq_zigzag: torch.Tensor = torch.cat( + _zigzag_split_seq(seq_natural, cp_size=2), dim=0 + ) + + result = _reorder_zigzag_to_natural_thd( + seq_zigzag, dim=0, cp_size=2, seq_lens=[100] + ) + assert torch.equal(result, seq_natural) + + def test_with_leading_batch_dim(self) -> None: + """THD reorder with leading batch dim: shape [B, T, H], t is dim=1.""" + torch.manual_seed(42) + batch: int = 2 + hidden: int = 4 + seq_a_natural = torch.randn(batch, 100, hidden) + seq_b_natural = torch.randn(batch, 64, hidden) + full_natural: torch.Tensor = torch.cat([seq_a_natural, seq_b_natural], dim=1) + + # Zigzag each seq along dim=1 + def zigzag_along_dim1(t: torch.Tensor) -> torch.Tensor: + num_chunks: int = 2 * 2 # cp_size=2 + chunks: list[torch.Tensor] = list(t.chunk(num_chunks, dim=1)) + order: list[int] = [0, 3, 1, 2] # zigzag for cp_size=2 + return torch.cat([chunks[i] for i in order], dim=1) + + seq_a_zigzag: torch.Tensor = zigzag_along_dim1(seq_a_natural) + seq_b_zigzag: torch.Tensor = zigzag_along_dim1(seq_b_natural) + combined_zigzag: torch.Tensor = torch.cat([seq_a_zigzag, seq_b_zigzag], dim=1) + + result = _reorder_zigzag_to_natural_thd( + combined_zigzag, dim=1, cp_size=2, seq_lens=[100, 64] + ) + assert torch.equal(result, full_natural) + + +class TestThdCpZigzagE2E: + """End-to-end unshard + reorder tests for THD CP zigzag format. + + Simulates Miles/Megatron forward data splitting: + + cp_size=2, batch with 2 seqs: seqA(100 tokens), seqB(61→pad to 64) + + Forward: + seqA(100): chunk_size=25, 4 chunks → rank0=[chunk0+chunk3](50), rank1=[chunk1+chunk2](50) + seqB(64): chunk_size=16, 4 chunks → rank0=[chunk0+chunk3](32), rank1=[chunk1+chunk2](32) + global pad → align to 128 + rank0: [seqA_r0(50) | seqB_r0(32) | pad(46)] = 128 tokens + rank1: [seqA_r1(50) | seqB_r1(32) | pad(46)] = 128 tokens + global cu_seqlens: [0, 100, 164, 256] + + Comparator undo: + Step 1 THD unshard: per-seq cross-rank concat → [seqA_zigzag(100) | seqB_zigzag(64) | pad(92)] + Step 2 THD reorder: per-seq zigzag→natural → [seqA_natural(100) | seqB_natural(64) | pad(92)] + """ + + def test_thd_cp2_two_seqs(self) -> None: + """cp_size=2, 2 seqs (100, 61→64) + global pad.""" + torch.manual_seed(42) + cp_size: int = 2 + total_per_rank: int = 128 + + seq_a_natural = torch.randn(100) + seq_b_natural_raw = torch.randn(61) + seq_b_padded = torch.cat([seq_b_natural_raw, torch.zeros(3)]) # pad 61→64 + + seq_a_ranks: list[torch.Tensor] = _zigzag_split_seq( + seq_a_natural, cp_size=cp_size + ) + seq_b_ranks: list[torch.Tensor] = _zigzag_split_seq( + seq_b_padded, cp_size=cp_size + ) + + # Build per-rank tensors: [seqA_r | seqB_r | pad_r] + rank_tensors: list[torch.Tensor] = [] + for rank in range(cp_size): + used: int = seq_a_ranks[rank].shape[0] + seq_b_ranks[rank].shape[0] + pad_len: int = total_per_rank - used + rank_tensor: torch.Tensor = torch.cat( + [seq_a_ranks[rank], seq_b_ranks[rank], torch.zeros(pad_len)] + ).refine_names("t") + rank_tensors.append(rank_tensor) + + # Step 1: THD unshard + seq_lens_per_rank: list[int] = [50, 32, 46] + unshard_plan = UnsharderPlan( + axis=ParallelAxis.CP, + params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=seq_lens_per_rank), + groups=[[0, 1]], + ) + with warning_sink.context(): + unsharded: list[torch.Tensor] = execute_unsharder_plan( + unshard_plan, rank_tensors + ) + assert len(unsharded) == 1 + + # Step 2: THD reorder + reorder_seq_lens: list[int] = [s * cp_size for s in seq_lens_per_rank] + reorder_plan = ReordererPlan( + params=ZigzagToNaturalThdParams( + dim_name="t", cp_size=cp_size, seq_lens=reorder_seq_lens + ) + ) + reordered: list[torch.Tensor] = execute_reorderer_plan(reorder_plan, unsharded) + assert len(reordered) == 1 + + result: torch.Tensor = reordered[0].rename(None) + assert torch.equal(result[:100], seq_a_natural) + assert torch.equal(result[100:164], seq_b_padded) + + def test_thd_cp3_single_seq(self) -> None: + """cp_size=3, single seq (120 tokens).""" + torch.manual_seed(42) + cp_size: int = 3 + seq_natural = torch.randn(120) + + seq_ranks: list[torch.Tensor] = _zigzag_split_seq(seq_natural, cp_size=cp_size) + + rank_tensors: list[torch.Tensor] = [t.refine_names("t") for t in seq_ranks] + + # Step 1: THD unshard + seq_len_per_rank: int = 120 // cp_size # 40 + unshard_plan = UnsharderPlan( + axis=ParallelAxis.CP, + params=CpThdConcatParams( + dim_name="t", seq_lens_per_rank=[seq_len_per_rank] + ), + groups=[list(range(cp_size))], + ) + with warning_sink.context(): + unsharded: list[torch.Tensor] = execute_unsharder_plan( + unshard_plan, rank_tensors + ) + assert len(unsharded) == 1 + + # Step 2: THD reorder + reorder_plan = ReordererPlan( + params=ZigzagToNaturalThdParams( + dim_name="t", cp_size=cp_size, seq_lens=[120] + ) + ) + reordered: list[torch.Tensor] = execute_reorderer_plan(reorder_plan, unsharded) + assert len(reordered) == 1 + + result: torch.Tensor = reordered[0].rename(None) + assert torch.equal(result, seq_natural) + + if __name__ == "__main__": sys.exit(pytest.main([__file__])) diff --git a/test/registered/debug_utils/comparator/aligner/reorderer/test_planner.py b/test/registered/debug_utils/comparator/aligner/reorderer/test_planner.py index fa3723fbc..e8aed9ad5 100644 --- a/test/registered/debug_utils/comparator/aligner/reorderer/test_planner.py +++ b/test/registered/debug_utils/comparator/aligner/reorderer/test_planner.py @@ -43,8 +43,8 @@ class TestComputeReordererPlans: assert plans[0].params.dim_name == "s" assert plans[0].params.cp_size == 2 - def test_compute_reorderer_plans_non_seq_dim_raises(self) -> None: - """Zigzag on non-sequence dim (e.g. t(cp,zigzag)) raises ValueError.""" + def test_compute_reorderer_plans_thd_zigzag(self) -> None: + """t(cp,zigzag) produces a ZigzagToNaturalThdParams plan.""" dim_specs = parse_dims("t(cp,zigzag) h(tp)") parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ { @@ -52,9 +52,54 @@ class TestComputeReordererPlans: ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), }, ] + thd_global_seq_lens: list[int] = [100, 64, 92] + plans = compute_reorderer_plans( + dim_specs=dim_specs, + parallel_infos=parallel_infos, + thd_global_seq_lens=thd_global_seq_lens, + ) + + assert len(plans) == 1 + assert plans[0].params.op == "zigzag_to_natural_thd" + assert plans[0].params.cp_size == 2 + assert plans[0].params.seq_lens == [100, 64, 92] + + def test_non_seq_dim_still_raises(self) -> None: + """Zigzag on non-sequence/non-token dim (e.g. h(cp,zigzag)) raises ValueError.""" + dim_specs = parse_dims("h(cp,zigzag) d") + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + {ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2)}, + ] with pytest.raises(ValueError, match="only supported on sequence dims"): compute_reorderer_plans(dim_specs=dim_specs, parallel_infos=parallel_infos) + def test_thd_zigzag_without_seq_lens_raises(self) -> None: + """t(cp,zigzag) without thd_global_seq_lens raises ValueError.""" + dim_specs = parse_dims("t(cp,zigzag) h(tp)") + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + { + ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + }, + ] + with pytest.raises(ValueError, match="thd_global_seq_lens is required"): + compute_reorderer_plans(dim_specs=dim_specs, parallel_infos=parallel_infos) + + def test_thd_natural_no_reorder(self) -> None: + """t(cp,natural) and t(cp) produce no reorder plans.""" + for dims_str in ["t(cp,natural) h(tp)", "t(cp) h(tp)"]: + dim_specs = parse_dims(dims_str) + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + { + ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + }, + ] + plans = compute_reorderer_plans( + dim_specs=dim_specs, parallel_infos=parallel_infos + ) + assert plans == [] + def test_compute_reorderer_plans_natural(self) -> None: """s(cp) and s(cp,natural) produce no reorder plans.""" for dims_str in ["b s(cp) h(tp)", "b s(cp,natural) h(tp)"]: