Support singleton dimension squeezing in dump comparator (#19566)

This commit is contained in:
fzyzcjy
2026-02-28 18:11:46 +08:00
committed by GitHub
parent 80bbd30909
commit 5705e02d28
16 changed files with 841 additions and 26 deletions

View File

@@ -0,0 +1,109 @@
from __future__ import annotations
from typing import Optional
import torch
from einops import rearrange
from sglang.srt.debug_utils.comparator.dims import (
_SingletonDimUtil,
parse_dims,
)
from sglang.srt.debug_utils.comparator.utils import Pair, _FrozenBase
from sglang.srt.debug_utils.comparator.warning_sink import warning_sink
# --- types ---
class AxisAlignerPlan(_FrozenBase):
pattern: Pair[Optional[str]] # einops pattern per side, None = no-op
# --- planner ---
def compute_axis_aligner_plan(
dims_str_pair: Pair[Optional[str]],
) -> Optional[AxisAlignerPlan]:
if dims_str_pair.x is None or dims_str_pair.y is None:
return None
# Pair[str] after None check
dims_pair: Pair[str] = Pair(x=dims_str_pair.x, y=dims_str_pair.y)
raw_names: Pair[list[str]] = dims_pair.map(
lambda s: [spec.name for spec in parse_dims(s)]
)
filtered_names: Pair[list[str]] = dims_pair.map(
lambda s: [spec.name for spec in _SingletonDimUtil.filter_out(parse_dims(s))]
)
target_order: Optional[list[str]] = _resolve_target_order(
x_names=filtered_names.x, y_names=filtered_names.y
)
if target_order is None:
return None
pattern: Pair[Optional[str]] = raw_names.map(
lambda names: _build_pattern(source=names, target=target_order)
)
if pattern.x is None and pattern.y is None:
return None
return AxisAlignerPlan(pattern=pattern)
def _resolve_target_order(
x_names: list[str], y_names: list[str]
) -> Optional[list[str]]:
"""Determine the canonical dim order both sides should align to.
Returns y_names (the target ordering) if name sets match, or None on mismatch.
If both sides are identical, returns the shared order.
"""
if x_names == y_names:
return y_names
if set(x_names) != set(y_names):
# Local import to avoid circular dependency:
# output_types -> aligner/entrypoint/types -> axis_aligner -> output_types
from sglang.srt.debug_utils.comparator.output_types import GeneralWarning
warning_sink.add(
GeneralWarning(
category="axis_aligner_dim_mismatch",
message=(
f"AxisAligner: dim name sets differ (x={x_names}, y={y_names}), "
f"skipping axis swap"
),
)
)
return None
return y_names
def _build_pattern(*, source: list[str], target: list[str]) -> Optional[str]:
"""Build an einops rearrange pattern from source dim names to target dim names.
Returns None if source already matches target (no rearrange needed).
"""
if source == target:
return None
return f"{' '.join(source)} -> {' '.join(target)}"
# --- executor ---
def execute_axis_aligner_plan(
tensor: torch.Tensor, plan: AxisAlignerPlan, *, side: str
) -> torch.Tensor:
pattern: Optional[str] = plan.pattern.x if side == "x" else plan.pattern.y
if pattern is not None:
tensor = rearrange(tensor.rename(None), pattern)
return tensor

View File

@@ -5,8 +5,8 @@ from typing import Optional
import torch
from sglang.srt.debug_utils.comparator.aligner.axis_swapper import (
execute_axis_swapper_plan,
from sglang.srt.debug_utils.comparator.aligner.axis_aligner import (
execute_axis_aligner_plan,
)
from sglang.srt.debug_utils.comparator.aligner.entrypoint.types import (
AlignerPerStepPlan,
@@ -65,11 +65,11 @@ def execute_aligner_plan(
y=list(step_tensors_y.values())[0],
)
# Cross-side: axis swap (rearrange x to match y's dim order)
if (swap_plan := plan.axis_swapper_plan) is not None:
# Cross-side: axis alignment (squeeze singletons + rearrange dim order)
if (aligner_plan := plan.axis_aligner_plan) is not None:
combined = Pair(
x=execute_axis_swapper_plan(tensor=combined.x, plan=swap_plan),
y=combined.y,
x=execute_axis_aligner_plan(tensor=combined.x, plan=aligner_plan, side="x"),
y=execute_axis_aligner_plan(tensor=combined.y, plan=aligner_plan, side="y"),
)
return AlignerResult(tensors=combined, failed_side_xy=None)

View File

@@ -2,9 +2,9 @@ from __future__ import annotations
from typing import Any, Optional
from sglang.srt.debug_utils.comparator.aligner.axis_swapper import (
AxisSwapperPlan,
compute_axis_swapper_plan,
from sglang.srt.debug_utils.comparator.aligner.axis_aligner import (
AxisAlignerPlan,
compute_axis_aligner_plan,
)
from sglang.srt.debug_utils.comparator.aligner.entrypoint.types import (
AlignerPerStepPlan,
@@ -23,7 +23,11 @@ from sglang.srt.debug_utils.comparator.aligner.unsharder.parallel_info import (
from sglang.srt.debug_utils.comparator.aligner.unsharder.planner import (
compute_unsharder_plan,
)
from sglang.srt.debug_utils.comparator.dims import DimSpec, parse_dims
from sglang.srt.debug_utils.comparator.dims import (
DimSpec,
_SingletonDimUtil,
parse_dims,
)
from sglang.srt.debug_utils.comparator.utils import Pair
@@ -38,7 +42,7 @@ def compute_aligner_plan(
dims_str_pair: Pair[Optional[str]] = metas_pair.map(
lambda metas: metas[0].get("dims") if metas else None
)
axis_swapper_plan: Optional[AxisSwapperPlan] = compute_axis_swapper_plan(
axis_aligner_plan: Optional[AxisAlignerPlan] = compute_axis_aligner_plan(
dims_str_pair=dims_str_pair
)
@@ -54,7 +58,7 @@ def compute_aligner_plan(
),
),
token_aligner_plan=token_aligner_plan,
axis_swapper_plan=axis_swapper_plan,
axis_aligner_plan=axis_aligner_plan,
)
@@ -100,7 +104,7 @@ def compute_per_step_sub_plans(
if dims_str is None:
return []
dim_specs: list[DimSpec] = parse_dims(dims_str)
dim_specs: list[DimSpec] = _SingletonDimUtil.filter_out(parse_dims(dims_str))
parallel_infos = [normalize_parallel_info(meta) for meta in metas]
unsharder_plans = compute_unsharder_plan(

View File

@@ -4,7 +4,7 @@ from typing import Annotated, Optional, Union
from pydantic import Discriminator
from sglang.srt.debug_utils.comparator.aligner.axis_swapper import AxisSwapperPlan
from sglang.srt.debug_utils.comparator.aligner.axis_aligner import AxisAlignerPlan
from sglang.srt.debug_utils.comparator.aligner.reorderer.types import ReordererPlan
from sglang.srt.debug_utils.comparator.aligner.token_aligner.types import (
TokenAlignerPlan,
@@ -27,4 +27,4 @@ class AlignerPerStepPlan(_FrozenBase):
class AlignerPlan(_FrozenBase):
per_step_plans: Pair[list[AlignerPerStepPlan]]
token_aligner_plan: Optional[TokenAlignerPlan] = None
axis_swapper_plan: Optional[AxisSwapperPlan] = None
axis_aligner_plan: Optional[AxisAlignerPlan] = None

View File

@@ -28,7 +28,7 @@ from sglang.srt.debug_utils.comparator.dims import (
ParallelAxis,
TokenLayout,
apply_dim_names,
parse_dim_names,
resolve_dim_names,
)
from sglang.srt.debug_utils.comparator.output_types import GeneralWarning
from sglang.srt.debug_utils.comparator.warning_sink import warning_sink
@@ -230,7 +230,7 @@ def _load_and_align_aux_tensor(
if sub_plans:
dims_str: Optional[str] = metas[0].get("dims")
if dims_str is not None:
dim_names: list[str] = parse_dim_names(dims_str)
dim_names: list[str] = resolve_dim_names(dims_str)
tensors = [apply_dim_names(t, dim_names) for t in tensors]
result = execute_sub_plans(tensors=tensors, plans=sub_plans)

View File

@@ -18,7 +18,10 @@ from sglang.srt.debug_utils.comparator.aligner.entrypoint.types import AlignerPl
from sglang.srt.debug_utils.comparator.aligner.token_aligner.types import (
TokenAlignerPlan,
)
from sglang.srt.debug_utils.comparator.dims import apply_dim_names, parse_dim_names
from sglang.srt.debug_utils.comparator.dims import (
apply_dim_names,
resolve_dim_names,
)
from sglang.srt.debug_utils.comparator.output_types import (
ComparisonRecord,
GeneralWarning,
@@ -241,7 +244,7 @@ def _apply_dim_names_from_meta(
if dims_str is None:
return tensors
dim_names: list[str] = parse_dim_names(dims_str)
dim_names: list[str] = resolve_dim_names(dims_str)
return [apply_dim_names(t, dim_names) for t in tensors]

View File

@@ -8,6 +8,7 @@ import torch
TOKEN_DIM_NAME: str = "t"
BATCH_DIM_NAME: str = "b"
SEQ_DIM_NAME: str = "s"
SQUEEZE_DIM_NAME: str = "1"
class TokenLayout(Enum):
@@ -40,6 +41,46 @@ class DimSpec:
reduction: Optional[Reduction] = None
class _SingletonDimUtil:
"""Utilities for squeeze dims (name="1") and their singleton tensor-name mapping."""
PREFIX: str = "singleton"
@staticmethod
def is_squeeze(spec: DimSpec) -> bool:
return spec.name == SQUEEZE_DIM_NAME
@staticmethod
def filter_out(dim_specs: list[DimSpec]) -> list[DimSpec]:
return [s for s in dim_specs if not _SingletonDimUtil.is_squeeze(s)]
@staticmethod
def make_name(index: int) -> str:
return f"{_SingletonDimUtil.PREFIX}{index}"
@staticmethod
def is_singleton_name(name: str) -> bool:
return (
name.startswith(_SingletonDimUtil.PREFIX)
and name[len(_SingletonDimUtil.PREFIX) :].isdigit()
)
@staticmethod
def sanitize_names(names: list[str]) -> list[str]:
"""Replace '1' with 'singleton0', 'singleton1', ... for named tensor compatibility."""
result: list[str] = []
sq_idx: int = 0
for name in names:
if name == SQUEEZE_DIM_NAME:
result.append(_SingletonDimUtil.make_name(sq_idx))
sq_idx += 1
else:
result.append(name)
return result
_DIM_PATTERN = re.compile(r"^(?P<name>[a-zA-Z_]\w*)(?:\((?P<modifiers>[^)]+)\))?$")
_MODIFIER_FIELDS: list[tuple[type[Enum], str]] = [
@@ -55,6 +96,9 @@ for _enum_cls, _field in _MODIFIER_FIELDS:
def parse_dim(token: str) -> DimSpec:
if token == SQUEEZE_DIM_NAME:
return DimSpec(name=SQUEEZE_DIM_NAME)
match = _DIM_PATTERN.match(token)
if match is None:
raise ValueError(f"Invalid dim token: {token!r}")
@@ -84,16 +128,22 @@ def parse_dims(dims_str: str) -> list[DimSpec]:
result = [parse_dim(token) for token in dims_str.strip().split()]
names = [spec.name for spec in result]
if len(names) != len(set(names)):
duplicates = sorted({n for n in names if names.count(n) > 1})
non_squeeze_names: list[str] = [
spec.name for spec in result if not _SingletonDimUtil.is_squeeze(spec)
]
if len(non_squeeze_names) != len(set(non_squeeze_names)):
duplicates = sorted(
{n for n in non_squeeze_names if non_squeeze_names.count(n) > 1}
)
raise ValueError(f"Duplicate dim names: {duplicates}")
return result
def parse_dim_names(dims_str: str) -> list[str]:
return [spec.name for spec in parse_dims(dims_str)]
def resolve_dim_names(dims_str: str) -> list[str]:
"""Parse dims string and return tensor-compatible names ('1''singleton0', ...)."""
names: list[str] = [spec.name for spec in parse_dims(dims_str)]
return _SingletonDimUtil.sanitize_names(names)
def find_dim_index(dim_specs: list[DimSpec], name: str) -> Optional[int]:

View File

@@ -219,8 +219,13 @@ def _format_aligner_plan(plan: AlignerPlan) -> str:
num_tokens: int = len(plan.token_aligner_plan.locators.x.steps)
lines.append(f" token_aligner: {num_tokens} tokens aligned")
if plan.axis_swapper_plan is not None:
lines.append(f" axis_swapper: {plan.axis_swapper_plan.pattern}")
if plan.axis_aligner_plan is not None:
parts: list[str] = []
if plan.axis_aligner_plan.pattern.x:
parts.append(f"x: {plan.axis_aligner_plan.pattern.x}")
if plan.axis_aligner_plan.pattern.y:
parts.append(f"y: {plan.axis_aligner_plan.pattern.y}")
lines.append(f" axis_aligner: {', '.join(parts)}")
return "\n".join(lines)