From 2f01950a0e947e6d7e8c53a3ed783ef3f6949141 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Fri, 20 Mar 2026 22:01:26 +0800 Subject: [PATCH] Support jointly-determined axes inference in dump comparator (#21025) --- .../comparator/aligner/unsharder/planner.py | 50 ++ .../aligner/unsharder/test_planner.py | 522 ++++++++++++++++++ .../debug_utils/comparator/test_display.py | 1 + .../debug_utils/comparator/test_entrypoint.py | 173 ++++++ 4 files changed, 746 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 382a536be..27c789b17 100644 --- a/python/sglang/srt/debug_utils/comparator/aligner/unsharder/planner.py +++ b/python/sglang/srt/debug_utils/comparator/aligner/unsharder/planner.py @@ -161,6 +161,15 @@ def _validate_explicit_replicated( ) undeclared: set[ParallelAxis] = all_axes - declared_axes + jointly_determined: frozenset[ParallelAxis] = frozenset( + child + for child in undeclared + if _is_jointly_determined( + parallel_infos, parent_axes=declared_axes, child=child + ) + ) + undeclared -= jointly_determined + if undeclared: undeclared_names: str = ", ".join(sorted(a.value for a in undeclared)) raise ValueError( @@ -238,6 +247,47 @@ def _is_dependent_axis( return True +def _is_jointly_determined( + parallel_infos: list[dict[ParallelAxis, AxisInfo]], + *, + parent_axes: frozenset[ParallelAxis], + child: ParallelAxis, +) -> bool: + """True if child's rank is uniquely determined by the joint tuple of parent ranks. + + Unlike ``_is_dependent_axis`` which checks single-parent dependency, this + checks whether the *combination* of all parent axes jointly determines the + child. For example, ``edp_rank`` may not be a function of ``tp_rank`` alone + or ``cp_rank`` alone, but it *is* a function of ``(tp_rank, cp_rank)``. + + Parent axes that are absent from *every* info are ignored (they carry no + information — e.g. DP with size 1 filtered by ``normalize_parallel_info``). + However, a parent axis present in *some* infos but missing from an info + that contains the child makes the determination incomplete → ``False``. + """ + if not parent_axes: + return False + + active_parents: frozenset[ParallelAxis] = frozenset( + ax for ax in parent_axes if any(ax in info for info in parallel_infos) + ) + if not active_parents: + return False + + mapping: dict[frozenset, int] = {} + for info in parallel_infos: + if child not in info: + continue + if not active_parents.issubset(info): + return False + parent_key = frozenset((ax, info[ax].axis_rank) for ax in active_parents) + child_rank: int = info[child].axis_rank + if mapping.setdefault(parent_key, child_rank) != child_rank: + return False + + return bool(mapping) + + def _group_and_project( *, current_coords: _CoordsList, 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 01d4f0d6b..7c449e0af 100644 --- a/test/registered/debug_utils/comparator/aligner/unsharder/test_planner.py +++ b/test/registered/debug_utils/comparator/aligner/unsharder/test_planner.py @@ -3,6 +3,7 @@ import pytest from sglang.srt.debug_utils.comparator.aligner.unsharder.planner import ( _compute_dependent_axes, _is_dependent_axis, + _is_jointly_determined, _validate_explicit_replicated, compute_unsharder_plan, ) @@ -418,6 +419,48 @@ class TestComputeUnsharderPlan: assert ParallelAxis.EP in axes_in_plan assert ParallelAxis.ETP not in axes_in_plan + def test_edp_jointly_determined_by_tp_and_cp(self) -> None: + """dims=t[cp:zigzag,sp] h # tp:replicated, EDP determined by (TP,CP) jointly → plan succeeds. + + Simulates tp=2, cp=2, ep=1, etp=1 on 4 GPUs. + """ + dim_specs = parse_dims("t[cp:zigzag,sp] h # tp:replicated").dims + 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.SP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=0, axis_size=4), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.SP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=1, axis_size=4), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.SP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=2, axis_size=4), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.SP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=3, axis_size=4), + }, + ] + plans = compute_unsharder_plan( + dim_specs, + parallel_infos, + explicit_replicated_axes=frozenset({ParallelAxis.TP}), + ) + axes_in_plan = [p.axis for p in plans] + assert ParallelAxis.CP in axes_in_plan + assert ParallelAxis.TP in axes_in_plan + assert ParallelAxis.EDP not in axes_in_plan + class TestExplicitReplicatedAxes: def test_replicated_tp_with_sharded_cp(self) -> None: @@ -1353,3 +1396,482 @@ class TestValidateExplicitReplicated: all_axes=set(), parallel_infos=[{}], ) + + def test_jointly_determined_axis_passes(self) -> None: + """EDP determined by (TP, CP) jointly but not by either alone → no error. + + Simulates tp=2, cp=2, ep=1, etp=1 on 4 GPUs where edp_size=4 + and edp_rank = unique per (tp_rank, cp_rank) combination. + """ + 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.SP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=0, axis_size=4), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.SP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=1, axis_size=4), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.SP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=2, axis_size=4), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.SP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=3, axis_size=4), + }, + ] + _validate_explicit_replicated( + explicit_replicated_axes=frozenset({ParallelAxis.TP}), + sharded_axes={ParallelAxis.CP, ParallelAxis.SP}, + all_axes={ + ParallelAxis.TP, + ParallelAxis.CP, + ParallelAxis.SP, + ParallelAxis.EDP, + }, + parallel_infos=parallel_infos, + ) + + def test_jointly_undetermined_axis_still_raises(self) -> None: + """Axis not determined even by the combination of all declared axes → raises. + + DP is orthogonal to TP (each TP rank pairs with both DP ranks), + so (TP,) cannot determine DP. + """ + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + { + ParallelAxis.TP: AxisInfo(axis_rank=tp_rank, axis_size=2), + ParallelAxis.DP: AxisInfo(axis_rank=dp_rank, axis_size=2), + } + for tp_rank in range(2) + for dp_rank in range(2) + ] + with pytest.raises(ValueError, match="dp.*not declared"): + _validate_explicit_replicated( + explicit_replicated_axes=frozenset(), + sharded_axes={ParallelAxis.TP}, + all_axes={ParallelAxis.TP, ParallelAxis.DP}, + parallel_infos=parallel_infos, + ) + + +class TestIsJointlyDetermined: + def test_edp_determined_by_tp_and_cp(self) -> None: + """EDP rank = unique per (TP, CP) combination → True.""" + 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.EDP: AxisInfo(axis_rank=0, axis_size=4), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=1, axis_size=4), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=2, axis_size=4), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=3, axis_size=4), + }, + ] + assert _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP, ParallelAxis.CP}), + child=ParallelAxis.EDP, + ) + + def test_dp_not_determined_by_tp_alone(self) -> None: + """DP is orthogonal to TP → False.""" + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + { + ParallelAxis.TP: AxisInfo(axis_rank=tp_rank, axis_size=2), + ParallelAxis.DP: AxisInfo(axis_rank=dp_rank, axis_size=2), + } + for tp_rank in range(2) + for dp_rank in range(2) + ] + assert not _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP}), + child=ParallelAxis.DP, + ) + + def test_empty_parallel_infos_returns_false(self) -> None: + """No parallel_info entries → False (no evidence).""" + assert not _is_jointly_determined( + [], + parent_axes=frozenset({ParallelAxis.TP}), + child=ParallelAxis.EDP, + ) + + def test_child_absent_from_infos_returns_false(self) -> None: + """Child axis not present in any info → False.""" + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + {ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2)}, + ] + assert not _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP}), + child=ParallelAxis.EDP, + ) + + def test_empty_parent_axes_returns_false(self) -> None: + """Empty parent_axes → False (no parents to determine child).""" + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + { + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=0, axis_size=2), + }, + ] + assert not _is_jointly_determined( + parallel_infos, + parent_axes=frozenset(), + child=ParallelAxis.EDP, + ) + + def test_single_parent_determines_child(self) -> None: + """Single parent tp_rank uniquely maps to edp_rank → True (degenerate joint case).""" + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + { + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=0, axis_size=2), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=1, axis_size=2), + }, + ] + assert _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP}), + child=ParallelAxis.EDP, + ) + + def test_conflict_returns_false(self) -> None: + """Same (tp_rank, cp_rank) maps to different edp_rank → False.""" + 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.EDP: AxisInfo(axis_rank=0, axis_size=4), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=1, axis_size=4), + }, + ] + assert not _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP, ParallelAxis.CP}), + child=ParallelAxis.EDP, + ) + + def test_two_parents_jointly_determine_child(self) -> None: + """(tp_rank, cp_rank) tuple uniquely determines edp_rank → True.""" + 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.EDP: AxisInfo(axis_rank=0, axis_size=4), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=1, axis_size=4), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=2, axis_size=4), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=3, axis_size=4), + }, + ] + assert _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP, ParallelAxis.CP}), + child=ParallelAxis.EDP, + ) + + def test_three_parents_jointly_determine_child(self) -> None: + """(tp, cp, ep) triple uniquely determines edp → True.""" + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + { + ParallelAxis.TP: AxisInfo(axis_rank=tp, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=cp, axis_size=2), + ParallelAxis.EP: AxisInfo(axis_rank=ep, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=tp * 4 + cp * 2 + ep, axis_size=8), + } + for tp in range(2) + for cp in range(2) + for ep in range(2) + ] + assert _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP, ParallelAxis.CP, ParallelAxis.EP}), + child=ParallelAxis.EDP, + ) + + def test_parent_partially_absent_causes_ambiguity(self) -> None: + """Some infos lack a parent axis → False, even if child values differ. + + When cp is missing from some infos, the joint determination is + incomplete because we cannot construct a full parent key. + """ + 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.EDP: AxisInfo(axis_rank=0, axis_size=4), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + # cp absent — parent key is incomplete + ParallelAxis.EDP: AxisInfo(axis_rank=1, axis_size=4), + }, + ] + assert not _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP, ParallelAxis.CP}), + child=ParallelAxis.EDP, + ) + + def test_partial_parent_first_info_missing_returns_false(self) -> None: + """First info lacks a parent axis; second info has all parents → False.""" + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + { + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + # cp absent + ParallelAxis.EDP: 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.EDP: AxisInfo(axis_rank=1, axis_size=2), + }, + ] + assert not _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP, ParallelAxis.CP}), + child=ParallelAxis.EDP, + ) + + def test_universally_absent_parent_ignored_remaining_determines(self) -> None: + """Parent axis absent from ALL infos is ignored; remaining parent determines child → True. + + Models the real scenario where DP (size 1) is in declared_axes but + filtered out of all parallel_infos by normalize_parallel_info. + """ + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + { + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=0, axis_size=2), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=1, axis_size=2), + }, + ] + assert _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP, ParallelAxis.CP}), + child=ParallelAxis.EDP, + ) + + def test_all_parents_universally_absent_returns_false(self) -> None: + """Every parent axis absent from ALL infos → no active parents → False.""" + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + { + ParallelAxis.EDP: AxisInfo(axis_rank=0, axis_size=2), + }, + { + ParallelAxis.EDP: AxisInfo(axis_rank=1, axis_size=2), + }, + ] + assert not _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP, ParallelAxis.CP}), + child=ParallelAxis.EDP, + ) + + def test_universally_absent_parent_remaining_conflict_returns_false(self) -> None: + """Parent axis absent from ALL infos ignored, but remaining parent has conflict → False.""" + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + { + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=0, axis_size=2), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=1, axis_size=2), + }, + ] + assert not _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP, ParallelAxis.CP}), + child=ParallelAxis.EDP, + ) + + def test_partial_parent_matching_child_still_returns_false(self) -> None: + """Even when child values match across infos, incomplete parent → False. + + Ensures the check is about parent completeness, not child conflict. + """ + 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.EDP: AxisInfo(axis_rank=0, axis_size=2), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + # cp absent — but edp_rank is SAME as first info + ParallelAxis.EDP: AxisInfo(axis_rank=0, axis_size=2), + }, + ] + assert not _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP, ParallelAxis.CP}), + child=ParallelAxis.EDP, + ) + + def test_many_infos_consistent_joint_mapping(self) -> None: + """8 ranks with (tp, cp) consistently mapping to edp → True.""" + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + { + ParallelAxis.TP: AxisInfo(axis_rank=tp, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=cp, axis_size=2), + ParallelAxis.EP: AxisInfo(axis_rank=ep, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=tp * 2 + cp, axis_size=4), + } + for tp in range(2) + for cp in range(2) + for ep in range(2) + ] + assert _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP, ParallelAxis.CP}), + child=ParallelAxis.EDP, + ) + + def test_partial_parent_middle_info_missing_returns_false(self) -> None: + """Middle info in a 3-info list lacks a parent → False.""" + 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.EDP: AxisInfo(axis_rank=0, axis_size=3), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), + # cp absent + ParallelAxis.EDP: AxisInfo(axis_rank=1, axis_size=3), + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=2, axis_size=3), + }, + ] + assert not _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP, ParallelAxis.CP}), + child=ParallelAxis.EDP, + ) + + def test_child_absent_from_some_infos_still_true(self) -> None: + """Child absent from some infos but consistent where present → True. + + Infos without the child are skipped; no parent completeness issue. + """ + 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.EDP: 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), + # edp absent — this info is skipped + }, + { + ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=1, axis_size=2), + }, + ] + assert _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP, ParallelAxis.CP}), + child=ParallelAxis.EDP, + ) + + def test_child_absent_from_all_infos_returns_false(self) -> None: + """Child not present in any info → mapping is empty → False.""" + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + {ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2)}, + {ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2)}, + ] + assert not _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP}), + child=ParallelAxis.CP, + ) + + def test_parent_present_in_some_but_missing_with_child_returns_false(self) -> None: + """Parent present in some infos but absent in an info that has child. + + This is the potential false-positive scenario: an info has child but + not all active parents, so the parent key cannot be fully constructed. + """ + 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.EDP: AxisInfo(axis_rank=0, axis_size=2), + }, + { + # TP present in first info so it's active, but absent here + ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2), + ParallelAxis.EDP: AxisInfo(axis_rank=1, axis_size=2), + }, + ] + assert not _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP, ParallelAxis.CP}), + child=ParallelAxis.EDP, + ) + + def test_single_info_with_all_axes_returns_true(self) -> None: + """Single info entry with parent and child → trivially determined → True.""" + parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ + { + ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=1), + ParallelAxis.EDP: AxisInfo(axis_rank=0, axis_size=1), + }, + ] + assert _is_jointly_determined( + parallel_infos, + parent_axes=frozenset({ParallelAxis.TP}), + child=ParallelAxis.EDP, + ) diff --git a/test/registered/debug_utils/comparator/test_display.py b/test/registered/debug_utils/comparator/test_display.py index c0a8e09a3..9a557a5de 100644 --- a/test/registered/debug_utils/comparator/test_display.py +++ b/test/registered/debug_utils/comparator/test_display.py @@ -13,6 +13,7 @@ from sglang.srt.debug_utils.comparator.display import ( _collect_rank_info, _extract_parallel_info, _render_polars_as_text, + _extract_parallel_info, ) from sglang.srt.debug_utils.comparator.output_types import ( InputIdsRecord, diff --git a/test/registered/debug_utils/comparator/test_entrypoint.py b/test/registered/debug_utils/comparator/test_entrypoint.py index d9987c14e..3e7c09459 100644 --- a/test/registered/debug_utils/comparator/test_entrypoint.py +++ b/test/registered/debug_utils/comparator/test_entrypoint.py @@ -4511,6 +4511,179 @@ class TestEntrypointAutoDescend: run(parse_args(argv)) +class TestPartialParallelInfo: + """Regression tests for _is_jointly_determined with incomplete parallel_info. + + When some ranks lack a parallel axis that other ranks have, the unsharder + planner must detect the inconsistency and report the axis as undeclared + rather than silently accepting it as jointly determined. + """ + + def test_missing_parent_axis_triggers_undeclared_error( + self, tmp_path: Path, capsys: pytest.CaptureFixture + ) -> None: + """Ranks with inconsistent parallel_info → undeclared axis error. + + # Step 1: Create 4 target ranks where moe_tp is absent from ranks 2-3. + # This makes moe_tp implicitly-sharded (dependent on tp for ranks 0-1), + # but edp is NOT dependent on tp alone (tp=0 maps to edp=0 AND edp=2). + # Step 2: _is_jointly_determined is called with parent_axes={tp, moe_tp} + # for child=edp. Ranks 2-3 lack moe_tp → returns False. + # Step 3: edp remains undeclared → ValueError emitted as error record. + """ + torch.manual_seed(42) + full_tensor = torch.randn(2, 8) + shard0 = full_tensor[:, :4] + shard1 = full_tensor[:, 4:] + + baseline_dir = tmp_path / "baseline" + target_dir = tmp_path / "target" + + _create_rank_dump( + baseline_dir, + rank=0, + name="hidden", + tensor=full_tensor, + dims="b h", + ) + + # Ranks 0-1: have tp + moe_tp + edp + _create_rank_dump( + target_dir, + rank=0, + name="hidden", + tensor=shard0, + dims="b h[tp]", + parallel_info={ + "tp_rank": 0, + "tp_size": 2, + "moe_tp_rank": 0, + "moe_tp_size": 2, + "edp_rank": 0, + "edp_size": 4, + }, + ) + _create_rank_dump( + target_dir, + rank=1, + name="hidden", + tensor=shard1, + dims="b h[tp]", + parallel_info={ + "tp_rank": 1, + "tp_size": 2, + "moe_tp_rank": 1, + "moe_tp_size": 2, + "edp_rank": 1, + "edp_size": 4, + }, + ) + + # Ranks 2-3: have tp + edp but NO moe_tp + _create_rank_dump( + target_dir, + rank=2, + name="hidden", + tensor=shard0, + dims="b h[tp]", + parallel_info={ + "tp_rank": 0, + "tp_size": 2, + "edp_rank": 2, + "edp_size": 4, + }, + ) + _create_rank_dump( + target_dir, + rank=3, + name="hidden", + tensor=shard1, + dims="b h[tp]", + parallel_info={ + "tp_rank": 1, + "tp_size": 2, + "edp_rank": 3, + "edp_size": 4, + }, + ) + + 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) + + assert exit_code == 1 + + errors = [r for r in records if isinstance(r, ComparisonErrorRecord)] + assert len(errors) >= 1 + assert any("not declared" in e.traceback_str for e in errors) + + def test_consistent_parallel_info_allows_joint_determination( + self, tmp_path: Path, capsys: pytest.CaptureFixture + ) -> None: + """All ranks have complete parallel_info → edp is jointly determined, comparison succeeds. + + # Step 1: 4 target ranks with TP=2, CP=2 (replicated), EDP=4. + # edp is NOT dependent on tp alone (tp=0→edp=0,2) or cp alone (cp=0→edp=0,1). + # Step 2: _is_jointly_determined is called with parent_axes={tp, cp}, child=edp. + # All infos have both tp and cp → joint mapping is consistent → True. + # Step 3: CP replicated picks one rank per tp group → TP concat → correct shape. + """ + torch.manual_seed(42) + full_tensor = torch.randn(2, 8) + shard0 = full_tensor[:, :4] + shard1 = full_tensor[:, 4:] + + baseline_dir = tmp_path / "baseline" + target_dir = tmp_path / "target" + + _create_rank_dump( + baseline_dir, + rank=0, + name="hidden", + tensor=full_tensor, + dims="b h", + ) + + # CP=replicated → ranks with different cp_rank have same tensor shard + for rank, tp, cp, edp, shard in [ + (0, 0, 0, 0, shard0), + (1, 1, 0, 1, shard1), + (2, 0, 1, 2, shard0), + (3, 1, 1, 3, shard1), + ]: + _create_rank_dump( + target_dir, + rank=rank, + name="hidden", + tensor=shard, + dims="b h[tp] # cp:replicated", + parallel_info={ + "tp_rank": tp, + "tp_size": 2, + "cp_rank": cp, + "cp_size": 2, + "edp_rank": edp, + "edp_size": 4, + }, + ) + + 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) + + assert exit_code == 0 + comp = _assert_single_comparison_passed(records) + assert comp.name == "hidden" + + class TestErrorResilience: """Bundle comparison exception → continue with remaining bundles."""