Support jointly-determined axes inference in dump comparator (#21025)

This commit is contained in:
fzyzcjy
2026-03-20 22:01:26 +08:00
committed by GitHub
parent ecd7e40d20
commit 2f01950a0e
4 changed files with 746 additions and 0 deletions
@@ -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,
@@ -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,
)
@@ -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,
@@ -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."""