Support presets and arbitrary skipping keys in dump comparator (#19676)

This commit is contained in:
fzyzcjy
2026-03-02 18:41:49 +08:00
committed by GitHub
parent 2e15c015c0
commit ec44bc82ab
13 changed files with 802 additions and 391 deletions

View File

@@ -22,11 +22,11 @@ from sglang.srt.debug_utils.comparator.aligner.unsharder.types import (
)
from sglang.srt.debug_utils.comparator.dims import ParallelAxis, TokenLayout
from sglang.srt.debug_utils.comparator.output_types import (
ComparisonRecord,
GeneralWarning,
NonTensorRecord,
SkipRecord,
NonTensorComparisonRecord,
SkipComparisonRecord,
SummaryRecord,
TensorComparisonRecord,
parse_record_json,
)
from sglang.srt.debug_utils.comparator.tensor_comparator.types import (
@@ -208,9 +208,9 @@ def _make_comparison_record(
*,
diff: DiffInfo | None,
warnings: list | None = None,
) -> ComparisonRecord:
) -> TensorComparisonRecord:
ti: TensorInfo = _make_tensor_info()
return ComparisonRecord(
return TensorComparisonRecord(
name="t",
baseline=ti,
target=ti,
@@ -223,7 +223,7 @@ def _make_comparison_record(
class TestOutputRecordCategories:
def test_skip_record_with_warnings_is_failed(self) -> None:
record = SkipRecord(
record = SkipComparisonRecord(
name="t",
reason="test",
warnings=[GeneralWarning(category="c", message="m")],
@@ -231,28 +231,28 @@ class TestOutputRecordCategories:
assert record.category == "failed"
def test_skip_record_no_warnings_is_skipped(self) -> None:
record = SkipRecord(name="t", reason="test")
record = SkipComparisonRecord(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)
record: TensorComparisonRecord = _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(
record: TensorComparisonRecord = _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(
record: TensorComparisonRecord = _make_comparison_record(
diff=_make_diff_info(passed=True),
)
assert record.category == "passed"
def test_non_tensor_record_equal_is_passed(self) -> None:
record = NonTensorRecord(
record = NonTensorComparisonRecord(
name="sm_scale",
baseline_value="0.125",
target_value="0.125",
@@ -263,7 +263,7 @@ class TestOutputRecordCategories:
assert record.category == "passed"
def test_non_tensor_record_different_is_failed(self) -> None:
record = NonTensorRecord(
record = NonTensorComparisonRecord(
name="sm_scale",
baseline_value="0.125",
target_value="0.25",
@@ -274,7 +274,7 @@ class TestOutputRecordCategories:
assert record.category == "failed"
def test_non_tensor_record_with_warnings_is_failed(self) -> None:
record = NonTensorRecord(
record = NonTensorComparisonRecord(
name="sm_scale",
baseline_value="0.125",
target_value="0.125",
@@ -286,7 +286,7 @@ class TestOutputRecordCategories:
assert record.category == "failed"
def test_non_tensor_record_json_roundtrip(self) -> None:
record = NonTensorRecord(
record = NonTensorComparisonRecord(
name="sm_scale",
baseline_value="0.125",
target_value="0.25",
@@ -296,14 +296,14 @@ class TestOutputRecordCategories:
)
json_str: str = record.model_dump_json()
roundtripped = parse_record_json(json_str)
assert isinstance(roundtripped, NonTensorRecord)
assert isinstance(roundtripped, NonTensorComparisonRecord)
assert roundtripped.name == "sm_scale"
assert roundtripped.values_equal is False
assert roundtripped.baseline_value == "0.125"
assert roundtripped.target_value == "0.25"
def test_non_tensor_record_text_format_equal(self) -> None:
record = NonTensorRecord(
record = NonTensorComparisonRecord(
name="sm_scale",
baseline_value="0.125",
target_value="0.125",
@@ -316,7 +316,7 @@ class TestOutputRecordCategories:
assert "[equal]" in text
def test_non_tensor_record_text_format_different(self) -> None:
record = NonTensorRecord(
record = NonTensorComparisonRecord(
name="sm_scale",
baseline_value="0.125",
target_value="0.25",
@@ -351,10 +351,10 @@ def _make_aligner_plan() -> AlignerPlan:
)
class TestAlignerPlanInComparisonRecord:
class TestAlignerPlanInTensorComparisonRecord:
def test_comparison_record_with_aligner_plan(self) -> None:
plan: AlignerPlan = _make_aligner_plan()
record: ComparisonRecord = _make_comparison_record(
record: TensorComparisonRecord = _make_comparison_record(
diff=_make_diff_info(passed=True),
)
record_with_plan = record.model_copy(update={"aligner_plan": plan})
@@ -363,7 +363,7 @@ class TestAlignerPlanInComparisonRecord:
def test_aligner_plan_json_roundtrip(self) -> None:
plan: AlignerPlan = _make_aligner_plan()
record: ComparisonRecord = _make_comparison_record(
record: TensorComparisonRecord = _make_comparison_record(
diff=_make_diff_info(passed=True),
)
record_with_plan = record.model_copy(update={"aligner_plan": plan})
@@ -376,7 +376,7 @@ class TestAlignerPlanInComparisonRecord:
== "unsharder"
)
roundtripped: ComparisonRecord = parse_record_json(json_str)
roundtripped: TensorComparisonRecord = parse_record_json(json_str)
assert roundtripped.aligner_plan is not None
assert (
roundtripped.aligner_plan.per_step_plans.x[0].sub_plans[0].type
@@ -384,16 +384,16 @@ class TestAlignerPlanInComparisonRecord:
)
def test_comparison_record_without_aligner_plan(self) -> None:
record: ComparisonRecord = _make_comparison_record(
record: TensorComparisonRecord = _make_comparison_record(
diff=_make_diff_info(passed=True),
)
json_str: str = record.model_dump_json()
roundtripped: ComparisonRecord = parse_record_json(json_str)
roundtripped: TensorComparisonRecord = parse_record_json(json_str)
assert roundtripped.aligner_plan is None
def test_aligner_plan_text_format(self) -> None:
plan: AlignerPlan = _make_aligner_plan()
record: ComparisonRecord = _make_comparison_record(
record: TensorComparisonRecord = _make_comparison_record(
diff=_make_diff_info(passed=True),
)
record_with_plan = record.model_copy(update={"aligner_plan": plan})