Support non orthogonal parallel axes and explicit replication annotation in dump comparator (#19679)

This commit is contained in:
fzyzcjy
2026-03-02 18:44:33 +08:00
committed by GitHub
parent a70dd11011
commit 6980416149
40 changed files with 1451 additions and 1012 deletions

View File

@@ -5,7 +5,7 @@ from typing import Optional
import torch
from einops import rearrange
from sglang.srt.debug_utils.comparator.dims import (
from sglang.srt.debug_utils.comparator.dims_spec import (
_FUSED_NAME_SEP,
DimSpec,
_SingletonDimUtil,

View File

@@ -23,8 +23,10 @@ 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 (
from sglang.srt.debug_utils.comparator.dims_spec import (
DimSpec,
DimsSpec,
ParallelAxis,
_SingletonDimUtil,
parse_dims,
)
@@ -106,12 +108,15 @@ def compute_per_step_sub_plans(
if dims_str is None:
return []
dim_specs: list[DimSpec] = _SingletonDimUtil.filter_out(parse_dims(dims_str).dims)
dims_spec: DimsSpec = parse_dims(dims_str)
dim_specs: list[DimSpec] = _SingletonDimUtil.filter_out(dims_spec.dims)
replicated_axes: frozenset[ParallelAxis] = dims_spec.replicated_axes
parallel_infos = [normalize_parallel_info(meta) for meta in metas]
unsharder_plans = compute_unsharder_plan(
dim_specs=dim_specs,
parallel_infos=parallel_infos,
explicit_replicated_axes=replicated_axes,
thd_global_seq_lens=thd_global_seq_lens,
)
reorderer_plans = compute_reorderer_plans(

View File

@@ -7,7 +7,7 @@ from sglang.srt.debug_utils.comparator.aligner.reorderer.types import (
ZigzagToNaturalParams,
ZigzagToNaturalThdParams,
)
from sglang.srt.debug_utils.comparator.dims import (
from sglang.srt.debug_utils.comparator.dims_spec import (
resolve_dim_by_name,
strip_dim_names,
)

View File

@@ -6,7 +6,7 @@ from sglang.srt.debug_utils.comparator.aligner.reorderer.types import (
ZigzagToNaturalThdParams,
)
from sglang.srt.debug_utils.comparator.aligner.unsharder.types import AxisInfo
from sglang.srt.debug_utils.comparator.dims import (
from sglang.srt.debug_utils.comparator.dims_spec import (
SEQ_DIM_NAME,
TOKEN_DIM_NAME,
DimSpec,

View File

@@ -4,7 +4,7 @@ from typing import Optional
import torch
from sglang.srt.debug_utils.comparator.dims import (
from sglang.srt.debug_utils.comparator.dims_spec import (
SEQ_DIM_NAME,
TOKEN_DIM_NAME,
)

View File

@@ -24,7 +24,7 @@ from sglang.srt.debug_utils.comparator.aligner.token_aligner.smart.types import
from sglang.srt.debug_utils.comparator.aligner.unsharder.parallel_info import (
normalize_parallel_info,
)
from sglang.srt.debug_utils.comparator.dims import (
from sglang.srt.debug_utils.comparator.dims_spec import (
ParallelAxis,
TokenLayout,
apply_dim_names,

View File

@@ -11,7 +11,7 @@ from sglang.srt.debug_utils.comparator.aligner.token_aligner.smart.types import
SGLangSeqId,
TokenAlignerStepAux,
)
from sglang.srt.debug_utils.comparator.dims import TokenLayout
from sglang.srt.debug_utils.comparator.dims_spec import TokenLayout
from sglang.srt.debug_utils.comparator.log_sink import log_sink
from sglang.srt.debug_utils.comparator.output_types import InfoLog

View File

@@ -7,7 +7,7 @@ from sglang.srt.debug_utils.comparator.aligner.token_aligner.smart.types import
TokenAlignerPlan,
TokenLocator,
)
from sglang.srt.debug_utils.comparator.dims import (
from sglang.srt.debug_utils.comparator.dims_spec import (
BATCH_DIM_NAME,
SEQ_DIM_NAME,
TOKEN_DIM_NAME,

View File

@@ -5,7 +5,7 @@ from typing import NamedTuple, Optional, Union
from pydantic import model_validator
from sglang.srt.debug_utils.comparator.dims import TokenLayout
from sglang.srt.debug_utils.comparator.dims_spec import TokenLayout
from sglang.srt.debug_utils.comparator.utils import (
Pair,
_check_equal_lengths,

View File

@@ -11,7 +11,7 @@ from sglang.srt.debug_utils.comparator.aligner.unsharder.types import (
UnsharderParams,
UnsharderPlan,
)
from sglang.srt.debug_utils.comparator.dims import (
from sglang.srt.debug_utils.comparator.dims_spec import (
ParallelAxis,
resolve_dim_by_name,
)

View File

@@ -1,7 +1,7 @@
from typing import Optional
from sglang.srt.debug_utils.comparator.aligner.unsharder.types import AxisInfo
from sglang.srt.debug_utils.comparator.dims import ParallelAxis
from sglang.srt.debug_utils.comparator.dims_spec import ParallelAxis
_PARALLEL_INFO_KEYS = ("sglang_parallel_info", "megatron_parallel_info")

View File

@@ -10,7 +10,7 @@ from sglang.srt.debug_utils.comparator.aligner.unsharder.types import (
UnsharderParams,
UnsharderPlan,
)
from sglang.srt.debug_utils.comparator.dims import (
from sglang.srt.debug_utils.comparator.dims_spec import (
TOKEN_DIM_NAME,
DimSpec,
ParallelAxis,
@@ -32,6 +32,7 @@ def compute_unsharder_plan(
dim_specs: list[DimSpec],
parallel_infos: list[dict[ParallelAxis, AxisInfo]],
*,
explicit_replicated_axes: frozenset[ParallelAxis] = frozenset(),
thd_global_seq_lens: Optional[list[int]] = None,
) -> list[UnsharderPlan]:
if not parallel_infos:
@@ -54,7 +55,21 @@ def compute_unsharder_plan(
reversed_sharded_modifiers = [
(name, m) for name, m in reversed_sharded_modifiers if m.axis in sharded_axes
]
replicated_axes: set[ParallelAxis] = all_axes - sharded_axes
# RECOMPUTE_PSEUDO is always implicitly replicated (system-injected, not user-facing)
auto_replicated: frozenset[ParallelAxis] = frozenset(
{ParallelAxis.RECOMPUTE_PSEUDO} & all_axes
)
effective_replicated: frozenset[ParallelAxis] = (
explicit_replicated_axes | auto_replicated
)
_validate_explicit_replicated(
explicit_replicated_axes=effective_replicated,
sharded_axes=sharded_axes,
all_axes=all_axes,
)
replicated_axes: frozenset[ParallelAxis] = effective_replicated
if not sharded_axes and not replicated_axes:
return []
@@ -96,6 +111,37 @@ def compute_unsharder_plan(
return plans
def _validate_explicit_replicated(
*,
explicit_replicated_axes: frozenset[ParallelAxis],
sharded_axes: set[ParallelAxis],
all_axes: set[ParallelAxis],
) -> None:
"""Validate explicit replicated declarations against sharded axes and parallel_infos."""
invalid: frozenset[ParallelAxis] = explicit_replicated_axes - all_axes
if invalid:
invalid_names: str = ", ".join(sorted(a.value for a in invalid))
raise ValueError(
f"Declared replicated axes {{{invalid_names}}} not found in parallel_infos "
f"(active axes: {{{', '.join(sorted(a.value for a in all_axes))}}})"
)
conflict: set[ParallelAxis] = explicit_replicated_axes & sharded_axes
if conflict:
conflict_names: str = ", ".join(sorted(a.value for a in conflict))
raise ValueError(
f"Axes {{{conflict_names}}} declared as both sharded and replicated"
)
undeclared: set[ParallelAxis] = all_axes - sharded_axes - explicit_replicated_axes
if undeclared:
undeclared_names: str = ", ".join(sorted(a.value for a in undeclared))
raise ValueError(
f"Axes {{{undeclared_names}}} are active (axis_size > 1) but not declared "
f"in dims. Annotate as sharded in dim spec or as '# axis:replicated'."
)
def _validate(
*,
axes_to_validate: set[ParallelAxis],

View File

@@ -4,7 +4,7 @@ from typing import Annotated, Literal, Union
from pydantic import Field, model_validator
from sglang.srt.debug_utils.comparator.dims import ParallelAxis
from sglang.srt.debug_utils.comparator.dims_spec import ParallelAxis
from sglang.srt.debug_utils.comparator.utils import _FrozenBase

View File

@@ -18,7 +18,7 @@ from sglang.srt.debug_utils.comparator.aligner.entrypoint.types import AlignerPl
from sglang.srt.debug_utils.comparator.aligner.token_aligner.smart.types import (
TokenAlignerPlan,
)
from sglang.srt.debug_utils.comparator.dims import (
from sglang.srt.debug_utils.comparator.dims_spec import (
SEQ_DIM_NAME,
TOKEN_DIM_NAME,
apply_dim_names,

View File

@@ -1,338 +0,0 @@
from __future__ import annotations
import re
from enum import Enum
from typing import Optional
import torch
from sglang.srt.debug_utils.comparator.utils import _FrozenBase
TOKEN_DIM_NAME: str = "t"
BATCH_DIM_NAME: str = "b"
SEQ_DIM_NAME: str = "s"
SQUEEZE_DIM_NAME: str = "1"
class TokenLayout(Enum):
T = "t" # single flat token dim
BS = "bs" # separate batch + seq dims, need collapse
class ParallelAxis(Enum):
TP = "tp"
CP = "cp"
EP = "ep"
SP = "sp"
RECOMPUTE_PSEUDO = "recompute_pseudo"
class Ordering(Enum):
ZIGZAG = "zigzag"
NATURAL = "natural"
class Reduction(Enum):
PARTIAL = "partial"
class ParallelModifier(_FrozenBase):
axis: ParallelAxis
ordering: Optional[Ordering] = None
reduction: Optional[Reduction] = None
_FUSED_NAME_SEP: str = "___"
class DimSpec(_FrozenBase):
name: str
parallel_modifiers: list[ParallelModifier] = []
@property
def sub_dims(self) -> list[str]:
"""Sub-dim names. Fused: ``["num_heads", "head_dim"]``; plain: ``["h"]``."""
return self.name.split("*")
@property
def is_fused(self) -> bool:
return len(self.sub_dims) > 1
@property
def sanitized_name(self) -> str:
"""Name safe for PyTorch named tensors (``*`` → ``___``)."""
if self.is_fused:
return _FUSED_NAME_SEP.join(self.sub_dims)
return self.name
class DimsSpec(_FrozenBase):
"""Parsed result of a full dims string like ``"b s h[tp] # dp:=moe_dp"``."""
dims: list[DimSpec]
dp_group_alias: Optional[str] = None
class DimsSpec(_FrozenBase):
"""Parsed result of a full dims string like ``"b s h(tp) # dp:=moe_dp"``."""
dims: list[DimSpec]
dp_group_alias: Optional[str] = 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>[^\]]+)\])?$")
_FUSED_DIM_PATTERN = re.compile(r"^\((?P<inner>[^)]+)\)(?:\[(?P<modifiers>[^\]]+)\])?$")
_SUB_DIM_NAME_PATTERN = re.compile(r"^[a-zA-Z_]\w*$")
_AXIS_LOOKUP: dict[str, ParallelAxis] = {m.value: m for m in ParallelAxis}
_QUALIFIER_LOOKUP: dict[str, Ordering | Reduction] = {
**{m.value: m for m in Ordering},
**{m.value: m for m in Reduction},
}
def _parse_modifier_token(modifier_token: str, dim_token: str) -> ParallelModifier:
"""Parse 'sp', 'cp:zigzag', 'tp:partial', or 'cp:zigzag+partial' → ParallelModifier.
Format: ``axis`` or ``axis:qual`` or ``axis:qual+qual``.
Colon separates axis from qualifiers; ``+`` separates multiple qualifiers.
"""
axis_str: str
qualifiers_str: str
if ":" in modifier_token:
axis_str, qualifiers_str = modifier_token.split(":", maxsplit=1)
else:
axis_str, qualifiers_str = modifier_token, ""
axis_str = axis_str.strip()
axis: Optional[ParallelAxis] = _AXIS_LOOKUP.get(axis_str)
if axis is None:
raise ValueError(
f"Unknown axis {axis_str!r} in modifier {modifier_token!r} "
f"of dim spec: {dim_token!r}"
)
ordering: Optional[Ordering] = None
reduction: Optional[Reduction] = None
for q_str in (q.strip() for q in qualifiers_str.split("+") if q.strip()):
qualifier: Optional[Ordering | Reduction] = _QUALIFIER_LOOKUP.get(q_str)
if qualifier is None:
raise ValueError(
f"Unknown qualifier {q_str!r} in modifier "
f"{modifier_token!r} of dim spec: {dim_token!r}"
)
if isinstance(qualifier, Ordering):
if ordering is not None:
raise ValueError(
f"Multiple ordering values in modifier "
f"{modifier_token!r} of dim spec: {dim_token!r}"
)
ordering = qualifier
else:
if reduction is not None:
raise ValueError(
f"Multiple reduction values in modifier "
f"{modifier_token!r} of dim spec: {dim_token!r}"
)
reduction = qualifier
return ParallelModifier(axis=axis, ordering=ordering, reduction=reduction)
def parse_dim(token: str) -> DimSpec:
if token == SQUEEZE_DIM_NAME:
return DimSpec(name=SQUEEZE_DIM_NAME)
fused_match = _FUSED_DIM_PATTERN.match(token)
if fused_match is not None:
return _parse_fused_dim(token=token, fused_match=fused_match)
return _parse_single_dim(token)
def _parse_single_dim(token: str) -> DimSpec:
match = _DIM_PATTERN.match(token)
if match is None:
raise ValueError(f"Invalid dim token: {token!r}")
name: str = match.group("name")
modifiers: list[ParallelModifier] = _parse_modifiers(
modifiers_str=match.group("modifiers"), dim_token=token
)
return DimSpec(name=name, parallel_modifiers=modifiers)
def _parse_fused_dim(*, token: str, fused_match: re.Match[str]) -> DimSpec:
inner: str = fused_match.group("inner")
modifiers_str: Optional[str] = fused_match.group("modifiers")
sub_names: list[str] = [s.strip() for s in inner.split("*")]
for sub_name in sub_names:
if not _SUB_DIM_NAME_PATTERN.match(sub_name):
raise ValueError(
f"Invalid sub-dim {sub_name!r} in fused dim token: {token!r}"
)
if len(sub_names) != len(set(sub_names)):
raise ValueError(f"Duplicate sub-dim names in fused dim token: {token!r}")
if len(sub_names) < 2:
raise ValueError(
f"Fused dim must have at least 2 sub-dims, got {len(sub_names)} in: {token!r}"
)
fused_name: str = "*".join(sub_names)
modifiers: list[ParallelModifier] = _parse_modifiers(
modifiers_str=modifiers_str, dim_token=token
)
return DimSpec(name=fused_name, parallel_modifiers=modifiers)
def _parse_modifiers(
*, modifiers_str: Optional[str], dim_token: str
) -> list[ParallelModifier]:
if modifiers_str is None:
return []
modifiers: list[ParallelModifier] = []
seen_axes: set[ParallelAxis] = set()
for modifier_token in (p.strip() for p in modifiers_str.split(",")):
modifier: ParallelModifier = _parse_modifier_token(modifier_token, dim_token)
if modifier.axis in seen_axes:
raise ValueError(
f"Duplicate axis {modifier.axis.value!r} in dim spec: {dim_token!r}"
)
seen_axes.add(modifier.axis)
modifiers.append(modifier)
return modifiers
def parse_dims(dims_str: str) -> DimsSpec:
"""Parse ``"b s[cp:zigzag] h[tp] d # dp:=moe_dp"`` → :class:`DimsSpec`.
The shape part (before ``#``) produces :pyattr:`DimsSpec.dims`.
The declaration part (after ``#``) is scanned for ``dp:=<group>``
which populates :pyattr:`DimsSpec.dp_group_alias`.
"""
parts: list[str] = dims_str.split("#", maxsplit=1)
raw: str = parts[0]
if not raw.strip():
raise ValueError("dims string must not be empty")
dims: list[DimSpec] = [parse_dim(token) for token in raw.strip().split()]
# Collect all semantic names (expanding fused sub-dims) for duplicate detection
semantic_names: list[str] = []
for spec in dims:
if _SingletonDimUtil.is_squeeze(spec):
continue
semantic_names.extend(spec.sub_dims)
if len(semantic_names) != len(set(semantic_names)):
duplicates = sorted({n for n in semantic_names if semantic_names.count(n) > 1})
raise ValueError(f"Duplicate dim names: {duplicates}")
dp_group_alias: Optional[str] = (
_extract_dp_group_alias(parts[1]) if len(parts) > 1 else None
)
return DimsSpec(dims=dims, dp_group_alias=dp_group_alias)
def resolve_dim_names(dims_str: str) -> list[str]:
"""Parse dims string and return tensor-compatible names ('1''singleton0', ...)."""
specs: list[DimSpec] = parse_dims(dims_str).dims
names: list[str] = [spec.sanitized_name for spec in specs]
return _SingletonDimUtil.sanitize_names(names)
def find_dim_index(dim_specs: list[DimSpec], name: str) -> Optional[int]:
"""Find index by name. Accepts both ``*``-form and ``___``-form for fused dims."""
for i, spec in enumerate(dim_specs):
if spec.name == name or spec.sanitized_name == name:
return i
return None
def resolve_dim_by_name(tensor: torch.Tensor, name: str) -> int:
if tensor.names[0] is None:
raise ValueError(f"Tensor has no names, cannot resolve {name!r}")
names: tuple[Optional[str], ...] = tensor.names
try:
return list(names).index(name)
except ValueError:
raise ValueError(f"Dim name {name!r} not in tensor names {names}")
def apply_dim_names(tensor: torch.Tensor, dim_names: list[str]) -> torch.Tensor:
if tensor.ndim != len(dim_names):
raise ValueError(
f"dims metadata mismatch: tensor has {tensor.ndim} dims (shape {list(tensor.shape)}) "
f"but dims string specifies {len(dim_names)} names {dim_names}. "
f"Please fix the dims string in the dumper.dump() call to match the actual tensor shape."
)
return tensor.refine_names(*dim_names)
def strip_dim_names(tensor: torch.Tensor) -> torch.Tensor:
return tensor.rename(None)
_DP_ALIAS_PATTERN = re.compile(r"^dp:=(\w+)$")
def _extract_dp_group_alias(declaration_part: str) -> Optional[str]:
"""Scan the ``#`` declaration section for a ``dp:=<group>`` token."""
for token in declaration_part.strip().split():
match = _DP_ALIAS_PATTERN.match(token)
if match is not None:
return match.group(1)
return None

View File

@@ -0,0 +1,49 @@
from sglang.srt.debug_utils.comparator.dims_spec.dim_parser import parse_dim
from sglang.srt.debug_utils.comparator.dims_spec.dims_parser import (
_SingletonDimUtil,
parse_dims,
resolve_dim_names,
)
from sglang.srt.debug_utils.comparator.dims_spec.tensor_naming import (
apply_dim_names,
find_dim_index,
resolve_dim_by_name,
strip_dim_names,
)
from sglang.srt.debug_utils.comparator.dims_spec.types import (
_FUSED_NAME_SEP,
BATCH_DIM_NAME,
SEQ_DIM_NAME,
SQUEEZE_DIM_NAME,
TOKEN_DIM_NAME,
DimSpec,
DimsSpec,
Ordering,
ParallelAxis,
ParallelModifier,
Reduction,
TokenLayout,
)
__all__ = [
"BATCH_DIM_NAME",
"SEQ_DIM_NAME",
"SQUEEZE_DIM_NAME",
"TOKEN_DIM_NAME",
"DimsSpec",
"DimSpec",
"Ordering",
"ParallelAxis",
"ParallelModifier",
"Reduction",
"TokenLayout",
"_FUSED_NAME_SEP",
"_SingletonDimUtil",
"apply_dim_names",
"find_dim_index",
"parse_dim",
"parse_dims",
"resolve_dim_by_name",
"resolve_dim_names",
"strip_dim_names",
]

View File

@@ -0,0 +1,59 @@
from __future__ import annotations
import re
from typing import NamedTuple, Optional
from sglang.srt.debug_utils.comparator.dims_spec.types import (
_AXIS_LOOKUP,
ParallelAxis,
)
_DP_ALIAS_PATTERN = re.compile(r"^dp:=(\w+)$")
_REPLICATED_PATTERN = re.compile(r"^(\w+):replicated$")
class _CommentSuffix(NamedTuple):
dp_group_alias: Optional[str] = None
replicated_axes: frozenset[ParallelAxis] = frozenset()
def _parse_comment_suffix(declaration_part: str) -> _CommentSuffix:
"""Parse the ``#`` comment section for dp alias and replicated declarations."""
dp_group_alias: Optional[str] = None
replicated_axes: set[ParallelAxis] = set()
for token in declaration_part.strip().split():
dp_match = _DP_ALIAS_PATTERN.match(token)
if dp_match is not None:
if dp_group_alias is not None:
raise ValueError(
f"Duplicate dp alias declaration: already have {dp_group_alias!r}, "
f"got {dp_match.group(1)!r}"
)
dp_group_alias = dp_match.group(1)
continue
repl_match = _REPLICATED_PATTERN.match(token)
if repl_match is not None:
axis_str: str = repl_match.group(1)
axis: Optional[ParallelAxis] = _AXIS_LOOKUP.get(axis_str)
if axis is None:
raise ValueError(
f"Unknown axis {axis_str!r} in replicated declaration: {token!r}"
)
if axis in replicated_axes:
raise ValueError(
f"Duplicate replicated declaration for axis {axis_str!r}"
)
replicated_axes.add(axis)
continue
raise ValueError(
f"Unrecognized token {token!r} in # comment section. "
f"Expected 'dp:=<group>' or '<axis>:replicated'."
)
return _CommentSuffix(
dp_group_alias=dp_group_alias,
replicated_axes=frozenset(replicated_axes),
)

View File

@@ -0,0 +1,68 @@
from __future__ import annotations
import re
from typing import Optional
from sglang.srt.debug_utils.comparator.dims_spec.modifier_parser import (
_parse_modifiers,
)
from sglang.srt.debug_utils.comparator.dims_spec.types import (
SQUEEZE_DIM_NAME,
DimSpec,
ParallelModifier,
)
_DIM_PATTERN = re.compile(r"^(?P<name>[a-zA-Z_]\w*)(?:\[(?P<modifiers>[^\]]+)\])?$")
_FUSED_DIM_PATTERN = re.compile(r"^\((?P<inner>[^)]+)\)(?:\[(?P<modifiers>[^\]]+)\])?$")
_SUB_DIM_NAME_PATTERN = re.compile(r"^[a-zA-Z_]\w*$")
def parse_dim(token: str) -> DimSpec:
if token == SQUEEZE_DIM_NAME:
return DimSpec(name=SQUEEZE_DIM_NAME)
fused_match = _FUSED_DIM_PATTERN.match(token)
if fused_match is not None:
return _parse_fused_dim(token=token, fused_match=fused_match)
return _parse_single_dim(token)
def _parse_single_dim(token: str) -> DimSpec:
match = _DIM_PATTERN.match(token)
if match is None:
raise ValueError(f"Invalid dim token: {token!r}")
name: str = match.group("name")
modifiers: list[ParallelModifier] = _parse_modifiers(
modifiers_str=match.group("modifiers"), dim_token=token
)
return DimSpec(name=name, parallel_modifiers=modifiers)
def _parse_fused_dim(*, token: str, fused_match: re.Match[str]) -> DimSpec:
inner: str = fused_match.group("inner")
modifiers_str: Optional[str] = fused_match.group("modifiers")
sub_names: list[str] = [s.strip() for s in inner.split("*")]
for sub_name in sub_names:
if not _SUB_DIM_NAME_PATTERN.match(sub_name):
raise ValueError(
f"Invalid sub-dim {sub_name!r} in fused dim token: {token!r}"
)
if len(sub_names) != len(set(sub_names)):
raise ValueError(f"Duplicate sub-dim names in fused dim token: {token!r}")
if len(sub_names) < 2:
raise ValueError(
f"Fused dim must have at least 2 sub-dims, got {len(sub_names)} in: {token!r}"
)
fused_name: str = "*".join(sub_names)
modifiers: list[ParallelModifier] = _parse_modifiers(
modifiers_str=modifiers_str, dim_token=token
)
return DimSpec(name=fused_name, parallel_modifiers=modifiers)

View File

@@ -0,0 +1,113 @@
from __future__ import annotations
from typing import Optional
from sglang.srt.debug_utils.comparator.dims_spec.comment_parser import (
_CommentSuffix,
_parse_comment_suffix,
)
from sglang.srt.debug_utils.comparator.dims_spec.dim_parser import parse_dim
from sglang.srt.debug_utils.comparator.dims_spec.types import (
SQUEEZE_DIM_NAME,
DimSpec,
DimsSpec,
ParallelAxis,
)
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
def parse_dims(dims_str: str) -> DimsSpec:
"""Parse ``"b s[cp:zigzag] h[tp] d # dp:=moe_dp ep:replicated"`` → :class:`DimsSpec`.
The shape part (before ``#``) produces :pyattr:`DimsSpec.dims`.
The declaration part (after ``#``) is scanned for:
- ``dp:=<group>`` → :pyattr:`DimsSpec.dp_group_alias`
- ``axis:replicated`` → :pyattr:`DimsSpec.replicated_axes`
"""
parts: list[str] = dims_str.split("#", maxsplit=1)
raw: str = parts[0]
if not raw.strip():
raise ValueError("dims string must not be empty")
dims: list[DimSpec] = [parse_dim(token) for token in raw.strip().split()]
# Collect all semantic names (expanding fused sub-dims) for duplicate detection
semantic_names: list[str] = []
for spec in dims:
if _SingletonDimUtil.is_squeeze(spec):
continue
semantic_names.extend(spec.sub_dims)
if len(semantic_names) != len(set(semantic_names)):
duplicates = sorted({n for n in semantic_names if semantic_names.count(n) > 1})
raise ValueError(f"Duplicate dim names: {duplicates}")
comment_suffix: _CommentSuffix = (
_parse_comment_suffix(parts[1]) if len(parts) > 1 else _CommentSuffix()
)
dp_group_alias: Optional[str] = comment_suffix.dp_group_alias
replicated_axes: frozenset[ParallelAxis] = comment_suffix.replicated_axes
sharded_axes: set[ParallelAxis] = {
m.axis for spec in dims for m in spec.parallel_modifiers
}
conflict: frozenset[ParallelAxis] = replicated_axes & sharded_axes
if conflict:
conflict_names: str = ", ".join(sorted(a.value for a in conflict))
raise ValueError(
f"Axes declared as both sharded (in dim spec) and replicated "
f"(in # declaration): {conflict_names}"
)
return DimsSpec(
dims=dims,
dp_group_alias=dp_group_alias,
replicated_axes=replicated_axes,
)
def resolve_dim_names(dims_str: str) -> list[str]:
"""Parse dims string and return tensor-compatible names ('1''singleton0', ...)."""
specs: list[DimSpec] = parse_dims(dims_str).dims
names: list[str] = [spec.sanitized_name for spec in specs]
return _SingletonDimUtil.sanitize_names(names)

View File

@@ -0,0 +1,84 @@
from __future__ import annotations
from typing import Optional
from sglang.srt.debug_utils.comparator.dims_spec.types import (
_AXIS_LOOKUP,
_QUALIFIER_LOOKUP,
Ordering,
ParallelAxis,
ParallelModifier,
Reduction,
)
def _parse_modifier_token(modifier_token: str, dim_token: str) -> ParallelModifier:
"""Parse 'sp', 'cp:zigzag', 'tp:partial', or 'cp:zigzag+partial' → ParallelModifier.
Format: ``axis`` or ``axis:qual`` or ``axis:qual+qual``.
Colon separates axis from qualifiers; ``+`` separates multiple qualifiers.
"""
axis_str: str
qualifiers_str: str
if ":" in modifier_token:
axis_str, qualifiers_str = modifier_token.split(":", maxsplit=1)
else:
axis_str, qualifiers_str = modifier_token, ""
axis_str = axis_str.strip()
axis: Optional[ParallelAxis] = _AXIS_LOOKUP.get(axis_str)
if axis is None:
raise ValueError(
f"Unknown axis {axis_str!r} in modifier {modifier_token!r} "
f"of dim spec: {dim_token!r}"
)
ordering: Optional[Ordering] = None
reduction: Optional[Reduction] = None
for q_str in (q.strip() for q in qualifiers_str.split("+") if q.strip()):
if q_str == "sharded":
continue
qualifier: Optional[Ordering | Reduction] = _QUALIFIER_LOOKUP.get(q_str)
if qualifier is None:
raise ValueError(
f"Unknown qualifier {q_str!r} in modifier "
f"{modifier_token!r} of dim spec: {dim_token!r}"
)
if isinstance(qualifier, Ordering):
if ordering is not None:
raise ValueError(
f"Multiple ordering values in modifier "
f"{modifier_token!r} of dim spec: {dim_token!r}"
)
ordering = qualifier
else:
if reduction is not None:
raise ValueError(
f"Multiple reduction values in modifier "
f"{modifier_token!r} of dim spec: {dim_token!r}"
)
reduction = qualifier
return ParallelModifier(axis=axis, ordering=ordering, reduction=reduction)
def _parse_modifiers(
*, modifiers_str: Optional[str], dim_token: str
) -> list[ParallelModifier]:
if modifiers_str is None:
return []
modifiers: list[ParallelModifier] = []
seen_axes: set[ParallelAxis] = set()
for modifier_token in (p.strip() for p in modifiers_str.split(",")):
modifier: ParallelModifier = _parse_modifier_token(modifier_token, dim_token)
if modifier.axis in seen_axes:
raise ValueError(
f"Duplicate axis {modifier.axis.value!r} in dim spec: {dim_token!r}"
)
seen_axes.add(modifier.axis)
modifiers.append(modifier)
return modifiers

View File

@@ -0,0 +1,40 @@
from __future__ import annotations
from typing import Optional
import torch
from sglang.srt.debug_utils.comparator.dims_spec.types import DimSpec
def find_dim_index(dim_specs: list[DimSpec], name: str) -> Optional[int]:
"""Find index by name. Accepts both ``*``-form and ``___``-form for fused dims."""
for i, spec in enumerate(dim_specs):
if spec.name == name or spec.sanitized_name == name:
return i
return None
def resolve_dim_by_name(tensor: torch.Tensor, name: str) -> int:
if tensor.names[0] is None:
raise ValueError(f"Tensor has no names, cannot resolve {name!r}")
names: tuple[Optional[str], ...] = tensor.names
try:
return list(names).index(name)
except ValueError:
raise ValueError(f"Dim name {name!r} not in tensor names {names}")
def apply_dim_names(tensor: torch.Tensor, dim_names: list[str]) -> torch.Tensor:
if tensor.ndim != len(dim_names):
raise ValueError(
f"dims metadata mismatch: tensor has {tensor.ndim} dims (shape {list(tensor.shape)}) "
f"but dims string specifies {len(dim_names)} names {dim_names}. "
f"Please fix the dims string in the dumper.dump() call to match the actual tensor shape."
)
return tensor.refine_names(*dim_names)
def strip_dim_names(tensor: torch.Tensor) -> torch.Tensor:
return tensor.rename(None)

View File

@@ -0,0 +1,77 @@
from __future__ import annotations
from enum import Enum
from typing import Optional
from sglang.srt.debug_utils.comparator.utils import _FrozenBase
TOKEN_DIM_NAME: str = "t"
BATCH_DIM_NAME: str = "b"
SEQ_DIM_NAME: str = "s"
SQUEEZE_DIM_NAME: str = "1"
class TokenLayout(Enum):
T = "t" # single flat token dim
BS = "bs" # separate batch + seq dims, need collapse
class ParallelAxis(Enum):
TP = "tp"
CP = "cp"
EP = "ep"
SP = "sp"
RECOMPUTE_PSEUDO = "recompute_pseudo"
class Ordering(Enum):
ZIGZAG = "zigzag"
NATURAL = "natural"
class Reduction(Enum):
PARTIAL = "partial"
class ParallelModifier(_FrozenBase):
axis: ParallelAxis
ordering: Optional[Ordering] = None
reduction: Optional[Reduction] = None
_AXIS_LOOKUP: dict[str, ParallelAxis] = {m.value: m for m in ParallelAxis}
_QUALIFIER_LOOKUP: dict[str, Ordering | Reduction] = {
**{m.value: m for m in Ordering},
**{m.value: m for m in Reduction},
}
_FUSED_NAME_SEP: str = "___"
class DimSpec(_FrozenBase):
name: str
parallel_modifiers: list[ParallelModifier] = []
@property
def sub_dims(self) -> list[str]:
"""Sub-dim names. Fused: ``["num_heads", "head_dim"]``; plain: ``["h"]``."""
return self.name.split("*")
@property
def is_fused(self) -> bool:
return len(self.sub_dims) > 1
@property
def sanitized_name(self) -> str:
"""Name safe for PyTorch named tensors (``*`` → ``___``)."""
if self.is_fused:
return _FUSED_NAME_SEP.join(self.sub_dims)
return self.name
class DimsSpec(_FrozenBase):
"""Parsed result of a full dims string like ``"b s h[tp] # dp:=moe_dp"``."""
dims: list[DimSpec]
dp_group_alias: Optional[str] = None
replicated_axes: frozenset[ParallelAxis] = frozenset()