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
@@ -124,7 +124,7 @@ def compute_per_step_sub_plans(
parallel_infos=parallel_infos, parallel_infos=parallel_infos,
explicit_replicated_axes=replicated_axes, explicit_replicated_axes=replicated_axes,
thd_global_seq_lens=thd_global_seq_lens, thd_global_seq_lens=thd_global_seq_lens,
dp_filtered_axis=dp_axis, dp_filtered_axis=dims_spec.dp_axis,
) )
reorderer_plans = compute_reorderer_plans( reorderer_plans = compute_reorderer_plans(
dim_specs=dim_specs, dim_specs=dim_specs,
@@ -175,7 +175,7 @@ def _load_non_tensor_aux(
loaded: list[ValueWithMeta] = [ loaded: list[ValueWithMeta] = [
ValueWithMeta.load(dump_path / r["filename"]) for r in rows ValueWithMeta.load(dump_path / r["filename"]) for r in rows
] ]
loaded = filter_to_non_empty_dp_rank(loaded) loaded = filter_to_non_empty_dp_rank(loaded, dp_axis=ParallelAxis.DP)
if len(loaded) > 1: if len(loaded) > 1:
first_value = loaded[0].value first_value = loaded[0].value
@@ -212,7 +212,7 @@ def _load_and_align_aux_tensor(
loaded: list[ValueWithMeta] = [ loaded: list[ValueWithMeta] = [
ValueWithMeta.load(dump_path / r["filename"]) for r in rows ValueWithMeta.load(dump_path / r["filename"]) for r in rows
] ]
loaded = filter_to_non_empty_dp_rank(loaded) loaded = filter_to_non_empty_dp_rank(loaded, dp_axis=ParallelAxis.DP)
tensors: list[torch.Tensor] = [ tensors: list[torch.Tensor] = [
item.value for item in loaded if isinstance(item.value, torch.Tensor) item.value for item in loaded if isinstance(item.value, torch.Tensor)
@@ -84,3 +84,11 @@ class DimsSpec(_FrozenBase):
dims: list[DimSpec] dims: list[DimSpec]
dp_group_alias: Optional[str] = None dp_group_alias: Optional[str] = None
replicated_axes: frozenset[ParallelAxis] = frozenset() replicated_axes: frozenset[ParallelAxis] = frozenset()
@property
def dp_axis(self) -> ParallelAxis:
return (
ParallelAxis(self.dp_group_alias)
if self.dp_group_alias
else ParallelAxis.DP
)
@@ -7,18 +7,16 @@ from typing import Optional
import torch import torch
from sglang.srt.debug_utils.comparator.dims_spec import ParallelAxis
from sglang.srt.debug_utils.dump_loader import ValueWithMeta from sglang.srt.debug_utils.dump_loader import ValueWithMeta
_PARALLEL_INFO_KEYS = ("sglang_parallel_info", "megatron_parallel_info") _PARALLEL_INFO_KEYS = ("sglang_parallel_info", "megatron_parallel_info")
_DP_RANK_FIELD = "dp_rank"
_DP_SIZE_FIELD = "dp_size"
def filter_to_non_empty_dp_rank( def filter_to_non_empty_dp_rank(
items: list[ValueWithMeta], items: list[ValueWithMeta],
*, *,
dp_group_alias: Optional[str] = None, dp_axis: ParallelAxis,
) -> list[ValueWithMeta]: ) -> list[ValueWithMeta]:
"""Filter items to the single non-empty dp_rank. """Filter items to the single non-empty dp_rank.
@@ -26,16 +24,15 @@ def filter_to_non_empty_dp_rank(
- dp_size > 1: group by dp_rank, assert exactly one group has non-empty - dp_size > 1: group by dp_rank, assert exactly one group has non-empty
tensors, return that group. tensors, return that group.
When *dp_group_alias* is set (e.g. ``"moe_dp"``), the function looks *dp_axis* determines which rank/size fields to look up (e.g.
for ``<alias>_rank`` / ``<alias>_size`` instead of the default ``ParallelAxis.MOE_DP`` → ``moe_dp_rank`` / ``moe_dp_size``).
``dp_rank`` / ``dp_size``. If the aliased fields are absent the If the fields are absent the filter is a noop (items returned unchanged).
filter is a noop (items returned unchanged).
""" """
if not items: if not items:
return items return items
dp_info: Optional[tuple[int, int]] = _extract_dp_info( dp_info: Optional[tuple[int, int]] = _extract_dp_info(
items[0].meta, dp_group_alias=dp_group_alias items[0].meta, dp_axis=dp_axis
) )
if dp_info is None: if dp_info is None:
return items return items
@@ -51,7 +48,7 @@ def filter_to_non_empty_dp_rank(
groups: dict[int, list[ValueWithMeta]] = defaultdict(list) groups: dict[int, list[ValueWithMeta]] = defaultdict(list)
for item in items: for item in items:
item_dp: Optional[tuple[int, int]] = _extract_dp_info( item_dp: Optional[tuple[int, int]] = _extract_dp_info(
item.meta, dp_group_alias=dp_group_alias item.meta, dp_axis=dp_axis
) )
rank: int = item_dp[0] if item_dp is not None else 0 rank: int = item_dp[0] if item_dp is not None else 0
groups[rank].append(item) groups[rank].append(item)
@@ -71,15 +68,16 @@ def filter_to_non_empty_dp_rank(
def _extract_dp_info( def _extract_dp_info(
meta: dict, meta: dict,
*, *,
dp_group_alias: Optional[str] = None, dp_axis: ParallelAxis,
) -> Optional[tuple[int, int]]: ) -> Optional[tuple[int, int]]:
"""Extract (dp_rank, dp_size) from meta's parallel_info block. """Extract (dp_rank, dp_size) from meta's parallel_info block.
When *dp_group_alias* is given, look for ``<alias>_rank``/``<alias>_size`` *dp_axis* determines which fields to look up: e.g.
instead of the default ``dp_rank``/``dp_size``. ``ParallelAxis.DP`` → ``dp_rank``/``dp_size``,
``ParallelAxis.MOE_DP`` → ``moe_dp_rank``/``moe_dp_size``.
""" """
rank_field: str = f"{dp_group_alias}_rank" if dp_group_alias else _DP_RANK_FIELD rank_field: str = f"{dp_axis.value}_rank"
size_field: str = f"{dp_group_alias}_size" if dp_group_alias else _DP_SIZE_FIELD size_field: str = f"{dp_axis.value}_size"
for key in _PARALLEL_INFO_KEYS: for key in _PARALLEL_INFO_KEYS:
info = meta.get(key) info = meta.get(key)
@@ -3,6 +3,7 @@ import sys
import pytest import pytest
import torch import torch
from sglang.srt.debug_utils.comparator.dims_spec import ParallelAxis
from sglang.srt.debug_utils.comparator.dp_utils import ( from sglang.srt.debug_utils.comparator.dp_utils import (
_extract_dp_info, _extract_dp_info,
_group_has_data, _group_has_data,
@@ -52,18 +53,18 @@ def _make_item(value: object, meta: dict) -> ValueWithMeta:
class TestExtractDpInfo: class TestExtractDpInfo:
def test_sglang_dp(self) -> None: def test_sglang_dp(self) -> None:
meta: dict = _make_sglang_meta(dp_rank=1, dp_size=4) 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: def test_megatron_dp(self) -> None:
meta: dict = _make_megatron_meta(dp_rank=2, dp_size=8) 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: 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: def test_no_dp_fields(self) -> None:
meta: dict = {"sglang_parallel_info": {"tp_rank": 0, "tp_size": 2}} 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), 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 assert result is items
def test_no_parallel_info_returns_unchanged(self) -> None: def test_no_parallel_info_returns_unchanged(self) -> None:
items: list[ValueWithMeta] = [ items: list[ValueWithMeta] = [
_make_item(value=torch.tensor([1.0]), meta={}), _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 assert result is items
def test_empty_list_returns_empty(self) -> None: 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 == [] assert result == []
def test_dp2_all_non_tensor_returns_unchanged(self) -> None: 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 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 len(result) == 1
assert torch.equal(result[0].value, torch.tensor([1.0, 2.0])) 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 len(result) == 1
assert torch.equal(result[0].value, torch.tensor([3.0, 4.0])) assert torch.equal(result[0].value, torch.tensor([3.0, 4.0]))
@@ -184,7 +197,7 @@ class TestFilterToNonEmptyDpRank:
with pytest.raises( with pytest.raises(
AssertionError, match="Expected exactly 1 non-empty dp_rank" 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: def test_dp2_with_tp2_filters_correctly(self) -> None:
"""DP=2 x TP=2: 4 items total, 2 non-empty from dp_rank=0.""" """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 len(result) == 2
assert torch.equal(result[0].value, torch.tensor([1.0])) 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: class TestExtractDpInfoWithAxis:
def test_alias_found(self) -> None: def test_moe_dp_axis_found(self) -> None:
meta: dict = { meta: dict = {
"sglang_parallel_info": { "sglang_parallel_info": {
"dp_rank": 0, "dp_rank": 0,
@@ -229,20 +244,20 @@ class TestExtractDpInfoWithAlias:
"moe_dp_size": 4, "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) 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) 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: class TestFilterToNonEmptyDpRankWithAxis:
def test_alias_none_unchanged_behavior(self) -> None: def test_dp_axis_unchanged_behavior(self) -> None:
"""dp_group_alias=None → same behavior as before (regression).""" """dp_axis=ParallelAxis.DP → same behavior as default (regression)."""
items: list[ValueWithMeta] = [ items: list[ValueWithMeta] = [
_make_item( _make_item(
value=torch.tensor([1.0, 2.0]), value=torch.tensor([1.0, 2.0]),
@@ -255,14 +270,14 @@ class TestFilterToNonEmptyDpRankWithAlias:
] ]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank( result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
items, dp_group_alias=None items, dp_axis=ParallelAxis.DP
) )
assert len(result) == 1 assert len(result) == 1
assert torch.equal(result[0].value, torch.tensor([1.0, 2.0])) assert torch.equal(result[0].value, torch.tensor([1.0, 2.0]))
def test_alias_group_absent_noop(self) -> None: def test_moe_dp_axis_absent_noop(self) -> None:
"""Alias group not in metadata → noop, return items unchanged.""" """MOE_DP axis fields not in metadata → noop, return items unchanged."""
items: list[ValueWithMeta] = [ items: list[ValueWithMeta] = [
_make_item( _make_item(
value=torch.tensor([1.0]), value=torch.tensor([1.0]),
@@ -275,13 +290,13 @@ class TestFilterToNonEmptyDpRankWithAlias:
] ]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank( 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 assert result is items
def test_alias_size_1_noop(self) -> None: def test_moe_dp_axis_size_1_noop(self) -> None:
"""Alias group present but size=1 → noop.""" """MOE_DP axis present but size=1 → noop."""
meta: dict = { meta: dict = {
"sglang_parallel_info": { "sglang_parallel_info": {
"dp_rank": 0, "dp_rank": 0,
@@ -295,13 +310,13 @@ class TestFilterToNonEmptyDpRankWithAlias:
] ]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank( 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 assert result is items
def test_alias_filters_correctly(self) -> None: def test_moe_dp_axis_filters_correctly(self) -> None:
"""Alias group size=2, one empty rank → correctly filters.""" """MOE_DP axis size=2, one empty rank → correctly filters."""
meta_rank0: dict = { meta_rank0: dict = {
"sglang_parallel_info": { "sglang_parallel_info": {
"dp_rank": 0, "dp_rank": 0,
@@ -324,7 +339,7 @@ class TestFilterToNonEmptyDpRankWithAlias:
] ]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank( 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 assert len(result) == 1