Validate replicated axes orthogonality in dump comparator (#21026)

This commit is contained in:
fzyzcjy
2026-03-20 22:02:40 +08:00
committed by GitHub
parent 2f01950a0e
commit cc22601d28
3 changed files with 261 additions and 0 deletions

View File

@@ -138,6 +138,11 @@ def _validate_explicit_replicated(
f"Axes {{{conflict_names}}} declared as both sharded and replicated"
)
_validate_replicated_axes_orthogonal(
explicit_replicated_axes=explicit_replicated_axes,
parallel_infos=parallel_infos,
)
candidate_axes: set[ParallelAxis] = (
all_axes - sharded_axes - explicit_replicated_axes
)
@@ -178,6 +183,33 @@ def _validate_explicit_replicated(
)
def _validate_replicated_axes_orthogonal(
*,
explicit_replicated_axes: frozenset[ParallelAxis],
parallel_infos: list[dict[ParallelAxis, AxisInfo]],
) -> None:
"""Every pair of explicitly replicated axes must be fully orthogonal (no dependency)."""
axes: list[ParallelAxis] = sorted(explicit_replicated_axes, key=lambda a: a.value)
if len(axes) < 2:
return
violations: list[str] = []
for i, axis_a in enumerate(axes):
for axis_b in axes[i + 1 :]:
for parent, child in [(axis_a, axis_b), (axis_b, axis_a)]:
if _is_dependent_axis(parallel_infos, parent=parent, child=child):
violations.append(
f"'{parent.value}' determines '{child.value}'"
f"remove '{child.value}:replicated'"
)
if violations:
details = "; ".join(violations)
raise ValueError(
f"Explicitly-replicated axes overlap (not orthogonal): {details}"
)
def _validate(
*,
axes_to_validate: set[ParallelAxis],

View File

@@ -5,6 +5,7 @@ from sglang.srt.debug_utils.comparator.aligner.unsharder.planner import (
_is_dependent_axis,
_is_jointly_determined,
_validate_explicit_replicated,
_validate_replicated_axes_orthogonal,
compute_unsharder_plan,
)
from sglang.srt.debug_utils.comparator.aligner.unsharder.types import (
@@ -829,6 +830,41 @@ class TestAxisContainment:
dim_specs, parallel_infos, explicit_replicated_axes=replicated
)
def test_backward_compat_explicit_children(self) -> None:
"""Both tp:replicated and attn_tp:replicated → ValueError (not orthogonal)."""
dim_specs = parse_dims(
"t h # tp:replicated attn_tp:replicated moe_tp:replicated"
).dims
replicated = frozenset(
{ParallelAxis.TP, ParallelAxis.ATTN_TP, ParallelAxis.MOE_TP}
)
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=4),
ParallelAxis.ATTN_TP: AxisInfo(axis_rank=0, axis_size=2),
ParallelAxis.MOE_TP: AxisInfo(axis_rank=0, axis_size=2),
},
{
ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=4),
ParallelAxis.ATTN_TP: AxisInfo(axis_rank=1, axis_size=2),
ParallelAxis.MOE_TP: AxisInfo(axis_rank=0, axis_size=2),
},
{
ParallelAxis.TP: AxisInfo(axis_rank=2, axis_size=4),
ParallelAxis.ATTN_TP: AxisInfo(axis_rank=0, axis_size=2),
ParallelAxis.MOE_TP: AxisInfo(axis_rank=1, axis_size=2),
},
{
ParallelAxis.TP: AxisInfo(axis_rank=3, axis_size=4),
ParallelAxis.ATTN_TP: AxisInfo(axis_rank=1, axis_size=2),
ParallelAxis.MOE_TP: AxisInfo(axis_rank=1, axis_size=2),
},
]
with pytest.raises(ValueError, match="not orthogonal"):
compute_unsharder_plan(
dim_specs, parallel_infos, explicit_replicated_axes=replicated
)
class TestDpFilteredAxis:
"""Tests for dp_filtered_axis parameter: DP axis handled by upstream DP filter
@@ -1875,3 +1911,154 @@ class TestIsJointlyDetermined:
parent_axes=frozenset({ParallelAxis.TP}),
child=ParallelAxis.EDP,
)
class TestReplicatedAxesOrthogonality:
"""Tests for _validate_replicated_axes_orthogonal: every pair of explicitly
replicated axes must be fully orthogonal (no dependency relationship)."""
def test_tp_determines_moe_tp_raises(self) -> None:
"""TP4 + MOE_TP2 where tp_rank determines moe_tp_rank → ValueError."""
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=4),
ParallelAxis.MOE_TP: AxisInfo(axis_rank=0, axis_size=2),
},
{
ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=4),
ParallelAxis.MOE_TP: AxisInfo(axis_rank=1, axis_size=2),
},
{
ParallelAxis.TP: AxisInfo(axis_rank=2, axis_size=4),
ParallelAxis.MOE_TP: AxisInfo(axis_rank=0, axis_size=2),
},
{
ParallelAxis.TP: AxisInfo(axis_rank=3, axis_size=4),
ParallelAxis.MOE_TP: AxisInfo(axis_rank=1, axis_size=2),
},
]
with pytest.raises(ValueError, match="not orthogonal"):
_validate_replicated_axes_orthogonal(
explicit_replicated_axes=frozenset(
{ParallelAxis.TP, ParallelAxis.MOE_TP}
),
parallel_infos=parallel_infos,
)
def test_tp_determines_sp_identical_group_raises(self) -> None:
"""TP2 + SP2 where sp_rank == tp_rank → ValueError."""
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2),
ParallelAxis.SP: AxisInfo(axis_rank=0, axis_size=2),
},
{
ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2),
ParallelAxis.SP: AxisInfo(axis_rank=1, axis_size=2),
},
]
with pytest.raises(ValueError, match="not orthogonal"):
_validate_replicated_axes_orthogonal(
explicit_replicated_axes=frozenset({ParallelAxis.TP, ParallelAxis.SP}),
parallel_infos=parallel_infos,
)
def test_three_axes_two_overlapping_pairs_raises(self) -> None:
"""TP4 + ATTN_TP2 + MOE_TP2, TP determines both → error mentions two pairs."""
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=4),
ParallelAxis.ATTN_TP: AxisInfo(axis_rank=0, axis_size=2),
ParallelAxis.MOE_TP: AxisInfo(axis_rank=0, axis_size=2),
},
{
ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=4),
ParallelAxis.ATTN_TP: AxisInfo(axis_rank=1, axis_size=2),
ParallelAxis.MOE_TP: AxisInfo(axis_rank=0, axis_size=2),
},
{
ParallelAxis.TP: AxisInfo(axis_rank=2, axis_size=4),
ParallelAxis.ATTN_TP: AxisInfo(axis_rank=0, axis_size=2),
ParallelAxis.MOE_TP: AxisInfo(axis_rank=1, axis_size=2),
},
{
ParallelAxis.TP: AxisInfo(axis_rank=3, axis_size=4),
ParallelAxis.ATTN_TP: AxisInfo(axis_rank=1, axis_size=2),
ParallelAxis.MOE_TP: AxisInfo(axis_rank=1, axis_size=2),
},
]
with pytest.raises(ValueError, match="not orthogonal") as exc_info:
_validate_replicated_axes_orthogonal(
explicit_replicated_axes=frozenset(
{ParallelAxis.TP, ParallelAxis.ATTN_TP, ParallelAxis.MOE_TP}
),
parallel_infos=parallel_infos,
)
msg = str(exc_info.value)
assert "attn_tp" in msg
assert "moe_tp" in msg
def test_three_axes_one_overlap_one_orthogonal_raises(self) -> None:
"""TP4 + MOE_TP2 (dependent) + CP2 (independent) → only tp/moe_tp pair errors."""
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
for cp_rank in range(2):
for tp_rank in range(4):
parallel_infos.append(
{
ParallelAxis.TP: AxisInfo(axis_rank=tp_rank, axis_size=4),
ParallelAxis.MOE_TP: AxisInfo(
axis_rank=tp_rank % 2, axis_size=2
),
ParallelAxis.CP: AxisInfo(axis_rank=cp_rank, axis_size=2),
}
)
with pytest.raises(ValueError, match="not orthogonal") as exc_info:
_validate_replicated_axes_orthogonal(
explicit_replicated_axes=frozenset(
{ParallelAxis.TP, ParallelAxis.MOE_TP, ParallelAxis.CP}
),
parallel_infos=parallel_infos,
)
msg = str(exc_info.value)
assert "moe_tp" in msg
assert "cp" not in msg
def test_single_replicated_axis_no_check(self) -> None:
"""Only one replicated axis → no orthogonality check needed, passes."""
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2)},
{ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2)},
]
_validate_replicated_axes_orthogonal(
explicit_replicated_axes=frozenset({ParallelAxis.TP}),
parallel_infos=parallel_infos,
)
def test_two_independent_axes_ok(self) -> None:
"""TP2 + CP2 fully orthogonal → no error."""
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2),
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
},
{
ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2),
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
},
{
ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2),
ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2),
},
{
ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2),
ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2),
},
]
_validate_replicated_axes_orthogonal(
explicit_replicated_axes=frozenset({ParallelAxis.TP, ParallelAxis.CP}),
parallel_infos=parallel_infos,
)
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))

View File

@@ -1935,6 +1935,48 @@ class TestEntrypointReplicatedAxis:
assert isinstance(summary, SummaryRecord)
assert summary.failed == 1
def test_dependent_replicated_axes_error(self, tmp_path, capsys):
"""TP4 + MOE_TP2 both replicated, tp determines moe_tp → ComparisonErrorRecord."""
torch.manual_seed(42)
tensor = torch.randn(4, 8)
baseline_dir = tmp_path / "baseline"
target_dir = tmp_path / "target"
# TP4 with MOE_TP2: tp_rank determines moe_tp_rank (rank%2)
for side_dir in [baseline_dir, target_dir]:
for tp_rank in range(4):
_create_rank_dump(
side_dir,
rank=tp_rank,
name="gate_out",
tensor=tensor,
dims="b h # tp:replicated moe_tp:replicated",
parallel_info={
"tp_rank": tp_rank,
"tp_size": 4,
"moe_tp_rank": tp_rank % 2,
"moe_tp_size": 2,
},
)
argv = _make_argv(
baseline_dir / _FIXED_EXP_NAME,
target_dir / _FIXED_EXP_NAME,
diff_threshold=0.01,
)
records, exit_code = _run_and_parse(argv, capsys)
errors = [r for r in records if isinstance(r, ComparisonErrorRecord)]
assert len(errors) == 1
assert "not orthogonal" in errors[0].traceback_str
summary = records[-1]
assert isinstance(summary, SummaryRecord)
assert summary.errored == 1
assert exit_code == 1
def test_sharded_tp_with_dependent_etp_passes(self, tmp_path, capsys):
"""TP2 sharded + ETP2 dependent (etp=tp) + EP2 replicated → no undeclared error."""
torch.manual_seed(42)