Support context parallel zigzag reordering in dump comparator (#19281)
This commit is contained in:
@@ -0,0 +1,158 @@
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.srt.debug_utils.comparator.aligner.reorder import (
|
||||
ReorderPlan,
|
||||
_reorder_zigzag_to_natural,
|
||||
compute_reorder_plans,
|
||||
execute_reorder_plan,
|
||||
)
|
||||
from sglang.srt.debug_utils.comparator.aligner.unshard.executor import (
|
||||
execute_unshard_plan,
|
||||
)
|
||||
from sglang.srt.debug_utils.comparator.aligner.unshard.planner import (
|
||||
compute_unshard_plan,
|
||||
)
|
||||
from sglang.srt.debug_utils.comparator.aligner.unshard.types import AxisInfo
|
||||
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 TestZigzagToNatural:
|
||||
def test_zigzag_to_natural_cp2(self) -> None:
|
||||
"""cp_size=2: zigzag order [0,3,1,2] -> natural [0,1,2,3]."""
|
||||
natural = torch.arange(24).reshape(4, 6)
|
||||
chunks = list(natural.chunk(4, dim=0))
|
||||
|
||||
zigzag_order: list[int] = [0, 3, 1, 2]
|
||||
zigzagged = torch.cat([chunks[i] for i in zigzag_order], dim=0)
|
||||
|
||||
result = _reorder_zigzag_to_natural(zigzagged, dim=0, cp_size=2)
|
||||
assert torch.equal(result, natural)
|
||||
|
||||
def test_zigzag_to_natural_cp3(self) -> None:
|
||||
"""cp_size=3: zigzag 162534 -> natural 123456 (1-indexed)."""
|
||||
natural = torch.arange(60).reshape(6, 10)
|
||||
chunks = list(natural.chunk(6, dim=0))
|
||||
|
||||
zigzag_order: list[int] = [0, 5, 1, 4, 2, 3]
|
||||
zigzagged = torch.cat([chunks[i] for i in zigzag_order], dim=0)
|
||||
|
||||
result = _reorder_zigzag_to_natural(zigzagged, dim=0, cp_size=3)
|
||||
assert torch.equal(result, natural)
|
||||
|
||||
def test_zigzag_to_natural_arbitrary_dim(self) -> None:
|
||||
"""Reorder along dim=1 instead of dim=0."""
|
||||
natural = torch.arange(48).reshape(3, 4, 4)
|
||||
chunks = list(natural.chunk(4, dim=1))
|
||||
|
||||
zigzag_order: list[int] = [0, 3, 1, 2]
|
||||
zigzagged = torch.cat([chunks[i] for i in zigzag_order], dim=1)
|
||||
|
||||
result = _reorder_zigzag_to_natural(zigzagged, dim=1, cp_size=2)
|
||||
assert torch.equal(result, natural)
|
||||
|
||||
|
||||
class TestComputeReorderPlans:
|
||||
def test_compute_reorder_plans_zigzag(self) -> None:
|
||||
"""s(cp,zigzag) produces a ReorderPlan."""
|
||||
dim_specs = parse_dims("b s(cp,zigzag) 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),
|
||||
},
|
||||
]
|
||||
plans = compute_reorder_plans(
|
||||
dim_specs=dim_specs, parallel_infos=parallel_infos
|
||||
)
|
||||
|
||||
assert len(plans) == 1
|
||||
assert plans[0].params.op == "zigzag_to_natural"
|
||||
assert plans[0].params.dim == 1
|
||||
assert plans[0].params.cp_size == 2
|
||||
|
||||
def test_compute_reorder_plans_non_seq_dim_raises(self) -> None:
|
||||
"""Zigzag on non-sequence dim (e.g. t(cp,zigzag)) raises ValueError."""
|
||||
dim_specs = parse_dims("t(cp,zigzag) 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),
|
||||
},
|
||||
]
|
||||
with pytest.raises(ValueError, match="only supported on sequence dims"):
|
||||
compute_reorder_plans(dim_specs=dim_specs, parallel_infos=parallel_infos)
|
||||
|
||||
def test_compute_reorder_plans_natural(self) -> None:
|
||||
"""s(cp) and s(cp,natural) produce no reorder plans."""
|
||||
for dims_str in ["b s(cp) h(tp)", "b s(cp,natural) h(tp)"]:
|
||||
dim_specs = parse_dims(dims_str)
|
||||
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
|
||||
{
|
||||
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
|
||||
ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2),
|
||||
},
|
||||
]
|
||||
plans = compute_reorder_plans(
|
||||
dim_specs=dim_specs, parallel_infos=parallel_infos
|
||||
)
|
||||
assert plans == []
|
||||
|
||||
|
||||
class TestCpZigzagTpE2E:
|
||||
def test_cp_zigzag_tp_e2e(self) -> None:
|
||||
"""CP=2 zigzag + TP=2: full pipeline round-trip."""
|
||||
torch.manual_seed(42)
|
||||
full_tensor = torch.randn(4, 8, 16)
|
||||
|
||||
# Shard: first split seq dim (dim=1) into CP=2 with zigzag ordering,
|
||||
# then split hidden dim (dim=2) into TP=2.
|
||||
natural_cp_chunks = list(full_tensor.chunk(4, dim=1))
|
||||
zigzag_order: list[int] = [0, 3, 1, 2]
|
||||
zigzagged = torch.cat([natural_cp_chunks[i] for i in zigzag_order], dim=1)
|
||||
|
||||
cp_chunks = list(zigzagged.chunk(2, dim=1))
|
||||
tensors: list[torch.Tensor] = []
|
||||
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
|
||||
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,zigzag) h(tp)")
|
||||
|
||||
unshard_plans = compute_unshard_plan(
|
||||
dim_specs=dim_specs, parallel_infos=parallel_infos
|
||||
)
|
||||
reorder_plans = compute_reorder_plans(
|
||||
dim_specs=dim_specs, parallel_infos=parallel_infos
|
||||
)
|
||||
all_plans = [*unshard_plans, *reorder_plans]
|
||||
|
||||
assert len(unshard_plans) == 2
|
||||
assert len(reorder_plans) == 1
|
||||
|
||||
current: list[torch.Tensor] = tensors
|
||||
for plan in all_plans:
|
||||
if isinstance(plan, ReorderPlan):
|
||||
current = execute_reorder_plan(plan, current)
|
||||
else:
|
||||
current = execute_unshard_plan(plan, current)
|
||||
|
||||
assert len(current) == 1
|
||||
assert torch.allclose(current[0], full_tensor)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__]))
|
||||
@@ -0,0 +1,263 @@
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.srt.debug_utils.comparator.aligner.unshard.executor import (
|
||||
_apply_unshard,
|
||||
execute_unshard_plan,
|
||||
)
|
||||
from sglang.srt.debug_utils.comparator.aligner.unshard.planner import (
|
||||
compute_unshard_plan,
|
||||
)
|
||||
from sglang.srt.debug_utils.comparator.aligner.unshard.types import AxisInfo
|
||||
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 TestExecuteUnshardPlan:
|
||||
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_unshard_plan(dim_specs, parallel_infos)
|
||||
assert len(plans) == 1
|
||||
|
||||
result = execute_unshard_plan(plans[0], shards)
|
||||
assert len(result) == 1
|
||||
assert torch.allclose(result[0], full_tensor)
|
||||
|
||||
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_unshard_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
|
||||
]
|
||||
|
||||
result = execute_unshard_plan(plans[0], tensors_ordered_by_world_rank)
|
||||
assert len(result) == 1
|
||||
assert torch.allclose(result[0], full_tensor)
|
||||
|
||||
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_unshard_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])
|
||||
|
||||
intermediate = execute_unshard_plan(plans[0], tensors)
|
||||
assert len(intermediate) == 4
|
||||
|
||||
final = execute_unshard_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_unshard_plan(dim_specs, parallel_infos)
|
||||
assert len(plans) == 2
|
||||
|
||||
current = tensors
|
||||
for plan in plans:
|
||||
current = execute_unshard_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_unshard_plan(dim_specs, parallel_infos)
|
||||
assert len(plans) == 2
|
||||
|
||||
current = tensors
|
||||
for plan in plans:
|
||||
current = execute_unshard_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)])
|
||||
|
||||
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_unshard_plan(dim_specs, parallel_infos)
|
||||
assert len(plans) == 3
|
||||
|
||||
current = tensors
|
||||
for plan in plans:
|
||||
current = execute_unshard_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_unshard_plan(dim_specs, parallel_infos)
|
||||
assert len(plans) == 3
|
||||
|
||||
current = tensors
|
||||
for plan in plans:
|
||||
current = execute_unshard_plan(plan, current)
|
||||
|
||||
assert len(current) == 1
|
||||
assert torch.allclose(current[0], full_tensor)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__]))
|
||||
@@ -0,0 +1,70 @@
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
|
||||
from sglang.srt.debug_utils.comparator.aligner.unshard.parallel_info import (
|
||||
normalize_parallel_info,
|
||||
)
|
||||
from sglang.srt.debug_utils.comparator.aligner.unshard.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,249 @@
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
|
||||
from sglang.srt.debug_utils.comparator.aligner.unshard.planner import (
|
||||
compute_unshard_plan,
|
||||
)
|
||||
from sglang.srt.debug_utils.comparator.aligner.unshard.types import AxisInfo
|
||||
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 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_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_unshard_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_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__]))
|
||||
Reference in New Issue
Block a user