Handle recompute and verify closeness in dumper (#19564)

This commit is contained in:
fzyzcjy
2026-02-28 18:07:44 +08:00
committed by GitHub
parent 63a4778542
commit 40facdb28c
11 changed files with 557 additions and 13 deletions

View File

@@ -506,6 +506,138 @@ class TestVerifyReplicatedGroup:
assert len(warnings) == 1
assert warnings[0].differing_index == 1
def test_recompute_pseudo_mismatch_warns(self) -> None:
"""_verify_replicated_group produces warning for RECOMPUTE_PSEUDO axis mismatch."""
tensor_a = torch.ones(4)
tensor_b = torch.ones(4) + 0.1
with warning_sink.context() as warnings:
_verify_replicated_group(
[tensor_a, tensor_b],
axis=ParallelAxis.RECOMPUTE_PSEUDO,
group_index=0,
)
assert len(warnings) == 1
assert warnings[0].axis == "recompute_pseudo"
assert warnings[0].group_index == 0
assert warnings[0].differing_index == 1
assert warnings[0].baseline_index == 0
assert warnings[0].max_abs_diff == pytest.approx(0.1, abs=1e-5)
class TestThdCpConcat:
def test_single_seq(self) -> None:
"""Single seq THD unshard: 2 ranks → per-seq concat."""
rank0 = torch.tensor([1, 2, 3]).refine_names("t")
rank1 = torch.tensor([4, 5, 6]).refine_names("t")
plan = UnsharderPlan(
axis=ParallelAxis.CP,
params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[3]),
groups=[[0, 1]],
)
with warning_sink.context():
result = execute_unsharder_plan(plan, [rank0, rank1])
assert len(result) == 1
expected = torch.tensor([1, 2, 3, 4, 5, 6])
assert torch.equal(result[0].rename(None), expected)
def test_multi_seq(self) -> None:
"""Multi-seq THD unshard: 2 ranks, seq_lens=[50, 32, 46]."""
# rank0: [seqA_r0(50) | seqB_r0(32) | pad_r0(46)]
# rank1: [seqA_r1(50) | seqB_r1(32) | pad_r1(46)]
seq_a_r0 = torch.arange(0, 50)
seq_b_r0 = torch.arange(100, 132)
pad_r0 = torch.full((46,), -1)
rank0 = torch.cat([seq_a_r0, seq_b_r0, pad_r0]).refine_names("t")
seq_a_r1 = torch.arange(50, 100)
seq_b_r1 = torch.arange(132, 164)
pad_r1 = torch.full((46,), -2)
rank1 = torch.cat([seq_a_r1, seq_b_r1, pad_r1]).refine_names("t")
plan = UnsharderPlan(
axis=ParallelAxis.CP,
params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[50, 32, 46]),
groups=[[0, 1]],
)
with warning_sink.context():
result = execute_unsharder_plan(plan, [rank0, rank1])
assert len(result) == 1
unsharded: torch.Tensor = result[0].rename(None)
# seqA: r0(50) + r1(50) = 100 tokens, values 0..99
assert torch.equal(unsharded[:100], torch.cat([seq_a_r0, seq_a_r1]))
# seqB: r0(32) + r1(32) = 64 tokens
assert torch.equal(unsharded[100:164], torch.cat([seq_b_r0, seq_b_r1]))
# pad: r0(46) + r1(46) = 92 tokens
assert torch.equal(unsharded[164:256], torch.cat([pad_r0, pad_r1]))
def test_with_hidden_dim(self) -> None:
"""THD unshard with trailing hidden dim: shape [T, H]."""
torch.manual_seed(42)
hidden: int = 4
# rank0: [seqA_r0(3, 4) | seqB_r0(2, 4)]
# rank1: [seqA_r1(3, 4) | seqB_r1(2, 4)]
seq_a_r0 = torch.randn(3, hidden)
seq_b_r0 = torch.randn(2, hidden)
rank0 = torch.cat([seq_a_r0, seq_b_r0]).refine_names("t", "h")
seq_a_r1 = torch.randn(3, hidden)
seq_b_r1 = torch.randn(2, hidden)
rank1 = torch.cat([seq_a_r1, seq_b_r1]).refine_names("t", "h")
plan = UnsharderPlan(
axis=ParallelAxis.CP,
params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[3, 2]),
groups=[[0, 1]],
)
with warning_sink.context():
result = execute_unsharder_plan(plan, [rank0, rank1])
assert len(result) == 1
unsharded: torch.Tensor = result[0].rename(None)
assert unsharded.shape == (10, hidden)
assert torch.equal(unsharded[:6], torch.cat([seq_a_r0, seq_a_r1]))
assert torch.equal(unsharded[6:10], torch.cat([seq_b_r0, seq_b_r1]))
def test_with_leading_batch_dim(self) -> None:
"""THD unshard with leading batch dim: shape [B, T, H], t is dim=1."""
torch.manual_seed(42)
batch: int = 2
hidden: int = 4
# rank0: [seqA_r0(3) | seqB_r0(2)] per batch item
# rank1: [seqA_r1(3) | seqB_r1(2)] per batch item
seq_a_r0 = torch.randn(batch, 3, hidden)
seq_b_r0 = torch.randn(batch, 2, hidden)
rank0 = torch.cat([seq_a_r0, seq_b_r0], dim=1).refine_names("b", "t", "h")
seq_a_r1 = torch.randn(batch, 3, hidden)
seq_b_r1 = torch.randn(batch, 2, hidden)
rank1 = torch.cat([seq_a_r1, seq_b_r1], dim=1).refine_names("b", "t", "h")
plan = UnsharderPlan(
axis=ParallelAxis.CP,
params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[3, 2]),
groups=[[0, 1]],
)
with warning_sink.context():
result = execute_unsharder_plan(plan, [rank0, rank1])
assert len(result) == 1
unsharded: torch.Tensor = result[0].rename(None)
assert unsharded.shape == (batch, 10, hidden)
# seqA: r0(3) + r1(3) = 6 tokens per batch
assert torch.equal(unsharded[:, :6, :], torch.cat([seq_a_r0, seq_a_r1], dim=1))
# seqB: r0(2) + r1(2) = 4 tokens per batch
assert torch.equal(
unsharded[:, 6:10, :], torch.cat([seq_b_r0, seq_b_r1], dim=1)
)
class TestThdCpConcat:
def test_single_seq(self) -> None:

View File

@@ -81,6 +81,19 @@ class TestNormalizeParallelInfo:
}
assert normalize_parallel_info(meta) == {}
def test_recompute_pseudo_from_top_level_meta(self) -> None:
"""recompute_pseudo_rank/size at top-level meta is extracted alongside TP."""
meta = {
"recompute_pseudo_rank": 1,
"recompute_pseudo_size": 2,
"sglang_parallel_info": {"tp_rank": 0, "tp_size": 2},
}
result = normalize_parallel_info(meta)
assert result == {
ParallelAxis.RECOMPUTE_PSEUDO: AxisInfo(axis_rank=1, axis_size=2),
ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2),
}
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))

View File

@@ -421,6 +421,20 @@ class TestReplicatedAxes:
with pytest.raises(ValueError, match="missing parallel_info"):
compute_unsharder_plan(dim_specs, parallel_infos)
def test_recompute_pseudo_replicated(self) -> None:
"""RECOMPUTE_PSEUDO with no dim annotation → replicated → PickParams."""
dim_specs = parse_dims("h d")
parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [
{ParallelAxis.RECOMPUTE_PSEUDO: AxisInfo(axis_rank=0, axis_size=2)},
{ParallelAxis.RECOMPUTE_PSEUDO: AxisInfo(axis_rank=1, axis_size=2)},
]
plans = compute_unsharder_plan(dim_specs, parallel_infos)
assert len(plans) == 1
assert plans[0].axis == ParallelAxis.RECOMPUTE_PSEUDO
assert isinstance(plans[0].params, PickParams)
assert plans[0].groups == [[0, 1]]
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))

View File

@@ -13,13 +13,14 @@ from sglang.srt.debug_utils.comparator.output_types import (
ConfigRecord,
GeneralWarning,
NonTensorRecord,
ReplicatedMismatchWarning,
SkipRecord,
SummaryRecord,
WarningRecord,
_OutputRecord,
parse_record_json,
)
from sglang.srt.debug_utils.dumper import DumperConfig, _Dumper
from sglang.srt.debug_utils.dumper import DumperConfig, _Dumper, _RecomputeStatus
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=30, suite="default", nightly=True)
@@ -881,6 +882,142 @@ class TestEntrypointGroupingLogical:
comp = _assert_single_comparison_passed(records)
assert comp.name == "hidden"
def test_recompute_pseudo_replicated_verification(self, tmp_path, capsys):
"""Recompute pseudo-axis with identical original/recompute tensors → passed."""
torch.manual_seed(42)
tensor = torch.randn(4, 8)
baseline_dir = tmp_path / "baseline"
target_dir = tmp_path / "target"
for side_dir in [baseline_dir, target_dir]:
_create_recompute_rank_dump(
side_dir,
rank=0,
name="hidden",
original_tensor=tensor,
recompute_tensor=tensor.clone(),
)
args = _make_args(
baseline_dir / _FIXED_EXP_NAME,
target_dir / _FIXED_EXP_NAME,
diff_threshold=0.01,
)
records = _run_and_parse(args, capsys)
comp = _assert_single_comparison_passed(records)
assert comp.name == "hidden"
def test_recompute_pseudo_mismatch_warning(self, tmp_path, capsys):
"""Recompute pseudo-axis with differing original/recompute → ReplicatedMismatchWarning."""
torch.manual_seed(42)
tensor = torch.randn(4, 8)
mismatched_tensor = tensor + torch.randn(4, 8) * 10.0
baseline_dir = tmp_path / "baseline"
target_dir = tmp_path / "target"
for side_dir in [baseline_dir, target_dir]:
_create_recompute_rank_dump(
side_dir,
rank=0,
name="hidden",
original_tensor=tensor,
recompute_tensor=mismatched_tensor,
)
args = _make_args(
baseline_dir / _FIXED_EXP_NAME,
target_dir / _FIXED_EXP_NAME,
diff_threshold=0.01,
)
records = _run_and_parse(args, capsys)
comparisons = _get_comparisons(records)
assert len(comparisons) == 1
recompute_warnings = [
w
for w in comparisons[0].warnings
if isinstance(w, ReplicatedMismatchWarning) and w.axis == "recompute_pseudo"
]
assert len(recompute_warnings) > 0
class TestEntrypointAxisSwapper:
"""Test cross-framework dim reordering through the full entrypoint pipeline."""
def test_axis_swap_different_dim_order(self, tmp_path, capsys):
"""Baseline dims 'b h d' vs target dims 'b d h': axis swapper rearranges baseline to match."""
torch.manual_seed(42)
full_tensor = torch.randn(4, 8, 16)
baseline_dir = tmp_path / "baseline"
target_dir = tmp_path / "target"
_create_rank_dump(
baseline_dir,
rank=0,
name="hidden",
tensor=full_tensor,
dims="b h d",
)
_create_rank_dump(
target_dir,
rank=0,
name="hidden",
tensor=full_tensor.permute(0, 2, 1).contiguous(),
dims="b d h",
)
args = _make_args(
baseline_dir / _FIXED_EXP_NAME,
target_dir / _FIXED_EXP_NAME,
diff_threshold=1e-3,
)
records = _run_and_parse(args, capsys)
comp = _assert_single_comparison_passed(records)
assert comp.name == "hidden"
assert comp.baseline.shape == [4, 16, 8]
assert comp.target.shape == [4, 16, 8]
def test_axis_swap_with_tp_unshard(self, tmp_path, capsys):
"""Baseline TP=2 with dims 'b h(tp) d' vs target TP=2 with dims 'b d h(tp)': unshard + axis swap."""
torch.manual_seed(42)
full_tensor = torch.randn(4, 8, 16)
baseline_dir = tmp_path / "baseline"
target_dir = tmp_path / "target"
_create_tp_sharded_dumps(
baseline_dir,
full_tensor=full_tensor,
name="hidden",
tp_size=2,
shard_dim=1,
dims_str="b h(tp) d",
)
_create_tp_sharded_dumps(
target_dir,
full_tensor=full_tensor.permute(0, 2, 1).contiguous(),
name="hidden",
tp_size=2,
shard_dim=2,
dims_str="b d h(tp)",
)
args = _make_args(
baseline_dir / _FIXED_EXP_NAME,
target_dir / _FIXED_EXP_NAME,
diff_threshold=1e-3,
)
records = _run_and_parse(args, capsys)
comp = _assert_single_comparison_passed(records)
assert comp.name == "hidden"
class TestEntrypointAxisSwapper:
"""Test cross-framework dim reordering through the full entrypoint pipeline."""
@@ -1826,6 +1963,53 @@ def _create_tp_sharded_dumps(
return directory / _FIXED_EXP_NAME
def _create_recompute_rank_dump(
directory: Path,
*,
rank: int,
name: str,
original_tensor: torch.Tensor,
recompute_tensor: torch.Tensor,
dims: str = "h d",
) -> Path:
"""Create a dump with both original and recompute forward passes via monkeypatched dumper.
The dumper naturally produces recompute_pseudo_rank=0 for original and =1 for recompute,
plus recompute_pseudo_size=2.
"""
with pytest.MonkeyPatch.context() as mp:
mp.setattr(_dumper_module, "_get_rank", lambda: rank)
dumper = _Dumper(
config=DumperConfig(
enable=True,
dir=str(directory),
exp_name=_FIXED_EXP_NAME,
)
)
dumper.__dict__["_static_meta"] = {"world_rank": rank, "world_size": 1}
# dump original forward
mp.setattr(
_dumper_module,
"_detect_recompute_status",
lambda: _RecomputeStatus.ORIGINAL,
)
dumper.dump(name, original_tensor, dims=dims)
# dump recompute forward
mp.setattr(
_dumper_module,
"_detect_recompute_status",
lambda: _RecomputeStatus.RECOMPUTE,
)
dumper.dump(name, recompute_tensor, dims=dims)
dumper.step()
return directory / _FIXED_EXP_NAME
def _zigzag_split_seq(seq_natural: torch.Tensor, *, cp_size: int) -> list[torch.Tensor]:
"""Split a natural-order seq into per-rank zigzag segments."""
num_chunks: int = cp_size * 2

View File

@@ -5,10 +5,12 @@ import pytest
import torch
from sglang.srt.debug_utils.dump_loader import (
LOAD_FAILED,
ValueWithMeta,
_add_duplicate_index,
_cast_to_polars_dtype,
find_row,
parse_meta_from_filename,
read_meta,
)
from sglang.test.ci.ci_register import register_cpu_ci
@@ -92,9 +94,35 @@ class TestValueWithMeta:
path.write_text("not a valid pt file")
loaded = ValueWithMeta.load(path)
assert loaded.value is None
assert loaded.value is LOAD_FAILED
assert loaded.meta["name"] == "bad"
class TestRecomputeStatusParsing:
def test_parse_recompute_status_from_filename(self) -> None:
from pathlib import Path
meta_disabled = parse_meta_from_filename(
Path(
"step=0___rank=0___dump_index=1___name=x___recompute_status=disabled.pt"
)
)
assert meta_disabled["recompute_status"] == "disabled"
meta_recompute = parse_meta_from_filename(
Path(
"step=0___rank=0___dump_index=1___name=x___recompute_status=recompute.pt"
)
)
assert meta_recompute["recompute_status"] == "recompute"
meta_original = parse_meta_from_filename(
Path(
"step=0___rank=0___dump_index=1___name=x___recompute_status=original.pt"
)
)
assert meta_original["recompute_status"] == "original"
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))

View File

@@ -16,6 +16,7 @@ from sglang.srt.debug_utils.dumper import (
DumperConfig,
_collective_with_timeout,
_deepcopy_or_clone,
_detect_recompute_status,
_Dumper,
_format_tags,
_get_default_exp_name,
@@ -23,6 +24,7 @@ from sglang.srt.debug_utils.dumper import (
_materialize_value,
_MegatronPlugin,
_obj_to_dict,
_RecomputeStatus,
_register_forward_hook_or_replace_fn,
_SGLangPlugin,
_torch_save,
@@ -2400,5 +2402,113 @@ class TestCtxDecorator:
d.ctx()
class TestRecomputeStatus:
def test_disabled_by_default(self, tmp_path: Path) -> None:
d = _make_test_dumper(tmp_path)
tensor = torch.randn(3, 3)
d.dump("test_tensor", tensor)
filenames = _get_filenames(tmp_path)
_assert_files(filenames, exist=["recompute_status=disabled"])
def test_recompute_status_in_embedded_meta(self, tmp_path: Path) -> None:
d = _make_test_dumper(tmp_path)
tensor = torch.randn(3, 3)
d.dump("test_tensor", tensor)
path = _find_dump_file(tmp_path, rank=0, name="test_tensor")
raw = _load_dump(path)
assert raw["meta"]["recompute_status"] == "disabled"
def test_recompute_status_recompute(self, tmp_path: Path, monkeypatch) -> None:
import sglang.srt.debug_utils.dumper as dumper_mod
monkeypatch.setattr(
dumper_mod, "_detect_recompute_status", lambda: _RecomputeStatus.RECOMPUTE
)
d = _make_test_dumper(tmp_path)
tensor = torch.randn(3, 3)
d.dump("test_tensor", tensor)
filenames = _get_filenames(tmp_path)
_assert_files(filenames, exist=["recompute_status=recompute"])
path = _find_dump_file(tmp_path, rank=0, name="test_tensor")
raw = _load_dump(path)
assert raw["meta"]["recompute_status"] == "recompute"
assert raw["meta"]["recompute_pseudo_rank"] == 1
assert raw["meta"]["recompute_pseudo_size"] == 2
def test_recompute_status_original(self, tmp_path: Path, monkeypatch) -> None:
import sglang.srt.debug_utils.dumper as dumper_mod
monkeypatch.setattr(
dumper_mod,
"_detect_recompute_status",
lambda: _RecomputeStatus.ORIGINAL,
)
d = _make_test_dumper(tmp_path)
tensor = torch.randn(3, 3)
d.dump("test_tensor", tensor)
filenames = _get_filenames(tmp_path)
_assert_files(filenames, exist=["recompute_status=original"])
path = _find_dump_file(tmp_path, rank=0, name="test_tensor")
raw = _load_dump(path)
assert raw["meta"]["recompute_status"] == "original"
assert raw["meta"]["recompute_pseudo_rank"] == 0
assert raw["meta"]["recompute_pseudo_size"] == 2
def test_disabled_no_recompute_pseudo_fields(self, tmp_path: Path) -> None:
d = _make_test_dumper(tmp_path)
tensor = torch.randn(3, 3)
d.dump("test_tensor", tensor)
path = _find_dump_file(tmp_path, rank=0, name="test_tensor")
raw = _load_dump(path)
assert "recompute_pseudo_rank" not in raw["meta"]
assert "recompute_pseudo_size" not in raw["meta"]
def test_grad_hook_has_no_recompute_status(self, tmp_path: Path) -> None:
d = _make_test_dumper(tmp_path, enable_grad=True)
x = torch.randn(3, 3, requires_grad=True)
y = (x * 2).sum()
d.dump("test_tensor", x)
y.backward()
grad_files = [f for f in _get_filenames(tmp_path) if "grad__test_tensor" in f]
assert len(grad_files) == 1
assert "recompute_status" not in grad_files[0]
def test_non_intrusive_hooks_have_recompute_status(self, tmp_path: Path) -> None:
class Simple(torch.nn.Module):
def __init__(self):
super().__init__()
self.linear = torch.nn.Linear(4, 4)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.linear(x)
model = Simple()
d = _make_test_dumper(tmp_path, non_intrusive_mode="all")
d.register_non_intrusive_dumper(model)
with d.capture_output() as captured:
model(torch.randn(2, 4))
for key, data in captured.items():
assert (
"recompute_status" in data["meta"]
), f"missing recompute_status in {key}"
assert data["meta"]["recompute_status"] == "disabled"
def test_detect_recompute_status_default(self) -> None:
assert _detect_recompute_status() == _RecomputeStatus.DISABLED
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))