Support jointly-determined axes inference in dump comparator (#21025)
This commit is contained in:
@@ -161,6 +161,15 @@ def _validate_explicit_replicated(
|
|||||||
)
|
)
|
||||||
undeclared: set[ParallelAxis] = all_axes - declared_axes
|
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:
|
if undeclared:
|
||||||
undeclared_names: str = ", ".join(sorted(a.value for a in undeclared))
|
undeclared_names: str = ", ".join(sorted(a.value for a in undeclared))
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -238,6 +247,47 @@ def _is_dependent_axis(
|
|||||||
return True
|
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(
|
def _group_and_project(
|
||||||
*,
|
*,
|
||||||
current_coords: _CoordsList,
|
current_coords: _CoordsList,
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ import pytest
|
|||||||
from sglang.srt.debug_utils.comparator.aligner.unsharder.planner import (
|
from sglang.srt.debug_utils.comparator.aligner.unsharder.planner import (
|
||||||
_compute_dependent_axes,
|
_compute_dependent_axes,
|
||||||
_is_dependent_axis,
|
_is_dependent_axis,
|
||||||
|
_is_jointly_determined,
|
||||||
_validate_explicit_replicated,
|
_validate_explicit_replicated,
|
||||||
compute_unsharder_plan,
|
compute_unsharder_plan,
|
||||||
)
|
)
|
||||||
@@ -418,6 +419,48 @@ class TestComputeUnsharderPlan:
|
|||||||
assert ParallelAxis.EP in axes_in_plan
|
assert ParallelAxis.EP in axes_in_plan
|
||||||
assert ParallelAxis.ETP not 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:
|
class TestExplicitReplicatedAxes:
|
||||||
def test_replicated_tp_with_sharded_cp(self) -> None:
|
def test_replicated_tp_with_sharded_cp(self) -> None:
|
||||||
@@ -1353,3 +1396,482 @@ class TestValidateExplicitReplicated:
|
|||||||
all_axes=set(),
|
all_axes=set(),
|
||||||
parallel_infos=[{}],
|
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,
|
_collect_rank_info,
|
||||||
_extract_parallel_info,
|
_extract_parallel_info,
|
||||||
_render_polars_as_text,
|
_render_polars_as_text,
|
||||||
|
_extract_parallel_info,
|
||||||
)
|
)
|
||||||
from sglang.srt.debug_utils.comparator.output_types import (
|
from sglang.srt.debug_utils.comparator.output_types import (
|
||||||
InputIdsRecord,
|
InputIdsRecord,
|
||||||
|
|||||||
@@ -4511,6 +4511,179 @@ class TestEntrypointAutoDescend:
|
|||||||
run(parse_args(argv))
|
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:
|
class TestErrorResilience:
|
||||||
"""Bundle comparison exception → continue with remaining bundles."""
|
"""Bundle comparison exception → continue with remaining bundles."""
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user