Support dependent axis auto-resolution in dump comparator (#21024)
This commit is contained in:
@@ -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__]))
|
||||
|
||||
Reference in New Issue
Block a user