Support data parallel attention in dump comparator (#19602)

This commit is contained in:
fzyzcjy
2026-03-01 10:51:21 +08:00
committed by GitHub
parent ea6ff7b01f
commit e64095c3c7
19 changed files with 783 additions and 325 deletions

View File

@@ -26,7 +26,7 @@ register_cpu_ci(est_time=10, suite="default", nightly=True)
class TestComputeReordererPlans:
def test_compute_reorderer_plans_zigzag(self) -> None:
"""s(cp:zigzag) produces a ReordererPlan."""
dim_specs = parse_dims("b s(cp:zigzag) h(tp)")
dim_specs = parse_dims("b s(cp:zigzag) h(tp)").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
@@ -44,7 +44,7 @@ class TestComputeReordererPlans:
def test_compute_reorderer_plans_thd_zigzag(self) -> None:
"""t(cp:zigzag) produces a ZigzagToNaturalThdParams plan."""
dim_specs = parse_dims("t(cp:zigzag) h(tp)")
dim_specs = parse_dims("t(cp:zigzag) h(tp)").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
@@ -65,7 +65,7 @@ class TestComputeReordererPlans:
def test_non_seq_dim_still_raises(self) -> None:
"""Zigzag on non-sequence/non-token dim (e.g. h(cp:zigzag)) raises ValueError."""
dim_specs = parse_dims("h(cp:zigzag) d")
dim_specs = parse_dims("h(cp:zigzag) d").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2)},
]
@@ -74,7 +74,7 @@ class TestComputeReordererPlans:
def test_thd_zigzag_without_seq_lens_raises(self) -> None:
"""t(cp:zigzag) without thd_global_seq_lens raises ValueError."""
dim_specs = parse_dims("t(cp:zigzag) h(tp)")
dim_specs = parse_dims("t(cp:zigzag) h(tp)").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
@@ -87,7 +87,7 @@ class TestComputeReordererPlans:
def test_thd_natural_no_reorder(self) -> None:
"""t(cp:natural) and t(cp) produce no reorder plans."""
for dims_str in ["t(cp:natural) h(tp)", "t(cp) h(tp)"]:
dim_specs = parse_dims(dims_str)
dim_specs = parse_dims(dims_str).dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
@@ -102,7 +102,7 @@ class TestComputeReordererPlans:
def test_compute_reorderer_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)
dim_specs = parse_dims(dims_str).dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
@@ -141,7 +141,7 @@ class TestCpZigzagTpE2E:
}
)
dim_specs: list[DimSpec] = parse_dims("b s(cp:zigzag) h(tp)")
dim_specs: list[DimSpec] = parse_dims("b s(cp:zigzag) h(tp)").dims
dim_names: list[str] = [s.name for s in dim_specs]
unsharder_plans = compute_unsharder_plan(
@@ -215,7 +215,7 @@ class TestCpZigzagSpSameDimE2E:
}
)
dim_specs: list[DimSpec] = parse_dims("t(cp:zigzag,sp) h")
dim_specs: list[DimSpec] = parse_dims("t(cp:zigzag,sp) h").dims
dim_names: list[str] = [s.name for s in dim_specs]
unsharder_plans = compute_unsharder_plan(

View File

@@ -42,7 +42,7 @@ class TestExecuteUnsharderPlan:
full_tensor = torch.randn(2, 8, 16)
shards = list(full_tensor.chunk(4, dim=1))
dim_specs = parse_dims("b h(tp) d")
dim_specs = parse_dims("b h(tp) d").dims
parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=4)} for i in range(4)
]
@@ -67,7 +67,7 @@ class TestExecuteUnsharderPlan:
{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")
dim_specs = parse_dims("h(tp) d").dims
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 1
@@ -95,7 +95,7 @@ class TestExecuteUnsharderPlan:
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)")
dim_specs = parse_dims("s(cp) h(tp)").dims
parallel_infos = []
for cp_rank in range(2):
for tp_rank in range(4):
@@ -145,7 +145,7 @@ class TestExecuteUnsharderPlan:
}
)
dim_specs = parse_dims("b s(cp) h(tp)")
dim_specs = parse_dims("b s(cp) h(tp)").dims
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 2
@@ -187,7 +187,7 @@ class TestExecuteUnsharderPlan:
}
)
dim_specs = parse_dims("b s(cp) h(tp)")
dim_specs = parse_dims("b s(cp) h(tp)").dims
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 2
@@ -241,7 +241,7 @@ class TestExecuteUnsharderPlan:
}
)
dim_specs = parse_dims("b e(ep) s(cp) h(tp)")
dim_specs = parse_dims("b e(ep) s(cp) h(tp)").dims
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 3
@@ -290,7 +290,7 @@ class TestExecuteUnsharderPlan:
}
)
dim_specs = parse_dims("b e(ep) s(cp) h(tp)")
dim_specs = parse_dims("b e(ep) s(cp) h(tp)").dims
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 3
@@ -307,7 +307,7 @@ 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")
dim_specs = parse_dims("h d").dims
parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2)},
{ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2)},
@@ -326,7 +326,7 @@ class TestPickOperation:
def test_pick_multiple_groups(self) -> None:
"""PickParams with multiple groups picks one from each."""
dim_specs = parse_dims("h(tp)")
dim_specs = parse_dims("h(tp)").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
@@ -378,7 +378,7 @@ class TestPickOperation:
}
)
dim_specs = parse_dims("b s(cp) d")
dim_specs = parse_dims("b s(cp) d").dims
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 2
@@ -407,7 +407,7 @@ class TestPickOperation:
}
)
dim_specs = parse_dims("b h d")
dim_specs = parse_dims("b h d").dims
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 2
assert all(isinstance(p.params, PickParams) for p in plans)
@@ -472,7 +472,7 @@ class TestVerifyReplicatedGroup:
def test_execute_returns_replicated_checks(self) -> None:
"""execute_unsharder_plan returns replicated checks for mismatch."""
dim_specs = parse_dims("h d")
dim_specs = parse_dims("h d").dims
parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2)},
{ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2)},
@@ -646,6 +646,261 @@ class TestThdCpConcat:
)
class TestReduceSum:
def test_basic_tp2_reduce(self) -> None:
"""2 partial tensors sum to full tensor."""
torch.manual_seed(42)
full_tensor = torch.randn(4, 8)
part_a = full_tensor * 0.6
part_b = full_tensor * 0.4
dim_specs = parse_dims("h(tp:partial) d").dims
parallel_infos = [
{ParallelAxis.TP: 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 isinstance(plans[0].params, ReduceSumParams)
named_parts: list[torch.Tensor] = _name_tensors([part_a, part_b], dim_specs)
unsharder_result: UnsharderResult = execute_unsharder_plan(
plans[0], named_parts
)
assert len(unsharder_result.tensors) == 1
assert torch.allclose(unsharder_result.tensors[0].rename(None), full_tensor)
def test_tp4_reduce(self) -> None:
"""4 partial tensors sum to full tensor."""
torch.manual_seed(42)
full_tensor = torch.randn(4, 8)
parts: list[torch.Tensor] = [full_tensor * 0.25 for _ in range(4)]
dim_specs = parse_dims("h(tp:partial) d").dims
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
named_parts: list[torch.Tensor] = _name_tensors(parts, dim_specs)
unsharder_result: UnsharderResult = execute_unsharder_plan(
plans[0], named_parts
)
assert len(unsharder_result.tensors) == 1
assert torch.allclose(unsharder_result.tensors[0].rename(None), full_tensor)
def test_multi_axis_concat_then_reduce(self) -> None:
"""CP concat + TP reduce end-to-end."""
torch.manual_seed(42)
full_tensor = torch.randn(4, 8, 16)
cp_chunks = list(full_tensor.chunk(2, dim=1))
# Each CP chunk is held as partial sums across TP ranks
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] * 0.5)
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:partial)").dims
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 2
current: list[torch.Tensor] = _name_tensors(tensors, dim_specs)
for plan in plans:
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, current)
current = unsharder_result.tensors
assert len(current) == 1
assert torch.allclose(current[0].rename(None), full_tensor)
def test_reduce_scrambled_ranks(self) -> None:
"""Scrambled rank order — sum is commutative so result is the same."""
torch.manual_seed(42)
full_tensor = torch.randn(4, 8)
parts: list[torch.Tensor] = [
full_tensor * 0.1,
full_tensor * 0.2,
full_tensor * 0.3,
full_tensor * 0.4,
]
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:partial) d").dims
plans = compute_unsharder_plan(dim_specs, parallel_infos)
named_parts: list[torch.Tensor] = _name_tensors(parts, dim_specs)
unsharder_result: UnsharderResult = execute_unsharder_plan(
plans[0], named_parts
)
assert len(unsharder_result.tensors) == 1
assert torch.allclose(unsharder_result.tensors[0].rename(None), full_tensor)
def test_reduce_preserves_named_dims(self) -> None:
"""Named tensor dimensions are preserved through reduce_sum."""
dim_specs = parse_dims("h(tp:partial) d").dims
part_a = torch.randn(4, 8).refine_names("h", "d")
part_b = torch.randn(4, 8).refine_names("h", "d")
plan = UnsharderPlan(
axis=ParallelAxis.TP,
params=ReduceSumParams(),
groups=[[0, 1]],
)
unsharder_result: UnsharderResult = execute_unsharder_plan(
plan, [part_a, part_b]
)
assert len(unsharder_result.tensors) == 1
assert unsharder_result.tensors[0].names == ("h", "d")
expected = (part_a.rename(None) + part_b.rename(None)).refine_names("h", "d")
assert torch.allclose(
unsharder_result.tensors[0].rename(None), expected.rename(None)
)
def test_recompute_pseudo_mismatch(self) -> None:
"""_verify_replicated_group returns failed check for RECOMPUTE_PSEUDO axis mismatch."""
tensor_a = torch.ones(4)
tensor_b = torch.ones(4) + 0.1
checks: list[ReplicatedCheckResult] = _verify_replicated_group(
[tensor_a, tensor_b],
axis=ParallelAxis.RECOMPUTE_PSEUDO,
group_index=0,
)
assert len(checks) == 1
assert checks[0].axis == "recompute_pseudo"
assert checks[0].group_index == 0
assert checks[0].compared_index == 1
assert checks[0].baseline_index == 0
assert not checks[0].passed
assert checks[0].diff.max_abs_diff == pytest.approx(0.1, abs=1e-5)
class TestThdCpConcat:
def test_single_seq(self) -> None:
"""Single seq THD unshard: 2 ranks → per-seq concat."""
rank0 = torch.tensor([1, 2, 3]).refine_names("t")
rank1 = torch.tensor([4, 5, 6]).refine_names("t")
plan = UnsharderPlan(
axis=ParallelAxis.CP,
params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[3]),
groups=[[0, 1]],
)
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, [rank0, rank1])
assert len(unsharder_result.tensors) == 1
expected = torch.tensor([1, 2, 3, 4, 5, 6])
assert torch.equal(unsharder_result.tensors[0].rename(None), expected)
def test_multi_seq(self) -> None:
"""Multi-seq THD unshard: 2 ranks, seq_lens=[50, 32, 46]."""
# rank0: [seqA_r0(50) | seqB_r0(32) | pad_r0(46)]
# rank1: [seqA_r1(50) | seqB_r1(32) | pad_r1(46)]
seq_a_r0 = torch.arange(0, 50)
seq_b_r0 = torch.arange(100, 132)
pad_r0 = torch.full((46,), -1)
rank0 = torch.cat([seq_a_r0, seq_b_r0, pad_r0]).refine_names("t")
seq_a_r1 = torch.arange(50, 100)
seq_b_r1 = torch.arange(132, 164)
pad_r1 = torch.full((46,), -2)
rank1 = torch.cat([seq_a_r1, seq_b_r1, pad_r1]).refine_names("t")
plan = UnsharderPlan(
axis=ParallelAxis.CP,
params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[50, 32, 46]),
groups=[[0, 1]],
)
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, [rank0, rank1])
assert len(unsharder_result.tensors) == 1
unsharded: torch.Tensor = unsharder_result.tensors[0].rename(None)
# seqA: r0(50) + r1(50) = 100 tokens, values 0..99
assert torch.equal(unsharded[:100], torch.cat([seq_a_r0, seq_a_r1]))
# seqB: r0(32) + r1(32) = 64 tokens
assert torch.equal(unsharded[100:164], torch.cat([seq_b_r0, seq_b_r1]))
# pad: r0(46) + r1(46) = 92 tokens
assert torch.equal(unsharded[164:256], torch.cat([pad_r0, pad_r1]))
def test_with_hidden_dim(self) -> None:
"""THD unshard with trailing hidden dim: shape [T, H]."""
torch.manual_seed(42)
hidden: int = 4
# rank0: [seqA_r0(3, 4) | seqB_r0(2, 4)]
# rank1: [seqA_r1(3, 4) | seqB_r1(2, 4)]
seq_a_r0 = torch.randn(3, hidden)
seq_b_r0 = torch.randn(2, hidden)
rank0 = torch.cat([seq_a_r0, seq_b_r0]).refine_names("t", "h")
seq_a_r1 = torch.randn(3, hidden)
seq_b_r1 = torch.randn(2, hidden)
rank1 = torch.cat([seq_a_r1, seq_b_r1]).refine_names("t", "h")
plan = UnsharderPlan(
axis=ParallelAxis.CP,
params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[3, 2]),
groups=[[0, 1]],
)
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, [rank0, rank1])
assert len(unsharder_result.tensors) == 1
unsharded: torch.Tensor = unsharder_result.tensors[0].rename(None)
assert unsharded.shape == (10, hidden)
assert torch.equal(unsharded[:6], torch.cat([seq_a_r0, seq_a_r1]))
assert torch.equal(unsharded[6:10], torch.cat([seq_b_r0, seq_b_r1]))
def test_with_leading_batch_dim(self) -> None:
"""THD unshard with leading batch dim: shape [B, T, H], t is dim=1."""
torch.manual_seed(42)
batch: int = 2
hidden: int = 4
# rank0: [seqA_r0(3) | seqB_r0(2)] per batch item
# rank1: [seqA_r1(3) | seqB_r1(2)] per batch item
seq_a_r0 = torch.randn(batch, 3, hidden)
seq_b_r0 = torch.randn(batch, 2, hidden)
rank0 = torch.cat([seq_a_r0, seq_b_r0], dim=1).refine_names("b", "t", "h")
seq_a_r1 = torch.randn(batch, 3, hidden)
seq_b_r1 = torch.randn(batch, 2, hidden)
rank1 = torch.cat([seq_a_r1, seq_b_r1], dim=1).refine_names("b", "t", "h")
plan = UnsharderPlan(
axis=ParallelAxis.CP,
params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[3, 2]),
groups=[[0, 1]],
)
unsharder_result: UnsharderResult = execute_unsharder_plan(plan, [rank0, rank1])
assert len(unsharder_result.tensors) == 1
unsharded: torch.Tensor = unsharder_result.tensors[0].rename(None)
assert unsharded.shape == (batch, 10, hidden)
# seqA: r0(3) + r1(3) = 6 tokens per batch
assert torch.equal(unsharded[:, :6, :], torch.cat([seq_a_r0, seq_a_r1], dim=1))
# seqB: r0(2) + r1(2) = 4 tokens per batch
assert torch.equal(
unsharded[:, 6:10, :], torch.cat([seq_b_r0, seq_b_r1], dim=1)
)
class TestReduceSum:
def test_basic_tp2_reduce(self) -> None:
"""2 partial tensors sum to full tensor."""

View File

@@ -19,7 +19,7 @@ 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")
dim_specs = parse_dims("b s h(tp) d").dims
parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=4)} for i in range(4)
]
@@ -31,7 +31,7 @@ class TestComputeUnsharderPlan:
assert plans[0].groups == [[0, 1, 2, 3]]
def test_inconsistent_axis_size_raises(self) -> None:
dim_specs = parse_dims("h(tp)")
dim_specs = parse_dims("h(tp)").dims
parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=4)},
{ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2)},
@@ -41,7 +41,7 @@ class TestComputeUnsharderPlan:
def test_missing_axis_in_all_parallel_infos_skipped(self) -> None:
"""Axis in dims but absent from all parallel_infos -> axis_size=1, auto-skip."""
dim_specs = parse_dims("h(tp)")
dim_specs = parse_dims("h(tp)").dims
parallel_infos = [{ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2)}]
# TP not in any parallel_info → skipped; CP is replicated but only 1 rank
# with size=2 → incomplete coverage
@@ -49,13 +49,13 @@ class TestComputeUnsharderPlan:
compute_unsharder_plan(dim_specs, parallel_infos)
def test_empty_parallel_infos_raises(self) -> None:
dim_specs = parse_dims("h(tp)")
dim_specs = parse_dims("h(tp)").dims
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)")
dim_specs = parse_dims("h(tp)").dims
parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=2, axis_size=4)},
{ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=4)},
@@ -67,14 +67,14 @@ class TestComputeUnsharderPlan:
assert plans[0].groups == [[1, 3, 0, 2]]
def test_no_sharded_axes_returns_empty(self) -> None:
dim_specs = parse_dims("b s d")
dim_specs = parse_dims("b s d").dims
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)")
dim_specs = parse_dims("s(cp) h(tp)").dims
parallel_infos = [
{
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
@@ -101,7 +101,7 @@ class TestComputeUnsharderPlan:
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)")
dim_specs = parse_dims("s(cp) h(tp)").dims
parallel_infos = []
for cp_rank in range(2):
for tp_rank in range(4):
@@ -129,7 +129,7 @@ class TestComputeUnsharderPlan:
def test_cp_tp_scrambled_ranks(self) -> None:
"""Scrambled rank assignment still produces correct plan."""
dim_specs = parse_dims("s(cp) h(tp)")
dim_specs = parse_dims("s(cp) h(tp)").dims
parallel_infos = [
{
ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2),
@@ -165,7 +165,7 @@ class TestComputeUnsharderPlan:
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)")
dim_specs = parse_dims("h(tp)").dims
parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=4)},
{ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=4)},
@@ -175,7 +175,7 @@ class TestComputeUnsharderPlan:
compute_unsharder_plan(dim_specs, parallel_infos)
def test_reduction_partial_returns_reduce_sum(self) -> None:
dim_specs = parse_dims("h(tp:partial)")
dim_specs = parse_dims("h(tp:partial)").dims
parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2)
]
@@ -188,7 +188,7 @@ class TestComputeUnsharderPlan:
def test_reduction_partial_tp4(self) -> None:
"""TP=4 with partial reduction produces a single ReduceSumParams step."""
dim_specs = parse_dims("h(tp:partial)")
dim_specs = parse_dims("h(tp:partial)").dims
parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=4)} for i in range(4)
]
@@ -200,7 +200,7 @@ class TestComputeUnsharderPlan:
def test_multi_axis_with_reduction_on_one(self) -> None:
"""CP concat + TP reduce produces a 2-step plan."""
dim_specs = parse_dims("s(cp) h(tp:partial)")
dim_specs = parse_dims("s(cp) h(tp:partial)").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
for cp_rank in range(2):
for tp_rank in range(2):
@@ -221,7 +221,7 @@ class TestComputeUnsharderPlan:
def test_reduction_scrambled_ranks(self) -> None:
"""Scrambled world_rank order with partial reduction."""
dim_specs = parse_dims("h(tp:partial)")
dim_specs = parse_dims("h(tp:partial)").dims
parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=2, axis_size=4)},
{ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=4)},
@@ -235,7 +235,7 @@ class TestComputeUnsharderPlan:
assert plans[0].groups == [[1, 3, 0, 2]]
def test_ordering_zigzag_accepted(self) -> None:
dim_specs = parse_dims("s(cp:zigzag)")
dim_specs = parse_dims("s(cp:zigzag)").dims
parallel_infos = [
{ParallelAxis.CP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2)
]
@@ -244,7 +244,7 @@ class TestComputeUnsharderPlan:
assert plans[0].axis == ParallelAxis.CP
def test_ordering_natural_accepted(self) -> None:
dim_specs = parse_dims("s(cp:natural)")
dim_specs = parse_dims("s(cp:natural)").dims
parallel_infos = [
{ParallelAxis.CP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2)
]
@@ -254,7 +254,7 @@ class TestComputeUnsharderPlan:
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)")
dim_specs = parse_dims("b e(ep) s(cp) h(tp)").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
for ep_rank in range(2):
for cp_rank in range(2):
@@ -290,7 +290,7 @@ class TestComputeUnsharderPlan:
def test_same_dim_cp_sp_plan(self) -> None:
"""t(cp:zigzag,sp) with CP=2 SP=2: SP unshards first (inner), then CP."""
dim_specs = parse_dims("t(cp:zigzag,sp) 1 h")
dim_specs = parse_dims("t(cp:zigzag,sp) 1 h").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
for cp_rank in range(2):
for sp_rank in range(2):
@@ -328,7 +328,7 @@ class TestComputeUnsharderPlan:
CpThdConcatParams,
)
dim_specs = parse_dims("t(cp:zigzag,sp) h")
dim_specs = parse_dims("t(cp:zigzag,sp) h").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
for cp_rank in range(2):
for sp_rank in range(2):
@@ -361,7 +361,7 @@ class TestComputeUnsharderPlan:
def test_sp_in_dims_but_not_in_parallel_info(self) -> None:
"""s(sp) in dims but SP absent from parallel_info (SP disabled), should auto-skip."""
dim_specs = parse_dims("s(sp) b h(tp)")
dim_specs = parse_dims("s(sp) b h(tp)").dims
parallel_infos = [
{ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2)},
{ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2)},
@@ -372,14 +372,14 @@ class TestComputeUnsharderPlan:
def test_all_dims_sharded_but_single_gpu(self) -> None:
"""Single GPU (TP=1, CP=1), dims has s(cp) h(tp) but parallel_info is empty."""
dim_specs = parse_dims("b s(cp) h(tp) d")
dim_specs = parse_dims("b s(cp) h(tp) d").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [{}]
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert plans == []
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)")
dim_specs = parse_dims("s(cp) h(tp)").dims
parallel_infos = [
{
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
@@ -397,7 +397,7 @@ class TestComputeUnsharderPlan:
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")
dim_specs = parse_dims("b s(cp) d").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
@@ -431,7 +431,7 @@ class TestReplicatedAxes:
def test_fully_replicated(self) -> None:
"""CP2 TP2, dims='b h d' → PickPlan(CP) + PickPlan(TP)."""
dim_specs = parse_dims("b h d")
dim_specs = parse_dims("b h d").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
@@ -459,7 +459,7 @@ class TestReplicatedAxes:
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)")
dim_specs = parse_dims("h(tp)").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = []
for cp_rank in range(2):
for ep_rank in range(2):
@@ -486,7 +486,7 @@ class TestReplicatedAxes:
def test_replicated_scrambled_ranks(self) -> None:
"""Scrambled world_rank order with replicated axis."""
dim_specs = parse_dims("h(tp)")
dim_specs = parse_dims("h(tp)").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.CP: AxisInfo(axis_rank=1, axis_size=2),
@@ -515,7 +515,7 @@ class TestReplicatedAxes:
def test_replicated_axis_inconsistent_size_raises(self) -> None:
"""Replicated axis with inconsistent sizes raises ValueError."""
dim_specs = parse_dims("h(tp)")
dim_specs = parse_dims("h(tp)").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
@@ -531,7 +531,7 @@ class TestReplicatedAxes:
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)")
dim_specs = parse_dims("h(tp)").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{
ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2),
@@ -547,7 +547,7 @@ class TestReplicatedAxes:
def test_recompute_pseudo_replicated(self) -> None:
"""RECOMPUTE_PSEUDO with no dim annotation → replicated → PickParams."""
dim_specs = parse_dims("h d")
dim_specs = parse_dims("h d").dims
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{ParallelAxis.RECOMPUTE_PSEUDO: AxisInfo(axis_rank=0, axis_size=2)},
{ParallelAxis.RECOMPUTE_PSEUDO: AxisInfo(axis_rank=1, axis_size=2)},