Make reorderer support packed format with CP in dump comparator (#19462)
This commit is contained in:
@@ -1,6 +1,12 @@
|
|||||||
|
from typing import Optional
|
||||||
|
|
||||||
import torch
|
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 (
|
from sglang.srt.debug_utils.comparator.dims import (
|
||||||
resolve_dim_by_name,
|
resolve_dim_by_name,
|
||||||
strip_dim_names,
|
strip_dim_names,
|
||||||
@@ -11,12 +17,66 @@ def execute_reorderer_plan(
|
|||||||
plan: ReordererPlan,
|
plan: ReordererPlan,
|
||||||
tensors: list[torch.Tensor],
|
tensors: list[torch.Tensor],
|
||||||
) -> list[torch.Tensor]:
|
) -> list[torch.Tensor]:
|
||||||
|
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)
|
dim: int = resolve_dim_by_name(tensors[0], plan.params.dim_name)
|
||||||
return [
|
return [
|
||||||
_reorder_zigzag_to_natural(tensor, dim=dim, cp_size=plan.params.cp_size)
|
_reorder_zigzag_to_natural(tensor, dim=dim, cp_size=plan.params.cp_size)
|
||||||
for tensor in tensors
|
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(
|
def _reorder_zigzag_to_natural(
|
||||||
tensor: torch.Tensor, *, dim: int, cp_size: int
|
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).
|
(megatron/core/ssm/mamba_context_parallel.py:360-373).
|
||||||
"""
|
"""
|
||||||
stripped: torch.Tensor = strip_dim_names(tensor)
|
stripped: torch.Tensor = strip_dim_names(tensor)
|
||||||
names: tuple = tensor.names
|
names: tuple[Optional[str], ...] = tensor.names
|
||||||
|
|
||||||
num_chunks: int = cp_size * 2
|
num_chunks: int = cp_size * 2
|
||||||
chunks: tuple[torch.Tensor, ...] = stripped.chunk(num_chunks, dim=dim)
|
chunks: tuple[torch.Tensor, ...] = stripped.chunk(num_chunks, dim=dim)
|
||||||
|
|||||||
@@ -1,21 +1,27 @@
|
|||||||
|
from typing import Optional
|
||||||
|
|
||||||
from sglang.srt.debug_utils.comparator.aligner.reorderer.types import (
|
from sglang.srt.debug_utils.comparator.aligner.reorderer.types import (
|
||||||
ReordererPlan,
|
ReordererPlan,
|
||||||
ZigzagToNaturalParams,
|
ZigzagToNaturalParams,
|
||||||
|
ZigzagToNaturalThdParams,
|
||||||
)
|
)
|
||||||
from sglang.srt.debug_utils.comparator.aligner.unsharder.types import AxisInfo
|
from sglang.srt.debug_utils.comparator.aligner.unsharder.types import AxisInfo
|
||||||
from sglang.srt.debug_utils.comparator.dims import (
|
from sglang.srt.debug_utils.comparator.dims import (
|
||||||
SEQ_DIM_NAME,
|
SEQ_DIM_NAME,
|
||||||
|
TOKEN_DIM_NAME,
|
||||||
DimSpec,
|
DimSpec,
|
||||||
Ordering,
|
Ordering,
|
||||||
ParallelAxis,
|
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(
|
def compute_reorderer_plans(
|
||||||
dim_specs: list[DimSpec],
|
dim_specs: list[DimSpec],
|
||||||
parallel_infos: list[dict[ParallelAxis, AxisInfo]],
|
parallel_infos: list[dict[ParallelAxis, AxisInfo]],
|
||||||
|
*,
|
||||||
|
thd_global_seq_lens: Optional[list[int]] = None,
|
||||||
) -> list[ReordererPlan]:
|
) -> list[ReordererPlan]:
|
||||||
plans: list[ReordererPlan] = []
|
plans: list[ReordererPlan] = []
|
||||||
|
|
||||||
@@ -28,17 +34,35 @@ def compute_reorderer_plans(
|
|||||||
if spec.name not in _ALLOWED_ZIGZAG_DIM_NAMES:
|
if spec.name not in _ALLOWED_ZIGZAG_DIM_NAMES:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Zigzag ordering is only supported on sequence dims "
|
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"{sorted(_ALLOWED_ZIGZAG_DIM_NAMES)}), "
|
||||||
f"but got dim name {spec.name!r} in {spec}"
|
f"but got dim name {spec.name!r} in {spec}"
|
||||||
)
|
)
|
||||||
|
|
||||||
assert spec.ordering == Ordering.ZIGZAG
|
if spec.ordering != Ordering.ZIGZAG:
|
||||||
axis_size: int = parallel_infos[0][spec.parallel].axis_size
|
raise ValueError(
|
||||||
plans.append(
|
f"Unsupported ordering {spec.ordering!r} for dim {spec.name!r}"
|
||||||
ReordererPlan(
|
|
||||||
params=ZigzagToNaturalParams(dim_name=spec.name, cp_size=axis_size),
|
|
||||||
)
|
)
|
||||||
|
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
|
return plans
|
||||||
|
|||||||
@@ -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
|
from sglang.srt.debug_utils.comparator.utils import _FrozenBase
|
||||||
|
|
||||||
@@ -9,7 +11,17 @@ class ZigzagToNaturalParams(_FrozenBase):
|
|||||||
cp_size: int
|
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):
|
class ReordererPlan(_FrozenBase):
|
||||||
|
|||||||
@@ -5,12 +5,49 @@ import torch
|
|||||||
|
|
||||||
from sglang.srt.debug_utils.comparator.aligner.reorderer.executor import (
|
from sglang.srt.debug_utils.comparator.aligner.reorderer.executor import (
|
||||||
_reorder_zigzag_to_natural,
|
_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
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
register_cpu_ci(est_time=10, suite="default", nightly=True)
|
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:
|
class TestZigzagToNatural:
|
||||||
def test_zigzag_to_natural_cp2(self) -> None:
|
def test_zigzag_to_natural_cp2(self) -> None:
|
||||||
"""cp_size=2: zigzag order [0,3,1,2] -> natural [0,1,2,3]."""
|
"""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)
|
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__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__]))
|
sys.exit(pytest.main([__file__]))
|
||||||
|
|||||||
@@ -43,8 +43,8 @@ class TestComputeReordererPlans:
|
|||||||
assert plans[0].params.dim_name == "s"
|
assert plans[0].params.dim_name == "s"
|
||||||
assert plans[0].params.cp_size == 2
|
assert plans[0].params.cp_size == 2
|
||||||
|
|
||||||
def test_compute_reorderer_plans_non_seq_dim_raises(self) -> None:
|
def test_compute_reorderer_plans_thd_zigzag(self) -> None:
|
||||||
"""Zigzag on non-sequence dim (e.g. t(cp,zigzag)) raises ValueError."""
|
"""t(cp,zigzag) produces a ZigzagToNaturalThdParams plan."""
|
||||||
dim_specs = parse_dims("t(cp,zigzag) h(tp)")
|
dim_specs = parse_dims("t(cp,zigzag) h(tp)")
|
||||||
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
|
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
|
||||||
{
|
{
|
||||||
@@ -52,9 +52,54 @@ class TestComputeReordererPlans:
|
|||||||
ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2),
|
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"):
|
with pytest.raises(ValueError, match="only supported on sequence dims"):
|
||||||
compute_reorderer_plans(dim_specs=dim_specs, parallel_infos=parallel_infos)
|
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:
|
def test_compute_reorderer_plans_natural(self) -> None:
|
||||||
"""s(cp) and s(cp,natural) produce no reorder plans."""
|
"""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)"]:
|
for dims_str in ["b s(cp) h(tp)", "b s(cp,natural) h(tp)"]:
|
||||||
|
|||||||
Reference in New Issue
Block a user