Validate replicated axes orthogonality in dump comparator (#21026)
This commit is contained in:
@@ -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],
|
||||
|
||||
@@ -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__]))
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user