Reorganize modules and pipeline in dump comparator (#19374)

This commit is contained in:
fzyzcjy
2026-02-26 10:00:13 +08:00
committed by GitHub
parent 508b8e3387
commit 2739d7df62
35 changed files with 459 additions and 333 deletions
@@ -0,0 +1,493 @@
import sys
import pytest
import torch
from sglang.srt.debug_utils.comparator.aligner.unsharder.executor import (
_apply_unshard,
_verify_replicated_group,
execute_unsharder_plan,
)
from sglang.srt.debug_utils.comparator.aligner.unsharder.planner import (
compute_unsharder_plan,
)
from sglang.srt.debug_utils.comparator.aligner.unsharder.types import (
AxisInfo,
PickParams,
)
from sglang.srt.debug_utils.comparator.dims import ParallelAxis, parse_dims
from sglang.srt.debug_utils.comparator.warning_sink import warning_sink
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=10, suite="default", nightly=True)
class TestExecuteUnsharderPlan:
def test_tp4_concat(self) -> None:
full_tensor = torch.randn(2, 8, 16)
shards = list(full_tensor.chunk(4, dim=1))
dim_specs = parse_dims("b h(tp) d")
parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=4)} for i in range(4)
]
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 1
with warning_sink.context() as warnings:
result = execute_unsharder_plan(plans[0], shards)
assert len(result) == 1
assert torch.allclose(result[0], full_tensor)
assert warnings == []
def test_scrambled_world_ranks_correct_result(self) -> None:
full_tensor = torch.randn(4, 8)
shards = list(full_tensor.chunk(4, dim=0))
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)},
]
dim_specs = parse_dims("h(tp) d")
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 1
tensors_ordered_by_world_rank = [
shards[2], # world_rank=0, axis_rank=2
shards[0], # world_rank=1, axis_rank=0
shards[3], # world_rank=2, axis_rank=3
shards[1], # world_rank=3, axis_rank=1
]
with warning_sink.context() as warnings:
result = execute_unsharder_plan(plans[0], tensors_ordered_by_world_rank)
assert len(result) == 1
assert torch.allclose(result[0], full_tensor)
assert warnings == []
def test_single_step_reduces_tensor_count(self) -> None:
"""8 tensors with 2 groups of 4 produce 2 output tensors."""
full_a = torch.randn(4, 8)
full_b = torch.randn(4, 8)
shards_a = list(full_a.chunk(4, dim=0))
shards_b = list(full_b.chunk(4, dim=0))
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_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 2
tensors: list[torch.Tensor] = []
for cp_rank in range(2):
source = shards_a if cp_rank == 0 else shards_b
for tp_rank in range(4):
tensors.append(source[tp_rank])
with warning_sink.context() as _warnings:
intermediate = execute_unsharder_plan(plans[0], tensors)
assert len(intermediate) == 4
with warning_sink.context() as _warnings:
final = execute_unsharder_plan(plans[1], intermediate)
assert len(final) == 1
def test_cp_tp_concat(self) -> None:
"""CP=2 + TP=2: multi-step unshard reconstructs original tensor."""
torch.manual_seed(42)
full_tensor = torch.randn(4, 8, 16)
cp_chunks = list(full_tensor.chunk(2, dim=1))
tensors: list[torch.Tensor] = []
parallel_infos = []
for cp_rank in range(2):
tp_chunks = list(cp_chunks[cp_rank].chunk(2, dim=2))
for tp_rank in range(2):
tensors.append(tp_chunks[tp_rank])
parallel_infos.append(
{
ParallelAxis.CP: AxisInfo(axis_rank=cp_rank, axis_size=2),
ParallelAxis.TP: AxisInfo(axis_rank=tp_rank, axis_size=2),
}
)
dim_specs = parse_dims("b s(cp) h(tp)")
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 2
current = tensors
for plan in plans:
with warning_sink.context() as _warnings:
current = execute_unsharder_plan(plan, current)
assert len(current) == 1
assert torch.allclose(current[0], full_tensor)
def test_cp_tp_scrambled(self) -> None:
"""Scrambled world_ranks for CP=2 + TP=2 still reconstruct correctly."""
torch.manual_seed(42)
full_tensor = torch.randn(4, 8, 16)
cp_chunks = list(full_tensor.chunk(2, dim=1))
shard_map: dict[tuple[int, int], torch.Tensor] = {}
for cp_rank in range(2):
tp_chunks = list(cp_chunks[cp_rank].chunk(2, dim=2))
for tp_rank in range(2):
shard_map[(cp_rank, tp_rank)] = tp_chunks[tp_rank]
scrambled_assignment = [
(1, 1), # world_rank=0
(0, 0), # world_rank=1
(1, 0), # world_rank=2
(0, 1), # world_rank=3
]
tensors: list[torch.Tensor] = []
parallel_infos = []
for cp_rank, tp_rank in scrambled_assignment:
tensors.append(shard_map[(cp_rank, tp_rank)])
parallel_infos.append(
{
ParallelAxis.CP: AxisInfo(axis_rank=cp_rank, axis_size=2),
ParallelAxis.TP: AxisInfo(axis_rank=tp_rank, axis_size=2),
}
)
dim_specs = parse_dims("b s(cp) h(tp)")
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 2
current = tensors
for plan in plans:
with warning_sink.context() as _warnings:
current = execute_unsharder_plan(plan, current)
assert len(current) == 1
assert torch.allclose(current[0], full_tensor)
def test_unsupported_params_type_raises(self) -> None:
"""_apply_unshard raises ValueError for unknown params type."""
class _FakeParams:
pass
with pytest.raises(ValueError, match="Unsupported unshard"):
_apply_unshard(
_FakeParams(),
[torch.randn(2, 2)],
axis=ParallelAxis.TP,
group_index=0,
)
def test_cp_tp_ep_three_axis_concat(self) -> None:
"""CP=2 + TP=2 + EP=2: three-step unshard reconstructs original tensor."""
torch.manual_seed(42)
full_tensor = torch.randn(4, 8, 16, 32)
ep_chunks = list(full_tensor.chunk(2, dim=1))
shard_map: dict[tuple[int, int, int], torch.Tensor] = {}
for ep_rank in range(2):
cp_chunks = list(ep_chunks[ep_rank].chunk(2, dim=2))
for cp_rank in range(2):
tp_chunks = list(cp_chunks[cp_rank].chunk(2, dim=3))
for tp_rank in range(2):
shard_map[(ep_rank, cp_rank, tp_rank)] = tp_chunks[tp_rank]
tensors: list[torch.Tensor] = []
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
for ep_rank in range(2):
for cp_rank in range(2):
for tp_rank in range(2):
tensors.append(shard_map[(ep_rank, cp_rank, tp_rank)])
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),
}
)
dim_specs = parse_dims("b e(ep) s(cp) h(tp)")
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 3
current = tensors
for plan in plans:
with warning_sink.context() as _warnings:
current = execute_unsharder_plan(plan, current)
assert len(current) == 1
assert torch.allclose(current[0], full_tensor)
def test_cp_tp_ep_scrambled_three_axis(self) -> None:
"""Scrambled ranks for CP=2 + TP=2 + EP=2 still reconstruct correctly."""
torch.manual_seed(42)
full_tensor = torch.randn(4, 8, 16, 32)
ep_chunks = list(full_tensor.chunk(2, dim=1))
shard_map: dict[tuple[int, int, int], torch.Tensor] = {}
for ep_rank in range(2):
cp_chunks = list(ep_chunks[ep_rank].chunk(2, dim=2))
for cp_rank in range(2):
tp_chunks = list(cp_chunks[cp_rank].chunk(2, dim=3))
for tp_rank in range(2):
shard_map[(ep_rank, cp_rank, tp_rank)] = tp_chunks[tp_rank]
scrambled_assignment = [
(1, 0, 1), # world_rank=0
(0, 1, 0), # world_rank=1
(1, 1, 0), # world_rank=2
(0, 0, 0), # world_rank=3
(0, 1, 1), # world_rank=4
(1, 0, 0), # world_rank=5
(0, 0, 1), # world_rank=6
(1, 1, 1), # world_rank=7
]
tensors: list[torch.Tensor] = []
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
for ep_rank, cp_rank, tp_rank in scrambled_assignment:
tensors.append(shard_map[(ep_rank, cp_rank, tp_rank)])
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),
}
)
dim_specs = parse_dims("b e(ep) s(cp) h(tp)")
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 3
current = tensors
for plan in plans:
with warning_sink.context() as _warnings:
current = execute_unsharder_plan(plan, current)
assert len(current) == 1
assert torch.allclose(current[0], full_tensor)
class TestPickOperation:
def test_pick_single_group(self) -> None:
"""PickParams picks the first tensor from a single group."""
tensor = torch.randn(4, 8)
dim_specs = parse_dims("h d")
parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2)},
{ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2)},
]
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 1
assert isinstance(plans[0].params, PickParams)
with warning_sink.context() as warnings:
result = execute_unsharder_plan(plans[0], [tensor, tensor.clone()])
assert len(result) == 1
assert torch.allclose(result[0], tensor)
assert warnings == []
def test_pick_multiple_groups(self) -> None:
"""PickParams with multiple groups picks one from each."""
dim_specs = parse_dims("h(tp)")
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
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),
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=1, axis_size=2),
},
]
plans = compute_unsharder_plan(dim_specs, parallel_infos)
pick_plans = [p for p in plans if isinstance(p.params, PickParams)]
assert len(pick_plans) == 1
assert pick_plans[0].axis == ParallelAxis.CP
tensor = torch.randn(4)
tensors = [tensor.clone() for _ in range(4)]
with warning_sink.context() as warnings:
result = execute_unsharder_plan(pick_plans[0], tensors)
assert len(result) == 2
assert warnings == []
def test_replicated_tp_sharded_cp_e2e(self) -> None:
"""CP2 TP2, dims='b s(cp) d': replicated TP pick + sharded CP concat round-trip."""
torch.manual_seed(42)
full_tensor = torch.randn(4, 8, 16)
cp_chunks = list(full_tensor.chunk(2, dim=1))
tensors: list[torch.Tensor] = []
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
for cp_rank in range(2):
for tp_rank in range(2):
tensors.append(cp_chunks[cp_rank].clone())
parallel_infos.append(
{
ParallelAxis.CP: AxisInfo(axis_rank=cp_rank, axis_size=2),
ParallelAxis.TP: AxisInfo(axis_rank=tp_rank, axis_size=2),
}
)
dim_specs = parse_dims("b s(cp) d")
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 2
current = tensors
for plan in plans:
with warning_sink.context() as _warnings:
current = execute_unsharder_plan(plan, current)
assert len(current) == 1
assert torch.allclose(current[0], full_tensor)
def test_fully_replicated_e2e(self) -> None:
"""CP2 TP2, dims='b h d': fully replicated -> 2 pick steps -> 1 tensor."""
torch.manual_seed(42)
full_tensor = torch.randn(4, 8, 16)
tensors: list[torch.Tensor] = []
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
for cp_rank in range(2):
for tp_rank in range(2):
tensors.append(full_tensor.clone())
parallel_infos.append(
{
ParallelAxis.CP: AxisInfo(axis_rank=cp_rank, axis_size=2),
ParallelAxis.TP: AxisInfo(axis_rank=tp_rank, axis_size=2),
}
)
dim_specs = parse_dims("b h d")
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 2
assert all(isinstance(p.params, PickParams) for p in plans)
current = tensors
for plan in plans:
with warning_sink.context() as _warnings:
current = execute_unsharder_plan(plan, current)
assert len(current) == 1
assert torch.allclose(current[0], full_tensor)
class TestVerifyReplicatedGroup:
def test_warns_on_mismatch(self) -> None:
"""_verify_replicated_group produces warning when replicas differ."""
tensor_a = torch.ones(4)
tensor_b = torch.ones(4) + 0.1
with warning_sink.context() as warnings:
_verify_replicated_group(
[tensor_a, tensor_b],
axis=ParallelAxis.TP,
group_index=0,
)
assert len(warnings) == 1
assert warnings[0].axis == "tp"
assert warnings[0].group_index == 0
assert warnings[0].differing_index == 1
assert warnings[0].baseline_index == 0
assert warnings[0].max_abs_diff == pytest.approx(0.1, abs=1e-5)
def test_no_warn_when_identical(self) -> None:
"""_verify_replicated_group produces no warning for identical replicas."""
tensor = torch.randn(4, 8)
with warning_sink.context() as warnings:
_verify_replicated_group(
[tensor, tensor.clone()],
axis=ParallelAxis.TP,
group_index=0,
)
assert warnings == []
def test_multiple_mismatches(self) -> None:
"""_verify_replicated_group reports each differing replica."""
baseline = torch.zeros(4)
other_a = torch.ones(4)
other_b = torch.ones(4) * 2
with warning_sink.context() as warnings:
_verify_replicated_group(
[baseline, other_a, other_b],
axis=ParallelAxis.CP,
group_index=1,
)
assert len(warnings) == 2
assert warnings[0].differing_index == 1
assert warnings[1].differing_index == 2
assert warnings[1].max_abs_diff == pytest.approx(2.0, abs=1e-5)
def test_execute_returns_warnings(self) -> None:
"""execute_unsharder_plan emits warnings for replicated mismatch."""
dim_specs = parse_dims("h d")
parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2)},
{ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2)},
]
plans = compute_unsharder_plan(dim_specs, parallel_infos)
tensor_a = torch.zeros(4)
tensor_b = torch.ones(4)
with warning_sink.context() as warnings:
result = execute_unsharder_plan(plans[0], [tensor_a, tensor_b])
assert len(result) == 1
assert len(warnings) == 1
assert torch.allclose(result[0], tensor_a)
def test_atol_boundary_within(self) -> None:
"""Difference exactly at atol (1e-6) -> torch.allclose passes -> no warning."""
baseline = torch.zeros(4)
other = torch.full((4,), 1e-6)
with warning_sink.context() as warnings:
_verify_replicated_group(
[baseline, other],
axis=ParallelAxis.TP,
group_index=0,
)
assert warnings == []
def test_atol_boundary_exceeded(self) -> None:
"""Difference just above atol (1e-6 + 1e-9) -> torch.allclose fails -> warning."""
baseline = torch.zeros(4)
other = torch.full((4,), 1e-6 + 1e-9)
with warning_sink.context() as warnings:
_verify_replicated_group(
[baseline, other],
axis=ParallelAxis.TP,
group_index=0,
)
assert len(warnings) == 1
assert warnings[0].differing_index == 1
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))
@@ -0,0 +1,70 @@
import sys
import pytest
from sglang.srt.debug_utils.comparator.aligner.unsharder.parallel_info import (
normalize_parallel_info,
)
from sglang.srt.debug_utils.comparator.aligner.unsharder.types import AxisInfo
from sglang.srt.debug_utils.comparator.dims import ParallelAxis
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=10, suite="default", nightly=True)
class TestNormalizeParallelInfo:
def test_sglang_info(self) -> None:
meta = {
"sglang_parallel_info": {
"tp_rank": 2,
"tp_size": 4,
"pp_rank": 0,
"pp_size": 1,
}
}
result = normalize_parallel_info(meta)
assert result == {ParallelAxis.TP: AxisInfo(axis_rank=2, axis_size=4)}
def test_megatron_info(self) -> None:
meta = {
"megatron_parallel_info": {
"tp_rank": 1,
"tp_size": 2,
"cp_rank": 0,
"cp_size": 4,
"dp_rank": 0,
"dp_size": 1,
}
}
result = normalize_parallel_info(meta)
assert result == {
ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2),
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=4),
}
def test_no_parallel_info(self) -> None:
assert normalize_parallel_info({}) == {}
assert normalize_parallel_info({"other_key": 42}) == {}
def test_both_present_raises(self) -> None:
meta = {
"sglang_parallel_info": {"tp_rank": 0, "tp_size": 2},
"megatron_parallel_info": {"tp_rank": 0, "tp_size": 2},
}
with pytest.raises(ValueError, match="multiple parallel_info"):
normalize_parallel_info(meta)
def test_size_1_filtered(self) -> None:
meta = {
"sglang_parallel_info": {
"tp_rank": 0,
"tp_size": 1,
"cp_rank": 0,
"cp_size": 1,
}
}
assert normalize_parallel_info(meta) == {}
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))
@@ -0,0 +1,405 @@
import sys
import pytest
from sglang.srt.debug_utils.comparator.aligner.unsharder.planner import (
compute_unsharder_plan,
)
from sglang.srt.debug_utils.comparator.aligner.unsharder.types import (
AxisInfo,
ConcatParams,
PickParams,
)
from sglang.srt.debug_utils.comparator.dims import ParallelAxis, parse_dims
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=10, suite="default", nightly=True)
class TestComputeUnsharderPlan:
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_unsharder_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_unsharder_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_unsharder_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_unsharder_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_unsharder_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_unsharder_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_unsharder_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_unsharder_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_unsharder_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_unsharder_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_unsharder_plan(dim_specs, parallel_infos)
def test_ordering_zigzag_accepted(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)
]
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 1
assert plans[0].axis == ParallelAxis.CP
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_unsharder_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_unsharder_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_sharded_axis_missing_from_rank_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 — sharded axis absent from rank
},
]
with pytest.raises(ValueError, match="missing parallel_info"):
compute_unsharder_plan(dim_specs, parallel_infos)
class TestReplicatedAxes:
def test_replicated_tp_with_sharded_cp(self) -> None:
"""CP2 TP2, dims='b s(cp) d' → PickPlan(TP) + ConcatPlan(CP)."""
dim_specs = parse_dims("b s(cp) d")
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
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_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 2
assert plans[0].axis == ParallelAxis.TP
assert isinstance(plans[0].params, PickParams)
assert len(plans[0].groups) == 2
for group in plans[0].groups:
assert len(group) == 2
assert plans[1].axis == ParallelAxis.CP
assert isinstance(plans[1].params, ConcatParams)
assert plans[1].params.dim == 1
def test_fully_replicated(self) -> None:
"""CP2 TP2, dims='b h d' → PickPlan(CP) + PickPlan(TP)."""
dim_specs = parse_dims("b h d")
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
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_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 2
assert all(isinstance(p.params, PickParams) for p in plans)
axes = {p.axis for p in plans}
assert axes == {ParallelAxis.CP, ParallelAxis.TP}
def test_multiple_replicated_one_sharded(self) -> None:
"""CP2 TP2 EP2, dims='h(tp)' → PickPlan(CP) + PickPlan(EP) + ConcatPlan(TP)."""
dim_specs = parse_dims("h(tp)")
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
for cp_rank in range(2):
for ep_rank in range(2):
for tp_rank in range(2):
parallel_infos.append(
{
ParallelAxis.CP: AxisInfo(axis_rank=cp_rank, axis_size=2),
ParallelAxis.EP: AxisInfo(axis_rank=ep_rank, axis_size=2),
ParallelAxis.TP: AxisInfo(axis_rank=tp_rank, axis_size=2),
}
)
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 3
pick_plans = [p for p in plans if isinstance(p.params, PickParams)]
concat_plans = [p for p in plans if isinstance(p.params, ConcatParams)]
assert len(pick_plans) == 2
assert len(concat_plans) == 1
assert concat_plans[0].axis == ParallelAxis.TP
replicated_axes = {p.axis for p in pick_plans}
assert replicated_axes == {ParallelAxis.CP, ParallelAxis.EP}
def test_replicated_scrambled_ranks(self) -> None:
"""Scrambled world_rank order with replicated axis."""
dim_specs = parse_dims("h(tp)")
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.CP: AxisInfo(axis_rank=1, 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=0, 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_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 2
assert plans[0].axis == ParallelAxis.CP
assert isinstance(plans[0].params, PickParams)
assert plans[1].axis == ParallelAxis.TP
assert isinstance(plans[1].params, ConcatParams)
def test_replicated_axis_inconsistent_size_raises(self) -> None:
"""Replicated axis with inconsistent sizes raises ValueError."""
dim_specs = parse_dims("h(tp)")
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
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=4),
ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2),
},
]
with pytest.raises(ValueError, match="Inconsistent axis_size"):
compute_unsharder_plan(dim_specs, parallel_infos)
def test_replicated_axis_missing_from_rank_raises(self) -> None:
"""A rank missing a replicated axis that other ranks have raises ValueError."""
dim_specs = parse_dims("h(tp)")
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2),
},
{
# missing CP — replicated axis absent from this rank
ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2),
},
]
with pytest.raises(ValueError, match="missing parallel_info"):
compute_unsharder_plan(dim_specs, parallel_infos)
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))