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"
|
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] = (
|
candidate_axes: set[ParallelAxis] = (
|
||||||
all_axes - sharded_axes - explicit_replicated_axes
|
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(
|
def _validate(
|
||||||
*,
|
*,
|
||||||
axes_to_validate: set[ParallelAxis],
|
axes_to_validate: set[ParallelAxis],
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ from sglang.srt.debug_utils.comparator.aligner.unsharder.planner import (
|
|||||||
_is_dependent_axis,
|
_is_dependent_axis,
|
||||||
_is_jointly_determined,
|
_is_jointly_determined,
|
||||||
_validate_explicit_replicated,
|
_validate_explicit_replicated,
|
||||||
|
_validate_replicated_axes_orthogonal,
|
||||||
compute_unsharder_plan,
|
compute_unsharder_plan,
|
||||||
)
|
)
|
||||||
from sglang.srt.debug_utils.comparator.aligner.unsharder.types import (
|
from sglang.srt.debug_utils.comparator.aligner.unsharder.types import (
|
||||||
@@ -829,6 +830,41 @@ class TestAxisContainment:
|
|||||||
dim_specs, parallel_infos, explicit_replicated_axes=replicated
|
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:
|
class TestDpFilteredAxis:
|
||||||
"""Tests for dp_filtered_axis parameter: DP axis handled by upstream DP filter
|
"""Tests for dp_filtered_axis parameter: DP axis handled by upstream DP filter
|
||||||
@@ -1875,3 +1911,154 @@ class TestIsJointlyDetermined:
|
|||||||
parent_axes=frozenset({ParallelAxis.TP}),
|
parent_axes=frozenset({ParallelAxis.TP}),
|
||||||
child=ParallelAxis.EDP,
|
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 isinstance(summary, SummaryRecord)
|
||||||
assert summary.failed == 1
|
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):
|
def test_sharded_tp_with_dependent_etp_passes(self, tmp_path, capsys):
|
||||||
"""TP2 sharded + ETP2 dependent (etp=tp) + EP2 replicated → no undeclared error."""
|
"""TP2 sharded + ETP2 dependent (etp=tp) + EP2 replicated → no undeclared error."""
|
||||||
torch.manual_seed(42)
|
torch.manual_seed(42)
|
||||||
|
|||||||
Reference in New Issue
Block a user