Handle warnings via sink for structured output and add pair in dump comparator (#19373)
This commit is contained in:
@@ -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__]))
|
||||||
Reference in New Issue
Block a user