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
@@ -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(