Enhance replicated tensor checker in dump comparator (#19597)

This commit is contained in:
fzyzcjy
2026-03-01 10:34:34 +08:00
committed by GitHub
parent ec08240a6a
commit e41164af1c
15 changed files with 514 additions and 313 deletions
@@ -3,23 +3,20 @@ import sys
import pytest
from sglang.srt.debug_utils.comparator.output_types import ReplicatedMismatchWarning
from sglang.srt.debug_utils.comparator.output_types import GeneralWarning
from sglang.srt.debug_utils.comparator.warning_sink import WarningSink
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=10, suite="default", nightly=True)
def _make_warning(**overrides) -> ReplicatedMismatchWarning:
def _make_warning(**overrides) -> GeneralWarning:
defaults: dict = dict(
axis="tp",
group_index=0,
differing_index=1,
baseline_index=0,
max_abs_diff=0.1,
category="test",
message="test warning",
)
defaults.update(overrides)
return ReplicatedMismatchWarning(**defaults)
return GeneralWarning(**defaults)
class TestWarningSink:
@@ -35,8 +32,8 @@ class TestWarningSink:
def test_nested_contexts(self) -> None:
sink = WarningSink()
outer_warning = _make_warning(group_index=0)
inner_warning = _make_warning(group_index=1)
outer_warning = _make_warning(message="outer")
inner_warning = _make_warning(message="inner")
with sink.context() as outer:
sink.add(outer_warning)
@@ -61,7 +58,7 @@ class TestWarningSink:
sink.add(_make_warning())
captured = capsys.readouterr()
assert "Replicated along tp" in captured.out
assert "test warning" in captured.out
def test_context_captures_instead_of_printing(self, capsys) -> None:
sink = WarningSink()
@@ -96,9 +93,9 @@ class TestWarningSink:
assert len(collected) == 1
sink.add(_make_warning(group_index=99))
sink.add(_make_warning(message="after exception"))
captured = capsys.readouterr()
assert "Replicated along tp" in captured.out
assert "after exception" in captured.out
if __name__ == "__main__":