Support multi sharding group on the same dimension in dump comparator (#19601)

This commit is contained in:
fzyzcjy
2026-03-01 10:36:48 +08:00
committed by GitHub
parent 46960e65cf
commit ea6ff7b01f
14 changed files with 469 additions and 151 deletions
@@ -26,11 +26,10 @@ def compute_reorderer_plans(
plans: list[ReordererPlan] = [] plans: list[ReordererPlan] = []
for spec in dim_specs: for spec in dim_specs:
if ( for modifier in spec.parallel_modifiers:
spec.ordering is not None if modifier.ordering is None or modifier.ordering == Ordering.NATURAL:
and spec.ordering != Ordering.NATURAL continue
and spec.parallel is not None
):
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 "
@@ -39,11 +38,11 @@ def compute_reorderer_plans(
f"but got dim name {spec.name!r} in {spec}" f"but got dim name {spec.name!r} in {spec}"
) )
if spec.ordering != Ordering.ZIGZAG: if modifier.ordering != Ordering.ZIGZAG:
raise ValueError( raise ValueError(
f"Unsupported ordering {spec.ordering!r} for dim {spec.name!r}" f"Unsupported ordering {modifier.ordering!r} for dim {spec.name!r}"
) )
axis_size: int = parallel_infos[0][spec.parallel].axis_size axis_size: int = parallel_infos[0][modifier.axis].axis_size
if spec.name == TOKEN_DIM_NAME: if spec.name == TOKEN_DIM_NAME:
if thd_global_seq_lens is None: if thd_global_seq_lens is None:
@@ -122,7 +122,7 @@ class _SGLangPlugin(_AuxFrameworkPlugin):
will be mishandled. Callers should set dims explicitly for non-zigzag CP. will be mishandled. Callers should set dims explicitly for non-zigzag CP.
""" """
if ndim == 1: if ndim == 1:
return "t(cp,zigzag)" return "t(cp:zigzag)"
raise ValueError( raise ValueError(
f"SGLang: cannot infer dims for CP-sharded '{name}' with ndim={ndim}" f"SGLang: cannot infer dims for CP-sharded '{name}' with ndim={ndim}"
) )
@@ -208,9 +208,9 @@ class _MegatronPlugin(_AuxFrameworkPlugin):
will be mishandled. Callers should set dims explicitly for non-zigzag CP. will be mishandled. Callers should set dims explicitly for non-zigzag CP.
""" """
if ndim == 1: if ndim == 1:
return "t(cp,zigzag)" return "t(cp:zigzag)"
if ndim == 2: if ndim == 2:
return "b s(cp,zigzag)" return "b s(cp:zigzag)"
raise ValueError( raise ValueError(
f"Megatron: cannot infer dims for CP-sharded '{name}' with ndim={ndim}" f"Megatron: cannot infer dims for CP-sharded '{name}' with ndim={ndim}"
) )
@@ -14,6 +14,7 @@ from sglang.srt.debug_utils.comparator.dims import (
TOKEN_DIM_NAME, TOKEN_DIM_NAME,
DimSpec, DimSpec,
ParallelAxis, ParallelAxis,
ParallelModifier,
) )
# _CoordsList[tensor_index][axis] = # _CoordsList[tensor_index][axis] =
@@ -36,18 +37,21 @@ def compute_unsharder_plan(
if not parallel_infos: if not parallel_infos:
raise ValueError("parallel_infos must not be empty") raise ValueError("parallel_infos must not be empty")
sharded_axis_infos: dict[ParallelAxis, DimSpec] = { # Within each dim spec, reverse modifier order: innermost shard (rightmost) unshards first.
spec.parallel: spec for spec in dim_specs if spec.parallel is not None reversed_sharded_modifiers: list[tuple[str, ParallelModifier]] = [
} (spec.name, m) for spec in dim_specs for m in reversed(spec.parallel_modifiers)
sharded_axes_raw: set[ParallelAxis] = set(sharded_axis_infos) ]
sharded_axes_raw: set[ParallelAxis] = {
m.axis for _, m in reversed_sharded_modifiers
}
all_axes: set[ParallelAxis] = {axis for info in parallel_infos for axis in info} all_axes: set[ParallelAxis] = {axis for info in parallel_infos for axis in info}
# axis annotated in dims but absent from all parallel_infos -> axis_size=1, skip # axis annotated in dims but absent from all parallel_infos -> axis_size=1, skip
sharded_axes: set[ParallelAxis] = sharded_axes_raw & all_axes sharded_axes: set[ParallelAxis] = sharded_axes_raw & all_axes
sharded_axis_infos = { reversed_sharded_modifiers = [
k: v for k, v in sharded_axis_infos.items() if k in sharded_axes (name, m) for name, m in reversed_sharded_modifiers if m.axis in sharded_axes
} ]
replicated_axes: set[ParallelAxis] = all_axes - sharded_axes replicated_axes: set[ParallelAxis] = all_axes - sharded_axes
if not sharded_axes and not replicated_axes: if not sharded_axes and not replicated_axes:
@@ -67,14 +71,15 @@ def compute_unsharder_plan(
(axis, PickParams()) for axis in sorted(replicated_axes, key=lambda a: a.value) (axis, PickParams()) for axis in sorted(replicated_axes, key=lambda a: a.value)
] + [ ] + [
( (
axis, modifier.axis,
_resolve_unshard_params( _resolve_unshard_params(
spec=spec, modifier=modifier,
dim_name=dim_name,
parallel_infos=parallel_infos, parallel_infos=parallel_infos,
thd_global_seq_lens=thd_global_seq_lens, thd_global_seq_lens=thd_global_seq_lens,
), ),
) )
for axis, spec in sharded_axis_infos.items() for dim_name, modifier in reversed_sharded_modifiers
] ]
plans: list[UnsharderPlan] = [] plans: list[UnsharderPlan] = []
@@ -151,23 +156,20 @@ def _group_and_project(
def _resolve_unshard_params( def _resolve_unshard_params(
*, *,
spec: DimSpec, modifier: ParallelModifier,
dim_name: str,
parallel_infos: list[dict[ParallelAxis, AxisInfo]], parallel_infos: list[dict[ParallelAxis, AxisInfo]],
thd_global_seq_lens: Optional[list[int]] = None, thd_global_seq_lens: Optional[list[int]] = None,
) -> UnsharderParams: ) -> UnsharderParams:
if spec.reduction is not None: if modifier.reduction is not None:
return ReduceSumParams() return ReduceSumParams()
if ( if (
spec.name == TOKEN_DIM_NAME dim_name == TOKEN_DIM_NAME
and spec.parallel == ParallelAxis.CP and modifier.axis == ParallelAxis.CP
and thd_global_seq_lens is not None and thd_global_seq_lens is not None
): ):
if spec.parallel is None: axis_size: int = parallel_infos[0][modifier.axis].axis_size
raise ValueError(
f"THD unshard requires a parallel axis on dim '{spec.name}', but got None"
)
axis_size: int = parallel_infos[0][spec.parallel].axis_size
for s in thd_global_seq_lens: for s in thd_global_seq_lens:
if s % axis_size != 0: if s % axis_size != 0:
raise ValueError( raise ValueError(
@@ -175,8 +177,6 @@ def _resolve_unshard_params(
f"Sequences must be padded to a multiple of cp_size for CP zigzag." f"Sequences must be padded to a multiple of cp_size for CP zigzag."
) )
seq_lens_per_rank: list[int] = [s // axis_size for s in thd_global_seq_lens] seq_lens_per_rank: list[int] = [s // axis_size for s in thd_global_seq_lens]
return CpThdConcatParams( return CpThdConcatParams(dim_name=dim_name, seq_lens_per_rank=seq_lens_per_rank)
dim_name=spec.name, seq_lens_per_rank=seq_lens_per_rank
)
return ConcatParams(dim_name=spec.name) return ConcatParams(dim_name=dim_name)
@@ -1,10 +1,11 @@
import re import re
from dataclasses import dataclass
from enum import Enum from enum import Enum
from typing import Optional from typing import Optional
import torch import torch
from sglang.srt.debug_utils.comparator.utils import _FrozenBase
TOKEN_DIM_NAME: str = "t" TOKEN_DIM_NAME: str = "t"
BATCH_DIM_NAME: str = "b" BATCH_DIM_NAME: str = "b"
SEQ_DIM_NAME: str = "s" SEQ_DIM_NAME: str = "s"
@@ -33,14 +34,17 @@ class Reduction(Enum):
PARTIAL = "partial" PARTIAL = "partial"
@dataclass(frozen=True) class ParallelModifier(_FrozenBase):
class DimSpec: axis: ParallelAxis
name: str
parallel: Optional[ParallelAxis] = None
ordering: Optional[Ordering] = None ordering: Optional[Ordering] = None
reduction: Optional[Reduction] = None reduction: Optional[Reduction] = None
class DimSpec(_FrozenBase):
name: str
parallel_modifiers: list[ParallelModifier] = []
class _SingletonDimUtil: class _SingletonDimUtil:
"""Utilities for squeeze dims (name="1") and their singleton tensor-name mapping.""" """Utilities for squeeze dims (name="1") and their singleton tensor-name mapping."""
@@ -83,16 +87,60 @@ class _SingletonDimUtil:
_DIM_PATTERN = re.compile(r"^(?P<name>[a-zA-Z_]\w*)(?:\((?P<modifiers>[^)]+)\))?$") _DIM_PATTERN = re.compile(r"^(?P<name>[a-zA-Z_]\w*)(?:\((?P<modifiers>[^)]+)\))?$")
_MODIFIER_FIELDS: list[tuple[type[Enum], str]] = [ _AXIS_LOOKUP: dict[str, ParallelAxis] = {m.value: m for m in ParallelAxis}
(ParallelAxis, "parallel"), _QUALIFIER_LOOKUP: dict[str, Ordering | Reduction] = {
(Ordering, "ordering"), **{m.value: m for m in Ordering},
(Reduction, "reduction"), **{m.value: m for m in Reduction},
] }
_MODIFIER_LOOKUP: dict[str, tuple[str, Enum]] = {}
for _enum_cls, _field in _MODIFIER_FIELDS: def _parse_modifier_token(modifier_token: str, dim_token: str) -> ParallelModifier:
for _member in _enum_cls: """Parse 'sp', 'cp:zigzag', 'tp:partial', or 'cp:zigzag+partial' → ParallelModifier.
_MODIFIER_LOOKUP[_member.value] = (_field, _member)
Format: ``axis`` or ``axis:qual`` or ``axis:qual+qual``.
Colon separates axis from qualifiers; ``+`` separates multiple qualifiers.
"""
axis_str: str
qualifiers_str: str
if ":" in modifier_token:
axis_str, qualifiers_str = modifier_token.split(":", maxsplit=1)
else:
axis_str, qualifiers_str = modifier_token, ""
axis_str = axis_str.strip()
axis: Optional[ParallelAxis] = _AXIS_LOOKUP.get(axis_str)
if axis is None:
raise ValueError(
f"Unknown axis {axis_str!r} in modifier {modifier_token!r} "
f"of dim spec: {dim_token!r}"
)
ordering: Optional[Ordering] = None
reduction: Optional[Reduction] = None
for q_str in (q.strip() for q in qualifiers_str.split("+") if q.strip()):
qualifier: Optional[Ordering | Reduction] = _QUALIFIER_LOOKUP.get(q_str)
if qualifier is None:
raise ValueError(
f"Unknown qualifier {q_str!r} in modifier "
f"{modifier_token!r} of dim spec: {dim_token!r}"
)
if isinstance(qualifier, Ordering):
if ordering is not None:
raise ValueError(
f"Multiple ordering values in modifier "
f"{modifier_token!r} of dim spec: {dim_token!r}"
)
ordering = qualifier
else:
if reduction is not None:
raise ValueError(
f"Multiple reduction values in modifier "
f"{modifier_token!r} of dim spec: {dim_token!r}"
)
reduction = qualifier
return ParallelModifier(axis=axis, ordering=ordering, reduction=reduction)
def parse_dim(token: str) -> DimSpec: def parse_dim(token: str) -> DimSpec:
@@ -103,26 +151,29 @@ def parse_dim(token: str) -> DimSpec:
if match is None: if match is None:
raise ValueError(f"Invalid dim token: {token!r}") raise ValueError(f"Invalid dim token: {token!r}")
name = match.group("name") name: str = match.group("name")
modifiers_str = match.group("modifiers") modifiers_str: Optional[str] = match.group("modifiers")
if modifiers_str is None: if modifiers_str is None:
return DimSpec(name=name) return DimSpec(name=name)
fields: dict[str, Enum] = {} modifiers: list[ParallelModifier] = []
for part in (p.strip() for p in modifiers_str.split(",")): seen_axes: set[ParallelAxis] = set()
if part not in _MODIFIER_LOOKUP:
raise ValueError(f"Unknown modifier {part!r} in dim spec: {token!r}")
field_name, enum_value = _MODIFIER_LOOKUP[part]
if field_name in fields:
raise ValueError(f"Multiple {field_name} values in dim token: {token!r}")
fields[field_name] = enum_value
return DimSpec(name=name, **fields) for modifier_token in (p.strip() for p in modifiers_str.split(",")):
modifier: ParallelModifier = _parse_modifier_token(modifier_token, token)
if modifier.axis in seen_axes:
raise ValueError(
f"Duplicate axis {modifier.axis.value!r} in dim spec: {token!r}"
)
seen_axes.add(modifier.axis)
modifiers.append(modifier)
return DimSpec(name=name, parallel_modifiers=modifiers)
def parse_dims(dims_str: str) -> list[DimSpec]: def parse_dims(dims_str: str) -> list[DimSpec]:
"""Parse 'b s(cp,zigzag) h(tp) d' -> list[DimSpec].""" """Parse 'b s(cp:zigzag) h(tp) d' -> list[DimSpec]."""
if not dims_str.strip(): if not dims_str.strip():
raise ValueError("dims string must not be empty") raise ValueError("dims string must not be empty")
@@ -36,9 +36,6 @@ from sglang.srt.debug_utils.comparator.output_types import (
from sglang.srt.debug_utils.comparator.per_token_visualizer import ( from sglang.srt.debug_utils.comparator.per_token_visualizer import (
generate_per_token_heatmap, generate_per_token_heatmap,
) )
from sglang.srt.debug_utils.comparator.per_token_visualizer import (
generate_per_token_heatmap,
)
from sglang.srt.debug_utils.comparator.utils import Pair from sglang.srt.debug_utils.comparator.utils import Pair
from sglang.srt.debug_utils.dump_loader import read_meta, read_tokenizer_path from sglang.srt.debug_utils.dump_loader import read_meta, read_tokenizer_path
@@ -258,12 +255,6 @@ def _consume_comparison_records(
return summary, skipped_names return summary, skipped_names
if visualize_per_token is not None and collected_comparisons:
generate_per_token_heatmap(
records=collected_comparisons,
output_path=visualize_per_token,
)
def _parse_args() -> argparse.Namespace: def _parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser() parser = argparse.ArgumentParser()
@@ -85,8 +85,8 @@ class TestComputePerStepSubPlans:
def test_zigzag_returns_both_plans(self) -> None: def test_zigzag_returns_both_plans(self) -> None:
result: list[AlignerPerStepSubPlan] = compute_per_step_sub_plans( result: list[AlignerPerStepSubPlan] = compute_per_step_sub_plans(
metas=[ metas=[
_make_meta(dims="b s(cp,zigzag) h", cp_rank=0, cp_size=2), _make_meta(dims="b s(cp:zigzag) h", cp_rank=0, cp_size=2),
_make_meta(dims="b s(cp,zigzag) h", cp_rank=1, cp_size=2), _make_meta(dims="b s(cp:zigzag) h", cp_rank=1, cp_size=2),
] ]
) )
unsharder_plans: list[UnsharderPlan] = [ unsharder_plans: list[UnsharderPlan] = [
@@ -177,33 +177,33 @@ class TestComputeAlignerPlan:
class TestComputePerStepSubPlansThd: class TestComputePerStepSubPlansThd:
def test_thd_zigzag_returns_thd_plans(self) -> None: def test_thd_zigzag_returns_thd_plans(self) -> None:
"""t(cp,zigzag) h(tp) generates THD-typed unsharder + reorderer plans.""" """t(cp:zigzag) h(tp) generates THD-typed unsharder + reorderer plans."""
thd_global_seq_lens: list[int] = [100, 64, 92] thd_global_seq_lens: list[int] = [100, 64, 92]
result: list[AlignerPerStepSubPlan] = compute_per_step_sub_plans( result: list[AlignerPerStepSubPlan] = compute_per_step_sub_plans(
metas=[ metas=[
_make_meta( _make_meta(
dims="t(cp,zigzag) h(tp)", dims="t(cp:zigzag) h(tp)",
cp_rank=0, cp_rank=0,
cp_size=2, cp_size=2,
tp_rank=0, tp_rank=0,
tp_size=2, tp_size=2,
), ),
_make_meta( _make_meta(
dims="t(cp,zigzag) h(tp)", dims="t(cp:zigzag) h(tp)",
cp_rank=0, cp_rank=0,
cp_size=2, cp_size=2,
tp_rank=1, tp_rank=1,
tp_size=2, tp_size=2,
), ),
_make_meta( _make_meta(
dims="t(cp,zigzag) h(tp)", dims="t(cp:zigzag) h(tp)",
cp_rank=1, cp_rank=1,
cp_size=2, cp_size=2,
tp_rank=0, tp_rank=0,
tp_size=2, tp_size=2,
), ),
_make_meta( _make_meta(
dims="t(cp,zigzag) h(tp)", dims="t(cp:zigzag) h(tp)",
cp_rank=1, cp_rank=1,
cp_size=2, cp_size=2,
tp_rank=1, tp_rank=1,
@@ -25,8 +25,8 @@ register_cpu_ci(est_time=10, suite="default", nightly=True)
class TestComputeReordererPlans: class TestComputeReordererPlans:
def test_compute_reorderer_plans_zigzag(self) -> None: def test_compute_reorderer_plans_zigzag(self) -> None:
"""s(cp,zigzag) produces a ReordererPlan.""" """s(cp:zigzag) produces a ReordererPlan."""
dim_specs = parse_dims("b s(cp,zigzag) h(tp)") dim_specs = parse_dims("b s(cp:zigzag) h(tp)")
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{ {
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2), ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
@@ -43,8 +43,8 @@ class TestComputeReordererPlans:
assert plans[0].params.cp_size == 2 assert plans[0].params.cp_size == 2
def test_compute_reorderer_plans_thd_zigzag(self) -> None: def test_compute_reorderer_plans_thd_zigzag(self) -> None:
"""t(cp,zigzag) produces a ZigzagToNaturalThdParams plan.""" """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]] = [
{ {
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2), ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
@@ -64,8 +64,8 @@ class TestComputeReordererPlans:
assert plans[0].params.seq_lens == [100, 64, 92] assert plans[0].params.seq_lens == [100, 64, 92]
def test_non_seq_dim_still_raises(self) -> None: def test_non_seq_dim_still_raises(self) -> None:
"""Zigzag on non-sequence/non-token dim (e.g. h(cp,zigzag)) raises ValueError.""" """Zigzag on non-sequence/non-token dim (e.g. h(cp:zigzag)) raises ValueError."""
dim_specs = parse_dims("h(cp,zigzag) d") dim_specs = parse_dims("h(cp:zigzag) d")
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2)}, {ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2)},
] ]
@@ -73,8 +73,8 @@ class TestComputeReordererPlans:
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: def test_thd_zigzag_without_seq_lens_raises(self) -> None:
"""t(cp,zigzag) without thd_global_seq_lens raises ValueError.""" """t(cp:zigzag) without thd_global_seq_lens raises ValueError."""
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]] = [
{ {
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2), ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
@@ -85,8 +85,8 @@ class TestComputeReordererPlans:
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_natural_no_reorder(self) -> None: def test_thd_natural_no_reorder(self) -> None:
"""t(cp,natural) and t(cp) produce no reorder plans.""" """t(cp:natural) and t(cp) produce no reorder plans."""
for dims_str in ["t(cp,natural) h(tp)", "t(cp) h(tp)"]: for dims_str in ["t(cp:natural) h(tp)", "t(cp) h(tp)"]:
dim_specs = parse_dims(dims_str) dim_specs = parse_dims(dims_str)
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{ {
@@ -100,8 +100,8 @@ class TestComputeReordererPlans:
assert plans == [] 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)"]:
dim_specs = parse_dims(dims_str) dim_specs = parse_dims(dims_str)
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{ {
@@ -141,7 +141,7 @@ class TestCpZigzagTpE2E:
} }
) )
dim_specs: list[DimSpec] = parse_dims("b s(cp,zigzag) h(tp)") dim_specs: list[DimSpec] = parse_dims("b s(cp:zigzag) h(tp)")
dim_names: list[str] = [s.name for s in dim_specs] dim_names: list[str] = [s.name for s in dim_specs]
unsharder_plans = compute_unsharder_plan( unsharder_plans = compute_unsharder_plan(
@@ -166,5 +166,83 @@ class TestCpZigzagTpE2E:
assert torch.allclose(current[0].rename(None), full_tensor) assert torch.allclose(current[0].rename(None), full_tensor)
class TestCpZigzagSpSameDimE2E:
"""E2E test for t(cp:zigzag,sp) — two axes sharding the same token dim."""
def test_cp2_sp2_zigzag_e2e(self) -> None:
"""CP=2 zigzag + SP=2 on same token dim: full unshard + reorder round-trip.
Shard order (outer to inner, matching left-to-right in dims annotation):
1. CP zigzag splits token dim into 2 CP chunks (zigzag order)
2. SP splits each CP chunk into 2 SP chunks
Unshard order (inner to outer, right-to-left):
1. SP concat (inner): merge SP chunks back
2. CP concat (outer): merge CP chunks back
3. Zigzag reorder: restore natural token order
"""
torch.manual_seed(42)
total_tokens: int = 16
hidden: int = 8
full_tensor: torch.Tensor = torch.randn(total_tokens, hidden)
# Step 1: CP zigzag split — split into 2*cp_size=4 natural chunks, reorder by zigzag
cp_size: int = 2
sp_size: int = 2
n_natural_chunks: int = cp_size * 2
natural_chunks: list[torch.Tensor] = list(
full_tensor.chunk(n_natural_chunks, dim=0)
)
zigzag_order: list[int] = [0, 3, 1, 2]
zigzagged: torch.Tensor = torch.cat(
[natural_chunks[i] for i in zigzag_order], dim=0
)
cp_chunks: list[torch.Tensor] = list(zigzagged.chunk(cp_size, dim=0))
# Step 2: SP split within each CP chunk
tensors: list[torch.Tensor] = []
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
for cp_rank in range(cp_size):
sp_chunks: list[torch.Tensor] = list(
cp_chunks[cp_rank].chunk(sp_size, dim=0)
)
for sp_rank in range(sp_size):
tensors.append(sp_chunks[sp_rank])
parallel_infos.append(
{
ParallelAxis.CP: AxisInfo(axis_rank=cp_rank, axis_size=cp_size),
ParallelAxis.SP: AxisInfo(axis_rank=sp_rank, axis_size=sp_size),
}
)
dim_specs: list[DimSpec] = parse_dims("t(cp:zigzag,sp) h")
dim_names: list[str] = [s.name for s in dim_specs]
unsharder_plans = compute_unsharder_plan(
dim_specs=dim_specs, parallel_infos=parallel_infos
)
reorderer_plans = compute_reorderer_plans(
dim_specs=dim_specs,
parallel_infos=parallel_infos,
thd_global_seq_lens=[total_tokens],
)
all_plans = [*unsharder_plans, *reorderer_plans]
assert len(unsharder_plans) == 2 # SP concat, CP concat
assert unsharder_plans[0].axis == ParallelAxis.SP
assert unsharder_plans[1].axis == ParallelAxis.CP
assert len(reorderer_plans) == 1 # zigzag reorder
current: list[torch.Tensor] = [t.refine_names(*dim_names) for t in tensors]
for plan in all_plans:
if isinstance(plan, ReordererPlan):
current = execute_reorderer_plan(plan, current)
else:
current = execute_unsharder_plan(plan, current).tensors
assert len(current) == 1
assert torch.allclose(current[0].rename(None), full_tensor)
if __name__ == "__main__": if __name__ == "__main__":
sys.exit(pytest.main([__file__])) sys.exit(pytest.main([__file__]))
@@ -90,7 +90,7 @@ class TestEnsureDimsInMetas:
assert result is metas assert result is metas
def test_cp_sharded_sglang_input_ids_infers_dims(self): def test_cp_sharded_sglang_input_ids_infers_dims(self):
"""CP + input_ids in sglang infers dims 't(cp,zigzag)'.""" """CP + input_ids in sglang infers dims 't(cp:zigzag)'."""
metas: list[dict] = [ metas: list[dict] = [
self._make_meta(cp_size=2, cp_rank=0), self._make_meta(cp_size=2, cp_rank=0),
self._make_meta(cp_size=2, cp_rank=1), self._make_meta(cp_size=2, cp_rank=1),
@@ -99,11 +99,11 @@ class TestEnsureDimsInMetas:
name="input_ids", plugin=_sglang_plugin, metas=metas, ndim=1 name="input_ids", plugin=_sglang_plugin, metas=metas, ndim=1
) )
assert result is not metas assert result is not metas
assert result[0]["dims"] == "t(cp,zigzag)" assert result[0]["dims"] == "t(cp:zigzag)"
assert result[1]["dims"] == "t(cp,zigzag)" assert result[1]["dims"] == "t(cp:zigzag)"
def test_cp_sharded_sglang_positions_infers_dims(self): def test_cp_sharded_sglang_positions_infers_dims(self):
"""CP + positions in sglang infers dims 't(cp,zigzag)'.""" """CP + positions in sglang infers dims 't(cp:zigzag)'."""
metas: list[dict] = [ metas: list[dict] = [
self._make_meta(cp_size=2, cp_rank=0), self._make_meta(cp_size=2, cp_rank=0),
self._make_meta(cp_size=2, cp_rank=1), self._make_meta(cp_size=2, cp_rank=1),
@@ -111,10 +111,10 @@ class TestEnsureDimsInMetas:
result = _ensure_dims_in_metas( result = _ensure_dims_in_metas(
name="positions", plugin=_sglang_plugin, metas=metas, ndim=1 name="positions", plugin=_sglang_plugin, metas=metas, ndim=1
) )
assert result[0]["dims"] == "t(cp,zigzag)" assert result[0]["dims"] == "t(cp:zigzag)"
def test_cp_sharded_megatron_input_ids_infers_dims_1d(self): def test_cp_sharded_megatron_input_ids_infers_dims_1d(self):
"""CP + input_ids in megatron (1D) infers dims 't(cp,zigzag)'.""" """CP + input_ids in megatron (1D) infers dims 't(cp:zigzag)'."""
metas: list[dict] = [ metas: list[dict] = [
{"megatron_parallel_info": {"cp_rank": 0, "cp_size": 2}}, {"megatron_parallel_info": {"cp_rank": 0, "cp_size": 2}},
{"megatron_parallel_info": {"cp_rank": 1, "cp_size": 2}}, {"megatron_parallel_info": {"cp_rank": 1, "cp_size": 2}},
@@ -122,10 +122,10 @@ class TestEnsureDimsInMetas:
result = _ensure_dims_in_metas( result = _ensure_dims_in_metas(
name="input_ids", plugin=_megatron_plugin, metas=metas, ndim=1 name="input_ids", plugin=_megatron_plugin, metas=metas, ndim=1
) )
assert result[0]["dims"] == "t(cp,zigzag)" assert result[0]["dims"] == "t(cp:zigzag)"
def test_cp_sharded_megatron_input_ids_infers_dims_2d(self): def test_cp_sharded_megatron_input_ids_infers_dims_2d(self):
"""CP + input_ids in megatron (2D) infers dims 'b s(cp,zigzag)'.""" """CP + input_ids in megatron (2D) infers dims 'b s(cp:zigzag)'."""
metas: list[dict] = [ metas: list[dict] = [
{"megatron_parallel_info": {"cp_rank": 0, "cp_size": 2}}, {"megatron_parallel_info": {"cp_rank": 0, "cp_size": 2}},
{"megatron_parallel_info": {"cp_rank": 1, "cp_size": 2}}, {"megatron_parallel_info": {"cp_rank": 1, "cp_size": 2}},
@@ -133,7 +133,7 @@ class TestEnsureDimsInMetas:
result = _ensure_dims_in_metas( result = _ensure_dims_in_metas(
name="input_ids", plugin=_megatron_plugin, metas=metas, ndim=2 name="input_ids", plugin=_megatron_plugin, metas=metas, ndim=2
) )
assert result[0]["dims"] == "b s(cp,zigzag)" assert result[0]["dims"] == "b s(cp:zigzag)"
def test_cp_non_sharded_name_returns_metas_unchanged(self): def test_cp_non_sharded_name_returns_metas_unchanged(self):
"""CP + non-sharded tensor name (seq_lens) returns metas as-is.""" """CP + non-sharded tensor name (seq_lens) returns metas as-is."""
@@ -218,19 +218,19 @@ class TestInferCpShardedDims:
"""Tests for infer_cp_sharded_dims on each plugin.""" """Tests for infer_cp_sharded_dims on each plugin."""
def test_megatron_infer_1d(self) -> None: def test_megatron_infer_1d(self) -> None:
"""Megatron 1D → 't(cp,zigzag)'.""" """Megatron 1D → 't(cp:zigzag)'."""
result: str = _megatron_plugin.infer_cp_sharded_dims(name="input_ids", ndim=1) result: str = _megatron_plugin.infer_cp_sharded_dims(name="input_ids", ndim=1)
assert result == "t(cp,zigzag)" assert result == "t(cp:zigzag)"
def test_megatron_infer_2d(self) -> None: def test_megatron_infer_2d(self) -> None:
"""Megatron 2D → 'b s(cp,zigzag)'.""" """Megatron 2D → 'b s(cp:zigzag)'."""
result: str = _megatron_plugin.infer_cp_sharded_dims(name="input_ids", ndim=2) result: str = _megatron_plugin.infer_cp_sharded_dims(name="input_ids", ndim=2)
assert result == "b s(cp,zigzag)" assert result == "b s(cp:zigzag)"
def test_sglang_infer_1d(self) -> None: def test_sglang_infer_1d(self) -> None:
"""SGLang 1D → 't(cp,zigzag)'.""" """SGLang 1D → 't(cp:zigzag)'."""
result: str = _sglang_plugin.infer_cp_sharded_dims(name="input_ids", ndim=1) result: str = _sglang_plugin.infer_cp_sharded_dims(name="input_ids", ndim=1)
assert result == "t(cp,zigzag)" assert result == "t(cp:zigzag)"
def test_megatron_infer_3d_raises(self) -> None: def test_megatron_infer_3d_raises(self) -> None:
"""Megatron 3D raises ValueError.""" """Megatron 3D raises ValueError."""
@@ -654,7 +654,7 @@ class TestReduceSum:
part_a = full_tensor * 0.6 part_a = full_tensor * 0.6
part_b = full_tensor * 0.4 part_b = full_tensor * 0.4
dim_specs = parse_dims("h(tp,partial) d") dim_specs = parse_dims("h(tp:partial) d")
parallel_infos = [ parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2) {ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2)
] ]
@@ -676,7 +676,7 @@ class TestReduceSum:
full_tensor = torch.randn(4, 8) full_tensor = torch.randn(4, 8)
parts: list[torch.Tensor] = [full_tensor * 0.25 for _ in range(4)] parts: list[torch.Tensor] = [full_tensor * 0.25 for _ in range(4)]
dim_specs = parse_dims("h(tp,partial) d") dim_specs = parse_dims("h(tp:partial) d")
parallel_infos = [ parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=4)} for i in range(4) {ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=4)} for i in range(4)
] ]
@@ -710,7 +710,7 @@ class TestReduceSum:
} }
) )
dim_specs = parse_dims("b s(cp) h(tp,partial)") dim_specs = parse_dims("b s(cp) h(tp:partial)")
plans = compute_unsharder_plan(dim_specs, parallel_infos) plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 2 assert len(plans) == 2
@@ -739,7 +739,7 @@ class TestReduceSum:
{ParallelAxis.TP: AxisInfo(axis_rank=3, axis_size=4)}, {ParallelAxis.TP: AxisInfo(axis_rank=3, axis_size=4)},
{ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=4)}, {ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=4)},
] ]
dim_specs = parse_dims("h(tp,partial) d") dim_specs = parse_dims("h(tp:partial) d")
plans = compute_unsharder_plan(dim_specs, parallel_infos) plans = compute_unsharder_plan(dim_specs, parallel_infos)
named_parts: list[torch.Tensor] = _name_tensors(parts, dim_specs) named_parts: list[torch.Tensor] = _name_tensors(parts, dim_specs)
@@ -752,7 +752,7 @@ class TestReduceSum:
def test_reduce_preserves_named_dims(self) -> None: def test_reduce_preserves_named_dims(self) -> None:
"""Named tensor dimensions are preserved through reduce_sum.""" """Named tensor dimensions are preserved through reduce_sum."""
dim_specs = parse_dims("h(tp,partial) d") dim_specs = parse_dims("h(tp:partial) d")
part_a = torch.randn(4, 8).refine_names("h", "d") part_a = torch.randn(4, 8).refine_names("h", "d")
part_b = torch.randn(4, 8).refine_names("h", "d") part_b = torch.randn(4, 8).refine_names("h", "d")
@@ -175,7 +175,7 @@ class TestComputeUnsharderPlan:
compute_unsharder_plan(dim_specs, parallel_infos) compute_unsharder_plan(dim_specs, parallel_infos)
def test_reduction_partial_returns_reduce_sum(self) -> None: def test_reduction_partial_returns_reduce_sum(self) -> None:
dim_specs = parse_dims("h(tp,partial)") dim_specs = parse_dims("h(tp:partial)")
parallel_infos = [ parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2) {ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2)
] ]
@@ -188,7 +188,7 @@ class TestComputeUnsharderPlan:
def test_reduction_partial_tp4(self) -> None: def test_reduction_partial_tp4(self) -> None:
"""TP=4 with partial reduction produces a single ReduceSumParams step.""" """TP=4 with partial reduction produces a single ReduceSumParams step."""
dim_specs = parse_dims("h(tp,partial)") dim_specs = parse_dims("h(tp:partial)")
parallel_infos = [ parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=4)} for i in range(4) {ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=4)} for i in range(4)
] ]
@@ -200,7 +200,7 @@ class TestComputeUnsharderPlan:
def test_multi_axis_with_reduction_on_one(self) -> None: def test_multi_axis_with_reduction_on_one(self) -> None:
"""CP concat + TP reduce produces a 2-step plan.""" """CP concat + TP reduce produces a 2-step plan."""
dim_specs = parse_dims("s(cp) h(tp,partial)") dim_specs = parse_dims("s(cp) h(tp:partial)")
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [] parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
for cp_rank in range(2): for cp_rank in range(2):
for tp_rank in range(2): for tp_rank in range(2):
@@ -221,7 +221,7 @@ class TestComputeUnsharderPlan:
def test_reduction_scrambled_ranks(self) -> None: def test_reduction_scrambled_ranks(self) -> None:
"""Scrambled world_rank order with partial reduction.""" """Scrambled world_rank order with partial reduction."""
dim_specs = parse_dims("h(tp,partial)") dim_specs = parse_dims("h(tp:partial)")
parallel_infos = [ parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=2, axis_size=4)}, {ParallelAxis.TP: AxisInfo(axis_rank=2, axis_size=4)},
{ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=4)}, {ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=4)},
@@ -235,7 +235,7 @@ class TestComputeUnsharderPlan:
assert plans[0].groups == [[1, 3, 0, 2]] assert plans[0].groups == [[1, 3, 0, 2]]
def test_ordering_zigzag_accepted(self) -> None: def test_ordering_zigzag_accepted(self) -> None:
dim_specs = parse_dims("s(cp,zigzag)") dim_specs = parse_dims("s(cp:zigzag)")
parallel_infos = [ parallel_infos = [
{ParallelAxis.CP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2) {ParallelAxis.CP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2)
] ]
@@ -244,7 +244,7 @@ class TestComputeUnsharderPlan:
assert plans[0].axis == ParallelAxis.CP assert plans[0].axis == ParallelAxis.CP
def test_ordering_natural_accepted(self) -> None: def test_ordering_natural_accepted(self) -> None:
dim_specs = parse_dims("s(cp,natural)") dim_specs = parse_dims("s(cp:natural)")
parallel_infos = [ parallel_infos = [
{ParallelAxis.CP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2) {ParallelAxis.CP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2)
] ]
@@ -288,6 +288,77 @@ class TestComputeUnsharderPlan:
assert len(plans[2].groups) == 1 assert len(plans[2].groups) == 1
assert len(plans[2].groups[0]) == 2 assert len(plans[2].groups[0]) == 2
def test_same_dim_cp_sp_plan(self) -> None:
"""t(cp:zigzag,sp) with CP=2 SP=2: SP unshards first (inner), then CP."""
dim_specs = parse_dims("t(cp:zigzag,sp) 1 h")
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
for cp_rank in range(2):
for sp_rank in range(2):
parallel_infos.append(
{
ParallelAxis.CP: AxisInfo(axis_rank=cp_rank, axis_size=2),
ParallelAxis.SP: AxisInfo(axis_rank=sp_rank, axis_size=2),
}
)
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 2
# SP unshards first (rightmost modifier = innermost shard)
sp_plan = plans[0]
assert sp_plan.axis == ParallelAxis.SP
assert isinstance(sp_plan.params, ConcatParams)
assert sp_plan.params.dim_name == "t"
assert len(sp_plan.groups) == 2
for group in sp_plan.groups:
assert len(group) == 2
# CP unshards second (leftmost modifier = outermost shard)
cp_plan = plans[1]
assert cp_plan.axis == ParallelAxis.CP
assert isinstance(cp_plan.params, ConcatParams)
assert cp_plan.params.dim_name == "t"
assert len(cp_plan.groups) == 1
assert len(cp_plan.groups[0]) == 2
def test_same_dim_cp_sp_with_thd(self) -> None:
"""t(cp:zigzag,sp) with THD: SP → ConcatParams, CP → CpThdConcatParams."""
from sglang.srt.debug_utils.comparator.aligner.unsharder.types import (
CpThdConcatParams,
)
dim_specs = parse_dims("t(cp:zigzag,sp) h")
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
for cp_rank in range(2):
for sp_rank in range(2):
parallel_infos.append(
{
ParallelAxis.CP: AxisInfo(axis_rank=cp_rank, axis_size=2),
ParallelAxis.SP: AxisInfo(axis_rank=sp_rank, axis_size=2),
}
)
thd_global_seq_lens: list[int] = [100, 64]
plans = compute_unsharder_plan(
dim_specs, parallel_infos, thd_global_seq_lens=thd_global_seq_lens
)
assert len(plans) == 2
# SP unshards first: plain concat (SP is not CP, no THD special handling)
sp_plan = plans[0]
assert sp_plan.axis == ParallelAxis.SP
assert isinstance(sp_plan.params, ConcatParams)
assert sp_plan.params.dim_name == "t"
# CP unshards second: THD concat because dim is 't' + axis is CP + thd_global_seq_lens provided
cp_plan = plans[1]
assert cp_plan.axis == ParallelAxis.CP
assert isinstance(cp_plan.params, CpThdConcatParams)
assert cp_plan.params.dim_name == "t"
assert cp_plan.params.seq_lens_per_rank == [50, 32]
def test_sp_in_dims_but_not_in_parallel_info(self) -> None: def test_sp_in_dims_but_not_in_parallel_info(self) -> None:
"""s(sp) in dims but SP absent from parallel_info (SP disabled), should auto-skip.""" """s(sp) in dims but SP absent from parallel_info (SP disabled), should auto-skip."""
dim_specs = parse_dims("s(sp) b h(tp)") dim_specs = parse_dims("s(sp) b h(tp)")
@@ -11,6 +11,7 @@ from sglang.srt.debug_utils.comparator.dims import (
DimSpec, DimSpec,
Ordering, Ordering,
ParallelAxis, ParallelAxis,
ParallelModifier,
Reduction, Reduction,
_SingletonDimUtil, _SingletonDimUtil,
apply_dim_names, apply_dim_names,
@@ -32,48 +33,78 @@ class TestParseDim:
assert parse_dim("b") == DimSpec(name="b") assert parse_dim("b") == DimSpec(name="b")
def test_parallel_axis(self) -> None: def test_parallel_axis(self) -> None:
assert parse_dim("h(tp)") == DimSpec(name="h", parallel=ParallelAxis.TP) assert parse_dim("h(tp)") == DimSpec(
name="h",
parallel_modifiers=[ParallelModifier(axis=ParallelAxis.TP)],
)
def test_all_parallel_axes(self) -> None: def test_all_parallel_axes(self) -> None:
assert parse_dim("a(tp)").parallel == ParallelAxis.TP assert parse_dim("a(tp)").parallel_modifiers[0].axis == ParallelAxis.TP
assert parse_dim("a(cp)").parallel == ParallelAxis.CP assert parse_dim("a(cp)").parallel_modifiers[0].axis == ParallelAxis.CP
assert parse_dim("a(ep)").parallel == ParallelAxis.EP assert parse_dim("a(ep)").parallel_modifiers[0].axis == ParallelAxis.EP
assert parse_dim("a(sp)").parallel == ParallelAxis.SP assert parse_dim("a(sp)").parallel_modifiers[0].axis == ParallelAxis.SP
def test_ordering(self) -> None: def test_ordering(self) -> None:
assert parse_dim("s(cp,zigzag)").ordering == Ordering.ZIGZAG assert (
assert parse_dim("s(cp,natural)").ordering == Ordering.NATURAL parse_dim("s(cp:zigzag)").parallel_modifiers[0].ordering == Ordering.ZIGZAG
)
assert (
parse_dim("s(cp:natural)").parallel_modifiers[0].ordering
== Ordering.NATURAL
)
def test_reduction(self) -> None: def test_reduction(self) -> None:
assert parse_dim("h(tp,partial)").reduction == Reduction.PARTIAL assert (
parse_dim("h(tp:partial)").parallel_modifiers[0].reduction
== Reduction.PARTIAL
)
def test_all_modifiers(self) -> None: def test_all_qualifiers(self) -> None:
assert parse_dim("s(cp,zigzag,partial)") == DimSpec( assert parse_dim("s(cp:zigzag+partial)") == DimSpec(
name="s", name="s",
parallel=ParallelAxis.CP, parallel_modifiers=[
ParallelModifier(
axis=ParallelAxis.CP,
ordering=Ordering.ZIGZAG, ordering=Ordering.ZIGZAG,
reduction=Reduction.PARTIAL, reduction=Reduction.PARTIAL,
),
],
) )
def test_multi_axis(self) -> None:
result: DimSpec = parse_dim("t(cp:zigzag,sp)")
assert result.name == "t"
assert len(result.parallel_modifiers) == 2
assert result.parallel_modifiers[0] == ParallelModifier(
axis=ParallelAxis.CP, ordering=Ordering.ZIGZAG
)
assert result.parallel_modifiers[1] == ParallelModifier(axis=ParallelAxis.SP)
def test_invalid_token_raises(self) -> None: def test_invalid_token_raises(self) -> None:
with pytest.raises(ValueError, match="Invalid dim token"): with pytest.raises(ValueError, match="Invalid dim token"):
parse_dim("h()") parse_dim("h()")
with pytest.raises(ValueError, match="Invalid dim token"): with pytest.raises(ValueError, match="Invalid dim token"):
parse_dim("h(tp(x))") parse_dim("h(tp(x))")
def test_unknown_modifier_raises(self) -> None: def test_unknown_axis_raises(self) -> None:
with pytest.raises(ValueError, match="Unknown modifier"): with pytest.raises(ValueError, match="Unknown axis"):
parse_dim("h(xyz)") parse_dim("h(xyz)")
with pytest.raises(ValueError, match="Unknown modifier"):
parse_dim("h(tp,foobar)") def test_unknown_qualifier_raises(self) -> None:
with pytest.raises(ValueError, match="Unknown qualifier"):
parse_dim("h(tp:foobar)")
def test_multiple_ordering_raises(self) -> None: def test_multiple_ordering_raises(self) -> None:
with pytest.raises(ValueError, match="Multiple ordering"): with pytest.raises(ValueError, match="Multiple ordering"):
parse_dim("s(cp,zigzag,natural)") parse_dim("s(cp:zigzag+natural)")
def test_multiple_reduction_raises(self) -> None: def test_multiple_reduction_raises(self) -> None:
with pytest.raises(ValueError, match="Multiple reduction"): with pytest.raises(ValueError, match="Multiple reduction"):
parse_dim("h(tp,partial,partial)") parse_dim("h(tp:partial+partial)")
def test_duplicate_axis_raises(self) -> None:
with pytest.raises(ValueError, match="Duplicate axis"):
parse_dim("h(tp,tp)")
def test_squeeze_dim(self) -> None: def test_squeeze_dim(self) -> None:
assert parse_dim("1") == DimSpec(name="1") assert parse_dim("1") == DimSpec(name="1")
@@ -96,10 +127,18 @@ class TestParseDims:
assert parse_dims("t") == [DimSpec(name="t")] assert parse_dims("t") == [DimSpec(name="t")]
def test_mixed_annotated(self) -> None: def test_mixed_annotated(self) -> None:
assert parse_dims("b s(cp,zigzag) h(tp) d") == [ assert parse_dims("b s(cp:zigzag) h(tp) d") == [
DimSpec(name="b"), DimSpec(name="b"),
DimSpec(name="s", parallel=ParallelAxis.CP, ordering=Ordering.ZIGZAG), DimSpec(
DimSpec(name="h", parallel=ParallelAxis.TP), name="s",
parallel_modifiers=[
ParallelModifier(axis=ParallelAxis.CP, ordering=Ordering.ZIGZAG),
],
),
DimSpec(
name="h",
parallel_modifiers=[ParallelModifier(axis=ParallelAxis.TP)],
),
DimSpec(name="d"), DimSpec(name="d"),
] ]
@@ -158,7 +197,7 @@ class TestFindDimIndex:
assert find_dim_index(specs, "d") == 3 assert find_dim_index(specs, "d") == 3
def test_with_modifiers(self) -> None: def test_with_modifiers(self) -> None:
specs: list[DimSpec] = parse_dims("b s(cp,zigzag) h(tp) d") specs: list[DimSpec] = parse_dims("b s(cp:zigzag) h(tp) d")
assert find_dim_index(specs, "h") == 2 assert find_dim_index(specs, "h") == 2
def test_empty_list(self) -> None: def test_empty_list(self) -> None:
@@ -837,7 +837,7 @@ class TestEntrypointGroupingLogical:
tp_size=1, tp_size=1,
seq_dim=1, seq_dim=1,
head_dim=2, head_dim=2,
dims_str="b s(cp,zigzag) h", dims_str="b s(cp:zigzag) h",
) )
args = _make_args( args = _make_args(
@@ -871,7 +871,7 @@ class TestEntrypointGroupingLogical:
tp_size=2, tp_size=2,
seq_dim=1, seq_dim=1,
head_dim=2, head_dim=2,
dims_str="b s(cp,zigzag) h(tp)", dims_str="b s(cp:zigzag) h(tp)",
) )
args = _make_args( args = _make_args(
@@ -959,14 +959,14 @@ class TestEntrypointGroupingLogical:
full_tensor=full_baseline, full_tensor=full_baseline,
name="attn_out", name="attn_out",
tp_size=2, tp_size=2,
dims_str="b h(tp,partial)", dims_str="b h(tp:partial)",
) )
target_path = _create_tp_partial_dumps( target_path = _create_tp_partial_dumps(
target_dir, target_dir,
full_tensor=full_target, full_tensor=full_target,
name="attn_out", name="attn_out",
tp_size=2, tp_size=2,
dims_str="b h(tp,partial)", dims_str="b h(tp:partial)",
) )
args = _make_args(baseline_path, target_path, diff_threshold=0.01) args = _make_args(baseline_path, target_path, diff_threshold=0.01)
@@ -997,7 +997,7 @@ class TestEntrypointGroupingLogical:
full_tensor=target_full, full_tensor=target_full,
name="attn_out", name="attn_out",
tp_size=2, tp_size=2,
dims_str="b h(tp,partial)", dims_str="b h(tp:partial)",
) )
args = _make_args(baseline_path, target_path, diff_threshold=0.01) args = _make_args(baseline_path, target_path, diff_threshold=0.01)
@@ -1026,7 +1026,7 @@ class TestEntrypointGroupingLogical:
rank=rank, rank=rank,
name="hidden", name="hidden",
tensor=cp_chunks[cp_rank] / 2, tensor=cp_chunks[cp_rank] / 2,
dims="b s(cp) h(tp,partial)", dims="b s(cp) h(tp:partial)",
parallel_info={ parallel_info={
"cp_rank": cp_rank, "cp_rank": cp_rank,
"cp_size": 2, "cp_size": 2,
@@ -1046,6 +1046,38 @@ class TestEntrypointGroupingLogical:
comp = _assert_single_comparison_passed(records) comp = _assert_single_comparison_passed(records)
assert comp.name == "hidden" assert comp.name == "hidden"
def test_cp_zigzag_sp_same_dim_unshard(self, tmp_path, capsys):
"""CP=2 zigzag + SP=2 on same seq dim: multi-axis unshard + reorder."""
torch.manual_seed(42)
full_baseline = torch.randn(4, 8, 6)
full_target = full_baseline + torch.randn(4, 8, 6) * 0.001
baseline_dir = tmp_path / "baseline"
target_dir = tmp_path / "target"
for side_dir, full_tensor in [
(baseline_dir, full_baseline),
(target_dir, full_target),
]:
_create_cp_zigzag_sp_sharded_dumps(
side_dir,
full_tensor=full_tensor,
name="hidden",
cp_size=2,
sp_size=2,
dims_str="b s(cp:zigzag,sp) h",
)
args = _make_args(
baseline_dir / _FIXED_EXP_NAME,
target_dir / _FIXED_EXP_NAME,
diff_threshold=0.01,
)
records, _ = _run_and_parse(args, capsys)
comp = _assert_single_comparison_passed(records)
assert comp.name == "hidden"
class TestEntrypointConcatMode: class TestEntrypointConcatMode:
"""Test concat token-aligner mode through the full entrypoint pipeline.""" """Test concat token-aligner mode through the full entrypoint pipeline."""
@@ -2578,6 +2610,63 @@ def _create_cp_zigzag_tp_sharded_dumps(
return directory / _FIXED_EXP_NAME return directory / _FIXED_EXP_NAME
def _create_cp_zigzag_sp_sharded_dumps(
directory: Path,
*,
full_tensor: torch.Tensor,
name: str,
cp_size: int,
sp_size: int,
dims_str: str,
seq_dim: int = 1,
num_steps: int = 1,
) -> Path:
"""Create CP-zigzag + SP sharded dump files for a seq dim (b s h format).
Shard order (outer to inner, matching left-to-right in dims annotation):
1. CP zigzag splits seq dim into cp_size chunks (zigzag order)
2. SP splits each CP chunk into sp_size chunks
"""
num_chunks: int = cp_size * 2
natural_chunks: list[torch.Tensor] = list(
full_tensor.chunk(num_chunks, dim=seq_dim)
)
zigzag_order: list[int] = []
for i in range(cp_size):
zigzag_order.append(i)
zigzag_order.append(num_chunks - 1 - i)
zigzagged: torch.Tensor = torch.cat(
[natural_chunks[idx] for idx in zigzag_order], dim=seq_dim
)
cp_chunks: list[torch.Tensor] = list(zigzagged.chunk(cp_size, dim=seq_dim))
rank: int = 0
for cp_rank in range(cp_size):
sp_chunks: list[torch.Tensor] = list(
cp_chunks[cp_rank].chunk(sp_size, dim=seq_dim)
)
for sp_rank in range(sp_size):
_create_rank_dump(
directory,
rank=rank,
name=name,
tensor=sp_chunks[sp_rank],
dims=dims_str,
parallel_info={
"cp_rank": cp_rank,
"cp_size": cp_size,
"sp_rank": sp_rank,
"sp_size": sp_size,
},
num_steps=num_steps,
)
rank += 1
return directory / _FIXED_EXP_NAME
def _create_replicated_tp_sharded_cp_dumps( def _create_replicated_tp_sharded_cp_dumps(
directory: Path, directory: Path,
*, *,
@@ -2772,7 +2861,7 @@ def _create_thd_cp_zigzag_dumps(
seq_lens: list[int], seq_lens: list[int],
cp_size: int, cp_size: int,
total_per_rank: int, total_per_rank: int,
dims_str: str = "t(cp,zigzag)", dims_str: str = "t(cp:zigzag)",
num_steps: int = 1, num_steps: int = 1,
) -> Path: ) -> Path:
"""Create THD CP-zigzag sharded dump files simulating Megatron forward. """Create THD CP-zigzag sharded dump files simulating Megatron forward.
@@ -2981,7 +3070,7 @@ class TestEntrypointThdCpZigzag:
rank=cp_rank, rank=cp_rank,
name="hidden_states", name="hidden_states",
tensor=rank_hidden, tensor=rank_hidden,
dims="t(cp,zigzag) h", dims="t(cp:zigzag) h",
parallel_info={"cp_rank": cp_rank, "cp_size": cp_size}, parallel_info={"cp_rank": cp_rank, "cp_size": cp_size},
framework="megatron", framework="megatron",
extra_dumps=[ extra_dumps=[
+3 -3
View File
@@ -2348,7 +2348,7 @@ class TestDumperDims:
) )
tensor = torch.randn(4, 8, requires_grad=True) tensor = torch.randn(4, 8, requires_grad=True)
dumper.dump("hidden", tensor, dims="b h(tp)", dims_grad="b h(tp,partial)") dumper.dump("hidden", tensor, dims="b h(tp)", dims_grad="b h(tp:partial)")
dumper.step() dumper.step()
tensor.backward(torch.ones_like(tensor)) tensor.backward(torch.ones_like(tensor))
@@ -2362,10 +2362,10 @@ class TestDumperDims:
value_data = torch.load(value_file, weights_only=False) value_data = torch.load(value_file, weights_only=False)
assert value_data["meta"]["dims"] == "b h(tp)" assert value_data["meta"]["dims"] == "b h(tp)"
assert value_data["meta"]["dims_grad"] == "b h(tp,partial)" assert value_data["meta"]["dims_grad"] == "b h(tp:partial)"
grad_data = torch.load(grad_file, weights_only=False) grad_data = torch.load(grad_file, weights_only=False)
assert grad_data["meta"]["dims"] == "b h(tp,partial)" assert grad_data["meta"]["dims"] == "b h(tp:partial)"
def test_dims_grad_inherits(self, tmp_path) -> None: def test_dims_grad_inherits(self, tmp_path) -> None:
dumper = _Dumper( dumper = _Dumper(