Support loading token aligner data in dump comparator (#19376)

This commit is contained in:
fzyzcjy
2026-02-26 10:03:56 +08:00
committed by GitHub
parent e8dd14519d
commit d34d5aca07
10 changed files with 1182 additions and 43 deletions

View File

@@ -0,0 +1,224 @@
from __future__ import annotations
from pathlib import Path
from typing import Iterable, Optional, Tuple
import polars as pl
import torch
from sglang.srt.debug_utils.comparator.aligner.entrypoint.executor import (
execute_sub_plans,
)
from sglang.srt.debug_utils.comparator.aligner.entrypoint.planner import (
compute_per_step_sub_plans,
)
from sglang.srt.debug_utils.comparator.aligner.token_aligner.aux_plugins import (
AUX_NAMES,
_AuxFrameworkPlugin,
_plugins,
)
from sglang.srt.debug_utils.comparator.aligner.token_aligner.types import (
TokenAlignerGlobalAux,
TokenAlignerStepAux,
)
from sglang.srt.debug_utils.comparator.aligner.unsharder.parallel_info import (
normalize_parallel_info,
)
from sglang.srt.debug_utils.comparator.dims import ParallelAxis
from sglang.srt.debug_utils.comparator.output_types import GeneralWarning
from sglang.srt.debug_utils.comparator.warning_sink import warning_sink
from sglang.srt.debug_utils.dump_loader import ValueWithMeta, filter_rows
# re-export for existing callers
__all__ = ["AUX_NAMES", "has_aux_tensors", "load_and_normalize_aux"]
def load_and_normalize_aux(
dump_path: Path, df: pl.DataFrame
) -> Optional[TokenAlignerGlobalAux]:
"""Bootstrap: load, unshard, and normalize auxiliary tensors for one side."""
plugin: Optional[_AuxFrameworkPlugin] = _detect_plugin(df, dump_path=dump_path)
if plugin is None:
return None
available_names: set[str] = set(df["name"].unique().to_list()) & plugin.all_names
steps: list[int] = sorted(df["step"].unique().to_list())
assert len(steps) == 1, f"Multi-step not yet supported, got {len(steps)} steps"
tensor_names: set[str] = available_names & plugin.tensor_names
non_tensor_names: set[str] = available_names & plugin.non_tensor_names
steps_data: dict[int, dict[str, object]] = {}
for step in steps:
step_data = dict(
_load_step_data(
step=step,
tensor_names=tensor_names,
non_tensor_names=non_tensor_names,
df=df,
dump_path=dump_path,
plugin=plugin,
)
)
if step_data:
steps_data[step] = step_data
layout: str = plugin.detect_layout(steps_data)
step_auxs: dict[int, TokenAlignerStepAux] = {
step: plugin.compute_step_aux(step_data, layout=layout, step=step)
for step, step_data in steps_data.items()
}
return TokenAlignerGlobalAux(
step_auxs=step_auxs, framework=plugin.name, layout=layout
)
def has_aux_tensors(df: pl.DataFrame) -> bool:
"""Check if the DataFrame contains the minimum auxiliary tensors for alignment."""
names: set[str] = set(df["name"].unique().to_list())
return any(plugin.has_required_names(names) for plugin in _plugins)
def _detect_plugin(df: pl.DataFrame, dump_path: Path) -> Optional[_AuxFrameworkPlugin]:
names: set[str] = set(df["name"].unique().to_list())
for plugin in _plugins:
if names & plugin.discriminating_names:
return plugin
first_row: dict = df.row(0, named=True)
value: ValueWithMeta = ValueWithMeta.load(dump_path / first_row["filename"])
for plugin in _plugins:
if f"{plugin.name}_parallel_info" in value.meta:
return plugin
return None
def _load_step_data(
*,
step: int,
tensor_names: set[str],
non_tensor_names: set[str],
df: pl.DataFrame,
dump_path: Path,
plugin: _AuxFrameworkPlugin,
) -> Iterable[Tuple[str, object]]:
"""Load all tensor and non-tensor aux values for a single step."""
for name in non_tensor_names:
value = _load_non_tensor_aux(name=name, step=step, df=df, dump_path=dump_path)
if value is not None:
yield name, value
for name in tensor_names:
tensor = _load_and_align_aux_tensor(
name=name, step=step, df=df, dump_path=dump_path, plugin=plugin
)
if tensor is not None:
yield name, tensor
def _load_non_tensor_aux(
*, name: str, step: int, df: pl.DataFrame, dump_path: Path
) -> Optional[object]:
"""Load a non-tensor auxiliary value for a step, validating consistency across ranks."""
rows = filter_rows(df, conditions={"name": name, "step": step})
if not rows:
return None
loaded: list[ValueWithMeta] = [
ValueWithMeta.load(dump_path / r["filename"]) for r in rows
]
if len(loaded) > 1:
first_value = loaded[0].value
for i, item in enumerate(loaded[1:], start=1):
if item.value != first_value:
warning_sink.add(
GeneralWarning(
category=f"{name}_mismatch",
message=(
f"{name} mismatch across ranks: rank 0 has {first_value}, "
f"rank {i} has {item.value}"
),
)
)
break
return loaded[0].value
def _load_and_align_aux_tensor(
*,
name: str,
step: int,
df: pl.DataFrame,
dump_path: Path,
plugin: _AuxFrameworkPlugin,
) -> Optional[torch.Tensor]:
"""Load an auxiliary tensor for (name, step), align if needed."""
rows = filter_rows(df, conditions={"name": name, "step": step})
if not rows:
return None
loaded: list[ValueWithMeta] = [
ValueWithMeta.load(dump_path / r["filename"]) for r in rows
]
tensors: list[torch.Tensor] = [
item.value for item in loaded if isinstance(item.value, torch.Tensor)
]
if not tensors:
return None
if len(tensors) == 1:
return tensors[0]
metas: list[dict] = [item.meta for item in loaded]
metas = _ensure_dims_in_metas(name=name, plugin=plugin, metas=metas)
sub_plans = compute_per_step_sub_plans(metas=metas)
if sub_plans:
result = execute_sub_plans(tensors=tensors, plans=sub_plans)
assert result is not None
return result
warning_sink.add(
GeneralWarning(
category="aux_no_dims",
message=(
f"aux tensor '{name}' has {len(tensors)} ranks "
f"but no dims metadata, using rank 0 only"
),
)
)
return tensors[0]
def _ensure_dims_in_metas(
*, name: str, plugin: _AuxFrameworkPlugin, metas: list[dict]
) -> list[dict]:
"""Inject inferred dims into metas if not already present.
Returns metas unchanged if dims is already set, or a new list with dims
injected if inference succeeds. Raises if the tensor is CP-sharded
(not yet supported).
"""
if metas[0].get("dims") is not None:
return metas
parallel_infos = [normalize_parallel_info(m) for m in metas]
has_cp: bool = any(ParallelAxis.CP in info for info in parallel_infos)
if not has_cp:
return metas
if name in plugin.cp_sharded_names:
raise NotImplementedError(
f"Aux tensor '{name}' is CP-sharded but reorderer does not yet support "
f"zigzag reordering on the 't' dimension. "
f"Pass explicit dims= at dump time or wait for t-dim zigzag support."
)
return metas

View File

@@ -0,0 +1,222 @@
from __future__ import annotations
from abc import ABC, abstractmethod
import torch
from sglang.srt.debug_utils.comparator.aligner.token_aligner.types import (
PositionalSeqId,
SeqId,
SGLangSeqId,
TokenAlignerStepAux,
)
from sglang.srt.debug_utils.comparator.output_types import GeneralWarning
from sglang.srt.debug_utils.comparator.warning_sink import warning_sink
_BSHD_NOT_SUPPORTED_MSG: str = (
"BSHD layout is not currently supported. "
"Use aux_loader BSHD→THD conversion (planned)."
)
# ── plugin ABC ─────────────────────────────────────────────────────
class _AuxFrameworkPlugin(ABC):
@property
@abstractmethod
def name(self) -> str: ...
@property
@abstractmethod
def tensor_names(self) -> frozenset[str]: ...
@property
@abstractmethod
def non_tensor_names(self) -> frozenset[str]: ...
@property
def cp_sharded_names(self) -> frozenset[str]:
return frozenset()
@property
def discriminating_names(self) -> frozenset[str]:
"""Field names unique to this framework (excluding shared names like input_ids)."""
return frozenset()
@abstractmethod
def detect_layout(self, raw: dict[int, dict[str, object]]) -> str: ...
@abstractmethod
def compute_step_aux(
self, step_data: dict[str, object], *, layout: str, step: int
) -> TokenAlignerStepAux: ...
@abstractmethod
def has_required_names(self, names: set[str]) -> bool:
"""Whether the minimum set of aux names needed for alignment is present."""
...
@property
def all_names(self) -> frozenset[str]:
return self.tensor_names | self.non_tensor_names
# ── sglang plugin ─────────────────────────────────────────────────
class _SGLangPlugin(_AuxFrameworkPlugin):
@property
def name(self) -> str:
return "sglang"
@property
def tensor_names(self) -> frozenset[str]:
return frozenset({"input_ids", "positions", "seq_lens", "req_pool_indices"})
@property
def non_tensor_names(self) -> frozenset[str]:
return frozenset({"rids"})
@property
def cp_sharded_names(self) -> frozenset[str]:
return frozenset({"input_ids", "positions"})
@property
def discriminating_names(self) -> frozenset[str]:
return frozenset({"seq_lens", "positions", "req_pool_indices", "rids"})
def has_required_names(self, names: set[str]) -> bool:
return "input_ids" in names and "seq_lens" in names
def detect_layout(self, raw: dict[int, dict[str, object]]) -> str:
return "thd"
def compute_step_aux(
self, step_data: dict[str, object], *, layout: str, step: int
) -> TokenAlignerStepAux:
input_ids = step_data["input_ids"]
positions = step_data["positions"]
seq_lens = step_data["seq_lens"]
rids_raw = step_data.get("rids")
assert isinstance(
input_ids, torch.Tensor
), f"input_ids: expected Tensor, got {type(input_ids)}"
assert isinstance(
positions, torch.Tensor
), f"positions: expected Tensor, got {type(positions)}"
assert isinstance(
seq_lens, torch.Tensor
), f"seq_lens: expected Tensor, got {type(seq_lens)}"
seq_lens_list: list[int] = seq_lens.tolist()
num_seqs: int = len(seq_lens_list)
seq_ids: list[SeqId]
if rids_raw is not None and isinstance(rids_raw, (list, tuple)):
seq_ids = [SGLangSeqId(rid=str(r)) for r in rids_raw]
else:
seq_ids = [PositionalSeqId(step=step, seq_index=i) for i in range(num_seqs)]
return TokenAlignerStepAux(
input_ids=input_ids.tolist(),
positions=positions.tolist(),
seq_lens=seq_lens_list,
seq_ids=seq_ids,
)
# ── megatron plugin ───────────────────────────────────────────────
class _MegatronPlugin(_AuxFrameworkPlugin):
@property
def name(self) -> str:
return "megatron"
@property
def tensor_names(self) -> frozenset[str]:
return frozenset({"input_ids", "position_ids", "cu_seqlens_q", "cu_seqlens_kv"})
@property
def non_tensor_names(self) -> frozenset[str]:
return frozenset({"qkv_format"})
@property
def cp_sharded_names(self) -> frozenset[str]:
return frozenset({"input_ids", "position_ids"})
@property
def discriminating_names(self) -> frozenset[str]:
return frozenset({"cu_seqlens_q", "cu_seqlens_kv", "qkv_format"})
def has_required_names(self, names: set[str]) -> bool:
return "input_ids" in names and "cu_seqlens_q" in names
def detect_layout(self, raw: dict[int, dict[str, object]]) -> str:
for step_data in raw.values():
if (qkv_format := step_data.get("qkv_format")) is not None:
fmt = qkv_format if isinstance(qkv_format, str) else str(qkv_format)
if "bshd" in fmt.lower():
raise NotImplementedError(_BSHD_NOT_SUPPORTED_MSG)
return "thd"
input_ids = step_data.get("input_ids")
if isinstance(input_ids, torch.Tensor) and input_ids.ndim == 2:
raise NotImplementedError(_BSHD_NOT_SUPPORTED_MSG)
warning_sink.add(
GeneralWarning(
category="layout_detection_fallback",
message=(
"Megatron layout detection: no qkv_format or 2D input_ids found, "
"falling back to thd"
),
)
)
return "thd"
def compute_step_aux(
self, step_data: dict[str, object], *, layout: str, step: int
) -> TokenAlignerStepAux:
input_ids: torch.Tensor = step_data["input_ids"]
if (cu_seqlens_q := step_data.get("cu_seqlens_q")) is not None:
seq_lens: torch.Tensor = cu_seqlens_q[1:] - cu_seqlens_q[:-1]
else:
seq_lens = torch.tensor([input_ids.shape[0]], dtype=torch.long)
if (position_ids := step_data.get("position_ids")) is not None:
positions: torch.Tensor = position_ids
else:
positions = _infer_positions(seq_lens=seq_lens)
seq_lens_list: list[int] = seq_lens.tolist()
num_seqs: int = len(seq_lens_list)
seq_ids: list[SeqId] = [
PositionalSeqId(step=step, seq_index=seq_index)
for seq_index in range(num_seqs)
]
return TokenAlignerStepAux(
input_ids=input_ids.tolist(),
positions=positions.tolist(),
seq_lens=seq_lens_list,
seq_ids=seq_ids,
)
# ── plugin registry ───────────────────────────────────────────────
_plugins: list[_AuxFrameworkPlugin] = [_SGLangPlugin(), _MegatronPlugin()]
AUX_NAMES: frozenset[str] = frozenset().union(*(p.all_names for p in _plugins))
# ── helpers ────────────────────────────────────────────────────────
def _infer_positions(*, seq_lens: torch.Tensor) -> torch.Tensor:
"""Infer positions when position_ids is missing (THD only)."""
return torch.cat([torch.arange(int(slen.item())) for slen in seq_lens])

View File

@@ -0,0 +1,120 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import NamedTuple, Union
from pydantic import model_validator
from sglang.srt.debug_utils.comparator.utils import (
Pair,
_check_equal_lengths,
_FrozenBase,
)
class SGLangSeqId(NamedTuple):
rid: str
class PositionalSeqId(NamedTuple):
step: int
seq_index: int
SeqId = Union[SGLangSeqId, PositionalSeqId]
@dataclass(frozen=True)
class TokenAlignerStepAux:
"""Normalized auxiliary tensors for a single step (framework-agnostic)."""
input_ids: list[int] # [num_tokens]
positions: list[int] # [num_tokens]
seq_lens: list[int] # [num_seqs]
seq_ids: list[SeqId] # [num_seqs] — sequence identity
def __post_init__(self) -> None:
_check_equal_lengths(input_ids=self.input_ids, positions=self.positions)
_check_equal_lengths(seq_lens=self.seq_lens, seq_ids=self.seq_ids)
token_count: int = sum(self.seq_lens)
if token_count != len(self.input_ids):
raise ValueError(
f"sum(seq_lens)={token_count} != len(input_ids)={len(self.input_ids)}"
)
@dataclass(frozen=True)
class TokenAlignerGlobalAux:
"""Auxiliary tensors for one side across all steps + side-level metadata."""
step_auxs: dict[int, TokenAlignerStepAux]
framework: str # "sglang" | "megatron"
layout: str # "thd"
class TokenLocator(_FrozenBase):
"""Locates tokens within a single-step tensor.
token i is at tensor[token_index_in_step[i]].
"""
token_index_in_step: list[int]
def __add__(self, other: TokenLocator) -> TokenLocator:
return TokenLocator(
token_index_in_step=self.token_index_in_step + other.token_index_in_step,
)
class TokenAlignerSeqInfo(_FrozenBase):
"""Information for a sequence, containing information to locate all the tokens inside the sequence."""
# All these fields are of shape (num_tokens_in_seq,)
input_ids: list[int]
positions: list[int]
locator: TokenLocator
@model_validator(mode="after")
def _validate_fields(self) -> TokenAlignerSeqInfo:
n: int = len(self.input_ids)
_check_equal_lengths(
input_ids=self.input_ids,
positions=self.positions,
locator_token_index_in_step=self.locator.token_index_in_step,
)
if self.positions != list(range(n)):
raise ValueError(
f"positions must be [0, 1, ..., {n - 1}], got {self.positions}"
)
return self
def __add__(self, other: TokenAlignerSeqInfo) -> TokenAlignerSeqInfo:
return TokenAlignerSeqInfo(
input_ids=self.input_ids + other.input_ids,
positions=self.positions + other.positions,
locator=self.locator + other.locator,
)
class TokenAlignerSeqsInfo(_FrozenBase):
"""All sequences for one side across all steps."""
sequences: dict[SeqId, TokenAlignerSeqInfo]
layout: str
class TokenAlignerPlan(_FrozenBase):
"""Token alignment plan. locators.x[i] and locators.y[i] correspond to the same logical token."""
locators: Pair[TokenLocator]
@model_validator(mode="after")
def _validate_fields(self) -> TokenAlignerPlan:
_check_equal_lengths(
locators_x_token_index_in_step=self.locators.x.token_index_in_step,
locators_y_token_index_in_step=self.locators.y.token_index_in_step,
)
return self