Use named tensors in dump comparator (#19458)

This commit is contained in:
fzyzcjy
2026-02-27 08:09:55 +08:00
committed by GitHub
parent eb0e905fc3
commit e1e0cfd856
18 changed files with 248 additions and 96 deletions

View File

@@ -54,4 +54,4 @@ def compute_axis_swapper_plan(
def execute_axis_swapper_plan(
tensor: torch.Tensor, plan: AxisSwapperPlan
) -> torch.Tensor:
return rearrange(tensor, plan.pattern)
return rearrange(tensor.rename(None), plan.pattern)

View File

@@ -36,8 +36,6 @@ def compute_aligner_plan(
metas_pair: Pair[list[dict[str, Any]]],
token_aligner_plan: Optional[TokenAlignerPlan],
) -> AlignerPlan:
token_dims: Pair[int] = metas_pair.map(_compute_token_dim)
dims_str_pair: Pair[Optional[str]] = metas_pair.map(
lambda metas: metas[0].get("dims") if metas else None
)
@@ -50,7 +48,6 @@ def compute_aligner_plan(
lambda metas: _compute_per_step_plans(metas=metas)
),
token_aligner_plan=token_aligner_plan,
token_dims=token_dims,
axis_swapper_plan=axis_swapper_plan,
)

View File

@@ -25,5 +25,4 @@ class AlignerPerStepPlan:
class AlignerPlan:
per_step_plans: Pair[list[AlignerPerStepPlan]]
token_aligner_plan: Optional[TokenAlignerPlan]
token_dims: Pair[int]
axis_swapper_plan: Optional[AxisSwapperPlan] = None

View File

@@ -1,16 +1,19 @@
import torch
from sglang.srt.debug_utils.comparator.aligner.reorderer.types import ReordererPlan
from sglang.srt.debug_utils.comparator.dims import (
resolve_dim_by_name,
strip_dim_names,
)
def execute_reorderer_plan(
plan: ReordererPlan,
tensors: list[torch.Tensor],
) -> list[torch.Tensor]:
dim: int = resolve_dim_by_name(tensors[0], plan.params.dim_name)
return [
_reorder_zigzag_to_natural(
tensor, dim=plan.params.dim, cp_size=plan.params.cp_size
)
_reorder_zigzag_to_natural(tensor, dim=dim, cp_size=plan.params.cp_size)
for tensor in tensors
]
@@ -23,9 +26,16 @@ def _reorder_zigzag_to_natural(
Generalized from Megatron-LM _undo_attention_load_balancing
(megatron/core/ssm/mamba_context_parallel.py:360-373).
"""
stripped: torch.Tensor = strip_dim_names(tensor)
names: tuple = tensor.names
num_chunks: int = cp_size * 2
chunks: tuple[torch.Tensor, ...] = tensor.chunk(num_chunks, dim=dim)
chunks: tuple[torch.Tensor, ...] = stripped.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)
result: torch.Tensor = torch.cat([chunks[i] for i in order], dim=dim)
if names[0] is not None:
result = result.refine_names(*names)
return result

View File

@@ -19,7 +19,7 @@ def compute_reorderer_plans(
) -> list[ReordererPlan]:
plans: list[ReordererPlan] = []
for dim_index, spec in enumerate(dim_specs):
for spec in dim_specs:
if (
spec.ordering is not None
and spec.ordering != Ordering.NATURAL
@@ -37,7 +37,7 @@ def compute_reorderer_plans(
axis_size: int = parallel_infos[0][spec.parallel].axis_size
plans.append(
ReordererPlan(
params=ZigzagToNaturalParams(dim=dim_index, cp_size=axis_size),
params=ZigzagToNaturalParams(dim_name=spec.name, cp_size=axis_size),
)
)

View File

@@ -5,7 +5,7 @@ from sglang.srt.debug_utils.comparator.utils import _FrozenBase
class ZigzagToNaturalParams(_FrozenBase):
op: Literal["zigzag_to_natural"] = "zigzag_to_natural"
dim: int
dim_name: str
cp_size: int

View File

@@ -6,8 +6,21 @@ from sglang.srt.debug_utils.comparator.aligner.token_aligner.types import (
TokenAlignerPlan,
TokenLocator,
)
from sglang.srt.debug_utils.comparator.dims import (
TOKEN_DIM_NAME,
resolve_dim_by_name,
strip_dim_names,
)
from sglang.srt.debug_utils.comparator.utils import Pair
_UNNAMED_TOKEN_DIM_FALLBACK: int = 0
def _resolve_dim_or_fallback(tensor: torch.Tensor, name: str) -> int:
if tensor.names[0] is None:
return _UNNAMED_TOKEN_DIM_FALLBACK
return resolve_dim_by_name(tensor, name)
def execute_token_aligner(
plan: TokenAlignerPlan,
@@ -17,8 +30,8 @@ def execute_token_aligner(
) -> Pair[torch.Tensor]:
if not plan.locators.x.steps:
return Pair(
x=_make_empty(tensor_of_step=tensor_of_step_pair.x, token_dim=token_dims.x),
y=_make_empty(tensor_of_step=tensor_of_step_pair.y, token_dim=token_dims.y),
x=_make_empty(tensor_of_step=tensor_of_step_pair.x),
y=_make_empty(tensor_of_step=tensor_of_step_pair.y),
)
return Pair(
@@ -36,9 +49,11 @@ def execute_token_aligner(
def _make_empty(
*, tensor_of_step: dict[int, torch.Tensor], token_dim: int
*,
tensor_of_step: dict[int, torch.Tensor],
) -> torch.Tensor:
dummy: torch.Tensor = next(iter(tensor_of_step.values()))
token_dim: int = _resolve_dim_or_fallback(dummy, TOKEN_DIM_NAME)
shape: list[int] = list(dummy.shape)
shape[token_dim] = 0
return torch.empty(shape, dtype=dummy.dtype)
@@ -50,8 +65,11 @@ def _extract_and_stack_tokens(
locator: TokenLocator,
token_dim: int,
) -> torch.Tensor:
some_tensor: torch.Tensor = next(iter(tensor_of_step.values()))
token_dim: int = _resolve_dim_or_fallback(some_tensor, TOKEN_DIM_NAME)
tokens: list[torch.Tensor] = [
tensor_of_step[s].select(dim=token_dim, index=i)
strip_dim_names(tensor_of_step[s]).select(dim=token_dim, index=i)
for s, i in zip(locator.steps, locator.token_index_in_step)
]
return torch.stack(tokens, dim=token_dim)

View File

@@ -6,7 +6,7 @@ from sglang.srt.debug_utils.comparator.aligner.unsharder.types import (
UnsharderParams,
UnsharderPlan,
)
from sglang.srt.debug_utils.comparator.dims import ParallelAxis
from sglang.srt.debug_utils.comparator.dims import ParallelAxis, resolve_dim_by_name
from sglang.srt.debug_utils.comparator.output_types import ReplicatedMismatchWarning
from sglang.srt.debug_utils.comparator.warning_sink import warning_sink
@@ -46,7 +46,8 @@ def _apply_unshard(
return ordered_tensors[0]
if isinstance(params, ConcatParams):
return torch.cat(ordered_tensors, dim=params.dim)
dim: int = resolve_dim_by_name(ordered_tensors[0], params.dim_name)
return torch.cat(ordered_tensors, dim=dim)
# Phase 2: ReduceSumParams, CpZigzagParams
raise ValueError(f"Unsupported unshard operation: {type(params).__name__}")
@@ -58,10 +59,10 @@ def _verify_replicated_group(
axis: ParallelAxis,
group_index: int,
) -> None:
baseline = ordered_tensors[0]
baseline = ordered_tensors[0].rename(None)
for i in range(1, len(ordered_tensors)):
other = ordered_tensors[i]
other = ordered_tensors[i].rename(None)
if not torch.allclose(baseline, other, atol=1e-6):
warning_sink.add(
ReplicatedMismatchWarning(

View File

@@ -28,10 +28,8 @@ def compute_unsharder_plan(
if not parallel_infos:
raise ValueError("parallel_infos must not be empty")
sharded_axis_infos: dict[ParallelAxis, tuple[int, DimSpec]] = {
spec.parallel: (dim_idx, spec)
for dim_idx, spec in enumerate(dim_specs)
if spec.parallel is not None
sharded_axis_infos: dict[ParallelAxis, DimSpec] = {
spec.parallel: spec for spec in dim_specs if spec.parallel is not None
}
sharded_axes: set[ParallelAxis] = set(sharded_axis_infos)
@@ -54,8 +52,8 @@ def compute_unsharder_plan(
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()
(axis, _resolve_unshard_params(spec=spec))
for axis, spec in sharded_axis_infos.items()
]
plans: list[UnsharderPlan] = []
@@ -130,9 +128,9 @@ def _group_and_project(
return _GroupResult(groups=groups, projected_coords=projected)
def _resolve_unshard_params(*, spec: DimSpec, dim_index: int) -> UnsharderParams:
def _resolve_unshard_params(*, spec: DimSpec) -> UnsharderParams:
if spec.reduction is not None:
raise NotImplementedError(
f"Unshard for reduction={spec.reduction} not yet implemented (Phase 2)"
)
return ConcatParams(dim=dim_index)
return ConcatParams(dim_name=spec.name)

View File

@@ -25,7 +25,7 @@ class AxisInfo(_FrozenBase):
class ConcatParams(_FrozenBase):
op: Literal["concat"] = "concat"
dim: int
dim_name: str
class PickParams(_FrozenBase):

View File

@@ -18,6 +18,7 @@ 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.output_types import (
ComparisonRecord,
SkipRecord,
@@ -81,9 +82,16 @@ def _compare_bundle_pair_raw(
metas_pair=metas_pair, token_aligner_plan=token_aligner_plan
)
# 3. Execute (tensor + plan only, no meta)
tensors_pair: Pair[list[torch.Tensor]] = valid_pair.map(
lambda items: [it.value for it in items]
# 3. Apply dim names to tensors, then execute
tensors_pair: Pair[list[torch.Tensor]] = Pair(
x=_apply_dim_names_from_meta(
tensors=[it.value for it in valid_pair.x],
metas=metas_pair.x,
),
y=_apply_dim_names_from_meta(
tensors=[it.value for it in valid_pair.y],
metas=metas_pair.y,
),
)
aligner_result: AlignerResult = execute_aligner_plan(
tensors_pair=tensors_pair, plan=plan
@@ -97,14 +105,30 @@ def _compare_bundle_pair_raw(
# 4. Compare
info = compare_tensor_pair(
x_baseline=aligner_result.tensors.x,
x_target=aligner_result.tensors.y,
x_baseline=aligner_result.tensors.x.rename(None),
x_target=aligner_result.tensors.y.rename(None),
name=name,
diff_threshold=diff_threshold,
)
return ComparisonRecord(**info.model_dump())
def _apply_dim_names_from_meta(
*,
tensors: list[torch.Tensor],
metas: list[dict[str, Any]],
) -> list[torch.Tensor]:
if not metas:
return tensors
dims_str: Optional[str] = metas[0].get("dims")
if dims_str is None:
return tensors
dim_names: list[str] = parse_dim_names(dims_str)
return [apply_dim_names(t, dim_names) for t in tensors]
def _load_valid_tensors(filenames: list[str], base_path: Path) -> list[ValueWithMeta]:
return [
x

View File

@@ -3,6 +3,8 @@ from dataclasses import dataclass
from enum import Enum
from typing import Optional
import torch
TOKEN_DIM_NAME: str = "t"
BATCH_DIM_NAME: str = "b"
SEQ_DIM_NAME: str = "s"
@@ -89,6 +91,29 @@ def parse_dims(dims_str: str) -> list[DimSpec]:
return result
def parse_dim_names(dims_str: str) -> list[str]:
return [spec.name for spec in parse_dims(dims_str)]
def find_dim_index(dim_specs: list[DimSpec], name: str) -> Optional[int]:
names: list[str] = [spec.name for spec in dim_specs]
return names.index(name) if name in names else 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:
return tensor.refine_names(*dim_names)
def strip_dim_names(tensor: torch.Tensor) -> torch.Tensor:
return tensor.rename(None)