From 2739d7df622c929d587239acbf11a6db2238cde6 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Thu, 26 Feb 2026 10:00:13 +0800 Subject: [PATCH] Reorganize modules and pipeline in dump comparator (#19374) --- .../debug_utils/comparator/aligner/reorder.py | 82 ----------- .../{unshard => reorderer}/__init__.py | 0 .../comparator/aligner/reorderer/executor.py | 31 ++++ .../comparator/aligner/reorderer/planner.py | 39 +++++ .../comparator/aligner/reorderer/types.py | 16 ++ .../comparator/aligner/unsharder/__init__.py | 0 .../{unshard => unsharder}/executor.py | 41 +++--- .../{unshard => unsharder}/parallel_info.py | 2 +- .../aligner/{unshard => unsharder}/planner.py | 18 +-- .../aligner/{unshard => unsharder}/types.py | 18 ++- .../debug_utils/comparator/output_types.py | 4 +- .../srt/debug_utils/comparator/pipeline.py | 92 +++++++----- .../comparator/tensor_comparator/__init__.py | 3 + .../comparator.py} | 7 +- .../formatter.py | 2 +- .../printer.py | 2 +- .../types.py | 0 .../comparator/tensor_comparison/__init__.py | 1 - .../srt/debug_utils/comparator/utils.py | 4 +- .../comparator/aligner/conftest.py | 0 .../comparator/aligner/reorderer/conftest.py | 0 .../aligner/reorderer/test_executor.py | 50 +++++++ .../test_planner.py} | 94 ++++-------- .../comparator/aligner/unsharder/conftest.py | 0 .../test_executor.py} | 139 ++++++++++-------- .../test_parallel_info.py | 4 +- .../test_planner.py} | 50 +++---- .../debug_utils/comparator/conftest.py | 0 .../comparator/tensor_comparator/conftest.py | 0 .../test_comparator.py} | 16 +- .../test_formatter.py | 4 +- .../test_printer.py | 4 +- .../test_types.py | 4 +- .../comparator/test_model_validation.py | 29 +++- .../debug_utils/comparator/test_utils.py | 36 ++++- 35 files changed, 459 insertions(+), 333 deletions(-) delete mode 100644 python/sglang/srt/debug_utils/comparator/aligner/reorder.py rename python/sglang/srt/debug_utils/comparator/aligner/{unshard => reorderer}/__init__.py (100%) create mode 100644 python/sglang/srt/debug_utils/comparator/aligner/reorderer/executor.py create mode 100644 python/sglang/srt/debug_utils/comparator/aligner/reorderer/planner.py create mode 100644 python/sglang/srt/debug_utils/comparator/aligner/reorderer/types.py create mode 100644 python/sglang/srt/debug_utils/comparator/aligner/unsharder/__init__.py rename python/sglang/srt/debug_utils/comparator/aligner/{unshard => unsharder}/executor.py (64%) rename python/sglang/srt/debug_utils/comparator/aligner/{unshard => unsharder}/parallel_info.py (93%) rename python/sglang/srt/debug_utils/comparator/aligner/{unshard => unsharder}/planner.py (92%) rename python/sglang/srt/debug_utils/comparator/aligner/{unshard => unsharder}/types.py (62%) create mode 100644 python/sglang/srt/debug_utils/comparator/tensor_comparator/__init__.py rename python/sglang/srt/debug_utils/comparator/{tensor_comparison/compare.py => tensor_comparator/comparator.py} (95%) rename python/sglang/srt/debug_utils/comparator/{tensor_comparison => tensor_comparator}/formatter.py (97%) rename python/sglang/srt/debug_utils/comparator/{tensor_comparison => tensor_comparator}/printer.py (97%) rename python/sglang/srt/debug_utils/comparator/{tensor_comparison => tensor_comparator}/types.py (100%) delete mode 100644 python/sglang/srt/debug_utils/comparator/tensor_comparison/__init__.py create mode 100644 test/registered/debug_utils/comparator/aligner/conftest.py create mode 100644 test/registered/debug_utils/comparator/aligner/reorderer/conftest.py create mode 100644 test/registered/debug_utils/comparator/aligner/reorderer/test_executor.py rename test/registered/debug_utils/comparator/aligner/{test_reorder.py => reorderer/test_planner.py} (55%) create mode 100644 test/registered/debug_utils/comparator/aligner/unsharder/conftest.py rename test/registered/debug_utils/comparator/aligner/{unshard/test_execute.py => unsharder/test_executor.py} (78%) rename test/registered/debug_utils/comparator/aligner/{unshard => unsharder}/test_parallel_info.py (92%) rename test/registered/debug_utils/comparator/aligner/{unshard/test_plan.py => unsharder/test_planner.py} (90%) create mode 100644 test/registered/debug_utils/comparator/conftest.py create mode 100644 test/registered/debug_utils/comparator/tensor_comparator/conftest.py rename test/registered/debug_utils/comparator/{tensor_comparison/test_compare.py => tensor_comparator/test_comparator.py} (88%) rename test/registered/debug_utils/comparator/{tensor_comparison => tensor_comparator}/test_formatter.py (98%) rename test/registered/debug_utils/comparator/{tensor_comparison => tensor_comparator}/test_printer.py (98%) rename test/registered/debug_utils/comparator/{tensor_comparison => tensor_comparator}/test_types.py (98%) diff --git a/python/sglang/srt/debug_utils/comparator/aligner/reorder.py b/python/sglang/srt/debug_utils/comparator/aligner/reorder.py deleted file mode 100644 index fbebe6075..000000000 --- a/python/sglang/srt/debug_utils/comparator/aligner/reorder.py +++ /dev/null @@ -1,82 +0,0 @@ -from typing import Literal - -import torch - -from sglang.srt.debug_utils.comparator.aligner.unshard.types import AxisInfo -from sglang.srt.debug_utils.comparator.dims import DimSpec, Ordering, ParallelAxis -from sglang.srt.debug_utils.comparator.utils import _FrozenBase - - -class ZigzagToNaturalParams(_FrozenBase): - op: Literal["zigzag_to_natural"] = "zigzag_to_natural" - dim: int - cp_size: int - - -ReorderParams = ZigzagToNaturalParams - - -class ReorderPlan(_FrozenBase): - params: ReorderParams - - -_ALLOWED_ZIGZAG_DIM_NAMES: set[str] = {"s"} - - -def compute_reorder_plans( - dim_specs: list[DimSpec], - parallel_infos: list[dict[ParallelAxis, AxisInfo]], -) -> list[ReorderPlan]: - plans: list[ReorderPlan] = [] - - for dim_index, spec in enumerate(dim_specs): - if ( - spec.ordering is not None - and spec.ordering != Ordering.NATURAL - and spec.parallel is not None - ): - if spec.name not in _ALLOWED_ZIGZAG_DIM_NAMES: - raise ValueError( - f"Zigzag ordering is only supported on sequence dims " - f"(bshd/sbhd format, dim name must be one of " - f"{sorted(_ALLOWED_ZIGZAG_DIM_NAMES)}), " - f"but got dim name {spec.name!r} in {spec}" - ) - - assert spec.ordering == Ordering.ZIGZAG - axis_size: int = parallel_infos[0][spec.parallel].axis_size - plans.append( - ReorderPlan( - params=ZigzagToNaturalParams(dim=dim_index, cp_size=axis_size), - ) - ) - - return plans - - -def execute_reorder_plan( - plan: ReorderPlan, - tensors: list[torch.Tensor], -) -> list[torch.Tensor]: - return [ - _reorder_zigzag_to_natural( - tensor, dim=plan.params.dim, cp_size=plan.params.cp_size - ) - for tensor in tensors - ] - - -def _reorder_zigzag_to_natural( - tensor: torch.Tensor, *, dim: int, cp_size: int -) -> torch.Tensor: - """Undo CP zigzag interleaving, restoring natural chunk order. - - Generalized from Megatron-LM _undo_attention_load_balancing - (megatron/core/ssm/mamba_context_parallel.py:360-373). - """ - num_chunks: int = cp_size * 2 - chunks: tuple[torch.Tensor, ...] = tensor.chunk(num_chunks, dim=dim) - order: list[int] = [2 * i for i in range(cp_size)] + [ - num_chunks - 2 * i - 1 for i in range(cp_size) - ] - return torch.cat([chunks[i] for i in order], dim=dim) diff --git a/python/sglang/srt/debug_utils/comparator/aligner/unshard/__init__.py b/python/sglang/srt/debug_utils/comparator/aligner/reorderer/__init__.py similarity index 100% rename from python/sglang/srt/debug_utils/comparator/aligner/unshard/__init__.py rename to python/sglang/srt/debug_utils/comparator/aligner/reorderer/__init__.py diff --git a/python/sglang/srt/debug_utils/comparator/aligner/reorderer/executor.py b/python/sglang/srt/debug_utils/comparator/aligner/reorderer/executor.py new file mode 100644 index 000000000..0d261e4c1 --- /dev/null +++ b/python/sglang/srt/debug_utils/comparator/aligner/reorderer/executor.py @@ -0,0 +1,31 @@ +import torch + +from sglang.srt.debug_utils.comparator.aligner.reorderer.types import ReordererPlan + + +def execute_reorderer_plan( + plan: ReordererPlan, + tensors: list[torch.Tensor], +) -> list[torch.Tensor]: + return [ + _reorder_zigzag_to_natural( + tensor, dim=plan.params.dim, cp_size=plan.params.cp_size + ) + for tensor in tensors + ] + + +def _reorder_zigzag_to_natural( + tensor: torch.Tensor, *, dim: int, cp_size: int +) -> torch.Tensor: + """Undo CP zigzag interleaving, restoring natural chunk order. + + Generalized from Megatron-LM _undo_attention_load_balancing + (megatron/core/ssm/mamba_context_parallel.py:360-373). + """ + num_chunks: int = cp_size * 2 + chunks: tuple[torch.Tensor, ...] = tensor.chunk(num_chunks, dim=dim) + order: list[int] = [2 * i for i in range(cp_size)] + [ + num_chunks - 2 * i - 1 for i in range(cp_size) + ] + return torch.cat([chunks[i] for i in order], dim=dim) diff --git a/python/sglang/srt/debug_utils/comparator/aligner/reorderer/planner.py b/python/sglang/srt/debug_utils/comparator/aligner/reorderer/planner.py new file mode 100644 index 000000000..d792d3a4b --- /dev/null +++ b/python/sglang/srt/debug_utils/comparator/aligner/reorderer/planner.py @@ -0,0 +1,39 @@ +from sglang.srt.debug_utils.comparator.aligner.reorderer.types import ( + ReordererPlan, + ZigzagToNaturalParams, +) +from sglang.srt.debug_utils.comparator.aligner.unsharder.types import AxisInfo +from sglang.srt.debug_utils.comparator.dims import DimSpec, Ordering, ParallelAxis + +_ALLOWED_ZIGZAG_DIM_NAMES: set[str] = {"s"} + + +def compute_reorderer_plans( + dim_specs: list[DimSpec], + parallel_infos: list[dict[ParallelAxis, AxisInfo]], +) -> list[ReordererPlan]: + plans: list[ReordererPlan] = [] + + for dim_index, spec in enumerate(dim_specs): + if ( + spec.ordering is not None + and spec.ordering != Ordering.NATURAL + and spec.parallel is not None + ): + if spec.name not in _ALLOWED_ZIGZAG_DIM_NAMES: + raise ValueError( + f"Zigzag ordering is only supported on sequence dims " + f"(bshd/sbhd format, dim name must be one of " + f"{sorted(_ALLOWED_ZIGZAG_DIM_NAMES)}), " + f"but got dim name {spec.name!r} in {spec}" + ) + + assert spec.ordering == Ordering.ZIGZAG + axis_size: int = parallel_infos[0][spec.parallel].axis_size + plans.append( + ReordererPlan( + params=ZigzagToNaturalParams(dim=dim_index, cp_size=axis_size), + ) + ) + + return plans diff --git a/python/sglang/srt/debug_utils/comparator/aligner/reorderer/types.py b/python/sglang/srt/debug_utils/comparator/aligner/reorderer/types.py new file mode 100644 index 000000000..d11f7bbaa --- /dev/null +++ b/python/sglang/srt/debug_utils/comparator/aligner/reorderer/types.py @@ -0,0 +1,16 @@ +from typing import Literal + +from sglang.srt.debug_utils.comparator.utils import _FrozenBase + + +class ZigzagToNaturalParams(_FrozenBase): + op: Literal["zigzag_to_natural"] = "zigzag_to_natural" + dim: int + cp_size: int + + +ReordererParams = ZigzagToNaturalParams + + +class ReordererPlan(_FrozenBase): + params: ReordererParams diff --git a/python/sglang/srt/debug_utils/comparator/aligner/unsharder/__init__.py b/python/sglang/srt/debug_utils/comparator/aligner/unsharder/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/python/sglang/srt/debug_utils/comparator/aligner/unshard/executor.py b/python/sglang/srt/debug_utils/comparator/aligner/unsharder/executor.py similarity index 64% rename from python/sglang/srt/debug_utils/comparator/aligner/unshard/executor.py rename to python/sglang/srt/debug_utils/comparator/aligner/unsharder/executor.py index 209c047aa..77fa2aa0c 100644 --- a/python/sglang/srt/debug_utils/comparator/aligner/unshard/executor.py +++ b/python/sglang/srt/debug_utils/comparator/aligner/unsharder/executor.py @@ -1,56 +1,52 @@ import torch -from sglang.srt.debug_utils.comparator.aligner.unshard.types import ( +from sglang.srt.debug_utils.comparator.aligner.unsharder.types import ( ConcatParams, PickParams, - UnshardParams, - UnshardPlan, + UnsharderParams, + UnsharderPlan, ) from sglang.srt.debug_utils.comparator.dims import ParallelAxis -from sglang.srt.debug_utils.comparator.output_types import ( - AnyWarning, - ReplicatedMismatchWarning, -) +from sglang.srt.debug_utils.comparator.output_types import ReplicatedMismatchWarning +from sglang.srt.debug_utils.comparator.warning_sink import warning_sink -def execute_unshard_plan( - plan: UnshardPlan, +def execute_unsharder_plan( + plan: UnsharderPlan, tensors: list[torch.Tensor], -) -> tuple[list[torch.Tensor], list[AnyWarning]]: - all_warnings: list[AnyWarning] = [] +) -> list[torch.Tensor]: result: list[torch.Tensor] = [] for group_idx, group in enumerate(plan.groups): group_tensors = [tensors[i] for i in group] - tensor, warnings = _apply_unshard( + tensor = _apply_unshard( plan.params, group_tensors, axis=plan.axis, group_index=group_idx, ) result.append(tensor) - all_warnings.extend(warnings) - return result, all_warnings + return result def _apply_unshard( - params: UnshardParams, + params: UnsharderParams, ordered_tensors: list[torch.Tensor], *, axis: ParallelAxis, group_index: int, -) -> tuple[torch.Tensor, list[AnyWarning]]: +) -> torch.Tensor: if isinstance(params, PickParams): - warnings = _verify_replicated_group( + _verify_replicated_group( ordered_tensors, axis=axis, group_index=group_index, ) - return ordered_tensors[0], warnings + return ordered_tensors[0] if isinstance(params, ConcatParams): - return torch.cat(ordered_tensors, dim=params.dim), [] + return torch.cat(ordered_tensors, dim=params.dim) # Phase 2: ReduceSumParams, CpZigzagParams raise ValueError(f"Unsupported unshard operation: {type(params).__name__}") @@ -61,14 +57,13 @@ def _verify_replicated_group( *, axis: ParallelAxis, group_index: int, -) -> list[ReplicatedMismatchWarning]: - warnings: list[ReplicatedMismatchWarning] = [] +) -> None: baseline = ordered_tensors[0] for i in range(1, len(ordered_tensors)): other = ordered_tensors[i] if not torch.allclose(baseline, other, atol=1e-6): - warnings.append( + warning_sink.add( ReplicatedMismatchWarning( axis=axis.value, group_index=group_index, @@ -77,5 +72,3 @@ def _verify_replicated_group( max_abs_diff=(baseline - other).abs().max().item(), ) ) - - return warnings diff --git a/python/sglang/srt/debug_utils/comparator/aligner/unshard/parallel_info.py b/python/sglang/srt/debug_utils/comparator/aligner/unsharder/parallel_info.py similarity index 93% rename from python/sglang/srt/debug_utils/comparator/aligner/unshard/parallel_info.py rename to python/sglang/srt/debug_utils/comparator/aligner/unsharder/parallel_info.py index a12cb6c02..7efd79758 100644 --- a/python/sglang/srt/debug_utils/comparator/aligner/unshard/parallel_info.py +++ b/python/sglang/srt/debug_utils/comparator/aligner/unsharder/parallel_info.py @@ -1,6 +1,6 @@ from typing import Optional -from sglang.srt.debug_utils.comparator.aligner.unshard.types import AxisInfo +from sglang.srt.debug_utils.comparator.aligner.unsharder.types import AxisInfo from sglang.srt.debug_utils.comparator.dims import ParallelAxis _PARALLEL_INFO_KEYS = ("sglang_parallel_info", "megatron_parallel_info") diff --git a/python/sglang/srt/debug_utils/comparator/aligner/unshard/planner.py b/python/sglang/srt/debug_utils/comparator/aligner/unsharder/planner.py similarity index 92% rename from python/sglang/srt/debug_utils/comparator/aligner/unshard/planner.py rename to python/sglang/srt/debug_utils/comparator/aligner/unsharder/planner.py index f8badb588..6d34cda3f 100644 --- a/python/sglang/srt/debug_utils/comparator/aligner/unshard/planner.py +++ b/python/sglang/srt/debug_utils/comparator/aligner/unsharder/planner.py @@ -1,12 +1,12 @@ from collections import defaultdict from typing import NamedTuple -from sglang.srt.debug_utils.comparator.aligner.unshard.types import ( +from sglang.srt.debug_utils.comparator.aligner.unsharder.types import ( AxisInfo, ConcatParams, PickParams, - UnshardParams, - UnshardPlan, + UnsharderParams, + UnsharderPlan, ) from sglang.srt.debug_utils.comparator.dims import DimSpec, ParallelAxis @@ -21,10 +21,10 @@ class _GroupResult(NamedTuple): projected_coords: _CoordsList -def compute_unshard_plan( +def compute_unsharder_plan( dim_specs: list[DimSpec], parallel_infos: list[dict[ParallelAxis, AxisInfo]], -) -> list[UnshardPlan]: +) -> list[UnsharderPlan]: if not parallel_infos: raise ValueError("parallel_infos must not be empty") @@ -51,20 +51,20 @@ def compute_unshard_plan( for info in parallel_infos ] - axis_and_params: list[tuple[ParallelAxis, UnshardParams]] = [ + axis_and_params: list[tuple[ParallelAxis, UnsharderParams]] = [ (axis, PickParams()) for axis in sorted(replicated_axes, key=lambda a: a.value) ] + [ (axis, _resolve_unshard_params(spec=spec, dim_index=dim_index)) for axis, (dim_index, spec) in sharded_axis_infos.items() ] - plans: list[UnshardPlan] = [] + plans: list[UnsharderPlan] = [] for axis, params in axis_and_params: result = _group_and_project( current_coords=current_coords, target_axis=axis, ) - plans.append(UnshardPlan(axis=axis, params=params, groups=result.groups)) + plans.append(UnsharderPlan(axis=axis, params=params, groups=result.groups)) current_coords = result.projected_coords return plans @@ -130,7 +130,7 @@ def _group_and_project( return _GroupResult(groups=groups, projected_coords=projected) -def _resolve_unshard_params(*, spec: DimSpec, dim_index: int) -> UnshardParams: +def _resolve_unshard_params(*, spec: DimSpec, dim_index: int) -> UnsharderParams: if spec.reduction is not None: raise NotImplementedError( f"Unshard for reduction={spec.reduction} not yet implemented (Phase 2)" diff --git a/python/sglang/srt/debug_utils/comparator/aligner/unshard/types.py b/python/sglang/srt/debug_utils/comparator/aligner/unsharder/types.py similarity index 62% rename from python/sglang/srt/debug_utils/comparator/aligner/unshard/types.py rename to python/sglang/srt/debug_utils/comparator/aligner/unsharder/types.py index a8b860f26..306bbb4f4 100644 --- a/python/sglang/srt/debug_utils/comparator/aligner/unshard/types.py +++ b/python/sglang/srt/debug_utils/comparator/aligner/unsharder/types.py @@ -2,7 +2,7 @@ from __future__ import annotations from typing import Annotated, Literal, Union -from pydantic import Field +from pydantic import Field, model_validator from sglang.srt.debug_utils.comparator.dims import ParallelAxis from sglang.srt.debug_utils.comparator.utils import _FrozenBase @@ -12,6 +12,16 @@ class AxisInfo(_FrozenBase): axis_rank: int axis_size: int + @model_validator(mode="after") + def _validate_bounds(self) -> AxisInfo: + if self.axis_size <= 0: + raise ValueError(f"axis_size must be > 0, got {self.axis_size}") + if not (0 <= self.axis_rank < self.axis_size): + raise ValueError( + f"axis_rank must be in [0, {self.axis_size}), got {self.axis_rank}" + ) + return self + class ConcatParams(_FrozenBase): op: Literal["concat"] = "concat" @@ -22,15 +32,15 @@ class PickParams(_FrozenBase): op: Literal["pick"] = "pick" -UnshardParams = Annotated[ +UnsharderParams = Annotated[ Union[ConcatParams, PickParams], Field(discriminator="op"), ] -class UnshardPlan(_FrozenBase): +class UnsharderPlan(_FrozenBase): axis: ParallelAxis - params: UnshardParams + params: UnsharderParams # groups[i] = indices in the input tensor list, which will be operated (e.g. concat) into i-th output tensor. # # Multistep example (CP=2, TP=2, 4 input tensors): diff --git a/python/sglang/srt/debug_utils/comparator/output_types.py b/python/sglang/srt/debug_utils/comparator/output_types.py index df2624fc2..501e893cb 100644 --- a/python/sglang/srt/debug_utils/comparator/output_types.py +++ b/python/sglang/srt/debug_utils/comparator/output_types.py @@ -3,10 +3,10 @@ from typing import Annotated, Any, Literal, Union from pydantic import Discriminator, Field, TypeAdapter, model_validator -from sglang.srt.debug_utils.comparator.tensor_comparison.formatter import ( +from sglang.srt.debug_utils.comparator.tensor_comparator.formatter import ( format_comparison, ) -from sglang.srt.debug_utils.comparator.tensor_comparison.types import ( +from sglang.srt.debug_utils.comparator.tensor_comparator.types import ( TensorComparisonInfo, ) from sglang.srt.debug_utils.comparator.utils import _StrictBase diff --git a/python/sglang/srt/debug_utils/comparator/pipeline.py b/python/sglang/srt/debug_utils/comparator/pipeline.py index 291d7300f..6612b7f20 100644 --- a/python/sglang/srt/debug_utils/comparator/pipeline.py +++ b/python/sglang/srt/debug_utils/comparator/pipeline.py @@ -3,31 +3,36 @@ from typing import Any, Optional, Union import torch -from sglang.srt.debug_utils.comparator.aligner.reorder import ( - ReorderPlan, - compute_reorder_plans, - execute_reorder_plan, +from sglang.srt.debug_utils.comparator.aligner.reorderer.executor import ( + execute_reorderer_plan, ) -from sglang.srt.debug_utils.comparator.aligner.unshard.executor import ( - execute_unshard_plan, +from sglang.srt.debug_utils.comparator.aligner.reorderer.planner import ( + compute_reorderer_plans, ) -from sglang.srt.debug_utils.comparator.aligner.unshard.parallel_info import ( +from sglang.srt.debug_utils.comparator.aligner.reorderer.types import ReordererPlan +from sglang.srt.debug_utils.comparator.aligner.unsharder.executor import ( + execute_unsharder_plan, +) +from sglang.srt.debug_utils.comparator.aligner.unsharder.parallel_info import ( normalize_parallel_info, ) -from sglang.srt.debug_utils.comparator.aligner.unshard.planner import ( - compute_unshard_plan, +from sglang.srt.debug_utils.comparator.aligner.unsharder.planner import ( + compute_unsharder_plan, ) -from sglang.srt.debug_utils.comparator.aligner.unshard.types import UnshardPlan +from sglang.srt.debug_utils.comparator.aligner.unsharder.types import UnsharderPlan from sglang.srt.debug_utils.comparator.dims import parse_dims from sglang.srt.debug_utils.comparator.output_types import ( AnyWarning, ComparisonRecord, SkipRecord, ) -from sglang.srt.debug_utils.comparator.tensor_comparison.compare import compare_tensors +from sglang.srt.debug_utils.comparator.tensor_comparator.comparator import ( + compare_tensor_pair, +) +from sglang.srt.debug_utils.comparator.warning_sink import warning_sink from sglang.srt.debug_utils.dump_loader import ValueWithMeta -Plan = Union[UnshardPlan, ReorderPlan] +Plan = Union[UnsharderPlan, ReordererPlan] def process_tensor_group( @@ -38,6 +43,28 @@ def process_tensor_group( baseline_path: Path, target_path: Path, diff_threshold: float, +) -> ComparisonRecord | SkipRecord: + with warning_sink.context() as collected_warnings: + return _process_tensor_group_raw( + name=name, + baseline_filenames=baseline_filenames, + target_filenames=target_filenames, + baseline_path=baseline_path, + target_path=target_path, + diff_threshold=diff_threshold, + collected_warnings=collected_warnings, + ) + + +def _process_tensor_group_raw( + *, + name: str, + baseline_filenames: list[str], + target_filenames: list[str], + baseline_path: Path, + target_path: Path, + diff_threshold: float, + collected_warnings: list[AnyWarning], ) -> ComparisonRecord | SkipRecord: b_tensors = _load_tensors(baseline_filenames, baseline_path) t_tensors = _load_tensors(target_filenames, target_path) @@ -51,22 +78,21 @@ def process_tensor_group( t_extracted = _extract_tensors(t_tensors) del b_tensors, t_tensors - b_tensor, b_warns = _execute_plans(b_extracted, b_plans) - t_tensor, t_warns = _execute_plans(t_extracted, t_plans) - all_warnings: list[AnyWarning] = b_warns + t_warns + b_tensor = _execute_plans(b_extracted, b_plans) + t_tensor = _execute_plans(t_extracted, t_plans) if b_tensor is None or t_tensor is None: reason = "baseline_load_failed" if b_tensor is None else "target_load_failed" - return SkipRecord(name=name, reason=reason, warnings=all_warnings) + return SkipRecord(name=name, reason=reason, warnings=collected_warnings) - info = compare_tensors( + info = compare_tensor_pair( x_baseline=b_tensor, x_target=t_tensor, name=name, diff_threshold=diff_threshold, ) - return ComparisonRecord(**info.model_dump(), warnings=all_warnings) + return ComparisonRecord(**info.model_dump(), warnings=collected_warnings) def _load_tensors(filenames: list[str], base_path: Path) -> list[ValueWithMeta]: @@ -96,13 +122,13 @@ def _compute_plans_for_group(metas: list[dict[str, Any]]) -> list[Plan]: dim_specs = parse_dims(dims_str) parallel_infos = [normalize_parallel_info(meta) for meta in metas] - unshard_plans = compute_unshard_plan( + unsharder_plans = compute_unsharder_plan( dim_specs=dim_specs, parallel_infos=parallel_infos ) - reorder_plans = compute_reorder_plans( + reorderer_plans = compute_reorderer_plans( dim_specs=dim_specs, parallel_infos=parallel_infos ) - return [*unshard_plans, *reorder_plans] + return [*unsharder_plans, *reorderer_plans] def _extract_tensors( @@ -114,32 +140,30 @@ def _extract_tensors( def _execute_plans( tensors: list[torch.Tensor], plans: list[Plan], -) -> tuple[Optional[torch.Tensor], list[AnyWarning]]: +) -> Optional[torch.Tensor]: if not tensors: - return None, [] + return None if not plans: if len(tensors) != 1: - return None, [] - return tensors[0], [] + return None + return tensors[0] - warnings: list[AnyWarning] = [] current = tensors for plan in plans: - current, new_warnings = _execute_plan(current, plan) - warnings.extend(new_warnings) + current = _execute_plan(current, plan) assert len(current) == 1 - return current[0], warnings + return current[0] def _execute_plan( tensors: list[torch.Tensor], plan: Plan, -) -> tuple[list[torch.Tensor], list[AnyWarning]]: - if isinstance(plan, UnshardPlan): - return execute_unshard_plan(plan, tensors) - elif isinstance(plan, ReorderPlan): - return execute_reorder_plan(plan, tensors), [] +) -> list[torch.Tensor]: + if isinstance(plan, UnsharderPlan): + return execute_unsharder_plan(plan, tensors) + elif isinstance(plan, ReordererPlan): + return execute_reorderer_plan(plan, tensors) else: raise NotImplementedError(f"Unknown {plan=}") diff --git a/python/sglang/srt/debug_utils/comparator/tensor_comparator/__init__.py b/python/sglang/srt/debug_utils/comparator/tensor_comparator/__init__.py new file mode 100644 index 000000000..b9974802d --- /dev/null +++ b/python/sglang/srt/debug_utils/comparator/tensor_comparator/__init__.py @@ -0,0 +1,3 @@ +from sglang.srt.debug_utils.comparator.tensor_comparator.comparator import ( + compare_tensor_pair, +) diff --git a/python/sglang/srt/debug_utils/comparator/tensor_comparison/compare.py b/python/sglang/srt/debug_utils/comparator/tensor_comparator/comparator.py similarity index 95% rename from python/sglang/srt/debug_utils/comparator/tensor_comparison/compare.py rename to python/sglang/srt/debug_utils/comparator/tensor_comparator/comparator.py index 43c056a5a..c26db9aa2 100644 --- a/python/sglang/srt/debug_utils/comparator/tensor_comparison/compare.py +++ b/python/sglang/srt/debug_utils/comparator/tensor_comparator/comparator.py @@ -2,13 +2,14 @@ from typing import Optional import torch -from sglang.srt.debug_utils.comparator.tensor_comparison.types import ( +from sglang.srt.debug_utils.comparator.tensor_comparator.types import ( DiffInfo, TensorComparisonInfo, TensorInfo, TensorStats, ) from sglang.srt.debug_utils.comparator.utils import ( + Pair, argmax_coord, calc_rel_diff, compute_smaller_dtype, @@ -20,7 +21,7 @@ QUANTILE_NUMEL_THRESHOLD = 10_000_000 SAMPLE_DIFF_THRESHOLD = 1e-3 -def compare_tensors( +def compare_tensor_pair( x_baseline: torch.Tensor, x_target: torch.Tensor, name: str = "", @@ -66,7 +67,7 @@ def compare_tensors( if baseline_original_dtype != target_original_dtype: downcast_dtype = compute_smaller_dtype( - baseline_original_dtype, target_original_dtype + Pair(x=baseline_original_dtype, y=target_original_dtype) ) if downcast_dtype is not None: diff_downcast = _compute_diff( diff --git a/python/sglang/srt/debug_utils/comparator/tensor_comparison/formatter.py b/python/sglang/srt/debug_utils/comparator/tensor_comparator/formatter.py similarity index 97% rename from python/sglang/srt/debug_utils/comparator/tensor_comparison/formatter.py rename to python/sglang/srt/debug_utils/comparator/tensor_comparator/formatter.py index 57fbbd8e2..c300f85b3 100644 --- a/python/sglang/srt/debug_utils/comparator/tensor_comparison/formatter.py +++ b/python/sglang/srt/debug_utils/comparator/tensor_comparator/formatter.py @@ -1,4 +1,4 @@ -from sglang.srt.debug_utils.comparator.tensor_comparison.types import ( +from sglang.srt.debug_utils.comparator.tensor_comparator.types import ( DiffInfo, TensorComparisonInfo, TensorStats, diff --git a/python/sglang/srt/debug_utils/comparator/tensor_comparison/printer.py b/python/sglang/srt/debug_utils/comparator/tensor_comparator/printer.py similarity index 97% rename from python/sglang/srt/debug_utils/comparator/tensor_comparison/printer.py rename to python/sglang/srt/debug_utils/comparator/tensor_comparator/printer.py index c9ea05dd1..716474085 100644 --- a/python/sglang/srt/debug_utils/comparator/tensor_comparison/printer.py +++ b/python/sglang/srt/debug_utils/comparator/tensor_comparator/printer.py @@ -1,4 +1,4 @@ -from sglang.srt.debug_utils.comparator.tensor_comparison.types import ( +from sglang.srt.debug_utils.comparator.tensor_comparator.types import ( DiffInfo, TensorComparisonInfo, TensorStats, diff --git a/python/sglang/srt/debug_utils/comparator/tensor_comparison/types.py b/python/sglang/srt/debug_utils/comparator/tensor_comparator/types.py similarity index 100% rename from python/sglang/srt/debug_utils/comparator/tensor_comparison/types.py rename to python/sglang/srt/debug_utils/comparator/tensor_comparator/types.py diff --git a/python/sglang/srt/debug_utils/comparator/tensor_comparison/__init__.py b/python/sglang/srt/debug_utils/comparator/tensor_comparison/__init__.py deleted file mode 100644 index 14718ef55..000000000 --- a/python/sglang/srt/debug_utils/comparator/tensor_comparison/__init__.py +++ /dev/null @@ -1 +0,0 @@ -from sglang.srt.debug_utils.comparator.tensor_comparison.compare import compare_tensors diff --git a/python/sglang/srt/debug_utils/comparator/utils.py b/python/sglang/srt/debug_utils/comparator/utils.py index da805e886..c899d1c95 100644 --- a/python/sglang/srt/debug_utils/comparator/utils.py +++ b/python/sglang/srt/debug_utils/comparator/utils.py @@ -40,13 +40,13 @@ def argmax_coord(x: torch.Tensor) -> Tuple[int, ...]: def compute_smaller_dtype( - dtype_a: torch.dtype, dtype_b: torch.dtype + dtypes: Pair[torch.dtype], ) -> Optional[torch.dtype]: info_dict = { (torch.float32, torch.bfloat16): torch.bfloat16, # ... add more ... } - return info_dict.get((dtype_a, dtype_b)) or info_dict.get((dtype_b, dtype_a)) + return info_dict.get((dtypes.x, dtypes.y)) or info_dict.get((dtypes.y, dtypes.x)) def try_unify_shape(x: torch.Tensor, target_shape: torch.Size) -> torch.Tensor: diff --git a/test/registered/debug_utils/comparator/aligner/conftest.py b/test/registered/debug_utils/comparator/aligner/conftest.py new file mode 100644 index 000000000..e69de29bb diff --git a/test/registered/debug_utils/comparator/aligner/reorderer/conftest.py b/test/registered/debug_utils/comparator/aligner/reorderer/conftest.py new file mode 100644 index 000000000..e69de29bb diff --git a/test/registered/debug_utils/comparator/aligner/reorderer/test_executor.py b/test/registered/debug_utils/comparator/aligner/reorderer/test_executor.py new file mode 100644 index 000000000..1029e5475 --- /dev/null +++ b/test/registered/debug_utils/comparator/aligner/reorderer/test_executor.py @@ -0,0 +1,50 @@ +import sys + +import pytest +import torch + +from sglang.srt.debug_utils.comparator.aligner.reorderer.executor import ( + _reorder_zigzag_to_natural, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=10, suite="default", nightly=True) + + +class TestZigzagToNatural: + def test_zigzag_to_natural_cp2(self) -> None: + """cp_size=2: zigzag order [0,3,1,2] -> natural [0,1,2,3].""" + natural = torch.arange(24).reshape(4, 6) + chunks = list(natural.chunk(4, dim=0)) + + zigzag_order: list[int] = [0, 3, 1, 2] + zigzagged = torch.cat([chunks[i] for i in zigzag_order], dim=0) + + result = _reorder_zigzag_to_natural(zigzagged, dim=0, cp_size=2) + assert torch.equal(result, natural) + + def test_zigzag_to_natural_cp3(self) -> None: + """cp_size=3: zigzag 162534 -> natural 123456 (1-indexed).""" + natural = torch.arange(60).reshape(6, 10) + chunks = list(natural.chunk(6, dim=0)) + + zigzag_order: list[int] = [0, 5, 1, 4, 2, 3] + zigzagged = torch.cat([chunks[i] for i in zigzag_order], dim=0) + + result = _reorder_zigzag_to_natural(zigzagged, dim=0, cp_size=3) + assert torch.equal(result, natural) + + def test_zigzag_to_natural_arbitrary_dim(self) -> None: + """Reorder along dim=1 instead of dim=0.""" + natural = torch.arange(48).reshape(3, 4, 4) + chunks = list(natural.chunk(4, dim=1)) + + zigzag_order: list[int] = [0, 3, 1, 2] + zigzagged = torch.cat([chunks[i] for i in zigzag_order], dim=1) + + result = _reorder_zigzag_to_natural(zigzagged, dim=1, cp_size=2) + assert torch.equal(result, natural) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__])) diff --git a/test/registered/debug_utils/comparator/aligner/test_reorder.py b/test/registered/debug_utils/comparator/aligner/reorderer/test_planner.py similarity index 55% rename from test/registered/debug_utils/comparator/aligner/test_reorder.py rename to test/registered/debug_utils/comparator/aligner/reorderer/test_planner.py index a78c5e3d2..0bf15cdb3 100644 --- a/test/registered/debug_utils/comparator/aligner/test_reorder.py +++ b/test/registered/debug_utils/comparator/aligner/reorderer/test_planner.py @@ -3,63 +3,30 @@ import sys import pytest import torch -from sglang.srt.debug_utils.comparator.aligner.reorder import ( - ReorderPlan, - _reorder_zigzag_to_natural, - compute_reorder_plans, - execute_reorder_plan, +from sglang.srt.debug_utils.comparator.aligner.reorderer.executor import ( + execute_reorderer_plan, ) -from sglang.srt.debug_utils.comparator.aligner.unshard.executor import ( - execute_unshard_plan, +from sglang.srt.debug_utils.comparator.aligner.reorderer.planner import ( + compute_reorderer_plans, ) -from sglang.srt.debug_utils.comparator.aligner.unshard.planner import ( - compute_unshard_plan, +from sglang.srt.debug_utils.comparator.aligner.reorderer.types import ReordererPlan +from sglang.srt.debug_utils.comparator.aligner.unsharder.executor import ( + execute_unsharder_plan, ) -from sglang.srt.debug_utils.comparator.aligner.unshard.types import AxisInfo +from sglang.srt.debug_utils.comparator.aligner.unsharder.planner import ( + compute_unsharder_plan, +) +from sglang.srt.debug_utils.comparator.aligner.unsharder.types import AxisInfo from sglang.srt.debug_utils.comparator.dims import ParallelAxis, parse_dims +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=10, suite="default", nightly=True) -class TestZigzagToNatural: - def test_zigzag_to_natural_cp2(self) -> None: - """cp_size=2: zigzag order [0,3,1,2] -> natural [0,1,2,3].""" - natural = torch.arange(24).reshape(4, 6) - chunks = list(natural.chunk(4, dim=0)) - - zigzag_order: list[int] = [0, 3, 1, 2] - zigzagged = torch.cat([chunks[i] for i in zigzag_order], dim=0) - - result = _reorder_zigzag_to_natural(zigzagged, dim=0, cp_size=2) - assert torch.equal(result, natural) - - def test_zigzag_to_natural_cp3(self) -> None: - """cp_size=3: zigzag 162534 -> natural 123456 (1-indexed).""" - natural = torch.arange(60).reshape(6, 10) - chunks = list(natural.chunk(6, dim=0)) - - zigzag_order: list[int] = [0, 5, 1, 4, 2, 3] - zigzagged = torch.cat([chunks[i] for i in zigzag_order], dim=0) - - result = _reorder_zigzag_to_natural(zigzagged, dim=0, cp_size=3) - assert torch.equal(result, natural) - - def test_zigzag_to_natural_arbitrary_dim(self) -> None: - """Reorder along dim=1 instead of dim=0.""" - natural = torch.arange(48).reshape(3, 4, 4) - chunks = list(natural.chunk(4, dim=1)) - - zigzag_order: list[int] = [0, 3, 1, 2] - zigzagged = torch.cat([chunks[i] for i in zigzag_order], dim=1) - - result = _reorder_zigzag_to_natural(zigzagged, dim=1, cp_size=2) - assert torch.equal(result, natural) - - -class TestComputeReorderPlans: - def test_compute_reorder_plans_zigzag(self) -> None: - """s(cp,zigzag) produces a ReorderPlan.""" +class TestComputeReordererPlans: + def test_compute_reorderer_plans_zigzag(self) -> None: + """s(cp,zigzag) produces a ReordererPlan.""" dim_specs = parse_dims("b s(cp,zigzag) h(tp)") parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ { @@ -67,7 +34,7 @@ class TestComputeReorderPlans: ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), }, ] - plans = compute_reorder_plans( + plans = compute_reorderer_plans( dim_specs=dim_specs, parallel_infos=parallel_infos ) @@ -76,7 +43,7 @@ class TestComputeReorderPlans: assert plans[0].params.dim == 1 assert plans[0].params.cp_size == 2 - def test_compute_reorder_plans_non_seq_dim_raises(self) -> None: + def test_compute_reorderer_plans_non_seq_dim_raises(self) -> None: """Zigzag on non-sequence dim (e.g. t(cp,zigzag)) raises ValueError.""" dim_specs = parse_dims("t(cp,zigzag) h(tp)") parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [ @@ -86,9 +53,9 @@ class TestComputeReorderPlans: }, ] with pytest.raises(ValueError, match="only supported on sequence dims"): - compute_reorder_plans(dim_specs=dim_specs, parallel_infos=parallel_infos) + compute_reorderer_plans(dim_specs=dim_specs, parallel_infos=parallel_infos) - def test_compute_reorder_plans_natural(self) -> None: + def test_compute_reorderer_plans_natural(self) -> None: """s(cp) and s(cp,natural) produce no reorder plans.""" for dims_str in ["b s(cp) h(tp)", "b s(cp,natural) h(tp)"]: dim_specs = parse_dims(dims_str) @@ -98,7 +65,7 @@ class TestComputeReorderPlans: ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), }, ] - plans = compute_reorder_plans( + plans = compute_reorderer_plans( dim_specs=dim_specs, parallel_infos=parallel_infos ) assert plans == [] @@ -132,23 +99,24 @@ class TestCpZigzagTpE2E: dim_specs = parse_dims("b s(cp,zigzag) h(tp)") - unshard_plans = compute_unshard_plan( + unsharder_plans = compute_unsharder_plan( dim_specs=dim_specs, parallel_infos=parallel_infos ) - reorder_plans = compute_reorder_plans( + reorderer_plans = compute_reorderer_plans( dim_specs=dim_specs, parallel_infos=parallel_infos ) - all_plans = [*unshard_plans, *reorder_plans] + all_plans = [*unsharder_plans, *reorderer_plans] - assert len(unshard_plans) == 2 - assert len(reorder_plans) == 1 + assert len(unsharder_plans) == 2 + assert len(reorderer_plans) == 1 current: list[torch.Tensor] = tensors - for plan in all_plans: - if isinstance(plan, ReorderPlan): - current = execute_reorder_plan(plan, current) - else: - current, _ = execute_unshard_plan(plan, current) + with warning_sink.context(): + for plan in all_plans: + if isinstance(plan, ReordererPlan): + current = execute_reorderer_plan(plan, current) + else: + current = execute_unsharder_plan(plan, current) assert len(current) == 1 assert torch.allclose(current[0], full_tensor) diff --git a/test/registered/debug_utils/comparator/aligner/unsharder/conftest.py b/test/registered/debug_utils/comparator/aligner/unsharder/conftest.py new file mode 100644 index 000000000..e69de29bb diff --git a/test/registered/debug_utils/comparator/aligner/unshard/test_execute.py b/test/registered/debug_utils/comparator/aligner/unsharder/test_executor.py similarity index 78% rename from test/registered/debug_utils/comparator/aligner/unshard/test_execute.py rename to test/registered/debug_utils/comparator/aligner/unsharder/test_executor.py index e862ddbad..e52440cef 100644 --- a/test/registered/debug_utils/comparator/aligner/unshard/test_execute.py +++ b/test/registered/debug_utils/comparator/aligner/unsharder/test_executor.py @@ -3,25 +3,26 @@ import sys import pytest import torch -from sglang.srt.debug_utils.comparator.aligner.unshard.executor import ( +from sglang.srt.debug_utils.comparator.aligner.unsharder.executor import ( _apply_unshard, _verify_replicated_group, - execute_unshard_plan, + execute_unsharder_plan, ) -from sglang.srt.debug_utils.comparator.aligner.unshard.planner import ( - compute_unshard_plan, +from sglang.srt.debug_utils.comparator.aligner.unsharder.planner import ( + compute_unsharder_plan, ) -from sglang.srt.debug_utils.comparator.aligner.unshard.types import ( +from sglang.srt.debug_utils.comparator.aligner.unsharder.types import ( AxisInfo, PickParams, ) from sglang.srt.debug_utils.comparator.dims import ParallelAxis, parse_dims +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=10, suite="default", nightly=True) -class TestExecuteUnshardPlan: +class TestExecuteUnsharderPlan: def test_tp4_concat(self) -> None: full_tensor = torch.randn(2, 8, 16) shards = list(full_tensor.chunk(4, dim=1)) @@ -30,10 +31,11 @@ class TestExecuteUnshardPlan: parallel_infos = [ {ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=4)} for i in range(4) ] - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 1 - result, warnings = execute_unshard_plan(plans[0], shards) + with warning_sink.context() as warnings: + result = execute_unsharder_plan(plans[0], shards) assert len(result) == 1 assert torch.allclose(result[0], full_tensor) assert warnings == [] @@ -49,7 +51,7 @@ class TestExecuteUnshardPlan: {ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=4)}, ] dim_specs = parse_dims("h(tp) d") - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 1 tensors_ordered_by_world_rank = [ @@ -59,7 +61,8 @@ class TestExecuteUnshardPlan: shards[1], # world_rank=3, axis_rank=1 ] - result, warnings = execute_unshard_plan(plans[0], tensors_ordered_by_world_rank) + with warning_sink.context() as warnings: + result = execute_unsharder_plan(plans[0], tensors_ordered_by_world_rank) assert len(result) == 1 assert torch.allclose(result[0], full_tensor) assert warnings == [] @@ -82,7 +85,7 @@ class TestExecuteUnshardPlan: } ) - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 2 tensors: list[torch.Tensor] = [] @@ -91,10 +94,12 @@ class TestExecuteUnshardPlan: for tp_rank in range(4): tensors.append(source[tp_rank]) - intermediate, _ = execute_unshard_plan(plans[0], tensors) + with warning_sink.context() as _warnings: + intermediate = execute_unsharder_plan(plans[0], tensors) assert len(intermediate) == 4 - final, _ = execute_unshard_plan(plans[1], intermediate) + with warning_sink.context() as _warnings: + final = execute_unsharder_plan(plans[1], intermediate) assert len(final) == 1 def test_cp_tp_concat(self) -> None: @@ -117,12 +122,13 @@ class TestExecuteUnshardPlan: ) dim_specs = parse_dims("b s(cp) h(tp)") - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 2 current = tensors for plan in plans: - current, _ = execute_unshard_plan(plan, current) + with warning_sink.context() as _warnings: + current = execute_unsharder_plan(plan, current) assert len(current) == 1 assert torch.allclose(current[0], full_tensor) @@ -158,12 +164,13 @@ class TestExecuteUnshardPlan: ) dim_specs = parse_dims("b s(cp) h(tp)") - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 2 current = tensors for plan in plans: - current, _ = execute_unshard_plan(plan, current) + with warning_sink.context() as _warnings: + current = execute_unsharder_plan(plan, current) assert len(current) == 1 assert torch.allclose(current[0], full_tensor) @@ -211,12 +218,13 @@ class TestExecuteUnshardPlan: ) dim_specs = parse_dims("b e(ep) s(cp) h(tp)") - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 3 current = tensors for plan in plans: - current, _ = execute_unshard_plan(plan, current) + with warning_sink.context() as _warnings: + current = execute_unsharder_plan(plan, current) assert len(current) == 1 assert torch.allclose(current[0], full_tensor) @@ -259,12 +267,13 @@ class TestExecuteUnshardPlan: ) dim_specs = parse_dims("b e(ep) s(cp) h(tp)") - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 3 current = tensors for plan in plans: - current, _ = execute_unshard_plan(plan, current) + with warning_sink.context() as _warnings: + current = execute_unsharder_plan(plan, current) assert len(current) == 1 assert torch.allclose(current[0], full_tensor) @@ -280,11 +289,12 @@ class TestPickOperation: {ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2)}, ] - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 1 assert isinstance(plans[0].params, PickParams) - result, warnings = execute_unshard_plan(plans[0], [tensor, tensor.clone()]) + with warning_sink.context() as warnings: + result = execute_unsharder_plan(plans[0], [tensor, tensor.clone()]) assert len(result) == 1 assert torch.allclose(result[0], tensor) assert warnings == [] @@ -311,7 +321,7 @@ class TestPickOperation: }, ] - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) pick_plans = [p for p in plans if isinstance(p.params, PickParams)] assert len(pick_plans) == 1 assert pick_plans[0].axis == ParallelAxis.CP @@ -319,7 +329,8 @@ class TestPickOperation: tensor = torch.randn(4) tensors = [tensor.clone() for _ in range(4)] - result, warnings = execute_unshard_plan(pick_plans[0], tensors) + with warning_sink.context() as warnings: + result = execute_unsharder_plan(pick_plans[0], tensors) assert len(result) == 2 assert warnings == [] @@ -342,18 +353,19 @@ class TestPickOperation: ) dim_specs = parse_dims("b s(cp) d") - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 2 current = tensors for plan in plans: - current, _ = execute_unshard_plan(plan, current) + with warning_sink.context() as _warnings: + current = execute_unsharder_plan(plan, current) assert len(current) == 1 assert torch.allclose(current[0], full_tensor) def test_fully_replicated_e2e(self) -> None: - """CP2 TP2, dims='b h d': fully replicated → 2 pick steps → 1 tensor.""" + """CP2 TP2, dims='b h d': fully replicated -> 2 pick steps -> 1 tensor.""" torch.manual_seed(42) full_tensor = torch.randn(4, 8, 16) @@ -370,13 +382,14 @@ class TestPickOperation: ) dim_specs = parse_dims("b h d") - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 2 assert all(isinstance(p.params, PickParams) for p in plans) current = tensors for plan in plans: - current, _ = execute_unshard_plan(plan, current) + with warning_sink.context() as _warnings: + current = execute_unsharder_plan(plan, current) assert len(current) == 1 assert torch.allclose(current[0], full_tensor) @@ -388,11 +401,12 @@ class TestVerifyReplicatedGroup: tensor_a = torch.ones(4) tensor_b = torch.ones(4) + 0.1 - warnings = _verify_replicated_group( - [tensor_a, tensor_b], - axis=ParallelAxis.TP, - group_index=0, - ) + with warning_sink.context() as warnings: + _verify_replicated_group( + [tensor_a, tensor_b], + axis=ParallelAxis.TP, + group_index=0, + ) assert len(warnings) == 1 assert warnings[0].axis == "tp" assert warnings[0].group_index == 0 @@ -404,11 +418,12 @@ class TestVerifyReplicatedGroup: """_verify_replicated_group produces no warning for identical replicas.""" tensor = torch.randn(4, 8) - warnings = _verify_replicated_group( - [tensor, tensor.clone()], - axis=ParallelAxis.TP, - group_index=0, - ) + with warning_sink.context() as warnings: + _verify_replicated_group( + [tensor, tensor.clone()], + axis=ParallelAxis.TP, + group_index=0, + ) assert warnings == [] def test_multiple_mismatches(self) -> None: @@ -417,55 +432,59 @@ class TestVerifyReplicatedGroup: other_a = torch.ones(4) other_b = torch.ones(4) * 2 - warnings = _verify_replicated_group( - [baseline, other_a, other_b], - axis=ParallelAxis.CP, - group_index=1, - ) + with warning_sink.context() as warnings: + _verify_replicated_group( + [baseline, other_a, other_b], + axis=ParallelAxis.CP, + group_index=1, + ) assert len(warnings) == 2 assert warnings[0].differing_index == 1 assert warnings[1].differing_index == 2 assert warnings[1].max_abs_diff == pytest.approx(2.0, abs=1e-5) def test_execute_returns_warnings(self) -> None: - """execute_unshard_plan returns warnings for replicated mismatch.""" + """execute_unsharder_plan emits warnings for replicated mismatch.""" dim_specs = parse_dims("h d") parallel_infos = [ {ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2)}, {ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2)}, ] - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) tensor_a = torch.zeros(4) tensor_b = torch.ones(4) - result, warnings = execute_unshard_plan(plans[0], [tensor_a, tensor_b]) + with warning_sink.context() as warnings: + result = execute_unsharder_plan(plans[0], [tensor_a, tensor_b]) assert len(result) == 1 assert len(warnings) == 1 assert torch.allclose(result[0], tensor_a) def test_atol_boundary_within(self) -> None: - """Difference exactly at atol (1e-6) → torch.allclose passes → no warning.""" + """Difference exactly at atol (1e-6) -> torch.allclose passes -> no warning.""" baseline = torch.zeros(4) other = torch.full((4,), 1e-6) - warnings = _verify_replicated_group( - [baseline, other], - axis=ParallelAxis.TP, - group_index=0, - ) + with warning_sink.context() as warnings: + _verify_replicated_group( + [baseline, other], + axis=ParallelAxis.TP, + group_index=0, + ) assert warnings == [] def test_atol_boundary_exceeded(self) -> None: - """Difference just above atol (1e-6 + 1e-9) → torch.allclose fails → warning.""" + """Difference just above atol (1e-6 + 1e-9) -> torch.allclose fails -> warning.""" baseline = torch.zeros(4) other = torch.full((4,), 1e-6 + 1e-9) - warnings = _verify_replicated_group( - [baseline, other], - axis=ParallelAxis.TP, - group_index=0, - ) + with warning_sink.context() as warnings: + _verify_replicated_group( + [baseline, other], + axis=ParallelAxis.TP, + group_index=0, + ) assert len(warnings) == 1 assert warnings[0].differing_index == 1 diff --git a/test/registered/debug_utils/comparator/aligner/unshard/test_parallel_info.py b/test/registered/debug_utils/comparator/aligner/unsharder/test_parallel_info.py similarity index 92% rename from test/registered/debug_utils/comparator/aligner/unshard/test_parallel_info.py rename to test/registered/debug_utils/comparator/aligner/unsharder/test_parallel_info.py index 4b8181eb5..85060c5b0 100644 --- a/test/registered/debug_utils/comparator/aligner/unshard/test_parallel_info.py +++ b/test/registered/debug_utils/comparator/aligner/unsharder/test_parallel_info.py @@ -2,10 +2,10 @@ import sys import pytest -from sglang.srt.debug_utils.comparator.aligner.unshard.parallel_info import ( +from sglang.srt.debug_utils.comparator.aligner.unsharder.parallel_info import ( normalize_parallel_info, ) -from sglang.srt.debug_utils.comparator.aligner.unshard.types import AxisInfo +from sglang.srt.debug_utils.comparator.aligner.unsharder.types import AxisInfo from sglang.srt.debug_utils.comparator.dims import ParallelAxis from sglang.test.ci.ci_register import register_cpu_ci diff --git a/test/registered/debug_utils/comparator/aligner/unshard/test_plan.py b/test/registered/debug_utils/comparator/aligner/unsharder/test_planner.py similarity index 90% rename from test/registered/debug_utils/comparator/aligner/unshard/test_plan.py rename to test/registered/debug_utils/comparator/aligner/unsharder/test_planner.py index 0f5a6577f..d7c87e977 100644 --- a/test/registered/debug_utils/comparator/aligner/unshard/test_plan.py +++ b/test/registered/debug_utils/comparator/aligner/unsharder/test_planner.py @@ -2,10 +2,10 @@ import sys import pytest -from sglang.srt.debug_utils.comparator.aligner.unshard.planner import ( - compute_unshard_plan, +from sglang.srt.debug_utils.comparator.aligner.unsharder.planner import ( + compute_unsharder_plan, ) -from sglang.srt.debug_utils.comparator.aligner.unshard.types import ( +from sglang.srt.debug_utils.comparator.aligner.unsharder.types import ( AxisInfo, ConcatParams, PickParams, @@ -16,13 +16,13 @@ from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=10, suite="default", nightly=True) -class TestComputeUnshardPlan: +class TestComputeUnsharderPlan: def test_tp4_plan(self) -> None: dim_specs = parse_dims("b s h(tp) d") parallel_infos = [ {ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=4)} for i in range(4) ] - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 1 assert plans[0].axis == ParallelAxis.TP @@ -36,18 +36,18 @@ class TestComputeUnshardPlan: {ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2)}, ] with pytest.raises(ValueError, match="Inconsistent axis_size"): - compute_unshard_plan(dim_specs, parallel_infos) + compute_unsharder_plan(dim_specs, parallel_infos) def test_missing_axis_in_parallel_info_raises(self) -> None: dim_specs = parse_dims("h(tp)") parallel_infos = [{ParallelAxis.CP: AxisInfo(axis_rank=0, axis_size=2)}] with pytest.raises(ValueError, match="missing parallel_info"): - compute_unshard_plan(dim_specs, parallel_infos) + compute_unsharder_plan(dim_specs, parallel_infos) def test_empty_parallel_infos_raises(self) -> None: dim_specs = parse_dims("h(tp)") with pytest.raises(ValueError, match="must not be empty"): - compute_unshard_plan(dim_specs, []) + compute_unsharder_plan(dim_specs, []) def test_scrambled_world_ranks(self) -> None: """world_rank order != axis_rank order.""" @@ -58,14 +58,14 @@ class TestComputeUnshardPlan: {ParallelAxis.TP: AxisInfo(axis_rank=3, axis_size=4)}, {ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=4)}, ] - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 1 assert plans[0].groups == [[1, 3, 0, 2]] def test_no_sharded_axes_returns_empty(self) -> None: dim_specs = parse_dims("b s d") parallel_infos = [{}] - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert plans == [] def test_multi_axis_plan(self) -> None: @@ -89,7 +89,7 @@ class TestComputeUnshardPlan: ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), }, ] - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 2 assert plans[0].axis == ParallelAxis.CP @@ -108,7 +108,7 @@ class TestComputeUnshardPlan: } ) - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 2 @@ -144,7 +144,7 @@ class TestComputeUnshardPlan: ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=2), }, ] - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 2 @@ -168,7 +168,7 @@ class TestComputeUnshardPlan: {ParallelAxis.TP: AxisInfo(axis_rank=3, axis_size=4)}, ] with pytest.raises(ValueError, match="axis_rank coverage.*incomplete"): - compute_unshard_plan(dim_specs, parallel_infos) + compute_unsharder_plan(dim_specs, parallel_infos) def test_reduction_not_implemented_raises(self) -> None: dim_specs = parse_dims("h(tp,partial)") @@ -176,14 +176,14 @@ class TestComputeUnshardPlan: {ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2) ] with pytest.raises(NotImplementedError, match="reduction"): - compute_unshard_plan(dim_specs, parallel_infos) + compute_unsharder_plan(dim_specs, parallel_infos) def test_ordering_zigzag_accepted(self) -> None: dim_specs = parse_dims("s(cp,zigzag)") parallel_infos = [ {ParallelAxis.CP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2) ] - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 1 assert plans[0].axis == ParallelAxis.CP @@ -192,7 +192,7 @@ class TestComputeUnshardPlan: parallel_infos = [ {ParallelAxis.CP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2) ] - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 1 assert plans[0].axis == ParallelAxis.CP @@ -211,7 +211,7 @@ class TestComputeUnshardPlan: } ) - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 3 assert plans[0].axis == ParallelAxis.EP @@ -246,7 +246,7 @@ class TestComputeUnshardPlan: }, ] with pytest.raises(ValueError, match="missing parallel_info"): - compute_unshard_plan(dim_specs, parallel_infos) + compute_unsharder_plan(dim_specs, parallel_infos) class TestReplicatedAxes: @@ -271,7 +271,7 @@ class TestReplicatedAxes: ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), }, ] - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 2 assert plans[0].axis == ParallelAxis.TP @@ -305,7 +305,7 @@ class TestReplicatedAxes: ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), }, ] - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 2 assert all(isinstance(p.params, PickParams) for p in plans) @@ -327,7 +327,7 @@ class TestReplicatedAxes: } ) - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 3 pick_plans = [p for p in plans if isinstance(p.params, PickParams)] @@ -360,7 +360,7 @@ class TestReplicatedAxes: ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=2), }, ] - plans = compute_unshard_plan(dim_specs, parallel_infos) + plans = compute_unsharder_plan(dim_specs, parallel_infos) assert len(plans) == 2 assert plans[0].axis == ParallelAxis.CP @@ -382,7 +382,7 @@ class TestReplicatedAxes: }, ] with pytest.raises(ValueError, match="Inconsistent axis_size"): - compute_unshard_plan(dim_specs, parallel_infos) + compute_unsharder_plan(dim_specs, parallel_infos) def test_replicated_axis_missing_from_rank_raises(self) -> None: """A rank missing a replicated axis that other ranks have raises ValueError.""" @@ -398,7 +398,7 @@ class TestReplicatedAxes: }, ] with pytest.raises(ValueError, match="missing parallel_info"): - compute_unshard_plan(dim_specs, parallel_infos) + compute_unsharder_plan(dim_specs, parallel_infos) if __name__ == "__main__": diff --git a/test/registered/debug_utils/comparator/conftest.py b/test/registered/debug_utils/comparator/conftest.py new file mode 100644 index 000000000..e69de29bb diff --git a/test/registered/debug_utils/comparator/tensor_comparator/conftest.py b/test/registered/debug_utils/comparator/tensor_comparator/conftest.py new file mode 100644 index 000000000..e69de29bb diff --git a/test/registered/debug_utils/comparator/tensor_comparison/test_compare.py b/test/registered/debug_utils/comparator/tensor_comparator/test_comparator.py similarity index 88% rename from test/registered/debug_utils/comparator/tensor_comparison/test_compare.py rename to test/registered/debug_utils/comparator/tensor_comparator/test_comparator.py index e42bd20f2..716ce48cd 100644 --- a/test/registered/debug_utils/comparator/tensor_comparison/test_compare.py +++ b/test/registered/debug_utils/comparator/tensor_comparator/test_comparator.py @@ -3,12 +3,12 @@ import sys import pytest import torch -from sglang.srt.debug_utils.comparator.tensor_comparison.compare import ( +from sglang.srt.debug_utils.comparator.tensor_comparator.comparator import ( QUANTILE_NUMEL_THRESHOLD, SAMPLE_DIFF_THRESHOLD, _compute_diff, _compute_tensor_stats, - compare_tensors, + compare_tensor_pair, ) from sglang.test.ci.ci_register import register_cpu_ci @@ -83,7 +83,7 @@ class TestCompareTensors: x = torch.randn(5, 5) y = x + torch.randn(5, 5) * 0.001 - info = compare_tensors(x_baseline=x, x_target=y, name="test") + info = compare_tensor_pair(x_baseline=x, x_target=y, name="test") assert info.name == "test" assert info.baseline.shape == [5, 5] @@ -96,7 +96,7 @@ class TestCompareTensors: x = torch.randn(3, 4) y = torch.randn(5, 6) - info = compare_tensors(x_baseline=x, x_target=y, name="mismatch") + info = compare_tensor_pair(x_baseline=x, x_target=y, name="mismatch") assert info.shape_mismatch is True assert info.diff is None @@ -105,7 +105,7 @@ class TestCompareTensors: x = torch.randn(5, 5, dtype=torch.float32) y = torch.randn(5, 5, dtype=torch.bfloat16) - info = compare_tensors(x_baseline=x, x_target=y, name="dtype_test") + info = compare_tensor_pair(x_baseline=x, x_target=y, name="dtype_test") assert info.shape_mismatch is False assert info.diff is not None @@ -118,7 +118,7 @@ class TestCompareTensors: x = core.unsqueeze(0).unsqueeze(0) # [1, 1, 4, 8] y = core.clone() # [4, 8] - info = compare_tensors(x_baseline=x, x_target=y, name="unify") + info = compare_tensor_pair(x_baseline=x, x_target=y, name="unify") assert info.baseline.shape == [1, 1, 4, 8] assert info.unified_shape == [4, 8] @@ -130,7 +130,7 @@ class TestCompareTensors: x = torch.zeros(5, 5) y = torch.ones(5, 5) - info = compare_tensors(x_baseline=x, x_target=y, name="big_diff") + info = compare_tensor_pair(x_baseline=x, x_target=y, name="big_diff") assert info.diff is not None assert info.diff.max_abs_diff > SAMPLE_DIFF_THRESHOLD @@ -141,7 +141,7 @@ class TestCompareTensors: x = torch.ones(5, 5) y = x + 1e-5 - info = compare_tensors(x_baseline=x, x_target=y, name="tiny_diff") + info = compare_tensor_pair(x_baseline=x, x_target=y, name="tiny_diff") assert info.diff is not None assert info.diff.max_abs_diff < SAMPLE_DIFF_THRESHOLD diff --git a/test/registered/debug_utils/comparator/tensor_comparison/test_formatter.py b/test/registered/debug_utils/comparator/tensor_comparator/test_formatter.py similarity index 98% rename from test/registered/debug_utils/comparator/tensor_comparison/test_formatter.py rename to test/registered/debug_utils/comparator/tensor_comparator/test_formatter.py index 33c44920c..ac3d6ab06 100644 --- a/test/registered/debug_utils/comparator/tensor_comparison/test_formatter.py +++ b/test/registered/debug_utils/comparator/tensor_comparator/test_formatter.py @@ -2,10 +2,10 @@ import sys import pytest -from sglang.srt.debug_utils.comparator.tensor_comparison.formatter import ( +from sglang.srt.debug_utils.comparator.tensor_comparator.formatter import ( format_comparison, ) -from sglang.srt.debug_utils.comparator.tensor_comparison.types import ( +from sglang.srt.debug_utils.comparator.tensor_comparator.types import ( DiffInfo, TensorComparisonInfo, TensorInfo, diff --git a/test/registered/debug_utils/comparator/tensor_comparison/test_printer.py b/test/registered/debug_utils/comparator/tensor_comparator/test_printer.py similarity index 98% rename from test/registered/debug_utils/comparator/tensor_comparison/test_printer.py rename to test/registered/debug_utils/comparator/tensor_comparator/test_printer.py index 4e82f2670..f03a83462 100644 --- a/test/registered/debug_utils/comparator/tensor_comparison/test_printer.py +++ b/test/registered/debug_utils/comparator/tensor_comparator/test_printer.py @@ -2,10 +2,10 @@ import sys import pytest -from sglang.srt.debug_utils.comparator.tensor_comparison.printer import ( +from sglang.srt.debug_utils.comparator.tensor_comparator.printer import ( print_comparison, ) -from sglang.srt.debug_utils.comparator.tensor_comparison.types import ( +from sglang.srt.debug_utils.comparator.tensor_comparator.types import ( DiffInfo, TensorComparisonInfo, TensorInfo, diff --git a/test/registered/debug_utils/comparator/tensor_comparison/test_types.py b/test/registered/debug_utils/comparator/tensor_comparator/test_types.py similarity index 98% rename from test/registered/debug_utils/comparator/tensor_comparison/test_types.py rename to test/registered/debug_utils/comparator/tensor_comparator/test_types.py index 393501427..db13f7531 100644 --- a/test/registered/debug_utils/comparator/tensor_comparison/test_types.py +++ b/test/registered/debug_utils/comparator/tensor_comparator/test_types.py @@ -11,7 +11,7 @@ from sglang.srt.debug_utils.comparator.output_types import ( SummaryRecord, parse_record_json, ) -from sglang.srt.debug_utils.comparator.tensor_comparison.types import ( +from sglang.srt.debug_utils.comparator.tensor_comparator.types import ( DiffInfo, TensorInfo, TensorStats, @@ -133,7 +133,7 @@ def _make_warning(**overrides) -> ReplicatedMismatchWarning: return ReplicatedMismatchWarning(**defaults) -class TestAlignWarnings: +class TestWarnings: def test_comparison_record_failed_when_diff_passed_but_warnings(self): """ComparisonRecord with diff.passed=True but warnings → category=='failed'.""" record = ComparisonRecord( diff --git a/test/registered/debug_utils/comparator/test_model_validation.py b/test/registered/debug_utils/comparator/test_model_validation.py index 040f588c3..5fc722188 100644 --- a/test/registered/debug_utils/comparator/test_model_validation.py +++ b/test/registered/debug_utils/comparator/test_model_validation.py @@ -3,13 +3,14 @@ import sys import pytest from pydantic import ValidationError +from sglang.srt.debug_utils.comparator.aligner.unsharder.types import AxisInfo from sglang.srt.debug_utils.comparator.output_types import ( ComparisonRecord, GeneralWarning, SkipRecord, SummaryRecord, ) -from sglang.srt.debug_utils.comparator.tensor_comparison.types import ( +from sglang.srt.debug_utils.comparator.tensor_comparator.types import ( DiffInfo, TensorInfo, TensorStats, @@ -32,6 +33,32 @@ class TestCheckEqualLengths: _check_equal_lengths(a=[1, 2], b=[3]) +class TestAxisInfo: + def test_valid(self): + info = AxisInfo(axis_rank=0, axis_size=4) + assert info.axis_rank == 0 + + def test_axis_size_zero(self): + with pytest.raises(ValidationError, match="axis_size must be > 0"): + AxisInfo(axis_rank=0, axis_size=0) + + def test_axis_size_negative(self): + with pytest.raises(ValidationError, match="axis_size must be > 0"): + AxisInfo(axis_rank=0, axis_size=-1) + + def test_axis_rank_negative(self): + with pytest.raises(ValidationError, match="axis_rank must be in"): + AxisInfo(axis_rank=-1, axis_size=4) + + def test_axis_rank_too_large(self): + with pytest.raises(ValidationError, match="axis_rank must be in"): + AxisInfo(axis_rank=4, axis_size=4) + + def test_axis_rank_equals_size_minus_one(self): + info = AxisInfo(axis_rank=3, axis_size=4) + assert info.axis_rank == 3 + + class TestSummaryRecord: def test_valid(self): record = SummaryRecord(total=10, passed=7, failed=2, skipped=1) diff --git a/test/registered/debug_utils/comparator/test_utils.py b/test/registered/debug_utils/comparator/test_utils.py index 9d36c116a..89a2b9856 100644 --- a/test/registered/debug_utils/comparator/test_utils.py +++ b/test/registered/debug_utils/comparator/test_utils.py @@ -4,6 +4,7 @@ import pytest import torch from sglang.srt.debug_utils.comparator.utils import ( + Pair, argmax_coord, calc_rel_diff, compute_smaller_dtype, @@ -78,16 +79,43 @@ class TestTryUnifyShape: class TestComputeSmallerDtype: def test_float32_bfloat16(self): - assert compute_smaller_dtype(torch.float32, torch.bfloat16) == torch.bfloat16 + assert ( + compute_smaller_dtype(Pair(x=torch.float32, y=torch.bfloat16)) + == torch.bfloat16 + ) def test_reverse_order(self): - assert compute_smaller_dtype(torch.bfloat16, torch.float32) == torch.bfloat16 + assert ( + compute_smaller_dtype(Pair(x=torch.bfloat16, y=torch.float32)) + == torch.bfloat16 + ) def test_same_dtype_returns_none(self): - assert compute_smaller_dtype(torch.float32, torch.float32) is None + assert compute_smaller_dtype(Pair(x=torch.float32, y=torch.float32)) is None def test_unknown_pair_returns_none(self): - assert compute_smaller_dtype(torch.int32, torch.int64) is None + assert compute_smaller_dtype(Pair(x=torch.int32, y=torch.int64)) is None + + +class TestPairMap: + def test_map_basic(self): + pair = Pair(x=[1, 2, 3], y=[4, 5, 6]) + result = pair.map(lambda lst: sum(lst)) + assert result.x == 6 + assert result.y == 15 + + def test_map_type_change(self): + pair = Pair(x=[1, 2, 3], y=[10, 20]) + result = pair.map(len) + assert result.x == 3 + assert result.y == 2 + + def test_map_returns_new_pair(self): + pair = Pair(x="hello", y="world") + result = pair.map(str.upper) + assert result.x == "HELLO" + assert result.y == "WORLD" + assert result is not pair if __name__ == "__main__":