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

View File

@@ -214,5 +214,122 @@ class TestFilterToNonEmptyDpRank:
assert torch.equal(result[1].value, torch.tensor([2.0]))
# ---------------------------------------------------------------------------
# dp_group_alias tests
# ---------------------------------------------------------------------------
class TestExtractDpInfoWithAlias:
def test_alias_found(self) -> None:
meta: dict = {
"sglang_parallel_info": {
"dp_rank": 0,
"dp_size": 2,
"moe_dp_rank": 1,
"moe_dp_size": 4,
}
}
assert _extract_dp_info(meta, dp_group_alias="moe_dp") == (1, 4)
def test_alias_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
def test_alias_none_uses_default(self) -> None:
meta: dict = _make_sglang_meta(dp_rank=1, dp_size=4)
assert _extract_dp_info(meta, dp_group_alias=None) == (1, 4)
class TestFilterToNonEmptyDpRankWithAlias:
def test_alias_none_unchanged_behavior(self) -> None:
"""dp_group_alias=None → same behavior as before (regression)."""
items: list[ValueWithMeta] = [
_make_item(
value=torch.tensor([1.0, 2.0]),
meta=_make_sglang_meta(dp_rank=0, dp_size=2),
),
_make_item(
value=torch.tensor([]),
meta=_make_sglang_meta(dp_rank=1, dp_size=2),
),
]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
items, dp_group_alias=None
)
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."""
items: list[ValueWithMeta] = [
_make_item(
value=torch.tensor([1.0]),
meta=_make_sglang_meta(dp_rank=0, dp_size=2),
),
_make_item(
value=torch.tensor([2.0]),
meta=_make_sglang_meta(dp_rank=1, dp_size=2),
),
]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
items, dp_group_alias="moe_dp"
)
assert result is items
def test_alias_size_1_noop(self) -> None:
"""Alias group present but size=1 → noop."""
meta: dict = {
"sglang_parallel_info": {
"dp_rank": 0,
"dp_size": 2,
"moe_dp_rank": 0,
"moe_dp_size": 1,
}
}
items: list[ValueWithMeta] = [
_make_item(value=torch.tensor([1.0]), meta=meta),
]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
items, dp_group_alias="moe_dp"
)
assert result is items
def test_alias_filters_correctly(self) -> None:
"""Alias group size=2, one empty rank → correctly filters."""
meta_rank0: dict = {
"sglang_parallel_info": {
"dp_rank": 0,
"dp_size": 2,
"moe_dp_rank": 0,
"moe_dp_size": 2,
}
}
meta_rank1: dict = {
"sglang_parallel_info": {
"dp_rank": 0,
"dp_size": 2,
"moe_dp_rank": 1,
"moe_dp_size": 2,
}
}
items: list[ValueWithMeta] = [
_make_item(value=torch.tensor([1.0, 2.0]), meta=meta_rank0),
_make_item(value=torch.tensor([]), meta=meta_rank1),
]
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(
items, dp_group_alias="moe_dp"
)
assert len(result) == 1
assert torch.equal(result[0].value, torch.tensor([1.0, 2.0]))
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))