Support data parallel in dump comparator (#19596)
This commit is contained in:
@@ -30,6 +30,7 @@ from sglang.srt.debug_utils.comparator.dims import (
|
|||||||
apply_dim_names,
|
apply_dim_names,
|
||||||
resolve_dim_names,
|
resolve_dim_names,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.debug_utils.comparator.dp_utils import filter_to_non_empty_dp_rank
|
||||||
from sglang.srt.debug_utils.comparator.output_types import GeneralWarning
|
from sglang.srt.debug_utils.comparator.output_types import GeneralWarning
|
||||||
from sglang.srt.debug_utils.comparator.warning_sink import warning_sink
|
from sglang.srt.debug_utils.comparator.warning_sink import warning_sink
|
||||||
from sglang.srt.debug_utils.dump_loader import ValueWithMeta, filter_rows
|
from sglang.srt.debug_utils.dump_loader import ValueWithMeta, filter_rows
|
||||||
@@ -170,6 +171,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)
|
||||||
|
|
||||||
if len(loaded) > 1:
|
if len(loaded) > 1:
|
||||||
first_value = loaded[0].value
|
first_value = loaded[0].value
|
||||||
@@ -206,6 +208,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)
|
||||||
|
|
||||||
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)
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ from sglang.srt.debug_utils.comparator.dims import (
|
|||||||
apply_dim_names,
|
apply_dim_names,
|
||||||
resolve_dim_names,
|
resolve_dim_names,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.debug_utils.comparator.dp_utils import filter_to_non_empty_dp_rank
|
||||||
from sglang.srt.debug_utils.comparator.output_types import (
|
from sglang.srt.debug_utils.comparator.output_types import (
|
||||||
ComparisonRecord,
|
ComparisonRecord,
|
||||||
GeneralWarning,
|
GeneralWarning,
|
||||||
@@ -94,6 +95,12 @@ def _compare_bundle_pair_inner(
|
|||||||
reason = "baseline_load_failed" if not all_pair.x else "target_load_failed"
|
reason = "baseline_load_failed" if not all_pair.x else "target_load_failed"
|
||||||
return SkipRecord(name=name, reason=reason)
|
return SkipRecord(name=name, reason=reason)
|
||||||
|
|
||||||
|
# 1b. DP filter: keep only the non-empty dp_rank
|
||||||
|
all_pair = Pair(
|
||||||
|
x=filter_to_non_empty_dp_rank(all_pair.x),
|
||||||
|
y=filter_to_non_empty_dp_rank(all_pair.y),
|
||||||
|
)
|
||||||
|
|
||||||
# 2. Check if any side has non-tensor values → non-tensor display path
|
# 2. Check if any side has non-tensor values → non-tensor display path
|
||||||
has_non_tensor: bool = any(
|
has_non_tensor: bool = any(
|
||||||
not isinstance(it.value, torch.Tensor) for it in [*all_pair.x, *all_pair.y]
|
not isinstance(it.value, torch.Tensor) for it in [*all_pair.x, *all_pair.y]
|
||||||
|
|||||||
@@ -0,0 +1,78 @@
|
|||||||
|
"""DP filtering: keep only the non-empty dp_rank items."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections import defaultdict
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.debug_utils.dump_loader import ValueWithMeta
|
||||||
|
|
||||||
|
_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(items: list[ValueWithMeta]) -> list[ValueWithMeta]:
|
||||||
|
"""Filter items to the single non-empty dp_rank.
|
||||||
|
|
||||||
|
- dp_size <= 1: return items unchanged.
|
||||||
|
- dp_size > 1: group by dp_rank, assert exactly one group has non-empty
|
||||||
|
tensors, return that group.
|
||||||
|
"""
|
||||||
|
if not items:
|
||||||
|
return items
|
||||||
|
|
||||||
|
dp_info: Optional[tuple[int, int]] = _extract_dp_info(items[0].meta)
|
||||||
|
if dp_info is None:
|
||||||
|
return items
|
||||||
|
|
||||||
|
_dp_rank, dp_size = dp_info
|
||||||
|
if dp_size <= 1:
|
||||||
|
return items
|
||||||
|
|
||||||
|
has_any_tensor: bool = any(isinstance(item.value, torch.Tensor) for item in items)
|
||||||
|
if not has_any_tensor:
|
||||||
|
return items
|
||||||
|
|
||||||
|
groups: dict[int, list[ValueWithMeta]] = defaultdict(list)
|
||||||
|
for item in items:
|
||||||
|
item_dp: Optional[tuple[int, int]] = _extract_dp_info(item.meta)
|
||||||
|
rank: int = item_dp[0] if item_dp is not None else 0
|
||||||
|
groups[rank].append(item)
|
||||||
|
|
||||||
|
non_empty_ranks: list[int] = [
|
||||||
|
rank for rank, group in groups.items() if _group_has_data(group)
|
||||||
|
]
|
||||||
|
|
||||||
|
assert len(non_empty_ranks) == 1, (
|
||||||
|
f"Expected exactly 1 non-empty dp_rank, got {len(non_empty_ranks)}: "
|
||||||
|
f"ranks={non_empty_ranks}"
|
||||||
|
)
|
||||||
|
|
||||||
|
return groups[non_empty_ranks[0]]
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_dp_info(meta: dict) -> Optional[tuple[int, int]]:
|
||||||
|
"""Extract (dp_rank, dp_size) from meta's parallel_info block."""
|
||||||
|
for key in _PARALLEL_INFO_KEYS:
|
||||||
|
info = meta.get(key)
|
||||||
|
if not isinstance(info, dict) or not info:
|
||||||
|
continue
|
||||||
|
|
||||||
|
dp_rank = info.get(_DP_RANK_FIELD)
|
||||||
|
dp_size = info.get(_DP_SIZE_FIELD)
|
||||||
|
if dp_rank is not None and dp_size is not None:
|
||||||
|
return (int(dp_rank), int(dp_size))
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _group_has_data(group: list[ValueWithMeta]) -> bool:
|
||||||
|
"""Check if any tensor in the group is non-empty (numel > 0)."""
|
||||||
|
return any(
|
||||||
|
isinstance(item.value, torch.Tensor) and item.value.numel() > 0
|
||||||
|
for item in group
|
||||||
|
)
|
||||||
@@ -291,5 +291,98 @@ class TestLoadAndAlignAuxTensor:
|
|||||||
assert "aux_no_dims" in warnings[0].category
|
assert "aux_no_dims" in warnings[0].category
|
||||||
|
|
||||||
|
|
||||||
|
class TestLoadNonTensorAuxDp:
|
||||||
|
"""DP filtering in _load_non_tensor_aux."""
|
||||||
|
|
||||||
|
def test_dp2_non_tensor_returns_value(self, tmp_path: Path) -> None:
|
||||||
|
"""DP=2 non-tensor aux: both ranks have same value, filter keeps all (non-tensor)."""
|
||||||
|
fn0: str = _save_pt(
|
||||||
|
tmp_path,
|
||||||
|
name="rids",
|
||||||
|
step=0,
|
||||||
|
rank=0,
|
||||||
|
value=["req_A"],
|
||||||
|
meta={
|
||||||
|
"sglang_parallel_info": {
|
||||||
|
"dp_rank": 0,
|
||||||
|
"dp_size": 2,
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
fn1: str = _save_pt(
|
||||||
|
tmp_path,
|
||||||
|
name="rids",
|
||||||
|
step=0,
|
||||||
|
rank=1,
|
||||||
|
value=["req_A"],
|
||||||
|
meta={
|
||||||
|
"sglang_parallel_info": {
|
||||||
|
"dp_rank": 1,
|
||||||
|
"dp_size": 2,
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
df: pl.DataFrame = _make_df_from_filenames([fn0, fn1])
|
||||||
|
|
||||||
|
sink = WarningSink()
|
||||||
|
with sink.context():
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.debug_utils.comparator.aligner.token_aligner.aux_loader.warning_sink",
|
||||||
|
sink,
|
||||||
|
):
|
||||||
|
result = _load_non_tensor_aux(
|
||||||
|
name="rids", step=0, df=df, dump_path=tmp_path
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result == ["req_A"]
|
||||||
|
|
||||||
|
|
||||||
|
class TestLoadAndAlignAuxTensorDp:
|
||||||
|
"""DP filtering in _load_and_align_aux_tensor."""
|
||||||
|
|
||||||
|
def test_dp2_tensor_one_empty(self, tmp_path: Path) -> None:
|
||||||
|
"""DP=2 tensor aux: rank 0 has data, rank 1 empty -> returns rank 0 tensor."""
|
||||||
|
fn0: str = _save_pt(
|
||||||
|
tmp_path,
|
||||||
|
name="input_ids",
|
||||||
|
step=0,
|
||||||
|
rank=0,
|
||||||
|
value=torch.tensor([10, 20, 30]),
|
||||||
|
meta={
|
||||||
|
"sglang_parallel_info": {
|
||||||
|
"dp_rank": 0,
|
||||||
|
"dp_size": 2,
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
fn1: str = _save_pt(
|
||||||
|
tmp_path,
|
||||||
|
name="input_ids",
|
||||||
|
step=0,
|
||||||
|
rank=1,
|
||||||
|
value=torch.tensor([]),
|
||||||
|
meta={
|
||||||
|
"sglang_parallel_info": {
|
||||||
|
"dp_rank": 1,
|
||||||
|
"dp_size": 2,
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
df: pl.DataFrame = _make_df_from_filenames([fn0, fn1])
|
||||||
|
|
||||||
|
result = _load_and_align_aux_tensor(
|
||||||
|
name="input_ids",
|
||||||
|
step=0,
|
||||||
|
df=df,
|
||||||
|
dump_path=tmp_path,
|
||||||
|
plugin=_sglang_plugin,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result is not None
|
||||||
|
assert torch.equal(result, torch.tensor([10, 20, 30]))
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__]))
|
sys.exit(pytest.main([__file__]))
|
||||||
|
|||||||
@@ -0,0 +1,218 @@
|
|||||||
|
import sys
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.debug_utils.comparator.dp_utils import (
|
||||||
|
_extract_dp_info,
|
||||||
|
_group_has_data,
|
||||||
|
filter_to_non_empty_dp_rank,
|
||||||
|
)
|
||||||
|
from sglang.srt.debug_utils.dump_loader import ValueWithMeta
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=15, suite="default", nightly=True)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_sglang_meta(
|
||||||
|
*, tp_rank: int = 0, tp_size: int = 1, dp_rank: int = 0, dp_size: int = 1
|
||||||
|
) -> dict:
|
||||||
|
return {
|
||||||
|
"sglang_parallel_info": {
|
||||||
|
"tp_rank": tp_rank,
|
||||||
|
"tp_size": tp_size,
|
||||||
|
"dp_rank": dp_rank,
|
||||||
|
"dp_size": dp_size,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _make_megatron_meta(
|
||||||
|
*, tp_rank: int = 0, tp_size: int = 1, dp_rank: int = 0, dp_size: int = 1
|
||||||
|
) -> dict:
|
||||||
|
return {
|
||||||
|
"megatron_parallel_info": {
|
||||||
|
"tp_rank": tp_rank,
|
||||||
|
"tp_size": tp_size,
|
||||||
|
"dp_rank": dp_rank,
|
||||||
|
"dp_size": dp_size,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _make_item(value: object, meta: dict) -> ValueWithMeta:
|
||||||
|
return ValueWithMeta(value=value, meta=meta)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# _extract_dp_info
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
def test_megatron_dp(self) -> None:
|
||||||
|
meta: dict = _make_megatron_meta(dp_rank=2, dp_size=8)
|
||||||
|
assert _extract_dp_info(meta) == (2, 8)
|
||||||
|
|
||||||
|
def test_no_parallel_info(self) -> None:
|
||||||
|
assert _extract_dp_info({}) 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
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# _group_has_data
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
class TestGroupHasData:
|
||||||
|
def test_non_empty_tensor(self) -> None:
|
||||||
|
item: ValueWithMeta = _make_item(value=torch.tensor([1, 2, 3]), meta={})
|
||||||
|
assert _group_has_data([item]) is True
|
||||||
|
|
||||||
|
def test_empty_tensor(self) -> None:
|
||||||
|
item: ValueWithMeta = _make_item(value=torch.tensor([]), meta={})
|
||||||
|
assert _group_has_data([item]) is False
|
||||||
|
|
||||||
|
def test_non_tensor_value(self) -> None:
|
||||||
|
item: ValueWithMeta = _make_item(value="hello", meta={})
|
||||||
|
assert _group_has_data([item]) is False
|
||||||
|
|
||||||
|
def test_empty_group(self) -> None:
|
||||||
|
assert _group_has_data([]) is False
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# filter_to_non_empty_dp_rank
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
class TestFilterToNonEmptyDpRank:
|
||||||
|
def test_dp_size_1_returns_unchanged(self) -> None:
|
||||||
|
items: list[ValueWithMeta] = [
|
||||||
|
_make_item(
|
||||||
|
value=torch.tensor([1.0]),
|
||||||
|
meta=_make_sglang_meta(dp_size=1),
|
||||||
|
),
|
||||||
|
]
|
||||||
|
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(items)
|
||||||
|
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)
|
||||||
|
assert result is items
|
||||||
|
|
||||||
|
def test_empty_list_returns_empty(self) -> None:
|
||||||
|
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank([])
|
||||||
|
assert result == []
|
||||||
|
|
||||||
|
def test_dp2_all_non_tensor_returns_unchanged(self) -> None:
|
||||||
|
"""DP=2 with non-tensor values: skip filtering, return unchanged."""
|
||||||
|
items: list[ValueWithMeta] = [
|
||||||
|
_make_item(
|
||||||
|
value=["req_A"],
|
||||||
|
meta=_make_sglang_meta(dp_rank=0, dp_size=2),
|
||||||
|
),
|
||||||
|
_make_item(
|
||||||
|
value=["req_A"],
|
||||||
|
meta=_make_sglang_meta(dp_rank=1, dp_size=2),
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(items)
|
||||||
|
|
||||||
|
assert result is items
|
||||||
|
|
||||||
|
def test_dp2_one_empty_one_nonempty_sglang(self) -> None:
|
||||||
|
"""DP=2, rank 0 has data, rank 1 has empty tensor."""
|
||||||
|
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)
|
||||||
|
|
||||||
|
assert len(result) == 1
|
||||||
|
assert torch.equal(result[0].value, torch.tensor([1.0, 2.0]))
|
||||||
|
|
||||||
|
def test_dp2_one_empty_one_nonempty_megatron(self) -> None:
|
||||||
|
"""DP=2 megatron, rank 1 has data, rank 0 has empty tensor."""
|
||||||
|
items: list[ValueWithMeta] = [
|
||||||
|
_make_item(
|
||||||
|
value=torch.tensor([]),
|
||||||
|
meta=_make_megatron_meta(dp_rank=0, dp_size=2),
|
||||||
|
),
|
||||||
|
_make_item(
|
||||||
|
value=torch.tensor([3.0, 4.0]),
|
||||||
|
meta=_make_megatron_meta(dp_rank=1, dp_size=2),
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(items)
|
||||||
|
|
||||||
|
assert len(result) == 1
|
||||||
|
assert torch.equal(result[0].value, torch.tensor([3.0, 4.0]))
|
||||||
|
|
||||||
|
def test_dp2_both_nonempty_raises(self) -> None:
|
||||||
|
"""DP=2, both ranks have data: assertion error."""
|
||||||
|
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),
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
with pytest.raises(
|
||||||
|
AssertionError, match="Expected exactly 1 non-empty dp_rank"
|
||||||
|
):
|
||||||
|
filter_to_non_empty_dp_rank(items)
|
||||||
|
|
||||||
|
def test_dp2_with_tp2_filters_correctly(self) -> None:
|
||||||
|
"""DP=2 x TP=2: 4 items total, 2 non-empty from dp_rank=0."""
|
||||||
|
items: list[ValueWithMeta] = [
|
||||||
|
_make_item(
|
||||||
|
value=torch.tensor([1.0]),
|
||||||
|
meta=_make_sglang_meta(tp_rank=0, tp_size=2, dp_rank=0, dp_size=2),
|
||||||
|
),
|
||||||
|
_make_item(
|
||||||
|
value=torch.tensor([2.0]),
|
||||||
|
meta=_make_sglang_meta(tp_rank=1, tp_size=2, dp_rank=0, dp_size=2),
|
||||||
|
),
|
||||||
|
_make_item(
|
||||||
|
value=torch.tensor([]),
|
||||||
|
meta=_make_sglang_meta(tp_rank=0, tp_size=2, dp_rank=1, dp_size=2),
|
||||||
|
),
|
||||||
|
_make_item(
|
||||||
|
value=torch.tensor([]),
|
||||||
|
meta=_make_sglang_meta(tp_rank=1, tp_size=2, dp_rank=1, dp_size=2),
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
result: list[ValueWithMeta] = filter_to_non_empty_dp_rank(items)
|
||||||
|
|
||||||
|
assert len(result) == 2
|
||||||
|
assert torch.equal(result[0].value, torch.tensor([1.0]))
|
||||||
|
assert torch.equal(result[1].value, torch.tensor([2.0]))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(pytest.main([__file__]))
|
||||||
@@ -2525,5 +2525,216 @@ class TestEntrypointThdCpZigzag:
|
|||||||
assert all(c.diff is not None and c.diff.passed for c in hidden_comparisons)
|
assert all(c.diff is not None and c.diff.passed for c in hidden_comparisons)
|
||||||
|
|
||||||
|
|
||||||
|
class TestEntrypointDpFilter:
|
||||||
|
"""E2E tests for DP (data parallel) filtering.
|
||||||
|
|
||||||
|
When DP > 1, only one dp_rank has non-empty tensors; the others
|
||||||
|
dump empty (numel=0) tensors. The comparator should filter out the
|
||||||
|
empty dp_rank items and produce correct comparison results.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def test_dp2_sglang_both_sides(self, tmp_path: Path, capsys) -> None:
|
||||||
|
"""DP=2 sglang: both baseline and target have 1 non-empty + 1 empty dp_rank."""
|
||||||
|
torch.manual_seed(42)
|
||||||
|
tensor_data: torch.Tensor = torch.randn(10, 8)
|
||||||
|
target_data: torch.Tensor = tensor_data + torch.randn(10, 8) * 0.001
|
||||||
|
|
||||||
|
for side, side_dir_name, data in [
|
||||||
|
("baseline", "baseline", tensor_data),
|
||||||
|
("target", "target", target_data),
|
||||||
|
]:
|
||||||
|
side_dir: Path = tmp_path / side_dir_name
|
||||||
|
side_dir.mkdir()
|
||||||
|
|
||||||
|
# dp_rank=0: non-empty tensor
|
||||||
|
_create_rank_dump(
|
||||||
|
side_dir,
|
||||||
|
rank=0,
|
||||||
|
name="hidden",
|
||||||
|
tensor=data,
|
||||||
|
dims="t h",
|
||||||
|
parallel_info={
|
||||||
|
"tp_rank": 0,
|
||||||
|
"tp_size": 1,
|
||||||
|
"dp_rank": 0,
|
||||||
|
"dp_size": 2,
|
||||||
|
},
|
||||||
|
framework="sglang",
|
||||||
|
)
|
||||||
|
|
||||||
|
# dp_rank=1: empty tensor
|
||||||
|
_create_rank_dump(
|
||||||
|
side_dir,
|
||||||
|
rank=1,
|
||||||
|
name="hidden",
|
||||||
|
tensor=torch.empty(0, 8),
|
||||||
|
dims="t h",
|
||||||
|
parallel_info={
|
||||||
|
"tp_rank": 0,
|
||||||
|
"tp_size": 1,
|
||||||
|
"dp_rank": 1,
|
||||||
|
"dp_size": 2,
|
||||||
|
},
|
||||||
|
framework="sglang",
|
||||||
|
)
|
||||||
|
|
||||||
|
args: Namespace = _make_args(
|
||||||
|
tmp_path / "baseline" / _FIXED_EXP_NAME,
|
||||||
|
tmp_path / "target" / _FIXED_EXP_NAME,
|
||||||
|
grouping="logical",
|
||||||
|
diff_threshold=1e-3,
|
||||||
|
)
|
||||||
|
records: list[AnyRecord] = _run_and_parse(args, capsys)
|
||||||
|
|
||||||
|
comparison: ComparisonRecord = _assert_single_comparison_passed(records)
|
||||||
|
assert comparison.name == "hidden"
|
||||||
|
|
||||||
|
def test_dp2_megatron_both_sides(self, tmp_path: Path, capsys) -> None:
|
||||||
|
"""DP=2 megatron: both baseline and target have 1 non-empty + 1 empty dp_rank."""
|
||||||
|
torch.manual_seed(42)
|
||||||
|
tensor_data: torch.Tensor = torch.randn(10, 8)
|
||||||
|
target_data: torch.Tensor = tensor_data + torch.randn(10, 8) * 0.001
|
||||||
|
|
||||||
|
for side, side_dir_name, data in [
|
||||||
|
("baseline", "baseline", tensor_data),
|
||||||
|
("target", "target", target_data),
|
||||||
|
]:
|
||||||
|
side_dir: Path = tmp_path / side_dir_name
|
||||||
|
side_dir.mkdir()
|
||||||
|
|
||||||
|
# dp_rank=0: non-empty tensor
|
||||||
|
_create_rank_dump(
|
||||||
|
side_dir,
|
||||||
|
rank=0,
|
||||||
|
name="hidden",
|
||||||
|
tensor=data,
|
||||||
|
dims="t h",
|
||||||
|
parallel_info={
|
||||||
|
"tp_rank": 0,
|
||||||
|
"tp_size": 1,
|
||||||
|
"dp_rank": 0,
|
||||||
|
"dp_size": 2,
|
||||||
|
},
|
||||||
|
framework="megatron",
|
||||||
|
)
|
||||||
|
|
||||||
|
# dp_rank=1: empty tensor
|
||||||
|
_create_rank_dump(
|
||||||
|
side_dir,
|
||||||
|
rank=1,
|
||||||
|
name="hidden",
|
||||||
|
tensor=torch.empty(0, 8),
|
||||||
|
dims="t h",
|
||||||
|
parallel_info={
|
||||||
|
"tp_rank": 0,
|
||||||
|
"tp_size": 1,
|
||||||
|
"dp_rank": 1,
|
||||||
|
"dp_size": 2,
|
||||||
|
},
|
||||||
|
framework="megatron",
|
||||||
|
)
|
||||||
|
|
||||||
|
args: Namespace = _make_args(
|
||||||
|
tmp_path / "baseline" / _FIXED_EXP_NAME,
|
||||||
|
tmp_path / "target" / _FIXED_EXP_NAME,
|
||||||
|
grouping="logical",
|
||||||
|
diff_threshold=1e-3,
|
||||||
|
)
|
||||||
|
records: list[AnyRecord] = _run_and_parse(args, capsys)
|
||||||
|
|
||||||
|
comparison: ComparisonRecord = _assert_single_comparison_passed(records)
|
||||||
|
assert comparison.name == "hidden"
|
||||||
|
|
||||||
|
def test_dp2_tp2_sglang(self, tmp_path: Path, capsys) -> None:
|
||||||
|
"""DP=2 x TP=2 sglang: 4 ranks, dp_rank=0 has data, dp_rank=1 empty."""
|
||||||
|
torch.manual_seed(42)
|
||||||
|
full_tensor: torch.Tensor = torch.randn(10, 8)
|
||||||
|
tp_chunks: list[torch.Tensor] = list(full_tensor.chunk(2, dim=1))
|
||||||
|
|
||||||
|
target_full: torch.Tensor = full_tensor + torch.randn(10, 8) * 0.001
|
||||||
|
target_tp_chunks: list[torch.Tensor] = list(target_full.chunk(2, dim=1))
|
||||||
|
|
||||||
|
for side, side_dir_name, chunks in [
|
||||||
|
("baseline", "baseline", tp_chunks),
|
||||||
|
("target", "target", target_tp_chunks),
|
||||||
|
]:
|
||||||
|
side_dir: Path = tmp_path / side_dir_name
|
||||||
|
side_dir.mkdir()
|
||||||
|
|
||||||
|
rank: int = 0
|
||||||
|
for dp_rank in range(2):
|
||||||
|
for tp_rank in range(2):
|
||||||
|
tensor: torch.Tensor = (
|
||||||
|
chunks[tp_rank] if dp_rank == 0 else torch.empty(0, 4)
|
||||||
|
)
|
||||||
|
_create_rank_dump(
|
||||||
|
side_dir,
|
||||||
|
rank=rank,
|
||||||
|
name="hidden",
|
||||||
|
tensor=tensor,
|
||||||
|
dims="t h(tp)",
|
||||||
|
parallel_info={
|
||||||
|
"tp_rank": tp_rank,
|
||||||
|
"tp_size": 2,
|
||||||
|
"dp_rank": dp_rank,
|
||||||
|
"dp_size": 2,
|
||||||
|
},
|
||||||
|
framework="sglang",
|
||||||
|
)
|
||||||
|
rank += 1
|
||||||
|
|
||||||
|
args: Namespace = _make_args(
|
||||||
|
tmp_path / "baseline" / _FIXED_EXP_NAME,
|
||||||
|
tmp_path / "target" / _FIXED_EXP_NAME,
|
||||||
|
grouping="logical",
|
||||||
|
diff_threshold=1e-3,
|
||||||
|
)
|
||||||
|
records: list[AnyRecord] = _run_and_parse(args, capsys)
|
||||||
|
|
||||||
|
comparison: ComparisonRecord = _assert_single_comparison_passed(records)
|
||||||
|
assert comparison.name == "hidden"
|
||||||
|
|
||||||
|
def test_dp2_both_nonempty_raises(self, tmp_path: Path, capsys) -> None:
|
||||||
|
"""DP=2 sglang: both dp_rank=0 and dp_rank=1 have non-empty tensors => AssertionError."""
|
||||||
|
torch.manual_seed(42)
|
||||||
|
tensor_data: torch.Tensor = torch.randn(10, 8)
|
||||||
|
target_data: torch.Tensor = tensor_data + torch.randn(10, 8) * 0.001
|
||||||
|
|
||||||
|
for side, side_dir_name, data in [
|
||||||
|
("baseline", "baseline", tensor_data),
|
||||||
|
("target", "target", target_data),
|
||||||
|
]:
|
||||||
|
side_dir: Path = tmp_path / side_dir_name
|
||||||
|
side_dir.mkdir()
|
||||||
|
|
||||||
|
for dp_rank in range(2):
|
||||||
|
_create_rank_dump(
|
||||||
|
side_dir,
|
||||||
|
rank=dp_rank,
|
||||||
|
name="hidden",
|
||||||
|
tensor=data,
|
||||||
|
dims="t h",
|
||||||
|
parallel_info={
|
||||||
|
"tp_rank": 0,
|
||||||
|
"tp_size": 1,
|
||||||
|
"dp_rank": dp_rank,
|
||||||
|
"dp_size": 2,
|
||||||
|
},
|
||||||
|
framework="sglang",
|
||||||
|
)
|
||||||
|
|
||||||
|
args: Namespace = _make_args(
|
||||||
|
tmp_path / "baseline" / _FIXED_EXP_NAME,
|
||||||
|
tmp_path / "target" / _FIXED_EXP_NAME,
|
||||||
|
grouping="logical",
|
||||||
|
diff_threshold=1e-3,
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(
|
||||||
|
AssertionError, match="Expected exactly 1 non-empty dp_rank"
|
||||||
|
):
|
||||||
|
_run_and_parse(args, capsys)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__]))
|
sys.exit(pytest.main([__file__]))
|
||||||
|
|||||||
Reference in New Issue
Block a user