Enhance replication check, matching pattern, logging in dump comparator (#19677)

This commit is contained in:
fzyzcjy
2026-03-02 18:42:27 +08:00
committed by GitHub
parent ec44bc82ab
commit 15e83eea61
23 changed files with 783 additions and 461 deletions
@@ -3,11 +3,13 @@ import sys
import pytest
import torch
from sglang.srt.debug_utils.comparator.output_types import SummaryRecord
from sglang.srt.debug_utils.comparator.utils import (
Pair,
argmax_coord,
calc_per_token_rel_diff,
calc_rel_diff,
compute_exit_code,
compute_smaller_dtype,
try_unify_shape,
)
@@ -165,5 +167,247 @@ class TestPairMap:
assert result is not pair
class TestComputeExitCode:
"""Unit tests for compute_exit_code logic."""
def test_all_passed(self):
"""All passed → exit 0."""
summary = SummaryRecord(total=3, passed=3, failed=0, skipped=0)
assert (
compute_exit_code(
summary,
allow_skipped_pattern=".*",
skipped_names=[],
allow_failed_pattern=None,
failed_names=[],
)
== 0
)
def test_has_failed_and_passed(self):
"""Has failed and passed → exit 1."""
summary = SummaryRecord(total=4, passed=2, failed=2, skipped=0)
assert (
compute_exit_code(
summary,
allow_skipped_pattern=".*",
skipped_names=[],
allow_failed_pattern=None,
failed_names=["a", "b"],
)
== 1
)
def test_all_failed(self):
"""All failed (0 passed) → exit 1."""
summary = SummaryRecord(total=3, passed=0, failed=3, skipped=0)
assert (
compute_exit_code(
summary,
allow_skipped_pattern=".*",
skipped_names=[],
allow_failed_pattern=None,
failed_names=["a", "b", "c"],
)
== 1
)
def test_all_skipped_allow_all(self):
"""All skipped + allow_skipped_pattern='.*' → exit 1 (nothing passed)."""
summary = SummaryRecord(total=2, passed=0, failed=0, skipped=2)
assert (
compute_exit_code(
summary,
allow_skipped_pattern=".*",
skipped_names=["a", "b"],
allow_failed_pattern=None,
failed_names=[],
)
== 1
)
def test_all_skipped_forbid_all(self):
"""All skipped + allow_skipped_pattern='^$' → exit 1."""
summary = SummaryRecord(total=2, passed=0, failed=0, skipped=2)
assert (
compute_exit_code(
summary,
allow_skipped_pattern="^$",
skipped_names=["a", "b"],
allow_failed_pattern=None,
failed_names=[],
)
== 1
)
def test_passed_and_skipped_allow_all(self):
"""Passed + skipped, allow all → exit 0."""
summary = SummaryRecord(total=3, passed=2, failed=0, skipped=1)
assert (
compute_exit_code(
summary,
allow_skipped_pattern=".*",
skipped_names=["a"],
allow_failed_pattern=None,
failed_names=[],
)
== 0
)
def test_passed_and_skipped_forbid_all(self):
"""Passed + skipped + forbid all → exit 1."""
summary = SummaryRecord(total=3, passed=2, failed=0, skipped=1)
assert (
compute_exit_code(
summary,
allow_skipped_pattern="^$",
skipped_names=["a"],
allow_failed_pattern=None,
failed_names=[],
)
== 1
)
def test_skip_pattern_matches_specific_name(self):
"""Pattern matching specific name allows that skip, forbids others."""
summary = SummaryRecord(total=4, passed=2, failed=0, skipped=2)
assert (
compute_exit_code(
summary,
allow_skipped_pattern="positions|seq_lens",
skipped_names=["positions", "seq_lens"],
allow_failed_pattern=None,
failed_names=[],
)
== 0
)
def test_skip_pattern_partial_match_forbidden(self):
"""Pattern matches some skips but not all → exit 1."""
summary = SummaryRecord(total=4, passed=1, failed=0, skipped=3)
assert (
compute_exit_code(
summary,
allow_skipped_pattern="positions|seq_lens",
skipped_names=["positions", "seq_lens", "hidden_states"],
allow_failed_pattern=None,
failed_names=[],
)
== 1
)
def test_allow_failed_pattern_matches_all(self):
"""allow_failed_pattern='.*' tolerates all failures → exit 0."""
summary = SummaryRecord(total=3, passed=1, failed=2, skipped=0)
assert (
compute_exit_code(
summary,
allow_skipped_pattern=".*",
skipped_names=[],
allow_failed_pattern=".*",
failed_names=["a", "b"],
)
== 0
)
def test_allow_failed_pattern_matches_specific(self):
"""Pattern matches all failed names → exit 0."""
summary = SummaryRecord(total=3, passed=1, failed=2, skipped=0)
assert (
compute_exit_code(
summary,
allow_skipped_pattern=".*",
skipped_names=[],
allow_failed_pattern="hidden_states|logits",
failed_names=["hidden_states", "logits"],
)
== 0
)
def test_allow_failed_pattern_partial_match(self):
"""Pattern matches some but not all failures → exit 1."""
summary = SummaryRecord(total=3, passed=0, failed=3, skipped=0)
assert (
compute_exit_code(
summary,
allow_skipped_pattern=".*",
skipped_names=[],
allow_failed_pattern="hidden_states",
failed_names=["hidden_states", "logits", "attn"],
)
== 1
)
def test_allow_failed_pattern_no_failures(self):
"""Pattern set but no failures → exit 0."""
summary = SummaryRecord(total=2, passed=2, failed=0, skipped=0)
assert (
compute_exit_code(
summary,
allow_skipped_pattern=".*",
skipped_names=[],
allow_failed_pattern=".*",
failed_names=[],
)
== 0
)
def test_both_failed_and_skipped_patterns(self):
"""Both patterns set, both satisfied → exit 0."""
summary = SummaryRecord(total=4, passed=1, failed=1, skipped=2)
assert (
compute_exit_code(
summary,
allow_skipped_pattern="positions|seq_lens",
skipped_names=["positions", "seq_lens"],
allow_failed_pattern="logits",
failed_names=["logits"],
)
== 0
)
def test_failed_pattern_satisfied_but_skipped_not(self):
"""Failed pattern OK but skipped pattern fails → exit 1."""
summary = SummaryRecord(total=3, passed=1, failed=1, skipped=1)
assert (
compute_exit_code(
summary,
allow_skipped_pattern="^$",
skipped_names=["a"],
allow_failed_pattern=".*",
failed_names=["b"],
)
== 1
)
def test_zero_passed_exits_one(self):
"""No tensors passed → exit 1, even when all failures are allowed."""
summary = SummaryRecord(total=2, passed=0, failed=2, skipped=0)
assert (
compute_exit_code(
summary,
allow_skipped_pattern=".*",
skipped_names=[],
allow_failed_pattern=".*",
failed_names=["a", "b"],
)
== 1
)
def test_zero_passed_all_skipped_exits_one(self):
"""All skipped, nothing passed → exit 1."""
summary = SummaryRecord(total=3, passed=0, failed=0, skipped=3)
assert (
compute_exit_code(
summary,
allow_skipped_pattern=".*",
skipped_names=["a", "b", "c"],
allow_failed_pattern=None,
failed_names=[],
)
== 1
)
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))