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
|
||||
|
||||
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."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user