diff --git a/python/sglang/srt/debug_utils/comparator/bundle_comparator.py b/python/sglang/srt/debug_utils/comparator/bundle_comparator.py index eca4dc8df..a0913efdc 100644 --- a/python/sglang/srt/debug_utils/comparator/bundle_comparator.py +++ b/python/sglang/srt/debug_utils/comparator/bundle_comparator.py @@ -25,6 +25,7 @@ from sglang.srt.debug_utils.comparator.dims import ( resolve_dim_names, ) from sglang.srt.debug_utils.comparator.dp_utils import filter_to_non_empty_dp_rank +from sglang.srt.debug_utils.comparator.meta_overrider import MetaOverrider from sglang.srt.debug_utils.comparator.output_types import ( ComparisonRecord, GeneralWarning, @@ -54,6 +55,7 @@ def compare_bundle_pair( ), viz_output_dir: Optional[Path] = None, compute_per_token: bool = False, + meta_overrider: Optional[MetaOverrider] = None, ) -> Union[ComparisonRecord, SkipRecord, NonTensorRecord]: with warning_sink.context() as collected_warnings: result = _compare_bundle_pair_inner( @@ -66,6 +68,7 @@ def compare_bundle_pair( thd_seq_lens_by_step_pair=thd_seq_lens_by_step_pair, viz_output_dir=viz_output_dir, compute_per_token=compute_per_token, + meta_overrider=meta_overrider, ) return result.model_copy(update={"warnings": collected_warnings}) @@ -84,6 +87,7 @@ def _compare_bundle_pair_inner( ), viz_output_dir: Optional[Path] = None, compute_per_token: bool = False, + meta_overrider: Optional[MetaOverrider] = None, ) -> Union[ComparisonRecord, SkipRecord, NonTensorRecord]: # 1. Load all successfully loaded values all_pair: Pair[list[ValueWithMeta]] = Pair( @@ -101,6 +105,24 @@ def _compare_bundle_pair_inner( y=filter_to_non_empty_dp_rank(all_pair.y), ) + # 1c. Dims override: patch meta["dims"] before downstream reads it + if meta_overrider is not None and not meta_overrider.is_empty: + _apply = meta_overrider.apply_to_meta + all_pair = Pair( + x=[ + ValueWithMeta( + value=v.value, meta=_apply(name=name, meta=v.meta, side="baseline") + ) + for v in all_pair.x + ], + y=[ + ValueWithMeta( + value=v.value, meta=_apply(name=name, meta=v.meta, side="target") + ) + for v in all_pair.y + ], + ) + # 2. Check if any side has non-tensor values → non-tensor display path has_non_tensor: bool = any( not isinstance(it.value, torch.Tensor) for it in [*all_pair.x, *all_pair.y] diff --git a/python/sglang/srt/debug_utils/comparator/entrypoint.py b/python/sglang/srt/debug_utils/comparator/entrypoint.py index 840f79da6..e84da0814 100644 --- a/python/sglang/srt/debug_utils/comparator/entrypoint.py +++ b/python/sglang/srt/debug_utils/comparator/entrypoint.py @@ -22,6 +22,7 @@ from sglang.srt.debug_utils.comparator.bundle_matcher import ( match_bundles, ) from sglang.srt.debug_utils.comparator.display import emit_display_records +from sglang.srt.debug_utils.comparator.meta_overrider import MetaOverrider from sglang.srt.debug_utils.comparator.output_types import ( ComparisonRecord, ConfigRecord, @@ -85,6 +86,13 @@ def run(args: argparse.Namespace) -> None: Path(args.visualize_per_token) if args.visualize_per_token else None ) + meta_overrider: MetaOverrider = MetaOverrider.from_args_and_config( + override_dims=args.override_dims, + override_baseline_dims=args.override_baseline_dims, + override_target_dims=args.override_target_dims, + override_config=Path(args.override_config) if args.override_config else None, + ) + comparison_records = _compare_bundle_pairs( bundle_info_pairs=bundle_info_pairs, baseline_path=Path(args.baseline_path), @@ -94,6 +102,7 @@ def run(args: argparse.Namespace) -> None: thd_seq_lens_by_step_pair=ta_result.thd_seq_lens_by_step_pair, viz_output_dir=viz_output_dir, compute_per_token=visualize_per_token is not None, + meta_overrider=meta_overrider, ) _consume_comparison_records( comparison_records=comparison_records, @@ -155,6 +164,7 @@ def _compare_bundle_pairs( thd_seq_lens_by_step_pair: Pair[Optional[dict[int, list[int]]]], viz_output_dir: Optional[Path] = None, compute_per_token: bool = False, + meta_overrider: Optional[MetaOverrider] = None, ) -> Iterator[Union[ComparisonRecord, SkipRecord, NonTensorRecord]]: for bundle_info_pair in bundle_info_pairs: if not bundle_info_pair.y: @@ -174,6 +184,7 @@ def _compare_bundle_pairs( thd_seq_lens_by_step_pair=thd_seq_lens_by_step_pair, viz_output_dir=viz_output_dir, compute_per_token=compute_per_token, + meta_overrider=meta_overrider, ) @@ -252,4 +263,31 @@ def _parse_args() -> argparse.Namespace: default=None, help="Output path for per-token relative difference heatmap PNG", ) + + # Dims override + parser.add_argument( + "--override-dims", + action="append", + default=[], + help="Override dims for both sides: 'name:dims_string' (repeatable)", + ) + parser.add_argument( + "--override-baseline-dims", + action="append", + default=[], + help="Override dims for baseline only: 'name:dims_string' (repeatable)", + ) + parser.add_argument( + "--override-target-dims", + action="append", + default=[], + help="Override dims for target only: 'name:dims_string' (repeatable)", + ) + parser.add_argument( + "--override-config", + type=str, + default=None, + help="Path to YAML override config file (dims overrides, etc.)", + ) + return parser.parse_args() diff --git a/python/sglang/srt/debug_utils/comparator/meta_overrider.py b/python/sglang/srt/debug_utils/comparator/meta_overrider.py new file mode 100644 index 000000000..c1ae48eb2 --- /dev/null +++ b/python/sglang/srt/debug_utils/comparator/meta_overrider.py @@ -0,0 +1,107 @@ +"""Meta overrider: replace metadata fields without re-running dumps. + +Currently only overrides 'dims', but the design supports overriding +additional meta fields (e.g. parallel_info) in the future. +""" + +from __future__ import annotations + +import re +from pathlib import Path +from typing import Any, Literal, Optional + +import yaml + +from sglang.srt.debug_utils.comparator.utils import _StrictBase + + +class MetaOverrideRule(_StrictBase): + """Single override rule: regex match on tensor name → replacement meta field(s). + + Currently only 'dims' is supported; more fields may be added in the future. + """ + + match: str + dims: str + side: Literal["both", "baseline", "target"] = "both" + + +class MetaOverrideConfig(_StrictBase): + """YAML top-level config for overriding comparator behavior.""" + + overrides: list[MetaOverrideRule] = [] + + +class MetaOverrider: + """Holds override rules and applies first-match-wins replacement.""" + + def __init__(self, rules: list[MetaOverrideRule]) -> None: + self._rules: list[MetaOverrideRule] = rules + + @property + def is_empty(self) -> bool: + return len(self._rules) == 0 + + @classmethod + def from_args_and_config( + cls, + *, + override_dims: list[str], + override_baseline_dims: list[str], + override_target_dims: list[str], + override_config: Optional[Path], + ) -> "MetaOverrider": + per_side_args: list[tuple[list[str], Literal["both", "baseline", "target"]]] = [ + (override_dims, "both"), + (override_baseline_dims, "baseline"), + (override_target_dims, "target"), + ] + cli_rules: list[MetaOverrideRule] = [ + MetaOverrideRule(match=name, dims=dims_str, side=side) + for raw_args, side in per_side_args + for name, dims_str in [_parse_cli_override_arg(raw) for raw in raw_args] + ] + + yaml_rules: list[MetaOverrideRule] = ( + _load_yaml_rules(override_config) if override_config is not None else [] + ) + + return cls(rules=cli_rules + yaml_rules) + + def apply_to_meta( + self, + *, + name: str, + meta: dict[str, Any], + side: Literal["baseline", "target"], + ) -> dict[str, Any]: + """First-match-wins: return meta with dims replaced by the first matching rule for this side.""" + for rule in self._rules: + if rule.side not in ("both", side): + continue + if re.search(rule.match, name): + return {**meta, "dims": rule.dims} + + return meta + + +def _parse_cli_override_arg(raw: str) -> tuple[str, str]: + """Parse 'name:dims_string' from a CLI --override-* argument.""" + parts: list[str] = raw.split(":", maxsplit=1) + if len(parts) != 2 or not parts[0].strip() or not parts[1].strip(): + raise ValueError( + f"Invalid override format: {raw!r}; expected 'name:dims_string'" + ) + return parts[0].strip(), parts[1].strip() + + +def _load_yaml_rules(path: Path) -> list[MetaOverrideRule]: + """Load override rules from a YAML config file.""" + with open(path) as f: + raw_data: Any = yaml.safe_load(f) + + if raw_data is None: + return [] + + config: MetaOverrideConfig = MetaOverrideConfig.model_validate(raw_data) + return config.overrides diff --git a/test/registered/debug_utils/comparator/test_entrypoint.py b/test/registered/debug_utils/comparator/test_entrypoint.py index 1fc285377..b25410dba 100644 --- a/test/registered/debug_utils/comparator/test_entrypoint.py +++ b/test/registered/debug_utils/comparator/test_entrypoint.py @@ -1,4 +1,5 @@ import sys +import textwrap from argparse import Namespace from pathlib import Path @@ -1044,107 +1045,6 @@ class TestEntrypointGroupingLogical: comp = _assert_single_comparison_passed(records) assert comp.name == "hidden" - def test_tp_partial_reduction_unshard(self, tmp_path, capsys): - """TP=2 with partial reduction: element-wise sum reconstructs full tensor.""" - torch.manual_seed(42) - full_baseline = torch.randn(4, 8) - full_target = full_baseline + torch.randn(4, 8) * 0.001 - - baseline_dir = tmp_path / "baseline" - target_dir = tmp_path / "target" - - baseline_path = _create_tp_partial_dumps( - baseline_dir, - full_tensor=full_baseline, - name="attn_out", - tp_size=2, - dims_str="b h(tp,partial)", - ) - target_path = _create_tp_partial_dumps( - target_dir, - full_tensor=full_target, - name="attn_out", - tp_size=2, - dims_str="b h(tp,partial)", - ) - - args = _make_args(baseline_path, target_path, diff_threshold=0.01) - - records = _run_and_parse(args, capsys) - comp = _assert_single_comparison_passed(records) - assert comp.name == "attn_out" - - summary = records[-1] - assert isinstance(summary, SummaryRecord) - assert summary.total == 1 - assert summary.passed == 1 - - def test_tp_partial_vs_single_rank(self, tmp_path, capsys): - """Baseline single rank vs target TP=2 partial: unshard target then compare.""" - torch.manual_seed(42) - full_tensor = torch.randn(4, 8) - target_full = full_tensor + torch.randn(4, 8) * 0.001 - - baseline_dir = tmp_path / "baseline" - target_dir = tmp_path / "target" - - baseline_path = _create_rank_dump( - baseline_dir, rank=0, name="attn_out", tensor=full_tensor - ) - target_path = _create_tp_partial_dumps( - target_dir, - full_tensor=target_full, - name="attn_out", - tp_size=2, - dims_str="b h(tp,partial)", - ) - - args = _make_args(baseline_path, target_path, diff_threshold=0.01) - - records = _run_and_parse(args, capsys) - comp = _assert_single_comparison_passed(records) - assert comp.name == "attn_out" - - def test_cp_concat_tp_partial_reduction(self, tmp_path, capsys): - """CP=2 concat + TP=2 partial reduction: multi-axis unshard.""" - torch.manual_seed(42) - full_baseline = torch.randn(4, 8, 16) - full_target = full_baseline + torch.randn(4, 8, 16) * 0.001 - - for side_dir, full_tensor in [ - (tmp_path / "baseline", full_baseline), - (tmp_path / "target", full_target), - ]: - side_dir.mkdir() - cp_chunks = list(full_tensor.chunk(2, dim=1)) - rank = 0 - for cp_rank in range(2): - for tp_rank in range(2): - _create_rank_dump( - side_dir, - rank=rank, - name="hidden", - tensor=cp_chunks[cp_rank] / 2, - dims="b s(cp) h(tp,partial)", - parallel_info={ - "cp_rank": cp_rank, - "cp_size": 2, - "tp_rank": tp_rank, - "tp_size": 2, - }, - ) - rank += 1 - - args = _make_args( - tmp_path / "baseline" / _FIXED_EXP_NAME, - tmp_path / "target" / _FIXED_EXP_NAME, - diff_threshold=0.01, - ) - - records = _run_and_parse(args, capsys) - comp = _assert_single_comparison_passed(records) - assert comp.name == "hidden" - class TestEntrypointAxisAligner: """Test cross-framework dim reordering through the full entrypoint pipeline.""" @@ -1988,6 +1888,10 @@ def _make_args(baseline_path: Path, target_path: Path, **overrides) -> Namespace viz_bundle_details=False, viz_output_dir="/tmp/comparator_viz/", visualize_per_token=None, + override_dims=[], + override_baseline_dims=[], + override_target_dims=[], + override_config=None, ) defaults.update(overrides) return Namespace(**defaults) @@ -2837,5 +2741,353 @@ class TestEntrypointDpFilter: _run_and_parse(args, capsys) +class TestEntrypointMetaOverride: + """E2E: dump with wrong dims → --override-dims / --override-config corrects at comparison time.""" + + @staticmethod + def _create_single_rank_pair( + tmp_path: Path, + *, + name: str = "hidden", + baseline_dims: str | None = "x y", + target_dims: str | None = "x y", + ) -> tuple[Path, Path]: + """Create single-rank baseline+target dumps with a close tensor pair.""" + torch.manual_seed(42) + tensor: torch.Tensor = torch.randn(10, 8) + target: torch.Tensor = tensor + torch.randn(10, 8) * 0.001 + + baseline_dir: Path = tmp_path / "baseline" + target_dir: Path = tmp_path / "target" + baseline_dir.mkdir() + target_dir.mkdir() + + _create_rank_dump( + baseline_dir, rank=0, name=name, tensor=tensor, dims=baseline_dims + ) + _create_rank_dump( + target_dir, rank=0, name=name, tensor=target, dims=target_dims + ) + + return baseline_dir / _FIXED_EXP_NAME, target_dir / _FIXED_EXP_NAME + + @staticmethod + def _assert_all_passed( + records: list[AnyRecord], *, expected_count: int = 1 + ) -> None: + """Assert that exactly expected_count comparisons exist and all passed.""" + comparisons: list[ComparisonRecord] = _get_comparisons(records) + assert len(comparisons) == expected_count + assert all(c.diff is not None and c.diff.passed for c in comparisons) + + def test_override_dims_fixes_wrong_dims(self, tmp_path: Path, capsys) -> None: + """Tensor dumped with wrong dims='h d' is fixed by --override-dims to 't h(tp)'.""" + torch.manual_seed(42) + + full_tensor: torch.Tensor = torch.randn(10, 8) + tp_chunks: list[torch.Tensor] = list(full_tensor.chunk(2, dim=1)) + + target_full: torch.Tensor = full_tensor + torch.randn(10, 8) * 0.001 + target_tp_chunks: list[torch.Tensor] = list(target_full.chunk(2, dim=1)) + + baseline_dir: Path = tmp_path / "baseline" + target_dir: Path = tmp_path / "target" + baseline_dir.mkdir() + target_dir.mkdir() + + # Dump with WRONG dims "h d" instead of correct "t h(tp)" + for tp_rank in range(2): + _create_rank_dump( + baseline_dir, + rank=tp_rank, + name="hidden", + tensor=tp_chunks[tp_rank], + dims="h d", + parallel_info={"tp_rank": tp_rank, "tp_size": 2}, + ) + _create_rank_dump( + target_dir, + rank=tp_rank, + name="hidden", + tensor=target_tp_chunks[tp_rank], + dims="h d", + parallel_info={"tp_rank": tp_rank, "tp_size": 2}, + ) + + args = _make_args( + baseline_dir / _FIXED_EXP_NAME, + target_dir / _FIXED_EXP_NAME, + grouping="logical", + override_dims=["hidden:t h(tp)"], + ) + self._assert_all_passed(_run_and_parse(args, capsys)) + + @pytest.mark.parametrize( + "baseline_dims, target_dims, override_kwarg", + [ + ("x y", "t h", {"override_baseline_dims": ["hidden:t h"]}), + ("t h", "x y", {"override_target_dims": ["hidden:t h"]}), + ("x y", "x y", {"override_dims": ["hidden:t h"]}), + ], + ids=["baseline_only", "target_only", "both_via_override_dims"], + ) + def test_single_side_override( + self, + tmp_path: Path, + capsys, + baseline_dims: str, + target_dims: str, + override_kwarg: dict, + ) -> None: + """Per-side override fixes the wrong dims on one or both sides.""" + baseline_path, target_path = self._create_single_rank_pair( + tmp_path, + baseline_dims=baseline_dims, + target_dims=target_dims, + ) + + args = _make_args(baseline_path, target_path, grouping="raw", **override_kwarg) + self._assert_all_passed(_run_and_parse(args, capsys)) + + def test_override_config_yaml(self, tmp_path: Path, capsys) -> None: + """--override-config YAML overrides dims.""" + baseline_path, target_path = self._create_single_rank_pair(tmp_path) + + yaml_path: Path = tmp_path / "override.yaml" + yaml_path.write_text(textwrap.dedent("""\ + overrides: + - match: "hidden" + dims: "t h" + """)) + + args = _make_args( + baseline_path, + target_path, + grouping="raw", + override_config=str(yaml_path), + ) + self._assert_all_passed(_run_and_parse(args, capsys)) + + 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.""" + baseline_path, target_path = self._create_single_rank_pair( + tmp_path, + baseline_dims="t h", + target_dims="t h", + ) + + args = _make_args( + baseline_path, + target_path, + grouping="raw", + override_dims=["no_match_pattern:b s d"], + ) + self._assert_all_passed(_run_and_parse(args, capsys)) + + def test_selective_match_multi_tensor(self, tmp_path: Path, capsys) -> None: + """Override matches only 'logits'; 'hidden' uses original dims.""" + torch.manual_seed(42) + + baseline_dir: Path = tmp_path / "baseline" + target_dir: Path = tmp_path / "target" + baseline_dir.mkdir() + target_dir.mkdir() + + hidden_b: torch.Tensor = torch.randn(10, 8) + hidden_t: torch.Tensor = hidden_b + torch.randn(10, 8) * 0.001 + logits_b: torch.Tensor = torch.randn(10, 4) + logits_t: torch.Tensor = logits_b + torch.randn(10, 4) * 0.001 + + for name, b_tensor, t_tensor, dims in [ + ("hidden", hidden_b, hidden_t, "t h"), + ("logits", logits_b, logits_t, "x y"), + ]: + _create_rank_dump( + baseline_dir, rank=0, name=name, tensor=b_tensor, dims=dims + ) + _create_rank_dump(target_dir, rank=0, name=name, tensor=t_tensor, dims=dims) + + args = _make_args( + baseline_dir / _FIXED_EXP_NAME, + target_dir / _FIXED_EXP_NAME, + grouping="raw", + override_dims=["logits:t v"], + ) + self._assert_all_passed(_run_and_parse(args, capsys), expected_count=2) + + def test_multiple_cli_override_dims(self, tmp_path: Path, capsys) -> None: + """Multiple --override-dims for different tensors.""" + torch.manual_seed(42) + + baseline_dir: Path = tmp_path / "baseline" + target_dir: Path = tmp_path / "target" + baseline_dir.mkdir() + target_dir.mkdir() + + hidden_b: torch.Tensor = torch.randn(10, 8) + hidden_t: torch.Tensor = hidden_b + torch.randn(10, 8) * 0.001 + logits_b: torch.Tensor = torch.randn(10, 4) + logits_t: torch.Tensor = logits_b + torch.randn(10, 4) * 0.001 + + for name, b_tensor, t_tensor in [ + ("hidden", hidden_b, hidden_t), + ("logits", logits_b, logits_t), + ]: + _create_rank_dump( + baseline_dir, rank=0, name=name, tensor=b_tensor, dims="x y" + ) + _create_rank_dump( + target_dir, rank=0, name=name, tensor=t_tensor, dims="x y" + ) + + args = _make_args( + baseline_dir / _FIXED_EXP_NAME, + target_dir / _FIXED_EXP_NAME, + grouping="raw", + override_dims=["hidden:t h", "logits:t v"], + ) + self._assert_all_passed(_run_and_parse(args, capsys), 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.""" + torch.manual_seed(42) + full_tensor: torch.Tensor = torch.randn(10, 8) + target_full: torch.Tensor = full_tensor + torch.randn(10, 8) * 0.001 + + baseline_dir: Path = tmp_path / "baseline" + target_dir: Path = tmp_path / "target" + baseline_dir.mkdir() + target_dir.mkdir() + + b_chunks: list[torch.Tensor] = list(full_tensor.chunk(2, dim=1)) + for tp_rank in range(2): + _create_rank_dump( + baseline_dir, + rank=tp_rank, + name="hidden", + tensor=b_chunks[tp_rank], + dims="x y", + parallel_info={"tp_rank": tp_rank, "tp_size": 2}, + ) + + t_chunks: list[torch.Tensor] = list(target_full.chunk(2, dim=1)) + for ep_rank in range(2): + _create_rank_dump( + target_dir, + rank=ep_rank, + name="hidden", + tensor=t_chunks[ep_rank], + dims="x y", + parallel_info={"ep_rank": ep_rank, "ep_size": 2}, + ) + + args = _make_args( + baseline_dir / _FIXED_EXP_NAME, + target_dir / _FIXED_EXP_NAME, + grouping="logical", + override_baseline_dims=["hidden:t h(tp)"], + override_target_dims=["hidden:t h(ep)"], + ) + self._assert_all_passed(_run_and_parse(args, capsys)) + + def test_yaml_first_match_wins_e2e(self, tmp_path: Path, capsys) -> None: + """YAML with two matching rules: first rule wins in real pipeline.""" + baseline_path, target_path = self._create_single_rank_pair(tmp_path) + + yaml_path: Path = tmp_path / "override.yaml" + yaml_path.write_text(textwrap.dedent("""\ + overrides: + - match: "hidden" + dims: "t h" + - match: "hidden" + dims: "a b" + """)) + + args = _make_args( + baseline_path, + target_path, + grouping="raw", + override_config=str(yaml_path), + ) + self._assert_all_passed(_run_and_parse(args, capsys)) + + def test_cli_overrides_yaml_e2e(self, tmp_path: Path, capsys) -> None: + """CLI --override-dims wins over YAML rule for the same tensor.""" + baseline_path, target_path = self._create_single_rank_pair(tmp_path) + + yaml_path: Path = tmp_path / "override.yaml" + yaml_path.write_text(textwrap.dedent("""\ + overrides: + - match: "hidden" + dims: "a b" + """)) + + args = _make_args( + baseline_path, + target_path, + grouping="raw", + override_dims=["hidden:t h"], + override_config=str(yaml_path), + ) + self._assert_all_passed(_run_and_parse(args, capsys)) + + 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.""" + baseline_path, target_path = self._create_single_rank_pair( + tmp_path, + baseline_dims=None, + target_dims=None, + ) + + args = _make_args( + baseline_path, + target_path, + grouping="raw", + override_dims=["hidden:t h"], + ) + self._assert_all_passed(_run_and_parse(args, capsys)) + + def test_non_tensor_unaffected_by_override(self, tmp_path: Path, capsys) -> None: + """Non-tensor values pass through without error even with active override.""" + torch.manual_seed(42) + tensor: torch.Tensor = torch.randn(4, 4) + + baseline_dir: Path = tmp_path / "baseline" + target_dir: Path = tmp_path / "target" + baseline_dir.mkdir() + target_dir.mkdir() + + for side_dir in [baseline_dir, target_dir]: + _create_non_tensor_rank_dump( + side_dir, + rank=0, + name="sm_scale", + value=0.125, + extra_tensor_dumps=[("hidden", tensor)], + ) + + args = _make_args( + baseline_dir / _FIXED_EXP_NAME, + target_dir / _FIXED_EXP_NAME, + grouping="raw", + override_dims=["hidden:x y"], + ) + records = _run_and_parse(args, capsys) + + non_tensors: list[NonTensorRecord] = [ + r for r in records if isinstance(r, NonTensorRecord) + ] + assert len(non_tensors) == 1 + assert non_tensors[0].name == "sm_scale" + assert non_tensors[0].values_equal + + comparisons: list[ComparisonRecord] = _get_comparisons(records) + assert len(comparisons) == 1 + assert comparisons[0].name == "hidden" + + summary: SummaryRecord = [r for r in records if isinstance(r, SummaryRecord)][0] + assert summary.failed == 0 + + if __name__ == "__main__": sys.exit(pytest.main([__file__])) diff --git a/test/registered/debug_utils/comparator/test_meta_overrider.py b/test/registered/debug_utils/comparator/test_meta_overrider.py new file mode 100644 index 000000000..90a1911a8 --- /dev/null +++ b/test/registered/debug_utils/comparator/test_meta_overrider.py @@ -0,0 +1,297 @@ +"""Tests for meta_overrider — unit tests.""" + +from __future__ import annotations + +import sys +import textwrap +from pathlib import Path + +import pytest + +from sglang.srt.debug_utils.comparator.meta_overrider import ( + MetaOverrider, + MetaOverrideRule, + _load_yaml_rules, + _parse_cli_override_arg, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=10, suite="default", nightly=True) + + +# ───────────────────── Unit: MetaOverrideRule ───────────────────── + + +class TestMetaOverrideRule: + """Pydantic validation for MetaOverrideRule.""" + + def test_shared_dims_both(self) -> None: + """Default side='both' applies dims to both sides.""" + rule = MetaOverrideRule(match="hidden", dims="b s h d") + assert rule.dims == "b s h d" + assert rule.side == "both" + + def test_side_baseline(self) -> None: + """side='baseline' is accepted.""" + rule = MetaOverrideRule(match="logits", dims="b s v(tp)", side="baseline") + assert rule.dims == "b s v(tp)" + assert rule.side == "baseline" + + def test_side_target(self) -> None: + """side='target' is accepted.""" + rule = MetaOverrideRule(match="logits", dims="b s v(ep)", side="target") + assert rule.dims == "b s v(ep)" + assert rule.side == "target" + + def test_invalid_side_rejected(self) -> None: + """Invalid side value is rejected.""" + with pytest.raises(Exception): + MetaOverrideRule(match="x", dims="b s", side="invalid") + + def test_dims_required(self) -> None: + """Must specify dims.""" + with pytest.raises(Exception): + MetaOverrideRule(match="x") + + def test_extra_field_rejected(self) -> None: + """Extra fields are rejected by _StrictBase.""" + with pytest.raises(Exception): + MetaOverrideRule(match="x", dims="b s", bogus="y") + + +# ──────────────────── Unit: _parse_cli_override_arg ──────────────────── + + +class TestParseCLIOverrideArg: + """CLI arg parsing for 'name:dims_string' format.""" + + def test_basic(self) -> None: + """Standard 'name:dims' parsing.""" + name, dims_str = _parse_cli_override_arg("hidden_states:b s h d") + assert name == "hidden_states" + assert dims_str == "b s h d" + + def test_colon_in_dims(self) -> None: + """Extra colons in dims are kept (maxsplit=1).""" + name, dims_str = _parse_cli_override_arg("x:a:b") + assert name == "x" + assert dims_str == "a:b" + + def test_whitespace_trimmed(self) -> None: + """Leading/trailing whitespace around name and dims is stripped.""" + name, dims_str = _parse_cli_override_arg(" foo : b s ") + assert name == "foo" + assert dims_str == "b s" + + def test_missing_colon(self) -> None: + """No colon raises ValueError.""" + with pytest.raises(ValueError, match="Invalid override format"): + _parse_cli_override_arg("no_colon_here") + + def test_empty_name(self) -> None: + """Empty name raises ValueError.""" + with pytest.raises(ValueError, match="Invalid override format"): + _parse_cli_override_arg(":b s h") + + def test_empty_dims(self) -> None: + """Empty dims raises ValueError.""" + with pytest.raises(ValueError, match="Invalid override format"): + _parse_cli_override_arg("foo:") + + +# ──────────────────── Unit: MetaOverrider ──────────────────── + + +class TestMetaOverrider: + """MetaOverrider logic: matching, priority, apply_to_meta.""" + + def test_first_match_wins(self) -> None: + """First matching rule takes effect; later rules ignored.""" + overrider = MetaOverrider( + rules=[ + MetaOverrideRule(match="hidden", dims="FIRST"), + MetaOverrideRule(match="hidden", dims="SECOND"), + ] + ) + result: dict = overrider.apply_to_meta( + name="hidden_states", + meta={"dims": "old"}, + side="baseline", + ) + assert result["dims"] == "FIRST" + + def test_regex_contains_match(self) -> None: + """match is a regex contains search, not exact match.""" + overrider = MetaOverrider( + rules=[MetaOverrideRule(match=r"\.q_proj\.", dims="h d")] + ) + result: dict = overrider.apply_to_meta( + name="layers.0.q_proj.weight", + meta={"dims": "old"}, + side="baseline", + ) + assert result["dims"] == "h d" + + def test_no_match_preserves_original(self) -> None: + """No matching rule leaves meta untouched.""" + overrider = MetaOverrider( + rules=[MetaOverrideRule(match="logits", dims="b s v")] + ) + result: dict = overrider.apply_to_meta( + name="hidden_states", + meta={"dims": "original"}, + side="baseline", + ) + assert result["dims"] == "original" + + @pytest.mark.parametrize( + "rule_side,apply_side,should_match", + [ + ("baseline", "baseline", True), + ("baseline", "target", False), + ("target", "target", True), + ("target", "baseline", False), + ("both", "baseline", True), + ("both", "target", True), + ], + ) + def test_side_filtering( + self, rule_side: str, apply_side: str, should_match: bool + ) -> None: + """Rule only applies when its side matches the apply side.""" + overrider = MetaOverrider( + rules=[MetaOverrideRule(match="logits", dims="NEW", side=rule_side)] + ) + result: dict = overrider.apply_to_meta( + name="logits", + meta={"dims": "old"}, + side=apply_side, + ) + assert result["dims"] == ("NEW" if should_match else "old") + + def test_is_empty(self) -> None: + """Empty overrider reports is_empty=True.""" + assert MetaOverrider(rules=[]).is_empty + assert not MetaOverrider(rules=[MetaOverrideRule(match="x", dims="d")]).is_empty + + def test_meta_without_dims_key(self) -> None: + """Override adds 'dims' even if original meta lacks it.""" + overrider = MetaOverrider(rules=[MetaOverrideRule(match="hidden", dims="NEW")]) + result: dict = overrider.apply_to_meta( + name="hidden", + meta={"other": "val"}, + side="baseline", + ) + assert result["dims"] == "NEW" + + +# ──────────────────── Unit: from_args_and_config ──────────────────── + + +class TestFromArgsAndConfig: + """MetaOverrider.from_args_and_config merges CLI + YAML rules.""" + + def test_cli_before_yaml(self, tmp_path: Path) -> None: + """CLI rules are ordered before YAML rules (CLI wins on conflict).""" + yaml_path = tmp_path / "override.yaml" + yaml_path.write_text(textwrap.dedent("""\ + overrides: + - match: "hidden" + dims: "FROM_YAML" + """)) + + overrider = MetaOverrider.from_args_and_config( + override_dims=["hidden:FROM_CLI"], + override_baseline_dims=[], + override_target_dims=[], + override_config=yaml_path, + ) + + result: dict = overrider.apply_to_meta( + name="hidden", + meta={"dims": "old"}, + side="baseline", + ) + assert result["dims"] == "FROM_CLI" + + def test_no_config_no_cli(self) -> None: + """Empty CLI + no YAML yields empty overrider.""" + overrider = MetaOverrider.from_args_and_config( + override_dims=[], + override_baseline_dims=[], + override_target_dims=[], + override_config=None, + ) + assert overrider.is_empty + + def test_per_side_cli_produces_separate_rules(self) -> None: + """--override-baseline-dims and --override-target-dims produce separate rules with side field.""" + overrider = MetaOverrider.from_args_and_config( + override_dims=[], + override_baseline_dims=["hidden:b s h(tp)"], + override_target_dims=["hidden:b s h(ep)"], + override_config=None, + ) + + baseline: dict = overrider.apply_to_meta( + name="hidden", + meta={"dims": "old"}, + side="baseline", + ) + target: dict = overrider.apply_to_meta( + name="hidden", + meta={"dims": "old"}, + side="target", + ) + assert baseline["dims"] == "b s h(tp)" + assert target["dims"] == "b s h(ep)" + + +# ──────────────────── Unit: _load_yaml_rules ──────────────────── + + +class TestLoadYamlRules: + """YAML loading and validation.""" + + def test_valid_yaml(self, tmp_path: Path) -> None: + """Valid YAML with override rules loads correctly.""" + yaml_path = tmp_path / "override.yaml" + yaml_path.write_text(textwrap.dedent("""\ + overrides: + - match: "hidden" + dims: "b s h d" + - match: "logits" + dims: "b s v(tp)" + side: baseline + """)) + rules = _load_yaml_rules(yaml_path) + assert len(rules) == 2 + assert rules[0].dims == "b s h d" + assert rules[0].side == "both" + assert rules[1].dims == "b s v(tp)" + assert rules[1].side == "baseline" + + def test_empty_yaml(self, tmp_path: Path) -> None: + """Empty YAML file returns no rules.""" + yaml_path = tmp_path / "empty.yaml" + yaml_path.write_text("") + rules = _load_yaml_rules(yaml_path) + assert rules == [] + + def test_unknown_top_key_rejected(self, tmp_path: Path) -> None: + """Unknown top-level key is rejected by OverrideConfig.""" + yaml_path = tmp_path / "bad.yaml" + yaml_path.write_text("unknown_key: 42\n") + with pytest.raises(Exception): + _load_yaml_rules(yaml_path) + + def test_overrides_empty_list(self, tmp_path: Path) -> None: + """Only 'overrides' key with no entries returns empty list.""" + yaml_path = tmp_path / "minimal.yaml" + yaml_path.write_text("overrides: []\n") + rules = _load_yaml_rules(yaml_path) + assert rules == [] + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__]))