247 lines
9.5 KiB
Python
247 lines
9.5 KiB
Python
import sys
|
|
|
|
import pytest
|
|
|
|
from sglang.srt.debug_utils.comparator.dims import ParallelAxis, parse_dims
|
|
from sglang.srt.debug_utils.comparator.unshard.planner import compute_unshard_plan
|
|
from sglang.srt.debug_utils.comparator.unshard.types import AxisInfo
|
|
from sglang.test.ci.ci_register import register_cpu_ci
|
|
|
|
register_cpu_ci(est_time=10, suite="default", nightly=True)
|
|
|
|
|
|
class TestComputeUnshardPlan:
|
|
def test_tp4_plan(self) -> None:
|
|
dim_specs = parse_dims("b s h(tp) d")
|
|
parallel_infos = [
|
|
{ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=4)} for i in range(4)
|
|
]
|
|
plans = compute_unshard_plan(dim_specs, parallel_infos)
|
|
|
|
assert len(plans) == 1
|
|
assert plans[0].axis == ParallelAxis.TP
|
|
assert plans[0].params.dim == 2
|
|
assert plans[0].groups == [[0, 1, 2, 3]]
|
|
|
|
def test_inconsistent_axis_size_raises(self) -> None:
|
|
dim_specs = parse_dims("h(tp)")
|
|
parallel_infos = [
|
|
{ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=4)},
|
|
{ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2)},
|
|
]
|
|
with pytest.raises(ValueError, match="Inconsistent axis_size"):
|
|
compute_unshard_plan(dim_specs, parallel_infos)
|
|
|
|
def test_missing_axis_in_parallel_info_raises(self) -> None:
|
|
dim_specs = parse_dims("h(tp)")
|
|
parallel_infos = [{ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2)}]
|
|
with pytest.raises(ValueError, match="missing parallel_info"):
|
|
compute_unshard_plan(dim_specs, parallel_infos)
|
|
|
|
def test_empty_parallel_infos_raises(self) -> None:
|
|
dim_specs = parse_dims("h(tp)")
|
|
with pytest.raises(ValueError, match="must not be empty"):
|
|
compute_unshard_plan(dim_specs, [])
|
|
|
|
def test_scrambled_world_ranks(self) -> None:
|
|
"""world_rank order != axis_rank order."""
|
|
dim_specs = parse_dims("h(tp)")
|
|
parallel_infos = [
|
|
{ParallelAxis.TP: AxisInfo(axis_rank=2, axis_size=4)},
|
|
{ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=4)},
|
|
{ParallelAxis.TP: AxisInfo(axis_rank=3, axis_size=4)},
|
|
{ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=4)},
|
|
]
|
|
plans = compute_unshard_plan(dim_specs, parallel_infos)
|
|
assert len(plans) == 1
|
|
assert plans[0].groups == [[1, 3, 0, 2]]
|
|
|
|
def test_no_sharded_axes_returns_empty(self) -> None:
|
|
dim_specs = parse_dims("b s d")
|
|
parallel_infos = [{}]
|
|
plans = compute_unshard_plan(dim_specs, parallel_infos)
|
|
assert plans == []
|
|
|
|
def test_multi_axis_plan(self) -> None:
|
|
"""Multi-axis (TP + CP) produces a 2-step plan."""
|
|
dim_specs = parse_dims("s(cp) h(tp)")
|
|
parallel_infos = [
|
|
{
|
|
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
|
|
ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2),
|
|
},
|
|
{
|
|
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
|
|
ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2),
|
|
},
|
|
{
|
|
ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2),
|
|
ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2),
|
|
},
|
|
{
|
|
ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2),
|
|
ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2),
|
|
},
|
|
]
|
|
plans = compute_unshard_plan(dim_specs, parallel_infos)
|
|
|
|
assert len(plans) == 2
|
|
assert plans[0].axis == ParallelAxis.CP
|
|
assert plans[1].axis == ParallelAxis.TP
|
|
|
|
def test_cp_tp_plan(self) -> None:
|
|
"""CP=2 + TP=4 produces correct 2-step plan with correct groups."""
|
|
dim_specs = parse_dims("s(cp) h(tp)")
|
|
parallel_infos = []
|
|
for cp_rank in range(2):
|
|
for tp_rank in range(4):
|
|
parallel_infos.append(
|
|
{
|
|
ParallelAxis.CP: AxisInfo(axis_rank=cp_rank, axis_size=2),
|
|
ParallelAxis.TP: AxisInfo(axis_rank=tp_rank, axis_size=4),
|
|
}
|
|
)
|
|
|
|
plans = compute_unshard_plan(dim_specs, parallel_infos)
|
|
|
|
assert len(plans) == 2
|
|
|
|
cp_plan = plans[0]
|
|
assert cp_plan.axis == ParallelAxis.CP
|
|
assert len(cp_plan.groups) == 4
|
|
for group in cp_plan.groups:
|
|
assert len(group) == 2
|
|
|
|
tp_plan = plans[1]
|
|
assert tp_plan.axis == ParallelAxis.TP
|
|
assert len(tp_plan.groups) == 1
|
|
assert len(tp_plan.groups[0]) == 4
|
|
|
|
def test_cp_tp_scrambled_ranks(self) -> None:
|
|
"""Scrambled rank assignment still produces correct plan."""
|
|
dim_specs = parse_dims("s(cp) h(tp)")
|
|
parallel_infos = [
|
|
{
|
|
ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2),
|
|
ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2),
|
|
},
|
|
{
|
|
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
|
|
ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2),
|
|
},
|
|
{
|
|
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
|
|
ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2),
|
|
},
|
|
{
|
|
ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2),
|
|
ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2),
|
|
},
|
|
]
|
|
plans = compute_unshard_plan(dim_specs, parallel_infos)
|
|
|
|
assert len(plans) == 2
|
|
|
|
cp_plan = plans[0]
|
|
assert cp_plan.axis == ParallelAxis.CP
|
|
assert len(cp_plan.groups) == 2
|
|
for group in cp_plan.groups:
|
|
assert len(group) == 2
|
|
|
|
tp_plan = plans[1]
|
|
assert tp_plan.axis == ParallelAxis.TP
|
|
assert len(tp_plan.groups) == 1
|
|
assert len(tp_plan.groups[0]) == 2
|
|
|
|
def test_axis_rank_coverage_incomplete_raises(self) -> None:
|
|
"""TP size=4 but only ranks 0,1,3 provided (missing rank 2)."""
|
|
dim_specs = parse_dims("h(tp)")
|
|
parallel_infos = [
|
|
{ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=4)},
|
|
{ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=4)},
|
|
{ParallelAxis.TP: AxisInfo(axis_rank=3, axis_size=4)},
|
|
]
|
|
with pytest.raises(ValueError, match="axis_rank coverage.*incomplete"):
|
|
compute_unshard_plan(dim_specs, parallel_infos)
|
|
|
|
def test_reduction_not_implemented_raises(self) -> None:
|
|
dim_specs = parse_dims("h(tp,partial)")
|
|
parallel_infos = [
|
|
{ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2)
|
|
]
|
|
with pytest.raises(NotImplementedError, match="reduction"):
|
|
compute_unshard_plan(dim_specs, parallel_infos)
|
|
|
|
def test_ordering_not_natural_raises(self) -> None:
|
|
dim_specs = parse_dims("s(cp,zigzag)")
|
|
parallel_infos = [
|
|
{ParallelAxis.CP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2)
|
|
]
|
|
with pytest.raises(NotImplementedError, match="ordering"):
|
|
compute_unshard_plan(dim_specs, parallel_infos)
|
|
|
|
def test_ordering_natural_accepted(self) -> None:
|
|
dim_specs = parse_dims("s(cp,natural)")
|
|
parallel_infos = [
|
|
{ParallelAxis.CP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2)
|
|
]
|
|
plans = compute_unshard_plan(dim_specs, parallel_infos)
|
|
assert len(plans) == 1
|
|
assert plans[0].axis == ParallelAxis.CP
|
|
|
|
def test_three_axis_plan(self) -> None:
|
|
"""EP=2 + CP=2 + TP=2 produces a 3-step plan."""
|
|
dim_specs = parse_dims("b e(ep) s(cp) h(tp)")
|
|
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
|
|
for ep_rank in range(2):
|
|
for cp_rank in range(2):
|
|
for tp_rank in range(2):
|
|
parallel_infos.append(
|
|
{
|
|
ParallelAxis.EP: AxisInfo(axis_rank=ep_rank, axis_size=2),
|
|
ParallelAxis.CP: AxisInfo(axis_rank=cp_rank, axis_size=2),
|
|
ParallelAxis.TP: AxisInfo(axis_rank=tp_rank, axis_size=2),
|
|
}
|
|
)
|
|
|
|
plans = compute_unshard_plan(dim_specs, parallel_infos)
|
|
|
|
assert len(plans) == 3
|
|
assert plans[0].axis == ParallelAxis.EP
|
|
assert plans[1].axis == ParallelAxis.CP
|
|
assert plans[2].axis == ParallelAxis.TP
|
|
|
|
# Step 0 (EP): 8 tensors → 4 (groups of 2)
|
|
assert len(plans[0].groups) == 4
|
|
for group in plans[0].groups:
|
|
assert len(group) == 2
|
|
|
|
# Step 1 (CP): 4 tensors → 2 (groups of 2)
|
|
assert len(plans[1].groups) == 2
|
|
for group in plans[1].groups:
|
|
assert len(group) == 2
|
|
|
|
# Step 2 (TP): 2 tensors → 1 (single group of 2)
|
|
assert len(plans[2].groups) == 1
|
|
assert len(plans[2].groups[0]) == 2
|
|
|
|
def test_replicated_axis_raises(self) -> None:
|
|
"""A world_rank missing a sharded axis raises ValueError."""
|
|
dim_specs = parse_dims("s(cp) h(tp)")
|
|
parallel_infos = [
|
|
{
|
|
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
|
|
ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2),
|
|
},
|
|
{
|
|
ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2),
|
|
# missing TP — replicated
|
|
},
|
|
]
|
|
with pytest.raises(ValueError, match="missing parallel_info"):
|
|
compute_unshard_plan(dim_specs, parallel_infos)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
sys.exit(pytest.main([__file__]))
|