@@ -1,3 +1,4 @@
import subprocess
import sys
import textwrap
from argparse import Namespace
@@ -7,7 +8,7 @@ import pytest
import torch
import sglang . srt . debug_utils . dumper as _dumper_module
from sglang . srt . debug_utils . comparator . entrypoint import run
from sglang . srt . debug_utils . comparator . entrypoint import _compute_exit_code , run
from sglang . srt . debug_utils . comparator . output_types import (
AnyRecord ,
ComparisonRecord ,
@@ -39,7 +40,7 @@ class TestEntrypointGroupingRaw:
baseline_path , target_path = _create_dumps ( tmp_path , [ " tensor_a " , " tensor_b " ] )
args = _make_args ( baseline_path , target_path , grouping = " raw " )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
assert isinstance ( records [ 0 ] , ConfigRecord )
assert len ( _get_comparisons ( records ) ) == 2
@@ -54,7 +55,7 @@ class TestEntrypointGroupingRaw:
baseline_path , target_path = _create_dumps ( tmp_path , [ " tensor_a " , " tensor_b " ] )
args = _make_args ( baseline_path , target_path , filter = " tensor_a " , grouping = " raw " )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
assert len ( _get_comparisons ( records ) ) == 1
def test_no_baseline_skip ( self , tmp_path , capsys ) :
@@ -66,7 +67,7 @@ class TestEntrypointGroupingRaw:
)
args = _make_args ( baseline_path , target_path , grouping = " raw " )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
skips = [ r for r in records if isinstance ( r , SkipRecord ) ]
assert len ( skips ) == 1
assert skips [ 0 ] . reason == " baseline_load_failed "
@@ -82,7 +83,7 @@ class TestEntrypointGroupingRaw:
baseline_path , target_path , start_step = 1 , end_step = 1 , grouping = " raw "
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
summary = records [ - 1 ]
assert isinstance ( summary , SummaryRecord )
assert summary . total == 1
@@ -92,7 +93,7 @@ class TestEntrypointGroupingRaw:
baseline_path , target_path = _create_dumps ( tmp_path , [ " t " ] , num_steps = 2 )
args = _make_args ( baseline_path , target_path , grouping = " raw " )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
assert all ( isinstance ( r , _OutputRecord ) for r in records )
def test_comparison_failed ( self , tmp_path , capsys ) :
@@ -111,7 +112,7 @@ class TestEntrypointGroupingRaw:
baseline_path , target_path , grouping = " raw " , diff_threshold = 1e-3
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 1
assert comparisons [ 0 ] . diff is not None
@@ -133,7 +134,7 @@ class TestEntrypointGroupingRaw:
)
args = _make_args ( baseline_path , target_path , grouping = " raw " )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 1
assert comparisons [ 0 ] . shape_mismatch is True
@@ -159,7 +160,7 @@ class TestEntrypointGroupingRaw:
)
args = _make_args ( baseline_path , target_path , grouping = " raw " )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 1
@@ -189,7 +190,7 @@ class TestEntrypointGroupingRaw:
baseline_path , target_path , grouping = " raw " , diff_threshold = 0.01
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 1
assert comparisons [ 0 ] . diff_downcast is not None
@@ -227,7 +228,7 @@ class TestEntrypointGroupingRaw:
diff_threshold = 1e-3 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
summary = records [ - 1 ]
assert isinstance ( summary , SummaryRecord )
assert summary . passed == 1
@@ -242,7 +243,7 @@ class TestEntrypointGroupingRaw:
baseline_path , target_path , filter = " nonexistent_pattern " , grouping = " raw "
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
summary = records [ - 1 ]
assert isinstance ( summary , SummaryRecord )
assert summary . total == 0
@@ -271,7 +272,7 @@ class TestEntrypointGroupingRaw:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 2
@@ -348,7 +349,7 @@ class TestEntrypointGroupingRaw:
grouping = " raw " ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 2
assert all ( c . diff is not None and c . diff . passed for c in comparisons )
@@ -367,7 +368,7 @@ class TestEntrypointGroupingLogical:
baseline_path , target_path = _create_dumps ( tmp_path , [ " tensor_a " , " tensor_b " ] )
args = _make_args ( baseline_path , target_path )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
assert len ( _get_comparisons ( records ) ) == 2
summary = records [ - 1 ]
assert isinstance ( summary , SummaryRecord )
@@ -402,7 +403,7 @@ class TestEntrypointGroupingLogical:
args = _make_args ( baseline_path , target_path , diff_threshold = 0.01 )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comp = _assert_single_comparison_passed ( records )
assert comp . name == " hidden "
@@ -439,7 +440,7 @@ class TestEntrypointGroupingLogical:
args = _make_args ( baseline_path , target_path , diff_threshold = 0.01 )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
_assert_single_comparison_passed ( records )
def test_one_side_dims_single_baseline ( self , tmp_path , capsys ) :
@@ -466,7 +467,7 @@ class TestEntrypointGroupingLogical:
args = _make_args ( baseline_path , target_path , diff_threshold = 0.01 )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
_assert_single_comparison_passed ( records )
@pytest.mark.parametrize (
@@ -496,7 +497,7 @@ class TestEntrypointGroupingLogical:
target_dir / _FIXED_EXP_NAME ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
skips = [ r for r in records if isinstance ( r , SkipRecord ) ]
assert len ( skips ) == 1
assert skips [ 0 ] . reason == expected_reason
@@ -535,7 +536,7 @@ class TestEntrypointGroupingLogical:
args = _make_args ( baseline_path , target_path , diff_threshold = 0.01 )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
summary = records [ - 1 ]
assert isinstance ( summary , SummaryRecord )
assert summary . total == 2
@@ -572,7 +573,7 @@ class TestEntrypointGroupingLogical:
args = _make_args ( baseline_path , target_path , diff_threshold = 0.01 )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 1
# concat along dim 0 (fallback, no token dim) → 2 steps × [4, 8] = [8, 8]
@@ -613,7 +614,7 @@ class TestEntrypointGroupingLogical:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comp = _assert_single_comparison_passed ( records )
assert comp . name == " attn_out "
@@ -651,7 +652,7 @@ class TestEntrypointGroupingLogical:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 1
assert comparisons [ 0 ] . name == " t_a "
@@ -696,7 +697,7 @@ class TestEntrypointGroupingLogical:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 2
assert all ( c . diff is not None and c . diff . passed for c in comparisons )
@@ -737,7 +738,7 @@ class TestEntrypointGroupingLogical:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comp = _assert_single_comparison_passed ( records )
assert comp . name == " hidden "
@@ -776,7 +777,7 @@ class TestEntrypointGroupingLogical:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
_assert_single_comparison_passed ( records )
def test_ep_cp_tp_three_axis_unshard ( self , tmp_path , capsys ) :
@@ -811,7 +812,7 @@ class TestEntrypointGroupingLogical:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comp = _assert_single_comparison_passed ( records )
assert comp . name == " hidden "
@@ -845,7 +846,7 @@ class TestEntrypointGroupingLogical:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comp = _assert_single_comparison_passed ( records )
assert comp . name == " attn_out "
@@ -879,7 +880,7 @@ class TestEntrypointGroupingLogical:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comp = _assert_single_comparison_passed ( records )
assert comp . name == " hidden "
@@ -906,7 +907,7 @@ class TestEntrypointGroupingLogical:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comp = _assert_single_comparison_passed ( records )
assert comp . name == " hidden "
@@ -934,7 +935,7 @@ class TestEntrypointGroupingLogical:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 1
@@ -970,7 +971,7 @@ class TestEntrypointGroupingLogical:
args = _make_args ( baseline_path , target_path , diff_threshold = 0.01 )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comp = _assert_single_comparison_passed ( records )
assert comp . name == " attn_out "
@@ -1001,7 +1002,7 @@ class TestEntrypointGroupingLogical:
args = _make_args ( baseline_path , target_path , diff_threshold = 0.01 )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comp = _assert_single_comparison_passed ( records )
assert comp . name == " attn_out "
@@ -1041,7 +1042,7 @@ class TestEntrypointGroupingLogical:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comp = _assert_single_comparison_passed ( records )
assert comp . name == " hidden "
@@ -1105,7 +1106,8 @@ class TestEntrypointConcatMode:
args : Namespace = _make_args (
baseline_path , target_path , diff_threshold = diff_threshold
)
return _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
return records
def test_concat_multi_step_different_data ( self , tmp_path , capsys ) :
""" Multi-step concat with different data per step + truncation. """
@@ -1168,7 +1170,7 @@ class TestEntrypointConcatMode:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 1
# 2 steps × [4, 8] concat along dim 0 (fallback) → [8, 8]
@@ -1330,7 +1332,7 @@ class TestEntrypointConcatMode:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
# all 3 tensors should be compared (not filtered out)
names = { c . name for c in comparisons }
@@ -1425,7 +1427,7 @@ class TestEntrypointConcatMode:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 1
# CP unshard: [4,4,6] × 2 ranks → [4,8,6] per step
@@ -1477,7 +1479,7 @@ class TestEntrypointConcatMode:
token_aligner = " concat_steps " ,
diff_threshold = 1e-3 ,
)
records : list [ AnyRecord ] = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons : list [ ComparisonRecord ] = _get_comparisons ( records )
hidden_comparisons : list [ ComparisonRecord ] = [
@@ -1519,7 +1521,7 @@ class TestEntrypointAxisAligner:
diff_threshold = 1e-3 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comp = _assert_single_comparison_passed ( records )
assert comp . name == " hidden "
assert comp . baseline . shape == [ 4 , 16 , 8 ]
@@ -1556,7 +1558,7 @@ class TestEntrypointAxisAligner:
diff_threshold = 1e-3 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comp = _assert_single_comparison_passed ( records )
assert comp . name == " hidden "
@@ -1589,7 +1591,7 @@ class TestEntrypointAxisAligner:
diff_threshold = 1e-3 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comp = _assert_single_comparison_passed ( records )
assert comp . name == " hidden "
assert comp . baseline . shape == [ 4 , 8 ]
@@ -1702,7 +1704,7 @@ class TestEntrypointReplicatedAxis:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comp = _assert_single_comparison_passed ( records )
assert comp . warnings == [ ]
assert all ( c . passed for c in comp . replicated_checks )
@@ -1741,7 +1743,7 @@ class TestEntrypointReplicatedAxis:
diff_threshold = 0.01 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 1
assert comparisons [ 0 ] . category == " failed "
@@ -1787,7 +1789,7 @@ class TestEntrypointReplicatedAxis:
diff_threshold = 0.5 ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 1
@@ -1850,7 +1852,7 @@ class TestEntrypointAlignment:
args = _make_args (
exp_paths [ 0 ] , exp_paths [ 1 ] , grouping = " logical " , token_aligner = " smart "
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
# AUX_NAMES are filtered out after plan computation → only hidden_states remains
@@ -1964,7 +1966,7 @@ class TestEntrypointAlignment:
token_aligner = " smart " ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
warning_records = [ r for r in records if isinstance ( r , WarningRecord ) ]
layout_warnings = [
@@ -2031,7 +2033,7 @@ class TestEntrypointNonTensorValues:
tmp_path , name = " sm_scale " , baseline_value = 0.125 , target_value = 0.125
)
args = _make_args ( baseline_path , target_path , grouping = " raw " )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
non_tensors = _get_non_tensors ( records )
assert len ( non_tensors ) == 1
@@ -2050,7 +2052,7 @@ class TestEntrypointNonTensorValues:
tmp_path , name = " sm_scale " , baseline_value = 0.125 , target_value = 0.25
)
args = _make_args ( baseline_path , target_path , grouping = " raw " )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
non_tensors = _get_non_tensors ( records )
assert len ( non_tensors ) == 1
@@ -2070,7 +2072,7 @@ class TestEntrypointNonTensorValues:
target_value = " flash_attn " ,
)
args = _make_args ( baseline_path , target_path , grouping = " raw " )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
non_tensors = _get_non_tensors ( records )
assert len ( non_tensors ) == 1
@@ -2100,7 +2102,7 @@ class TestEntrypointNonTensorValues:
target_dir / _FIXED_EXP_NAME ,
grouping = " raw " ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
non_tensors = _get_non_tensors ( records )
@@ -2121,7 +2123,7 @@ class TestEntrypointNonTensorValues:
tmp_path , name = " debug_info " , baseline_value = value , target_value = value
)
args = _make_args ( baseline_path , target_path , grouping = " raw " )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
non_tensors = _get_non_tensors ( records )
assert len ( non_tensors ) == 1
@@ -2135,7 +2137,7 @@ class TestEntrypointNonTensorValues:
tmp_path , name = " optional_param " , baseline_value = None , target_value = None
)
args = _make_args ( baseline_path , target_path , grouping = " raw " )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
non_tensors = _get_non_tensors ( records )
assert len ( non_tensors ) == 1
@@ -2151,7 +2153,7 @@ class TestEntrypointNonTensorValues:
tmp_path , name = " sm_scale " , baseline_value = 0.125 , target_value = 0.125
)
args = _make_args ( baseline_path , target_path , grouping = " raw " )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
non_tensors = _get_non_tensors ( records )
assert len ( non_tensors ) == 1
@@ -2186,7 +2188,7 @@ class TestEntrypointVisualize:
viz_output_dir = str ( viz_dir ) ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
assert len ( _get_comparisons ( records ) ) == 1
png_files = list ( viz_dir . glob ( " *.png " ) )
@@ -2341,15 +2343,19 @@ def _make_args(baseline_path: Path, target_path: Path, **overrides) -> Namespace
override_baseline_dims = [ ] ,
override_target_dims = [ ] ,
override_config = None ,
allow_skip_pattern = " .* " ,
report_path = " " ,
)
defaults . update ( overrides )
return Namespace ( * * defaults )
def _run_and_parse ( args : Namespace , capsys : pytest . CaptureFixture ) - > list [ AnyRecord ] :
def _run_and_parse (
args : Namespace , capsys : pytest . CaptureFixture
) - > tuple [ list [ AnyRecord ] , int ] :
capsys . readouterr ( )
run ( args )
return _parse_jsonl ( capsys . readouterr ( ) . out )
exit_code : int = run( args )
return _parse_jsonl ( capsys . readouterr ( ) . out ) , exit_code
def _parse_jsonl ( output : str ) - > list [ AnyRecord ] :
@@ -2875,7 +2881,7 @@ class TestEntrypointPerTokenVisualization:
grouping = " raw " ,
visualize_per_token = str ( output_png ) ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 2
@@ -2892,7 +2898,7 @@ class TestEntrypointPerTokenVisualization:
baseline_path , target_path = _create_dumps ( tmp_path , [ " tensor_a " ] )
args = _make_args ( baseline_path , target_path , grouping = " raw " )
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons = _get_comparisons ( records )
assert len ( comparisons ) == 1
@@ -2992,7 +2998,7 @@ class TestEntrypointThdCpZigzag:
token_aligner = " smart " ,
diff_threshold = 1e-3 ,
)
records : list [ AnyRecord ] = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparisons : list [ ComparisonRecord ] = _get_comparisons ( records )
hidden_comparisons : list [ ComparisonRecord ] = [
@@ -3043,7 +3049,7 @@ class TestEntrypointThdCpZigzag:
token_aligner = " smart " ,
diff_threshold = 1e-3 ,
)
records : list [ AnyRecord ] = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
# hidden_states should pass comparison (after unshard + reorder)
comparisons : list [ ComparisonRecord ] = _get_comparisons ( records )
@@ -3113,7 +3119,7 @@ class TestEntrypointDpFilter:
grouping = " logical " ,
diff_threshold = 1e-3 ,
)
records : list [ AnyRecord ] = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparison : ComparisonRecord = _assert_single_comparison_passed ( records )
assert comparison . name == " hidden "
@@ -3169,7 +3175,7 @@ class TestEntrypointDpFilter:
grouping = " logical " ,
diff_threshold = 1e-3 ,
)
records : list [ AnyRecord ] = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparison : ComparisonRecord = _assert_single_comparison_passed ( records )
assert comparison . name == " hidden "
@@ -3218,7 +3224,7 @@ class TestEntrypointDpFilter:
grouping = " logical " ,
diff_threshold = 1e-3 ,
)
records : list [ AnyRecord ] = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
comparison : ComparisonRecord = _assert_single_comparison_passed ( records )
assert comparison . name == " hidden "
@@ -3344,7 +3350,7 @@ class TestEntrypointMetaOverride:
grouping = " logical " ,
override_dims = [ " hidden:t h(tp) " ] ,
)
self . _assert_all_passed ( _run_and_parse ( args , capsys ) )
self . _assert_all_passed ( _run_and_parse ( args , capsys ) [ 0 ] )
@pytest.mark.parametrize (
" baseline_dims, target_dims, override_kwarg " ,
@@ -3371,7 +3377,7 @@ class TestEntrypointMetaOverride:
)
args = _make_args ( baseline_path , target_path , grouping = " raw " , * * override_kwarg )
self . _assert_all_passed ( _run_and_parse ( args , capsys ) )
self . _assert_all_passed ( _run_and_parse ( args , capsys ) [ 0 ] )
def test_override_config_yaml ( self , tmp_path : Path , capsys ) - > None :
""" --override-config YAML overrides dims. """
@@ -3390,7 +3396,7 @@ class TestEntrypointMetaOverride:
grouping = " raw " ,
override_config = str ( yaml_path ) ,
)
self . _assert_all_passed ( _run_and_parse ( args , capsys ) )
self . _assert_all_passed ( _run_and_parse ( args , capsys ) [ 0 ] )
def test_no_match_uses_original_dims ( self , tmp_path : Path , capsys ) - > None :
""" When override regex doesn ' t match, original dims from dump are used. """
@@ -3406,7 +3412,7 @@ class TestEntrypointMetaOverride:
grouping = " raw " ,
override_dims = [ " no_match_pattern:b s d " ] ,
)
self . _assert_all_passed ( _run_and_parse ( args , capsys ) )
self . _assert_all_passed ( _run_and_parse ( args , capsys ) [ 0 ] )
def test_selective_match_multi_tensor ( self , tmp_path : Path , capsys ) - > None :
""" Override matches only ' logits ' ; ' hidden ' uses original dims. """
@@ -3437,7 +3443,7 @@ class TestEntrypointMetaOverride:
grouping = " raw " ,
override_dims = [ " logits:t v " ] ,
)
self . _assert_all_passed ( _run_and_parse ( args , capsys ) , expected_count = 2 )
self . _assert_all_passed ( _run_and_parse ( args , capsys ) [ 0 ] , expected_count = 2 )
def test_multiple_cli_override_dims ( self , tmp_path : Path , capsys ) - > None :
""" Multiple --override-dims for different tensors. """
@@ -3470,7 +3476,7 @@ class TestEntrypointMetaOverride:
grouping = " raw " ,
override_dims = [ " hidden:t h " , " logits:t v " ] ,
)
self . _assert_all_passed ( _run_and_parse ( args , capsys ) , expected_count = 2 )
self . _assert_all_passed ( _run_and_parse ( args , capsys ) [ 0 ] , expected_count = 2 )
def test_per_side_dims_different_parallelism ( self , tmp_path : Path , capsys ) - > None :
""" baseline TP-sharded, target EP-sharded — per-side override fixes both. """
@@ -3512,7 +3518,7 @@ class TestEntrypointMetaOverride:
override_baseline_dims = [ " hidden:t h(tp) " ] ,
override_target_dims = [ " hidden:t h(ep) " ] ,
)
self . _assert_all_passed ( _run_and_parse ( args , capsys ) )
self . _assert_all_passed ( _run_and_parse ( args , capsys ) [ 0 ] )
def test_yaml_first_match_wins_e2e ( self , tmp_path : Path , capsys ) - > None :
""" YAML with two matching rules: first rule wins in real pipeline. """
@@ -3533,7 +3539,7 @@ class TestEntrypointMetaOverride:
grouping = " raw " ,
override_config = str ( yaml_path ) ,
)
self . _assert_all_passed ( _run_and_parse ( args , capsys ) )
self . _assert_all_passed ( _run_and_parse ( args , capsys ) [ 0 ] )
def test_cli_overrides_yaml_e2e ( self , tmp_path : Path , capsys ) - > None :
""" CLI --override-dims wins over YAML rule for the same tensor. """
@@ -3553,7 +3559,7 @@ class TestEntrypointMetaOverride:
override_dims = [ " hidden:t h " ] ,
override_config = str ( yaml_path ) ,
)
self . _assert_all_passed ( _run_and_parse ( args , capsys ) )
self . _assert_all_passed ( _run_and_parse ( args , capsys ) [ 0 ] )
def test_override_injects_dims_when_absent ( self , tmp_path : Path , capsys ) - > None :
""" Override injects dims into meta even when dump had no dims annotation. """
@@ -3569,7 +3575,7 @@ class TestEntrypointMetaOverride:
grouping = " raw " ,
override_dims = [ " hidden:t h " ] ,
)
self . _assert_all_passed ( _run_and_parse ( args , capsys ) )
self . _assert_all_passed ( _run_and_parse ( args , capsys ) [ 0 ] )
def test_non_tensor_unaffected_by_override ( self , tmp_path : Path , capsys ) - > None :
""" Non-tensor values pass through without error even with active override. """
@@ -3596,7 +3602,7 @@ class TestEntrypointMetaOverride:
grouping = " raw " ,
override_dims = [ " hidden:x y " ] ,
)
records = _run_and_parse ( args , capsys )
records , _ = _run_and_parse ( args , capsys )
non_tensors : list [ NonTensorRecord ] = [
r for r in records if isinstance ( r , NonTensorRecord )
@@ -3613,5 +3619,297 @@ class TestEntrypointMetaOverride:
assert summary . failed == 0
class TestExitCode :
""" Tests for exit code behavior based on comparison results. """
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_skip_pattern = " .* " , skipped_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_skip_pattern = " .* " , skipped_names = [ ] ) == 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_skip_pattern = " .* " , skipped_names = [ ] ) == 1
)
def test_all_skipped_allow_all ( self ) :
""" All skipped + allow_skip_pattern= ' .* ' → exit 0. """
summary = SummaryRecord ( total = 2 , passed = 0 , failed = 0 , skipped = 2 )
assert (
_compute_exit_code (
summary , allow_skip_pattern = " .* " , skipped_names = [ " a " , " b " ]
)
== 0
)
def test_all_skipped_forbid_all ( self ) :
""" All skipped + allow_skip_pattern= ' ^$ ' → exit 1. """
summary = SummaryRecord ( total = 2 , passed = 0 , failed = 0 , skipped = 2 )
assert (
_compute_exit_code (
summary , allow_skip_pattern = " ^$ " , skipped_names = [ " a " , " b " ]
)
== 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_skip_pattern = " .* " , skipped_names = [ " a " ] )
== 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_skip_pattern = " ^$ " , skipped_names = [ " a " ] )
== 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_skip_pattern = " positions|seq_lens " ,
skipped_names = [ " positions " , " seq_lens " ] ,
)
== 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_skip_pattern = " positions|seq_lens " ,
skipped_names = [ " positions " , " seq_lens " , " hidden_states " ] ,
)
== 1
)
def test_e2e_all_passed_exit_zero ( self , tmp_path , capsys ) :
""" Integration: all comparisons pass → run() returns 0. """
baseline_path , target_path = _create_dumps ( tmp_path , [ " tensor_a " , " tensor_b " ] )
args = _make_args ( baseline_path , target_path , grouping = " raw " )
records , exit_code = _run_and_parse ( args , capsys )
summary = records [ - 1 ]
assert isinstance ( summary , SummaryRecord )
assert summary . passed == 2
assert summary . failed == 0
assert exit_code == 0
def test_e2e_has_failed_exit_nonzero ( self , tmp_path , capsys ) :
""" Integration: a failed comparison → run() returns 1. """
torch . manual_seed ( 42 )
baseline_path = _create_rank_dump (
tmp_path / " baseline " , rank = 0 , name = " tensor_a " , tensor = torch . randn ( 10 , 10 )
)
target_path = _create_rank_dump (
tmp_path / " target " ,
rank = 0 ,
name = " tensor_a " ,
tensor = torch . randn ( 10 , 10 ) * 100 ,
)
args = _make_args (
baseline_path , target_path , grouping = " raw " , diff_threshold = 1e-3
)
records , exit_code = _run_and_parse ( args , capsys )
summary = records [ - 1 ]
assert isinstance ( summary , SummaryRecord )
assert summary . failed == 1
assert exit_code == 1
class TestExitCodeSubprocess :
""" E2E subprocess tests: invoke comparator as a child process and verify exit code. """
@staticmethod
def _run_comparator (
baseline_path : Path ,
target_path : Path ,
* ,
grouping : str = " raw " ,
allow_skip_pattern : str = " .* " ,
) - > subprocess . CompletedProcess [ str ] :
cmd : list [ str ] = [
sys . executable ,
" -m " ,
" sglang.srt.debug_utils.comparator " ,
" --baseline-path " ,
str ( baseline_path ) ,
" --target-path " ,
str ( target_path ) ,
" --grouping " ,
grouping ,
" --output-format " ,
" json " ,
" --allow-skip-pattern " ,
allow_skip_pattern ,
]
return subprocess . run ( cmd , capture_output = True , text = True )
def test_all_passed_exit_zero ( self , tmp_path ) :
""" Subprocess: all comparisons pass → exit 0. """
baseline_path , target_path = _create_dumps ( tmp_path , [ " tensor_a " ] )
result = self . _run_comparator ( baseline_path , target_path )
assert result . returncode == 0
def test_failed_exit_nonzero ( self , tmp_path ) :
""" Subprocess: failed comparison → exit 1. """
torch . manual_seed ( 42 )
baseline_path = _create_rank_dump (
tmp_path / " baseline " , rank = 0 , name = " t " , tensor = torch . randn ( 10 , 10 )
)
target_path = _create_rank_dump (
tmp_path / " target " , rank = 0 , name = " t " , tensor = torch . randn ( 10 , 10 ) * 100
)
result = self . _run_comparator ( baseline_path , target_path )
assert result . returncode == 1
def test_skipped_allow_all_exit_zero ( self , tmp_path ) :
""" Subprocess: skipped comparison with allow_skip_pattern= ' .* ' → exit 0. """
baseline_path , target_path = _create_dumps (
tmp_path ,
tensor_names = [ " tensor_a " , " tensor_extra " ] ,
baseline_names = [ " tensor_a " ] ,
)
result = self . _run_comparator (
baseline_path , target_path , allow_skip_pattern = " .* "
)
assert result . returncode == 0
def test_skipped_forbid_all_exit_nonzero ( self , tmp_path ) :
""" Subprocess: skipped comparison with allow_skip_pattern= ' ^$ ' → exit 1. """
baseline_path , target_path = _create_dumps (
tmp_path ,
tensor_names = [ " tensor_a " , " tensor_extra " ] ,
baseline_names = [ " tensor_a " ] ,
)
result = self . _run_comparator (
baseline_path , target_path , allow_skip_pattern = " ^$ "
)
assert result . returncode == 1
class TestReportOutput :
""" Test JSONL report file output via ReportSink. """
def test_default_report_path ( self , tmp_path , capsys ) :
""" Default writes to <target>/comparator_report.jsonl with ConfigRecord + SummaryRecord. """
baseline_path , target_path = _create_dumps ( tmp_path , [ " tensor_a " ] )
args = _make_args ( baseline_path , target_path , grouping = " raw " , report_path = None )
exit_code : int = run ( args )
report_file : Path = target_path / " comparator_report.jsonl "
assert report_file . exists ( )
report_records : list [ AnyRecord ] = _parse_jsonl ( report_file . read_text ( ) )
assert isinstance ( report_records [ 0 ] , ConfigRecord )
assert isinstance ( report_records [ - 1 ] , SummaryRecord )
assert exit_code == 0
def test_custom_report_path ( self , tmp_path , capsys ) :
""" --report-path writes to the specified location. """
baseline_path , target_path = _create_dumps ( tmp_path , [ " tensor_a " ] )
custom_path : Path = tmp_path / " custom " / " report.jsonl "
args = _make_args (
baseline_path ,
target_path ,
grouping = " raw " ,
report_path = str ( custom_path ) ,
)
run ( args )
assert custom_path . exists ( )
report_records : list [ AnyRecord ] = _parse_jsonl ( custom_path . read_text ( ) )
assert isinstance ( report_records [ 0 ] , ConfigRecord )
assert isinstance ( report_records [ - 1 ] , SummaryRecord )
def test_disabled_report ( self , tmp_path , capsys ) :
""" --report-path ' ' disables file generation. """
baseline_path , target_path = _create_dumps ( tmp_path , [ " tensor_a " ] )
args = _make_args ( baseline_path , target_path , grouping = " raw " , report_path = " " )
run ( args )
report_file : Path = target_path / " comparator_report.jsonl "
assert not report_file . exists ( )
def test_report_matches_stdout_json ( self , tmp_path , capsys ) :
""" In json mode, report content matches stdout output. """
baseline_path , target_path = _create_dumps ( tmp_path , [ " tensor_a " ] )
report_file : Path = tmp_path / " report.jsonl "
args = _make_args (
baseline_path ,
target_path ,
grouping = " raw " ,
output_format = " json " ,
report_path = str ( report_file ) ,
)
capsys . readouterr ( )
run ( args )
stdout_lines : list [ str ] = capsys . readouterr ( ) . out . strip ( ) . splitlines ( )
report_lines : list [ str ] = report_file . read_text ( ) . strip ( ) . splitlines ( )
assert stdout_lines == report_lines
def test_text_mode_also_writes_report ( self , tmp_path , capsys ) :
""" Text stdout mode still writes JSONL report. """
baseline_path , target_path = _create_dumps ( tmp_path , [ " tensor_a " ] )
report_file : Path = tmp_path / " report.jsonl "
args = _make_args (
baseline_path ,
target_path ,
grouping = " raw " ,
output_format = " text " ,
report_path = str ( report_file ) ,
)
run ( args )
assert report_file . exists ( )
report_records : list [ AnyRecord ] = _parse_jsonl ( report_file . read_text ( ) )
assert isinstance ( report_records [ 0 ] , ConfigRecord )
assert isinstance ( report_records [ - 1 ] , SummaryRecord )
def test_streaming_flush ( self , tmp_path , capsys ) :
""" Report file is flushed after each record (readable before close). """
from sglang . srt . debug_utils . comparator . output_types import report_sink
report_file : Path = tmp_path / " stream_report.jsonl "
report_sink . configure (
output_format = " json " ,
report_path = report_file ,
)
report_sink . add ( ConfigRecord ( config = { " test " : True } ) )
content : str = report_file . read_text ( )
assert len ( content . strip ( ) . splitlines ( ) ) == 1
parsed : AnyRecord = parse_record_json ( content . strip ( ) )
assert isinstance ( parsed , ConfigRecord )
if __name__ == " __main__ " :
sys . exit ( pytest . main ( [ __file__ ] ) )