Refactor dp_utils to use ParallelAxis enum in dump comparator (#21028)

This commit is contained in:
fzyzcjy
2026-03-20 22:04:20 +08:00
committed by GitHub
parent 154395ab7d
commit fdbcb8156e
5 changed files with 72 additions and 51 deletions

View File

@@ -3,6 +3,7 @@ import sys
import pytest
import torch
from sglang.srt.debug_utils.comparator.dims_spec import ParallelAxis
from sglang.srt.debug_utils.comparator.dp_utils import (
_extract_dp_info,
_group_has_data,
@@ -52,18 +53,18 @@ def _make_item(value: object, meta: dict) -> ValueWithMeta:
class TestExtractDpInfo:
def test_sglang_dp(self) -> None:
meta: dict = _make_sglang_meta(dp_rank=1, dp_size=4)
assert _extract_dp_info(meta) == (1, 4)
assert _extract_dp_info(meta, dp_axis=ParallelAxis.DP) == (1, 4)
def test_megatron_dp(self) -> None:
meta: dict = _make_megatron_meta(dp_rank=2, dp_size=8)
assert _extract_dp_info(meta) == (2, 8)
assert _extract_dp_info(meta, dp_axis=ParallelAxis.DP) == (2, 8)
def test_no_parallel_info(self) -> None:
assert _extract_dp_info({}) is None
assert _extract_dp_info({}, dp_axis=ParallelAxis.DP) is None
def test_no_dp_fields(self) -> None:
meta: dict = {"sglang_parallel_info": {"tp_rank": 0, "tp_size": 2}}
assert _extract_dp_info(meta) is None
assert _extract_dp_info(meta, dp_axis=ParallelAxis.DP) is None
# ---------------------------------------------------------------------------
@@ -101,18 +102,24 @@ class TestFilterToNonEmptyDpRank:
meta=_make_sglang_meta(dp_size=1),
),
]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(items)
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
items, dp_axis=ParallelAxis.DP
)
assert result is items
def test_no_parallel_info_returns_unchanged(self) -> None:
items: list[ValueWithMeta] = [
_make_item(value=torch.tensor([1.0]), meta={}),
]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(items)
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
items, dp_axis=ParallelAxis.DP
)
assert result is items
def test_empty_list_returns_empty(self) -> None:
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank([])
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
[], dp_axis=ParallelAxis.DP
)
assert result == []
def test_dp2_all_non_tensor_returns_unchanged(self) -> None:
@@ -128,7 +135,9 @@ class TestFilterToNonEmptyDpRank:
),
]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(items)
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
items, dp_axis=ParallelAxis.DP
)
assert result is items
@@ -145,7 +154,9 @@ class TestFilterToNonEmptyDpRank:
),
]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(items)
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
items, dp_axis=ParallelAxis.DP
)
assert len(result) == 1
assert torch.equal(result[0].value, torch.tensor([1.0, 2.0]))
@@ -163,7 +174,9 @@ class TestFilterToNonEmptyDpRank:
),
]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(items)
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
items, dp_axis=ParallelAxis.DP
)
assert len(result) == 1
assert torch.equal(result[0].value, torch.tensor([3.0, 4.0]))
@@ -184,7 +197,7 @@ class TestFilterToNonEmptyDpRank:
with pytest.raises(
AssertionError, match="Expected exactly 1 non-empty dp_rank"
):
filter_to_non_empty_dp_rank(items)
filter_to_non_empty_dp_rank(items, dp_axis=ParallelAxis.DP)
def test_dp2_with_tp2_filters_correctly(self) -> None:
"""DP=2 x TP=2: 4 items total, 2 non-empty from dp_rank=0."""
@@ -207,7 +220,9 @@ class TestFilterToNonEmptyDpRank:
),
]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(items)
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
items, dp_axis=ParallelAxis.DP
)
assert len(result) == 2
assert torch.equal(result[0].value, torch.tensor([1.0]))
@@ -215,12 +230,12 @@ class TestFilterToNonEmptyDpRank:
# ---------------------------------------------------------------------------
# dp_group_alias tests
# dp_axis tests (non-default axis)
# ---------------------------------------------------------------------------
class TestExtractDpInfoWithAlias:
def test_alias_found(self) -> None:
class TestExtractDpInfoWithAxis:
def test_moe_dp_axis_found(self) -> None:
meta: dict = {
"sglang_parallel_info": {
"dp_rank": 0,
@@ -229,20 +244,20 @@ class TestExtractDpInfoWithAlias:
"moe_dp_size": 4,
}
}
assert _extract_dp_info(meta, dp_group_alias="moe_dp") == (1, 4)
assert _extract_dp_info(meta, dp_axis=ParallelAxis.MOE_DP) == (1, 4)
def test_alias_not_found_returns_none(self) -> None:
def test_moe_dp_axis_not_found_returns_none(self) -> None:
meta: dict = _make_sglang_meta(dp_rank=0, dp_size=2)
assert _extract_dp_info(meta, dp_group_alias="moe_dp") is None
assert _extract_dp_info(meta, dp_axis=ParallelAxis.MOE_DP) is None
def test_alias_none_uses_default(self) -> None:
def test_dp_axis_uses_default_fields(self) -> None:
meta: dict = _make_sglang_meta(dp_rank=1, dp_size=4)
assert _extract_dp_info(meta, dp_group_alias=None) == (1, 4)
assert _extract_dp_info(meta, dp_axis=ParallelAxis.DP) == (1, 4)
class TestFilterToNonEmptyDpRankWithAlias:
def test_alias_none_unchanged_behavior(self) -> None:
"""dp_group_alias=None → same behavior as before (regression)."""
class TestFilterToNonEmptyDpRankWithAxis:
def test_dp_axis_unchanged_behavior(self) -> None:
"""dp_axis=ParallelAxis.DP → same behavior as default (regression)."""
items: list[ValueWithMeta] = [
_make_item(
value=torch.tensor([1.0, 2.0]),
@@ -255,14 +270,14 @@ class TestFilterToNonEmptyDpRankWithAlias:
]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
items, dp_group_alias=None
items, dp_axis=ParallelAxis.DP
)
assert len(result) == 1
assert torch.equal(result[0].value, torch.tensor([1.0, 2.0]))
def test_alias_group_absent_noop(self) -> None:
"""Alias group not in metadata → noop, return items unchanged."""
def test_moe_dp_axis_absent_noop(self) -> None:
"""MOE_DP axis fields not in metadata → noop, return items unchanged."""
items: list[ValueWithMeta] = [
_make_item(
value=torch.tensor([1.0]),
@@ -275,13 +290,13 @@ class TestFilterToNonEmptyDpRankWithAlias:
]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
items, dp_group_alias="moe_dp"
items, dp_axis=ParallelAxis.MOE_DP
)
assert result is items
def test_alias_size_1_noop(self) -> None:
"""Alias group present but size=1 → noop."""
def test_moe_dp_axis_size_1_noop(self) -> None:
"""MOE_DP axis present but size=1 → noop."""
meta: dict = {
"sglang_parallel_info": {
"dp_rank": 0,
@@ -295,13 +310,13 @@ class TestFilterToNonEmptyDpRankWithAlias:
]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
items, dp_group_alias="moe_dp"
items, dp_axis=ParallelAxis.MOE_DP
)
assert result is items
def test_alias_filters_correctly(self) -> None:
"""Alias group size=2, one empty rank → correctly filters."""
def test_moe_dp_axis_filters_correctly(self) -> None:
"""MOE_DP axis size=2, one empty rank → correctly filters."""
meta_rank0: dict = {
"sglang_parallel_info": {
"dp_rank": 0,
@@ -324,7 +339,7 @@ class TestFilterToNonEmptyDpRankWithAlias:
]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
items, dp_group_alias="moe_dp"
items, dp_axis=ParallelAxis.MOE_DP
)
assert len(result) == 1