From cc22601d28c2e8965443c63c29458af0e21f5a37 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Fri, 20 Mar 2026 22:02:40 +0800 Subject: [PATCH] Validate replicated axes orthogonality in dump comparator (#21026) --- .../comparator/aligner/unsharder/planner.py | 32 +++ .../aligner/unsharder/test_planner.py | 187 ++++++++++++++++++ .../debug_utils/comparator/test_entrypoint.py | 42 ++++ 3 files changed, 261 insertions(+) diff --git a/python/sglang/srt/debug_utils/comparator/aligner/unsharder/planner.py b/python/sglang/srt/debug_utils/comparator/aligner/unsharder/planner.py index 27c789b17..e98596aa6 100644 --- a/python/sglang/srt/debug_utils/comparator/aligner/unsharder/planner.py +++ b/python/sglang/srt/debug_utils/comparator/aligner/unsharder/planner.py @@ -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], diff --git a/test/registered/debug_utils/comparator/aligner/unsharder/test_planner.py b/test/registered/debug_utils/comparator/aligner/unsharder/test_planner.py index 7c449e0af..3db3ff90c 100644 --- a/test/registered/debug_utils/comparator/aligner/unsharder/test_planner.py +++ b/test/registered/debug_utils/comparator/aligner/unsharder/test_planner.py @@ -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__])) diff --git a/test/registered/debug_utils/comparator/test_entrypoint.py b/test/registered/debug_utils/comparator/test_entrypoint.py index 3e7c09459..cd8a0d4ae 100644 --- a/test/registered/debug_utils/comparator/test_entrypoint.py +++ b/test/registered/debug_utils/comparator/test_entrypoint.py @@ -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)