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
@@ -9,8 +9,8 @@ from sglang.srt.debug_utils.comparator.aligner.axis_aligner import (
compute_axis_aligner_plan,
execute_axis_aligner_plan,
)
from sglang.srt.debug_utils.comparator.log_sink import log_sink
from sglang.srt.debug_utils.comparator.utils import Pair
from sglang.srt.debug_utils.comparator.warning_sink import warning_sink
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=15, suite="default", nightly=True)
@@ -37,7 +37,7 @@ class TestComputeAxisAlignerPlan:
assert result.pattern.y is None
def test_name_mismatch_returns_none_with_warning(self) -> None:
with warning_sink.context() as warnings:
with log_sink.context() as warnings:
result: Optional[AxisAlignerPlan] = compute_axis_aligner_plan(
Pair(x="t h d", y="t h e")
)
@@ -15,8 +15,8 @@ from sglang.srt.debug_utils.comparator.aligner.token_aligner.smart.aux_plugins i
_MegatronPlugin,
_SGLangPlugin,
)
from sglang.srt.debug_utils.comparator.output_types import GeneralWarning
from sglang.srt.debug_utils.comparator.warning_sink import WarningSink
from sglang.srt.debug_utils.comparator.log_sink import LogSink
from sglang.srt.debug_utils.comparator.output_types import ErrorLog, InfoLog
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=15, suite="default", nightly=True)
@@ -209,12 +209,12 @@ class TestLoadNonTensorAux:
fn1: str = _save_pt(tmp_path, name="rids", step=0, rank=1, value=["req_B"])
df: pl.DataFrame = _make_df_from_filenames([fn0, fn1])
sink = WarningSink()
sink = LogSink()
with sink.context() as warnings:
from unittest.mock import patch
with patch(
"sglang.srt.debug_utils.comparator.aligner.token_aligner.smart.aux_loader.warning_sink",
"sglang.srt.debug_utils.comparator.aligner.token_aligner.smart.aux_loader.log_sink",
sink,
):
result = _load_non_tensor_aux(
@@ -223,7 +223,7 @@ class TestLoadNonTensorAux:
assert result == ["req_A"]
assert len(warnings) == 1
assert isinstance(warnings[0], GeneralWarning)
assert isinstance(warnings[0], ErrorLog)
assert "rids_mismatch" in warnings[0].category
def test_no_rows_returns_none(self, tmp_path: Path) -> None:
@@ -268,12 +268,12 @@ class TestLoadAndAlignAuxTensor:
)
df: pl.DataFrame = _make_df_from_filenames([fn0, fn1])
sink = WarningSink()
sink = LogSink()
with sink.context() as warnings:
from unittest.mock import patch
with patch(
"sglang.srt.debug_utils.comparator.aligner.token_aligner.smart.aux_loader.warning_sink",
"sglang.srt.debug_utils.comparator.aligner.token_aligner.smart.aux_loader.log_sink",
sink,
):
result = _load_and_align_aux_tensor(
@@ -287,7 +287,7 @@ class TestLoadAndAlignAuxTensor:
assert result is not None
assert torch.equal(result, torch.tensor([1, 2, 3]))
assert len(warnings) == 1
assert isinstance(warnings[0], GeneralWarning)
assert isinstance(warnings[0], InfoLog)
assert "aux_no_dims" in warnings[0].category
@@ -324,12 +324,12 @@ class TestLoadNonTensorAuxDp:
)
df: pl.DataFrame = _make_df_from_filenames([fn0, fn1])
sink = WarningSink()
sink = LogSink()
with sink.context():
from unittest.mock import patch
with patch(
"sglang.srt.debug_utils.comparator.aligner.token_aligner.smart.aux_loader.warning_sink",
"sglang.srt.debug_utils.comparator.aligner.token_aligner.smart.aux_loader.log_sink",
sink,
):
result = _load_non_tensor_aux(