Handle warnings via sink for structured output and add pair in dump comparator (#19373)

This commit is contained in:
fzyzcjy
2026-02-26 09:59:15 +08:00
committed by GitHub
parent 46321ee70e
commit 508b8e3387
10 changed files with 367 additions and 64 deletions
@@ -8,7 +8,7 @@ from sglang.srt.debug_utils.comparator.aligner.unshard.types import (
) )
from sglang.srt.debug_utils.comparator.dims import ParallelAxis from sglang.srt.debug_utils.comparator.dims import ParallelAxis
from sglang.srt.debug_utils.comparator.output_types import ( from sglang.srt.debug_utils.comparator.output_types import (
AlignWarning, AnyWarning,
ReplicatedMismatchWarning, ReplicatedMismatchWarning,
) )
@@ -16,8 +16,8 @@ from sglang.srt.debug_utils.comparator.output_types import (
def execute_unshard_plan( def execute_unshard_plan(
plan: UnshardPlan, plan: UnshardPlan,
tensors: list[torch.Tensor], tensors: list[torch.Tensor],
) -> tuple[list[torch.Tensor], list[AlignWarning]]: ) -> tuple[list[torch.Tensor], list[AnyWarning]]:
all_warnings: list[AlignWarning] = [] all_warnings: list[AnyWarning] = []
result: list[torch.Tensor] = [] result: list[torch.Tensor] = []
for group_idx, group in enumerate(plan.groups): for group_idx, group in enumerate(plan.groups):
@@ -40,7 +40,7 @@ def _apply_unshard(
*, *,
axis: ParallelAxis, axis: ParallelAxis,
group_index: int, group_index: int,
) -> tuple[torch.Tensor, list[AlignWarning]]: ) -> tuple[torch.Tensor, list[AnyWarning]]:
if isinstance(params, PickParams): if isinstance(params, PickParams):
warnings = _verify_replicated_group( warnings = _verify_replicated_group(
ordered_tensors, ordered_tensors,
@@ -31,13 +31,7 @@ def run(args: argparse.Namespace) -> None:
assert all(c in df_target.columns for c in ["rank", "step", "dump_index", "name"]) assert all(c in df_target.columns for c in ["rank", "step", "dump_index", "name"])
print_record( print_record(
ConfigRecord( ConfigRecord.from_args(args),
baseline_path=args.baseline_path,
target_path=args.target_path,
diff_threshold=args.diff_threshold,
start_step=args.start_step,
end_step=args.end_step,
),
output_format=args.output_format, output_format=args.output_format,
) )
@@ -1,7 +1,7 @@
from abc import abstractmethod from abc import abstractmethod
from typing import Annotated, Literal, Union from typing import Annotated, Any, Literal, Union
from pydantic import Discriminator, Field, TypeAdapter from pydantic import Discriminator, Field, TypeAdapter, model_validator
from sglang.srt.debug_utils.comparator.tensor_comparison.formatter import ( from sglang.srt.debug_utils.comparator.tensor_comparison.formatter import (
format_comparison, format_comparison,
@@ -28,38 +28,45 @@ class ReplicatedMismatchWarning(_StrictBase):
) )
AlignWarning = ( class GeneralWarning(_StrictBase):
ReplicatedMismatchWarning # future: Annotated[Union[...], Discriminator("kind")] kind: Literal["general"] = "general"
) category: str
message: str
def to_text(self) -> str:
return self.message
AnyWarning = Annotated[
Union[ReplicatedMismatchWarning, GeneralWarning],
Discriminator("kind"),
]
class _OutputRecord(_StrictBase): class _OutputRecord(_StrictBase):
align_warnings: list[AlignWarning] = Field(default_factory=list) warnings: list[AnyWarning] = Field(default_factory=list)
@abstractmethod @abstractmethod
def _format_body(self) -> str: ... def _format_body(self) -> str: ...
def to_text(self) -> str: def to_text(self) -> str:
body = self._format_body() body = self._format_body()
if self.align_warnings: if self.warnings:
body += "\n" + "\n".join(f" ⚠ {w.to_text()}" for w in self.align_warnings) body += "\n" + "\n".join(f" ⚠ {w.to_text()}" for w in self.warnings)
return body return body
class ConfigRecord(_OutputRecord): class ConfigRecord(_OutputRecord):
type: Literal["config"] = "config" type: Literal["config"] = "config"
baseline_path: str config: dict[str, Any]
target_path: str
diff_threshold: float @classmethod
start_step: int def from_args(cls, args) -> "ConfigRecord":
end_step: int """Create ConfigRecord from argparse.Namespace."""
return cls(config=vars(args))
def _format_body(self) -> str: def _format_body(self) -> str:
return ( return f"Config: {self.config}"
f"Config: baseline={self.baseline_path} target={self.target_path}\n"
f"diff_threshold={self.diff_threshold} "
f"steps=[{self.start_step}, {self.end_step}]"
)
class SkipRecord(_OutputRecord): class SkipRecord(_OutputRecord):
@@ -69,7 +76,7 @@ class SkipRecord(_OutputRecord):
@property @property
def category(self) -> str: def category(self) -> str:
if self.align_warnings: if self.warnings:
return "failed" return "failed"
return "skipped" return "skipped"
@@ -82,7 +89,7 @@ class ComparisonRecord(TensorComparisonInfo, _OutputRecord):
@property @property
def category(self) -> str: def category(self) -> str:
if self.align_warnings: if self.warnings:
return "failed" return "failed"
return "passed" if self.diff is not None and self.diff.passed else "failed" return "passed" if self.diff is not None and self.diff.passed else "failed"
@@ -97,6 +104,15 @@ class SummaryRecord(_OutputRecord):
failed: int failed: int
skipped: int skipped: int
@model_validator(mode="after")
def _validate_totals(self) -> "SummaryRecord":
expected: int = self.passed + self.failed + self.skipped
if self.total != expected:
raise ValueError(
f"total={self.total} != passed({self.passed}) + failed({self.failed}) + skipped({self.skipped}) = {expected}"
)
return self
def _format_body(self) -> str: def _format_body(self) -> str:
return ( return (
f"Summary: {self.passed} passed, {self.failed} failed, " f"Summary: {self.passed} passed, {self.failed} failed, "
@@ -104,8 +120,15 @@ class SummaryRecord(_OutputRecord):
) )
class WarningRecord(_OutputRecord):
type: Literal["warning"] = "warning"
def _format_body(self) -> str:
return ""
AnyRecord = Annotated[ AnyRecord = Annotated[
Union[ConfigRecord, SkipRecord, ComparisonRecord, SummaryRecord], Union[ConfigRecord, SkipRecord, ComparisonRecord, SummaryRecord, WarningRecord],
Discriminator("type"), Discriminator("type"),
] ]
@@ -20,7 +20,7 @@ from sglang.srt.debug_utils.comparator.aligner.unshard.planner import (
from sglang.srt.debug_utils.comparator.aligner.unshard.types import UnshardPlan from sglang.srt.debug_utils.comparator.aligner.unshard.types import UnshardPlan
from sglang.srt.debug_utils.comparator.dims import parse_dims from sglang.srt.debug_utils.comparator.dims import parse_dims
from sglang.srt.debug_utils.comparator.output_types import ( from sglang.srt.debug_utils.comparator.output_types import (
AlignWarning, AnyWarning,
ComparisonRecord, ComparisonRecord,
SkipRecord, SkipRecord,
) )
@@ -53,11 +53,11 @@ def process_tensor_group(
b_tensor, b_warns = _execute_plans(b_extracted, b_plans) b_tensor, b_warns = _execute_plans(b_extracted, b_plans)
t_tensor, t_warns = _execute_plans(t_extracted, t_plans) t_tensor, t_warns = _execute_plans(t_extracted, t_plans)
all_warnings: list[AlignWarning] = b_warns + t_warns all_warnings: list[AnyWarning] = b_warns + t_warns
if b_tensor is None or t_tensor is None: if b_tensor is None or t_tensor is None:
reason = "baseline_load_failed" if b_tensor is None else "target_load_failed" reason = "baseline_load_failed" if b_tensor is None else "target_load_failed"
return SkipRecord(name=name, reason=reason, align_warnings=all_warnings) return SkipRecord(name=name, reason=reason, warnings=all_warnings)
info = compare_tensors( info = compare_tensors(
x_baseline=b_tensor, x_baseline=b_tensor,
@@ -66,7 +66,7 @@ def process_tensor_group(
diff_threshold=diff_threshold, diff_threshold=diff_threshold,
) )
return ComparisonRecord(**info.model_dump(), align_warnings=all_warnings) return ComparisonRecord(**info.model_dump(), warnings=all_warnings)
def _load_tensors(filenames: list[str], base_path: Path) -> list[ValueWithMeta]: def _load_tensors(filenames: list[str], base_path: Path) -> list[ValueWithMeta]:
@@ -114,7 +114,7 @@ def _extract_tensors(
def _execute_plans( def _execute_plans(
tensors: list[torch.Tensor], tensors: list[torch.Tensor],
plans: list[Plan], plans: list[Plan],
) -> tuple[Optional[torch.Tensor], list[AlignWarning]]: ) -> tuple[Optional[torch.Tensor], list[AnyWarning]]:
if not tensors: if not tensors:
return None, [] return None, []
@@ -123,7 +123,7 @@ def _execute_plans(
return None, [] return None, []
return tensors[0], [] return tensors[0], []
warnings: list[AlignWarning] = [] warnings: list[AnyWarning] = []
current = tensors current = tensors
for plan in plans: for plan in plans:
current, new_warnings = _execute_plan(current, plan) current, new_warnings = _execute_plan(current, plan)
@@ -136,7 +136,7 @@ def _execute_plans(
def _execute_plan( def _execute_plan(
tensors: list[torch.Tensor], tensors: list[torch.Tensor],
plan: Plan, plan: Plan,
) -> tuple[list[torch.Tensor], list[AlignWarning]]: ) -> tuple[list[torch.Tensor], list[AnyWarning]]:
if isinstance(plan, UnshardPlan): if isinstance(plan, UnshardPlan):
return execute_unshard_plan(plan, tensors) return execute_unshard_plan(plan, tensors)
elif isinstance(plan, ReorderPlan): elif isinstance(plan, ReorderPlan):
@@ -1,9 +1,22 @@
from __future__ import annotations
import functools import functools
from typing import Optional, Tuple from typing import Callable, Generic, Optional, Tuple, TypeVar
import torch import torch
from pydantic import BaseModel, ConfigDict from pydantic import BaseModel, ConfigDict
_T = TypeVar("_T")
_U = TypeVar("_U")
def _check_equal_lengths(**named_lists: list) -> None:
lengths: dict[str, int] = {name: len(lst) for name, lst in named_lists.items()}
unique: set[int] = set(lengths.values())
if len(unique) > 1:
details: str = ", ".join(f"{name}={length}" for name, length in lengths.items())
raise ValueError(f"Length mismatch: {details}")
class _StrictBase(BaseModel): class _StrictBase(BaseModel):
model_config = ConfigDict(extra="forbid") model_config = ConfigDict(extra="forbid")
@@ -13,6 +26,14 @@ class _FrozenBase(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid") model_config = ConfigDict(frozen=True, extra="forbid")
class Pair(_FrozenBase, Generic[_T]):
x: _T
y: _T
def map(self, fn: Callable[[_T], _U]) -> Pair[_U]:
return Pair(x=fn(self.x), y=fn(self.y))
def argmax_coord(x: torch.Tensor) -> Tuple[int, ...]: def argmax_coord(x: torch.Tensor) -> Tuple[int, ...]:
flat_idx = x.argmax() flat_idx = x.argmax()
return tuple(idx.item() for idx in torch.unravel_index(flat_idx, x.shape)) return tuple(idx.item() for idx in torch.unravel_index(flat_idx, x.shape))
@@ -0,0 +1,42 @@
from __future__ import annotations
from contextlib import contextmanager
from typing import Generator
from sglang.srt.debug_utils.comparator.output_types import AnyWarning
class WarningSink:
def __init__(self) -> None:
self._stack: list[list[AnyWarning]] = []
self._output_format: str = "text"
def set_output_format(self, output_format: str) -> None:
self._output_format = output_format
@contextmanager
def context(self) -> Generator[list[AnyWarning], None, None]:
bucket: list[AnyWarning] = []
self._stack.append(bucket)
try:
yield bucket
finally:
popped = self._stack.pop()
assert popped is bucket
def add(self, warning: AnyWarning) -> None:
if self._stack:
self._stack[-1].append(warning)
else:
from sglang.srt.debug_utils.comparator.output_types import (
WarningRecord,
print_record,
)
print_record(
WarningRecord(warnings=[warning]),
output_format=self._output_format,
)
warning_sink = WarningSink()
@@ -98,11 +98,13 @@ class TestRecordTypes:
def test_discriminated_union_parsing(self): def test_discriminated_union_parsing(self):
for record in [ for record in [
ConfigRecord( ConfigRecord(
baseline_path="/a", config={
target_path="/b", "baseline_path": "/a",
diff_threshold=1e-3, "target_path": "/b",
start_step=0, "diff_threshold": 1e-3,
end_step=100, "start_step": 0,
"end_step": 100,
},
), ),
SkipRecord(name="attn", reason="no_baseline"), SkipRecord(name="attn", reason="no_baseline"),
ComparisonRecord( ComparisonRecord(
@@ -133,7 +135,7 @@ def _make_warning(**overrides) -> ReplicatedMismatchWarning:
class TestAlignWarnings: class TestAlignWarnings:
def test_comparison_record_failed_when_diff_passed_but_warnings(self): def test_comparison_record_failed_when_diff_passed_but_warnings(self):
"""ComparisonRecord with diff.passed=True but align_warnings → category=='failed'.""" """ComparisonRecord with diff.passed=True but warnings → category=='failed'."""
record = ComparisonRecord( record = ComparisonRecord(
name="hidden", name="hidden",
baseline=_make_tensor_info(), baseline=_make_tensor_info(),
@@ -141,21 +143,21 @@ class TestAlignWarnings:
unified_shape=[4, 8], unified_shape=[4, 8],
shape_mismatch=False, shape_mismatch=False,
diff=_make_diff(passed=True), diff=_make_diff(passed=True),
align_warnings=[_make_warning()], warnings=[_make_warning()],
) )
assert record.category == "failed" assert record.category == "failed"
def test_skip_record_failed_when_warnings(self): def test_skip_record_failed_when_warnings(self):
"""SkipRecord with align_warnings → category=='failed' instead of 'skipped'.""" """SkipRecord with warnings → category=='failed' instead of 'skipped'."""
record = SkipRecord( record = SkipRecord(
name="x", name="x",
reason="no_baseline", reason="no_baseline",
align_warnings=[_make_warning()], warnings=[_make_warning()],
) )
assert record.category == "failed" assert record.category == "failed"
def test_align_warnings_json_round_trip(self): def test_warnings_json_round_trip(self):
"""align_warnings survive model_dump_json → parse_record_json round-trip.""" """warnings survive model_dump_json → parse_record_json round-trip."""
warning = _make_warning( warning = _make_warning(
axis="cp", axis="cp",
group_index=2, group_index=2,
@@ -170,14 +172,14 @@ class TestAlignWarnings:
unified_shape=[4, 8], unified_shape=[4, 8],
shape_mismatch=False, shape_mismatch=False,
diff=_make_diff(), diff=_make_diff(),
align_warnings=[warning], warnings=[warning],
) )
restored = parse_record_json(record.model_dump_json()) restored = parse_record_json(record.model_dump_json())
assert isinstance(restored, ComparisonRecord) assert isinstance(restored, ComparisonRecord)
assert len(restored.align_warnings) == 1 assert len(restored.warnings) == 1
restored_warning = restored.align_warnings[0] restored_warning = restored.warnings[0]
assert restored_warning.axis == "cp" assert restored_warning.axis == "cp"
assert restored_warning.group_index == 2 assert restored_warning.group_index == 2
assert restored_warning.differing_index == 3 assert restored_warning.differing_index == 3
@@ -540,7 +540,7 @@ class TestEntrypointGroupingLogical:
assert summary.skipped == 0 assert summary.skipped == 0
def test_multi_step_tp(self, tmp_path, capsys): def test_multi_step_tp(self, tmp_path, capsys):
"""Two steps with TP=2 shards produce two logical groups (one per step).""" """Two steps with TP=2 shards produce two per-step comparisons (no aux → no alignment)."""
torch.manual_seed(42) torch.manual_seed(42)
full_tensor = torch.randn(4, 8) full_tensor = torch.randn(4, 8)
@@ -571,6 +571,8 @@ class TestEntrypointGroupingLogical:
records = _run_and_parse(args, capsys) records = _run_and_parse(args, capsys)
comparisons = _get_comparisons(records) comparisons = _get_comparisons(records)
assert len(comparisons) == 2 assert len(comparisons) == 2
assert comparisons[0].baseline.shape == [4, 8]
assert comparisons[1].baseline.shape == [4, 8]
summary = records[-1] summary = records[-1]
assert isinstance(summary, SummaryRecord) assert isinstance(summary, SummaryRecord)
@@ -612,7 +614,7 @@ class TestEntrypointGroupingLogical:
assert comp.name == "attn_out" assert comp.name == "attn_out"
def test_filter_logical(self, tmp_path, capsys): def test_filter_logical(self, tmp_path, capsys):
"""--filter in logical grouping selects only matching tensor groups.""" """--filter in logical grouping selects only matching tensor bundles."""
torch.manual_seed(42) torch.manual_seed(42)
full_a = torch.randn(4, 8) full_a = torch.randn(4, 8)
full_b = torch.randn(4, 8) full_b = torch.randn(4, 8)
@@ -736,7 +738,7 @@ class TestEntrypointGroupingLogical:
assert comp.name == "hidden" assert comp.name == "hidden"
def test_cp_tp_different_sizes(self, tmp_path, capsys): def test_cp_tp_different_sizes(self, tmp_path, capsys):
"""Baseline CP=2+TP=2 vs target CP=1+TP=4: both sides independently unshard.""" """Baseline CP=2+TP=2 vs target CP=1+TP=4: both sides independently unsharder."""
torch.manual_seed(42) torch.manual_seed(42)
full_baseline = torch.randn(4, 8, 16) full_baseline = torch.randn(4, 8, 16)
full_target = full_baseline + torch.randn(4, 8, 16) * 0.001 full_target = full_baseline + torch.randn(4, 8, 16) * 0.001
@@ -882,7 +884,7 @@ class TestEntrypointReplicatedAxis:
"""Test replicated-axis scenarios through the full entrypoint pipeline.""" """Test replicated-axis scenarios through the full entrypoint pipeline."""
def test_replicated_axis_identical_replicas_passed(self, tmp_path, capsys): def test_replicated_axis_identical_replicas_passed(self, tmp_path, capsys):
"""CP2 TP2, TP replicated and identical → passed, no align_warnings.""" """CP2 TP2, TP replicated and identical → passed, no warnings."""
torch.manual_seed(42) torch.manual_seed(42)
full_baseline = torch.randn(4, 8, 6) full_baseline = torch.randn(4, 8, 6)
full_target = full_baseline + torch.randn(4, 8, 6) * 0.0001 full_target = full_baseline + torch.randn(4, 8, 6) * 0.0001
@@ -912,14 +914,14 @@ class TestEntrypointReplicatedAxis:
records = _run_and_parse(args, capsys) records = _run_and_parse(args, capsys)
comp = _assert_single_comparison_passed(records) comp = _assert_single_comparison_passed(records)
assert comp.align_warnings == [] assert comp.warnings == []
summary = records[-1] summary = records[-1]
assert isinstance(summary, SummaryRecord) assert isinstance(summary, SummaryRecord)
assert summary.passed == 1 assert summary.passed == 1
def test_replicated_mismatch_fails(self, tmp_path, capsys): def test_replicated_mismatch_fails(self, tmp_path, capsys):
"""CP2 TP2, TP replicas differ (> atol) → failed with align_warnings.""" """CP2 TP2, TP replicas differ (> atol) → failed with warnings."""
torch.manual_seed(42) torch.manual_seed(42)
full_baseline = torch.randn(4, 8, 6) full_baseline = torch.randn(4, 8, 6)
full_target = full_baseline + torch.randn(4, 8, 6) * 0.0001 full_target = full_baseline + torch.randn(4, 8, 6) * 0.0001
@@ -952,14 +954,14 @@ class TestEntrypointReplicatedAxis:
comparisons = _get_comparisons(records) comparisons = _get_comparisons(records)
assert len(comparisons) == 1 assert len(comparisons) == 1
assert comparisons[0].category == "failed" assert comparisons[0].category == "failed"
assert len(comparisons[0].align_warnings) > 0 assert len(comparisons[0].warnings) > 0
summary = records[-1] summary = records[-1]
assert isinstance(summary, SummaryRecord) assert isinstance(summary, SummaryRecord)
assert summary.failed == 1 assert summary.failed == 1
def test_summary_counts_failed_from_align_warnings_only(self, tmp_path, capsys): def test_summary_counts_failed_from_warnings_only(self, tmp_path, capsys):
"""Diff itself passes but TP replicas differ → summary.failed=1 from align_warnings.""" """Diff itself passes but TP replicas differ → summary.failed=1 from warnings."""
torch.manual_seed(42) torch.manual_seed(42)
full_baseline = torch.randn(4, 8, 6) full_baseline = torch.randn(4, 8, 6)
full_target = full_baseline + torch.randn(4, 8, 6) * 0.0001 full_target = full_baseline + torch.randn(4, 8, 6) * 0.0001
@@ -1001,7 +1003,7 @@ class TestEntrypointReplicatedAxis:
comp = comparisons[0] comp = comparisons[0]
assert comp.diff is not None assert comp.diff is not None
assert comp.diff.passed assert comp.diff.passed
assert len(comp.align_warnings) > 0 assert len(comp.warnings) > 0
assert comp.category == "failed" assert comp.category == "failed"
summary = records[-1] summary = records[-1]
@@ -0,0 +1,114 @@
import sys
import pytest
from pydantic import ValidationError
from sglang.srt.debug_utils.comparator.output_types import (
ComparisonRecord,
GeneralWarning,
SkipRecord,
SummaryRecord,
)
from sglang.srt.debug_utils.comparator.tensor_comparison.types import (
DiffInfo,
TensorInfo,
TensorStats,
)
from sglang.srt.debug_utils.comparator.utils import _check_equal_lengths
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=10, suite="default", nightly=True)
class TestCheckEqualLengths:
def test_all_equal(self):
_check_equal_lengths(a=[1, 2], b=[3, 4])
def test_empty_lists(self):
_check_equal_lengths(a=[], b=[])
def test_mismatch_raises(self):
with pytest.raises(ValueError, match="Length mismatch"):
_check_equal_lengths(a=[1, 2], b=[3])
class TestSummaryRecord:
def test_valid(self):
record = SummaryRecord(total=10, passed=7, failed=2, skipped=1)
assert record.total == 10
def test_total_mismatch(self):
with pytest.raises(ValidationError, match="total=10"):
SummaryRecord(total=10, passed=5, failed=2, skipped=1)
def _make_tensor_info() -> TensorInfo:
return TensorInfo(
shape=[4, 4],
dtype="float32",
stats=TensorStats(mean=0.0, std=1.0, min=-2.0, max=2.0),
)
def _make_diff_info(*, passed: bool) -> DiffInfo:
return DiffInfo(
rel_diff=0.001,
max_abs_diff=0.01,
mean_abs_diff=0.005,
max_diff_coord=[0, 0],
baseline_at_max=1.0,
target_at_max=1.01,
passed=passed,
)
def _make_comparison_record(
*,
diff: DiffInfo | None,
warnings: list | None = None,
) -> ComparisonRecord:
ti: TensorInfo = _make_tensor_info()
return ComparisonRecord(
name="t",
baseline=ti,
target=ti,
unified_shape=[4, 4],
shape_mismatch=False,
diff=diff,
warnings=warnings or [],
)
class TestOutputRecordCategories:
def test_skip_record_with_warnings_is_failed(self) -> None:
record = SkipRecord(
name="t",
reason="test",
warnings=[GeneralWarning(category="c", message="m")],
)
assert record.category == "failed"
def test_skip_record_no_warnings_is_skipped(self) -> None:
record = SkipRecord(name="t", reason="test")
assert record.category == "skipped"
def test_comparison_record_diff_none_is_failed(self) -> None:
record: ComparisonRecord = _make_comparison_record(diff=None)
assert record.category == "failed"
def test_comparison_record_passed_with_warnings_is_failed(self) -> None:
record: ComparisonRecord = _make_comparison_record(
diff=_make_diff_info(passed=True),
warnings=[GeneralWarning(category="c", message="m")],
)
assert record.category == "failed"
def test_comparison_record_passed_no_warnings_is_passed(self) -> None:
record: ComparisonRecord = _make_comparison_record(
diff=_make_diff_info(passed=True),
)
assert record.category == "passed"
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))
@@ -0,0 +1,105 @@
import json
import sys
import pytest
from sglang.srt.debug_utils.comparator.output_types import ReplicatedMismatchWarning
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:
defaults: dict = dict(
axis="tp",
group_index=0,
differing_index=1,
baseline_index=0,
max_abs_diff=0.1,
)
defaults.update(overrides)
return ReplicatedMismatchWarning(**defaults)
class TestWarningSink:
def test_basic_collection(self) -> None:
sink = WarningSink()
warning = _make_warning()
with sink.context() as collected:
sink.add(warning)
assert len(collected) == 1
assert collected[0] is warning
def test_nested_contexts(self) -> None:
sink = WarningSink()
outer_warning = _make_warning(group_index=0)
inner_warning = _make_warning(group_index=1)
with sink.context() as outer:
sink.add(outer_warning)
with sink.context() as inner:
sink.add(inner_warning)
assert len(inner) == 1
assert inner[0] is inner_warning
assert len(outer) == 1
assert outer[0] is outer_warning
def test_empty_context(self) -> None:
sink = WarningSink()
with sink.context() as collected:
pass
assert collected == []
def test_add_outside_context_prints(self, capsys) -> None:
sink = WarningSink()
sink.set_output_format("text")
sink.add(_make_warning())
captured = capsys.readouterr()
assert "Replicated along tp" in captured.out
def test_context_captures_instead_of_printing(self, capsys) -> None:
sink = WarningSink()
sink.set_output_format("text")
with sink.context() as collected:
sink.add(_make_warning())
assert len(collected) == 1
captured = capsys.readouterr()
assert captured.out == ""
def test_json_output_outside_context(self, capsys) -> None:
sink = WarningSink()
sink.set_output_format("json")
sink.add(_make_warning())
captured = capsys.readouterr()
parsed: dict = json.loads(captured.out.strip())
assert "warnings" in parsed
assert len(parsed["warnings"]) == 1
def test_exception_in_context_cleans_stack(self, capsys) -> None:
sink = WarningSink()
sink.set_output_format("text")
with pytest.raises(RuntimeError):
with sink.context() as collected:
sink.add(_make_warning())
raise RuntimeError("boom")
assert len(collected) == 1
sink.add(_make_warning(group_index=99))
captured = capsys.readouterr()
assert "Replicated along tp" in captured.out
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))