Support unifying axis ordering in dump comparator (#19456)
This commit is contained in:
@@ -0,0 +1,57 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
|
||||
from sglang.srt.debug_utils.comparator.dims import parse_dims
|
||||
from sglang.srt.debug_utils.comparator.output_types import GeneralWarning
|
||||
from sglang.srt.debug_utils.comparator.utils import Pair, _FrozenBase
|
||||
from sglang.srt.debug_utils.comparator.warning_sink import warning_sink
|
||||
|
||||
# --- types ---
|
||||
|
||||
|
||||
class AxisSwapperPlan(_FrozenBase):
|
||||
pattern: str # einops pattern, e.g. "t h d -> t d h"
|
||||
|
||||
|
||||
# --- planner ---
|
||||
|
||||
|
||||
def compute_axis_swapper_plan(
|
||||
dims_str_pair: Pair[Optional[str]],
|
||||
) -> Optional[AxisSwapperPlan]:
|
||||
if dims_str_pair.x is None or dims_str_pair.y is None:
|
||||
return None
|
||||
|
||||
x_names: list[str] = [spec.name for spec in parse_dims(dims_str_pair.x)]
|
||||
y_names: list[str] = [spec.name for spec in parse_dims(dims_str_pair.y)]
|
||||
|
||||
if x_names == y_names:
|
||||
return None
|
||||
|
||||
if set(x_names) != set(y_names):
|
||||
warning_sink.add(
|
||||
GeneralWarning(
|
||||
category="axis_swapper_dim_mismatch",
|
||||
message=(
|
||||
f"AxisSwapper: dim name sets differ (x={x_names}, y={y_names}), "
|
||||
f"skipping axis swap"
|
||||
),
|
||||
)
|
||||
)
|
||||
return None
|
||||
|
||||
pattern: str = f"{' '.join(x_names)} -> {' '.join(y_names)}"
|
||||
return AxisSwapperPlan(pattern=pattern)
|
||||
|
||||
|
||||
# --- executor ---
|
||||
|
||||
|
||||
def execute_axis_swapper_plan(
|
||||
tensor: torch.Tensor, plan: AxisSwapperPlan
|
||||
) -> torch.Tensor:
|
||||
return rearrange(tensor, plan.pattern)
|
||||
@@ -5,6 +5,9 @@ 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.entrypoint.types import (
|
||||
AlignerPerStepPlan,
|
||||
AlignerPerStepSubPlan,
|
||||
@@ -63,6 +66,13 @@ 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:
|
||||
combined = Pair(
|
||||
x=execute_axis_swapper_plan(tensor=combined.x, plan=swap_plan),
|
||||
y=combined.y,
|
||||
)
|
||||
|
||||
return AlignerResult(tensors=combined, failed_side_xy=None)
|
||||
|
||||
|
||||
|
||||
@@ -2,6 +2,10 @@ 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.entrypoint.types import (
|
||||
AlignerPerStepPlan,
|
||||
AlignerPerStepSubPlan,
|
||||
@@ -33,12 +37,21 @@ def compute_aligner_plan(
|
||||
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
|
||||
)
|
||||
axis_swapper_plan: Optional[AxisSwapperPlan] = compute_axis_swapper_plan(
|
||||
dims_str_pair=dims_str_pair
|
||||
)
|
||||
|
||||
return AlignerPlan(
|
||||
per_step_plans=metas_pair.map(
|
||||
lambda metas: _compute_per_step_plans(metas=metas)
|
||||
),
|
||||
token_aligner_plan=token_aligner_plan,
|
||||
token_dims=token_dims,
|
||||
axis_swapper_plan=axis_swapper_plan,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -3,6 +3,7 @@ from __future__ import annotations
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Union
|
||||
|
||||
from sglang.srt.debug_utils.comparator.aligner.axis_swapper import AxisSwapperPlan
|
||||
from sglang.srt.debug_utils.comparator.aligner.reorderer.types import ReordererPlan
|
||||
from sglang.srt.debug_utils.comparator.aligner.token_aligner.types import (
|
||||
TokenAlignerPlan,
|
||||
@@ -25,3 +26,4 @@ class AlignerPlan:
|
||||
per_step_plans: Pair[list[AlignerPerStepPlan]]
|
||||
token_aligner_plan: Optional[TokenAlignerPlan]
|
||||
token_dims: Pair[int]
|
||||
axis_swapper_plan: Optional[AxisSwapperPlan] = None
|
||||
|
||||
Reference in New Issue
Block a user