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

View File

@@ -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.output_types import (
AlignWarning,
AnyWarning,
ReplicatedMismatchWarning,
)
@@ -16,8 +16,8 @@ from sglang.srt.debug_utils.comparator.output_types import (
def execute_unshard_plan(
plan: UnshardPlan,
tensors: list[torch.Tensor],
) -> tuple[list[torch.Tensor], list[AlignWarning]]:
all_warnings: list[AlignWarning] = []
) -> tuple[list[torch.Tensor], list[AnyWarning]]:
all_warnings: list[AnyWarning] = []
result: list[torch.Tensor] = []
for group_idx, group in enumerate(plan.groups):
@@ -40,7 +40,7 @@ def _apply_unshard(
*,
axis: ParallelAxis,
group_index: int,
) -> tuple[torch.Tensor, list[AlignWarning]]:
) -> tuple[torch.Tensor, list[AnyWarning]]:
if isinstance(params, PickParams):
warnings = _verify_replicated_group(
ordered_tensors,

View File

@@ -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"])
print_record(
ConfigRecord(
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,
),
ConfigRecord.from_args(args),
output_format=args.output_format,
)

View File

@@ -1,7 +1,7 @@
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 (
format_comparison,
@@ -28,38 +28,45 @@ class ReplicatedMismatchWarning(_StrictBase):
)
AlignWarning = (
ReplicatedMismatchWarning # future: Annotated[Union[...], Discriminator("kind")]
)
class GeneralWarning(_StrictBase):
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):
align_warnings: list[AlignWarning] = Field(default_factory=list)
warnings: list[AnyWarning] = Field(default_factory=list)
@abstractmethod
def _format_body(self) -> str: ...
def to_text(self) -> str:
body = self._format_body()
if self.align_warnings:
body += "\n" + "\n".join(f"{w.to_text()}" for w in self.align_warnings)
if self.warnings:
body += "\n" + "\n".join(f"{w.to_text()}" for w in self.warnings)
return body
class ConfigRecord(_OutputRecord):
type: Literal["config"] = "config"
baseline_path: str
target_path: str
diff_threshold: float
start_step: int
end_step: int
config: dict[str, Any]
@classmethod
def from_args(cls, args) -> "ConfigRecord":
"""Create ConfigRecord from argparse.Namespace."""
return cls(config=vars(args))
def _format_body(self) -> str:
return (
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}]"
)
return f"Config: {self.config}"
class SkipRecord(_OutputRecord):
@@ -69,7 +76,7 @@ class SkipRecord(_OutputRecord):
@property
def category(self) -> str:
if self.align_warnings:
if self.warnings:
return "failed"
return "skipped"
@@ -82,7 +89,7 @@ class ComparisonRecord(TensorComparisonInfo, _OutputRecord):
@property
def category(self) -> str:
if self.align_warnings:
if self.warnings:
return "failed"
return "passed" if self.diff is not None and self.diff.passed else "failed"
@@ -97,6 +104,15 @@ class SummaryRecord(_OutputRecord):
failed: 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:
return (
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[
Union[ConfigRecord, SkipRecord, ComparisonRecord, SummaryRecord],
Union[ConfigRecord, SkipRecord, ComparisonRecord, SummaryRecord, WarningRecord],
Discriminator("type"),
]

View File

@@ -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.dims import parse_dims
from sglang.srt.debug_utils.comparator.output_types import (
AlignWarning,
AnyWarning,
ComparisonRecord,
SkipRecord,
)
@@ -53,11 +53,11 @@ def process_tensor_group(
b_tensor, b_warns = _execute_plans(b_extracted, b_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:
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(
x_baseline=b_tensor,
@@ -66,7 +66,7 @@ def process_tensor_group(
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]:
@@ -114,7 +114,7 @@ def _extract_tensors(
def _execute_plans(
tensors: list[torch.Tensor],
plans: list[Plan],
) -> tuple[Optional[torch.Tensor], list[AlignWarning]]:
) -> tuple[Optional[torch.Tensor], list[AnyWarning]]:
if not tensors:
return None, []
@@ -123,7 +123,7 @@ def _execute_plans(
return None, []
return tensors[0], []
warnings: list[AlignWarning] = []
warnings: list[AnyWarning] = []
current = tensors
for plan in plans:
current, new_warnings = _execute_plan(current, plan)
@@ -136,7 +136,7 @@ def _execute_plans(
def _execute_plan(
tensors: list[torch.Tensor],
plan: Plan,
) -> tuple[list[torch.Tensor], list[AlignWarning]]:
) -> tuple[list[torch.Tensor], list[AnyWarning]]:
if isinstance(plan, UnshardPlan):
return execute_unshard_plan(plan, tensors)
elif isinstance(plan, ReorderPlan):

View File

@@ -1,9 +1,22 @@
from __future__ import annotations
import functools
from typing import Optional, Tuple
from typing import Callable, Generic, Optional, Tuple, TypeVar
import torch
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):
model_config = ConfigDict(extra="forbid")
@@ -13,6 +26,14 @@ class _FrozenBase(BaseModel):
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, ...]:
flat_idx = x.argmax()
return tuple(idx.item() for idx in torch.unravel_index(flat_idx, x.shape))

View File

@@ -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()

View File

@@ -98,11 +98,13 @@ class TestRecordTypes:
def test_discriminated_union_parsing(self):
for record in [
ConfigRecord(
baseline_path="/a",
target_path="/b",
diff_threshold=1e-3,
start_step=0,
end_step=100,
config={
"baseline_path": "/a",
"target_path": "/b",
"diff_threshold": 1e-3,
"start_step": 0,
"end_step": 100,
},
),
SkipRecord(name="attn", reason="no_baseline"),
ComparisonRecord(
@@ -133,7 +135,7 @@ def _make_warning(**overrides) -> ReplicatedMismatchWarning:
class TestAlignWarnings:
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(
name="hidden",
baseline=_make_tensor_info(),
@@ -141,21 +143,21 @@ class TestAlignWarnings:
unified_shape=[4, 8],
shape_mismatch=False,
diff=_make_diff(passed=True),
align_warnings=[_make_warning()],
warnings=[_make_warning()],
)
assert record.category == "failed"
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(
name="x",
reason="no_baseline",
align_warnings=[_make_warning()],
warnings=[_make_warning()],
)
assert record.category == "failed"
def test_align_warnings_json_round_trip(self):
"""align_warnings survive model_dump_json → parse_record_json round-trip."""
def test_warnings_json_round_trip(self):
"""warnings survive model_dump_json → parse_record_json round-trip."""
warning = _make_warning(
axis="cp",
group_index=2,
@@ -170,14 +172,14 @@ class TestAlignWarnings:
unified_shape=[4, 8],
shape_mismatch=False,
diff=_make_diff(),
align_warnings=[warning],
warnings=[warning],
)
restored = parse_record_json(record.model_dump_json())
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.group_index == 2
assert restored_warning.differing_index == 3

View File

@@ -540,7 +540,7 @@ class TestEntrypointGroupingLogical:
assert summary.skipped == 0
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)
full_tensor = torch.randn(4, 8)
@@ -571,6 +571,8 @@ class TestEntrypointGroupingLogical:
records = _run_and_parse(args, capsys)
comparisons = _get_comparisons(records)
assert len(comparisons) == 2
assert comparisons[0].baseline.shape == [4, 8]
assert comparisons[1].baseline.shape == [4, 8]
summary = records[-1]
assert isinstance(summary, SummaryRecord)
@@ -612,7 +614,7 @@ class TestEntrypointGroupingLogical:
assert comp.name == "attn_out"
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)
full_a = torch.randn(4, 8)
full_b = torch.randn(4, 8)
@@ -736,7 +738,7 @@ class TestEntrypointGroupingLogical:
assert comp.name == "hidden"
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)
full_baseline = torch.randn(4, 8, 16)
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."""
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)
full_baseline = torch.randn(4, 8, 6)
full_target = full_baseline + torch.randn(4, 8, 6) * 0.0001
@@ -912,14 +914,14 @@ class TestEntrypointReplicatedAxis:
records = _run_and_parse(args, capsys)
comp = _assert_single_comparison_passed(records)
assert comp.align_warnings == []
assert comp.warnings == []
summary = records[-1]
assert isinstance(summary, SummaryRecord)
assert summary.passed == 1
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)
full_baseline = torch.randn(4, 8, 6)
full_target = full_baseline + torch.randn(4, 8, 6) * 0.0001
@@ -952,14 +954,14 @@ class TestEntrypointReplicatedAxis:
comparisons = _get_comparisons(records)
assert len(comparisons) == 1
assert comparisons[0].category == "failed"
assert len(comparisons[0].align_warnings) > 0
assert len(comparisons[0].warnings) > 0
summary = records[-1]
assert isinstance(summary, SummaryRecord)
assert summary.failed == 1
def test_summary_counts_failed_from_align_warnings_only(self, tmp_path, capsys):
"""Diff itself passes but TP replicas differ → summary.failed=1 from align_warnings."""
def test_summary_counts_failed_from_warnings_only(self, tmp_path, capsys):
"""Diff itself passes but TP replicas differ → summary.failed=1 from warnings."""
torch.manual_seed(42)
full_baseline = torch.randn(4, 8, 6)
full_target = full_baseline + torch.randn(4, 8, 6) * 0.0001
@@ -1001,7 +1003,7 @@ class TestEntrypointReplicatedAxis:
comp = comparisons[0]
assert comp.diff is not None
assert comp.diff.passed
assert len(comp.align_warnings) > 0
assert len(comp.warnings) > 0
assert comp.category == "failed"
summary = records[-1]

View File

@@ -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__]))

View 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__]))