Enhance replicated tensor checker in dump comparator (#19597)
This commit is contained in:
@@ -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__":
|
||||
|
||||
Reference in New Issue
Block a user