Support data parallel attention in dump comparator (#19602)
This commit is contained in:
@@ -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__]))
|
||||
|
||||
Reference in New Issue
Block a user