Support dependent axis auto-resolution in dump comparator (#21024)

This commit is contained in:
fzyzcjy
2026-03-20 21:56:39 +08:00
committed by GitHub
parent 2d7a262ca3
commit ecd7e40d20
7 changed files with 888 additions and 11 deletions
@@ -36,16 +36,20 @@ def _make_meta(
tp_size: int = 1,
cp_rank: int = 0,
cp_size: int = 1,
extra_parallel_info: Optional[dict[str, int]] = None,
) -> dict[str, Any]:
meta: dict[str, Any] = {"step": step}
if dims is not None:
meta["dims"] = dims
meta["sglang_parallel_info"] = {
parallel_info: dict[str, int] = {
"tp_rank": tp_rank,
"tp_size": tp_size,
"cp_rank": cp_rank,
"cp_size": cp_size,
}
if extra_parallel_info is not None:
parallel_info.update(extra_parallel_info)
meta["sglang_parallel_info"] = parallel_info
return meta
@@ -235,5 +239,97 @@ class TestComputePerStepSubPlansThd:
assert reorderer_plans[0].params.seq_lens == [100, 64, 92]
class TestComputePerStepSubPlansDpFiltered:
"""Tests that compute_per_step_sub_plans passes dp_filtered_axis to unsharder,
so DP axes already handled by the upstream DP filter don't cause validation errors.
"""
def test_dp2_tp2_does_not_raise(self) -> None:
"""DP2 + TP2, dims='t h[tp]' → should not raise despite DP being active."""
result: list[AlignerPerStepSubPlan] = compute_per_step_sub_plans(
metas=[
_make_meta(
dims="t h[tp]",
tp_rank=0,
tp_size=2,
extra_parallel_info={"dp_rank": 0, "dp_size": 2},
),
_make_meta(
dims="t h[tp]",
tp_rank=1,
tp_size=2,
extra_parallel_info={"dp_rank": 0, "dp_size": 2},
),
]
)
unsharder_plans: list[UnsharderPlan] = [
p for p in result if isinstance(p, UnsharderPlan)
]
assert len(unsharder_plans) == 1
assert unsharder_plans[0].axis.value == "tp"
def test_dp2_only_no_sharding_does_not_raise(self) -> None:
"""DP2 only, dims='t h' → should not raise, no plans produced."""
result: list[AlignerPerStepSubPlan] = compute_per_step_sub_plans(
metas=[
_make_meta(
dims="t h",
extra_parallel_info={"dp_rank": 0, "dp_size": 2},
),
_make_meta(
dims="t h",
extra_parallel_info={"dp_rank": 0, "dp_size": 2},
),
]
)
assert result == []
def test_dp_alias_passes_correct_filtered_axis(self) -> None:
"""dims with '# dp:=moe_dp', metas have moe_dp → should not raise."""
result: list[AlignerPerStepSubPlan] = compute_per_step_sub_plans(
metas=[
_make_meta(
dims="t h[tp] # dp:=moe_dp",
tp_rank=0,
tp_size=2,
extra_parallel_info={"moe_dp_rank": 0, "moe_dp_size": 2},
),
_make_meta(
dims="t h[tp] # dp:=moe_dp",
tp_rank=1,
tp_size=2,
extra_parallel_info={"moe_dp_rank": 0, "moe_dp_size": 2},
),
]
)
unsharder_plans: list[UnsharderPlan] = [
p for p in result if isinstance(p, UnsharderPlan)
]
assert len(unsharder_plans) == 1
assert unsharder_plans[0].axis.value == "tp"
def test_dp2_tp2_cp2_does_not_raise(self) -> None:
"""DP2 + TP2 + CP2, dims='s[cp] h[tp]' → should not raise."""
metas = []
for cp_rank in range(2):
for tp_rank in range(2):
metas.append(
_make_meta(
dims="s[cp] h[tp]",
tp_rank=tp_rank,
tp_size=2,
cp_rank=cp_rank,
cp_size=2,
extra_parallel_info={"dp_rank": 0, "dp_size": 2},
)
)
result: list[AlignerPerStepSubPlan] = compute_per_step_sub_plans(metas=metas)
unsharder_plans: list[UnsharderPlan] = [
p for p in result if isinstance(p, UnsharderPlan)
]
axes = {p.axis.value for p in unsharder_plans}
assert axes == {"cp", "tp"}
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))