Support TP unification and enhance tests in dump comparator (#19278)

This commit is contained in:
fzyzcjy
2026-02-25 09:42:56 +08:00
committed by GitHub
parent 39ba9b5ab5
commit 94ca2ac5d7
4 changed files with 490 additions and 156 deletions

View File

@@ -1,19 +1,17 @@
import argparse
from pathlib import Path
from typing import Optional
import polars as pl
import torch
from sglang.srt.debug_utils.comparator.output_types import (
ComparisonRecord,
ConfigRecord,
SkipRecord,
SummaryRecord,
print_record,
)
from sglang.srt.debug_utils.comparator.tensor_comparison import compare_tensors
from sglang.srt.debug_utils.dump_loader import ValueWithMeta, find_row, read_meta
from sglang.srt.debug_utils.comparator.pipeline import process_tensor_group
from sglang.srt.debug_utils.dump_loader import filter_rows, read_meta
_NON_KEY_COLS = {"dump_index", "filename", "duplicate_index"}
def main() -> None:
@@ -22,6 +20,8 @@ def main() -> None:
def run(args: argparse.Namespace) -> None:
df_baseline = read_meta(args.baseline_path)
df_target = read_meta(args.target_path)
df_target = df_target.filter(
(pl.col("step") >= args.start_step) & (pl.col("step") <= args.end_step)
@@ -30,8 +30,6 @@ def run(args: argparse.Namespace) -> None:
df_target = df_target.filter(pl.col("filename").str.contains(args.filter))
assert all(c in df_target.columns for c in ["rank", "step", "dump_index", "name"])
df_baseline = read_meta(args.baseline_path)
print_record(
ConfigRecord(
baseline_path=args.baseline_path,
@@ -44,60 +42,27 @@ def run(args: argparse.Namespace) -> None:
)
counts: dict[str, int] = {"passed": 0, "failed": 0, "skipped": 0}
grouping: str = args.grouping
for row in df_target.iter_rows(named=True):
path_target = Path(args.target_path) / row["filename"]
baseline_step = row["step"]
non_key_cols = _NON_KEY_COLS | ({"rank"} if grouping == "logical" else set())
key_cols = [c for c in df_target.columns if c not in non_key_cols]
tensor_group_keys = df_target.unique(subset=key_cols)
row_baseline = find_row(
df_baseline,
conditions=dict(
step=baseline_step,
**{
k: v
for k, v in row.items()
if k not in ["step", "dump_index", "filename"]
},
),
)
for tensor_group_key in tensor_group_keys.iter_rows(named=True):
conditions = {k: tensor_group_key[k] for k in key_cols}
baseline_rows = filter_rows(df_baseline, conditions=conditions)
target_rows = filter_rows(df_target, conditions=conditions)
if row_baseline is None:
counts["skipped"] += 1
print_record(
SkipRecord(name=row["name"], reason="no_baseline"),
output_format=args.output_format,
)
continue
path_baseline = Path(args.baseline_path) / row_baseline["filename"]
x_baseline = _load_tensor(path_baseline)
x_target = _load_tensor(path_target)
if x_baseline is None or x_target is None:
counts["skipped"] += 1
print_record(
SkipRecord(name=row["name"], reason="load_failed"),
output_format=args.output_format,
)
continue
info = compare_tensors(
x_baseline=x_baseline,
x_target=x_target,
name=row["name"],
record = process_tensor_group(
name=tensor_group_key["name"],
baseline_filenames=[r["filename"] for r in baseline_rows],
target_filenames=[r["filename"] for r in target_rows],
baseline_path=Path(args.baseline_path),
target_path=Path(args.target_path),
diff_threshold=args.diff_threshold,
)
if info.diff is not None and info.diff.passed:
counts["passed"] += 1
else:
counts["failed"] += 1
print_record(
ComparisonRecord(**info.model_dump()),
output_format=args.output_format,
)
counts[record.category] += 1
print_record(record, output_format=args.output_format)
print_record(
SummaryRecord(total=sum(counts.values()), **counts),
@@ -105,15 +70,7 @@ def run(args: argparse.Namespace) -> None:
)
def _load_tensor(path: Path) -> Optional[torch.Tensor]:
loaded = ValueWithMeta.load(path)
if not isinstance(loaded.value, torch.Tensor):
return None
return loaded.value
def _parse_args() -> argparse.Namespace:
# python -m sglang.srt.debug_utils.comparator --baseline-path ... --target-path ...
parser = argparse.ArgumentParser()
parser.add_argument("--baseline-path", type=str)
parser.add_argument("--target-path", type=str)
@@ -130,4 +87,11 @@ def _parse_args() -> argparse.Namespace:
default="text",
help="Output format: text (default) or json (JSONL, one JSON object per line)",
)
parser.add_argument(
"--grouping",
type=str,
choices=["logical", "raw"],
default="logical",
help="Grouping mode: logical (cross-rank unshard) or raw (rank-by-rank)",
)
return parser.parse_args()

View File

@@ -38,6 +38,10 @@ class SkipRecord(_OutputRecord):
name: str
reason: str
@property
def category(self):
return "skipped"
def to_text(self) -> str:
return f"Skip: {self.name} ({self.reason})"
@@ -45,6 +49,10 @@ class SkipRecord(_OutputRecord):
class ComparisonRecord(TensorComparisonInfo, _OutputRecord):
type: Literal["comparison"] = "comparison"
@property
def category(self):
return "passed" if self.diff is not None and self.diff.passed else "failed"
def to_text(self) -> str:
return format_comparison(self)

View File

@@ -0,0 +1,119 @@
from pathlib import Path
from typing import Any, Optional
import torch
from sglang.srt.debug_utils.comparator.dims import parse_dims
from sglang.srt.debug_utils.comparator.output_types import (
ComparisonRecord,
SkipRecord,
)
from sglang.srt.debug_utils.comparator.tensor_comparison.compare import compare_tensors
from sglang.srt.debug_utils.comparator.unshard.executor import execute_unshard_plan
from sglang.srt.debug_utils.comparator.unshard.parallel_info import (
normalize_parallel_info,
)
from sglang.srt.debug_utils.comparator.unshard.planner import compute_unshard_plan
from sglang.srt.debug_utils.comparator.unshard.types import Plan, UnshardPlan
from sglang.srt.debug_utils.dump_loader import ValueWithMeta
def process_tensor_group(
*,
name: str,
baseline_filenames: list[str],
target_filenames: list[str],
baseline_path: Path,
target_path: Path,
diff_threshold: float,
) -> ComparisonRecord | SkipRecord:
b_tensors = _load_tensors(baseline_filenames, baseline_path)
t_tensors = _load_tensors(target_filenames, target_path)
b_plans, t_plans = _compute_plans(
baseline_metas=[item.meta for item in b_tensors],
target_metas=[item.meta for item in t_tensors],
)
b_extracted = _extract_tensors(b_tensors)
t_extracted = _extract_tensors(t_tensors)
del b_tensors, t_tensors
b_tensor = _execute_plans(b_extracted, b_plans)
t_tensor = _execute_plans(t_extracted, t_plans)
if b_tensor is None or t_tensor is None:
reason = "baseline_load_failed" if b_tensor is None else "target_load_failed"
return SkipRecord(name=name, reason=reason)
info = compare_tensors(
x_baseline=b_tensor,
x_target=t_tensor,
name=name,
diff_threshold=diff_threshold,
)
return ComparisonRecord(**info.model_dump())
def _load_tensors(filenames: list[str], base_path: Path) -> list[ValueWithMeta]:
return [ValueWithMeta.load(base_path / f) for f in filenames]
def _compute_plans(
*,
baseline_metas: list[dict[str, Any]],
target_metas: list[dict[str, Any]],
) -> tuple[list[Plan], list[Plan]]:
"""This function deliberately takes metadata, since plan computation must never inspect actual tensor data."""
return (
_compute_plans_for_group(baseline_metas),
_compute_plans_for_group(target_metas),
)
def _compute_plans_for_group(metas: list[dict[str, Any]]) -> list[Plan]:
if not metas or len(metas) == 1:
return []
dims_str = metas[0].get("dims")
if dims_str is None:
return []
dim_specs = parse_dims(dims_str)
parallel_infos = [normalize_parallel_info(meta) for meta in metas]
plan = compute_unshard_plan(dim_specs=dim_specs, parallel_infos=parallel_infos)
return [plan] if plan is not None else []
def _extract_tensors(
loaded: list[ValueWithMeta],
) -> Optional[list[torch.Tensor]]:
return [value for item in loaded if isinstance(value := item.value, torch.Tensor)]
def _execute_plans(
tensors: list[torch.Tensor],
plans: list[Plan],
) -> Optional[torch.Tensor]:
if not tensors:
return None
if not plans:
if len(tensors) != 1:
return None
return tensors[0]
assert len(plans) <= 1, "multi-plan not supported yet"
for plan in plans:
if isinstance(plan, UnshardPlan):
# TODO: incorrect `tensors_by_world_rank` if multi UnshardPlan
tensors = execute_unshard_plan(
plan, tensors_by_world_rank=dict(enumerate(tensors))
)
else:
raise NotImplementedError(f"Unknown {plan=}")
return tensors