Enhance replication check, matching pattern, logging in dump comparator (#19677)
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user