Add e2e demo test in dump comparator (#21031)

This commit is contained in:
fzyzcjy
2026-03-20 22:41:01 +08:00
committed by GitHub
parent 6703cc4484
commit 146700db68
7 changed files with 372 additions and 27 deletions

View File

@@ -261,24 +261,6 @@ def _compute_dependent_axes(
)
def _is_dependent_axis(
parallel_infos: list[dict[ParallelAxis, AxisInfo]],
*,
parent: ParallelAxis,
child: ParallelAxis,
) -> bool:
"""True if child's rank is uniquely determined by parent's rank."""
parent_rank_to_child_rank: dict[int, int] = {}
for info in parallel_infos:
if parent not in info or child not in info:
continue
parent_rank = info[parent].axis_rank
child_rank = info[child].axis_rank
if parent_rank_to_child_rank.setdefault(parent_rank, child_rank) != child_rank:
return False
return True
def _is_jointly_determined(
parallel_infos: list[dict[ParallelAxis, AxisInfo]],
*,
@@ -320,6 +302,24 @@ def _is_jointly_determined(
return bool(mapping)
def _is_dependent_axis(
parallel_infos: list[dict[ParallelAxis, AxisInfo]],
*,
parent: ParallelAxis,
child: ParallelAxis,
) -> bool:
"""True if child's rank is uniquely determined by parent's rank."""
parent_rank_to_child_rank: dict[int, int] = {}
for info in parallel_infos:
if parent not in info or child not in info:
continue
parent_rank = info[parent].axis_rank
child_rank = info[child].axis_rank
if parent_rank_to_child_rank.setdefault(parent_rank, child_rank) != child_rank:
return False
return True
def _group_and_project(
*,
current_coords: _CoordsList,

View File

@@ -1,3 +1,5 @@
import __future__
import importlib
import inspect
import textwrap
@@ -85,7 +87,12 @@ def patch_function(
if preamble.strip():
modified_source = _insert_preamble(source=modified_source, preamble=preamble)
code: types.CodeType = compile(modified_source, inspect.getfile(target), "exec")
code: types.CodeType = compile(
modified_source,
inspect.getfile(target),
"exec",
flags=__future__.annotations.compiler_flag,
)
temp_namespace: dict[str, Any] = {}
exec(code, target.__globals__, temp_namespace)

View File

@@ -1,3 +1,5 @@
import sys
import pytest
from sglang.srt.debug_utils.comparator.aligner.unsharder.planner import (

View File

@@ -17,6 +17,79 @@ from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=20, suite="stage-a-cpu-only", nightly=True)
class TestComputeTensorInfo:
def test_basic_tensor_returns_correct_shape_and_dtype(self) -> None:
tensor = torch.randn(2, 3)
info = compute_tensor_info(tensor)
assert info.shape == [2, 3]
assert info.dtype == "torch.float32"
assert info.stats.mean == pytest.approx(tensor.float().mean().item(), abs=1e-4)
def test_include_sample_false_returns_none_sample(self) -> None:
tensor = torch.randn(2, 3)
info = compute_tensor_info(tensor, include_sample=False)
assert info.sample is None
def test_include_sample_true_returns_string_sample(self) -> None:
tensor = torch.randn(2, 3)
info = compute_tensor_info(tensor, include_sample=True)
assert info.sample is not None
assert isinstance(info.sample, str)
def test_empty_tensor_stats_are_zero(self) -> None:
tensor = torch.tensor([])
info = compute_tensor_info(tensor)
assert info.stats.mean == 0.0
assert info.stats.std == 0.0
assert info.shape == [0]
def test_integer_tensor_converted_to_float_for_stats(self) -> None:
"""Integer tensors should be cast to float internally for stats computation."""
tensor = torch.tensor([1, 2, 3, 4], dtype=torch.int32)
info = compute_tensor_info(tensor)
assert info.dtype == "torch.int32"
assert info.stats.mean == pytest.approx(2.5, abs=1e-4)
assert info.stats.min == pytest.approx(1.0, abs=1e-4)
assert info.stats.max == pytest.approx(4.0, abs=1e-4)
def test_bfloat16_tensor_shape_and_stats(self) -> None:
"""bfloat16 tensors produce correct shape and dtype string."""
tensor = torch.ones(3, 4, dtype=torch.bfloat16)
info = compute_tensor_info(tensor)
assert info.shape == [3, 4]
assert info.dtype == "torch.bfloat16"
assert info.stats.mean == pytest.approx(1.0, abs=1e-2)
def test_multidimensional_shape(self) -> None:
"""Shape is preserved for high-rank tensors."""
tensor = torch.randn(2, 3, 4, 5)
info = compute_tensor_info(tensor)
assert info.shape == [2, 3, 4, 5]
def test_scalar_tensor(self) -> None:
"""Scalar (0-dim) tensor produces empty shape list."""
tensor = torch.tensor(3.14)
info = compute_tensor_info(tensor)
assert info.shape == []
assert info.stats.mean == pytest.approx(3.14, abs=1e-4)
assert info.stats.min == pytest.approx(3.14, abs=1e-4)
assert info.stats.max == pytest.approx(3.14, abs=1e-4)
def test_include_sample_true_contains_tensor_representation(self) -> None:
"""Sample string should contain some recognizable tensor content."""
tensor = torch.tensor([1.0, 2.0])
info = compute_tensor_info(tensor, include_sample=True)
assert info.sample is not None
assert "1." in info.sample or "2." in info.sample
def test_percentiles_present_for_small_tensor(self) -> None:
"""Small tensors (< threshold) should have percentile data."""
tensor = torch.randn(100)
info = compute_tensor_info(tensor)
assert len(info.stats.percentiles) > 0
assert 50 in info.stats.percentiles
class TestComputeTensorInfo:
def test_basic_tensor_returns_correct_shape_and_dtype(self) -> None:
tensor = torch.randn(2, 3)

View File

@@ -68,6 +68,14 @@ _DEFAULT_PERCENTILE_LINES: list[str] = [
" [blue]p99 [/] 1.8000 1.8000 [dim]+0.00e+00[/]",
]
_DEFAULT_PERCENTILE_LINES: list[str] = [
" [blue]p1 [/] -1.8000 -1.8000 [dim]+0.00e+00[/]",
" [blue]p5 [/] -1.5000 -1.5000 [dim]+0.00e+00[/]",
" [blue]p50 [/] 0.0000 0.0000 [dim]+0.00e+00[/]",
" [blue]p95 [/] 1.5000 1.5000 [dim]+0.00e+00[/]",
" [blue]p99 [/] 1.8000 1.8000 [dim]+0.00e+00[/]",
]
# Snapshot strings below are intentionally spelled out in full per test.
# The shared skeleton (stats block, diff block) looks duplicated, but keeping

View File

@@ -0,0 +1,249 @@
"""Minimal demo: run the comparator on synthetic data and print its output.
This is NOT a correctness test suite.
The sole purpose is to let a new user run ``pytest -s test_e2e_demo.py``
and immediately see what comparator text output looks like (passed, failed,
skipped in one shot). Correctness is verified via the JSONL report file.
"""
from __future__ import annotations
import sys
from pathlib import Path
from typing import Dict, List, Optional
import pytest
import torch
import sglang.srt.debug_utils.dumper as _dumper_module
from sglang.srt.debug_utils.comparator.entrypoint import parse_args, run
from sglang.srt.debug_utils.comparator.output_types import (
AnyRecord,
ComparisonErrorRecord,
SummaryRecord,
parse_record_json,
)
from sglang.srt.debug_utils.dumper import DumperConfig, _Dumper
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=10, suite="default", nightly=True)
_EXP_NAME = "demo_exp"
# This file has exactly ONE test. All demo scenarios go here — do not add separate tests.
def test_demo(tmp_path: Path) -> None:
"""Passed + failed + skipped + sharded + errored in a single demo file."""
torch.manual_seed(0)
good_tensor = torch.randn(4, 8)
sharded_full = torch.randn(2, 8, 16)
baseline_dir = tmp_path / "baseline"
target_dir = tmp_path / "target"
baseline_dir.mkdir()
target_dir.mkdir()
# Step 1: simple tensors (single rank, no parallelism)
_dump_single(baseline_dir, name="my_good_tensor", tensor=good_tensor)
_dump_single(baseline_dir, name="my_bad_tensor", tensor=torch.randn(4, 8))
_dump_single(
target_dir, name="my_good_tensor", tensor=good_tensor + torch.randn(4, 8) * 1e-5
)
_dump_single(target_dir, name="my_bad_tensor", tensor=torch.randn(4, 8) * 100)
_dump_single(target_dir, name="my_orphan_tensor", tensor=torch.randn(4, 8))
# Step 2: sharded tensor (BSHD) — baseline: TP=2 on h, target: CP=2 zigzag + SP=2 on s
sharded_target = sharded_full + torch.randn_like(sharded_full) * 1e-5
_dump_tp_sharded(
baseline_dir, name="my_sharded_tensor", full_tensor=sharded_full, tp_size=2
)
_dump_cp_zigzag_sp_sharded(
target_dir,
name="my_sharded_tensor",
full_tensor=sharded_target,
cp_size=2,
sp_size=2,
)
# Step 3: bad dims — target says h[cp] but parallel_info has tp → undeclared axis error
bad_dims_tensor = torch.randn(2, 8, 16)
for tp_rank, shard in enumerate(bad_dims_tensor.chunk(2, dim=-1)):
_dump_rank(
baseline_dir,
rank=tp_rank,
name="my_bad_dims_tensor",
tensor=shard,
dims="b s h[tp]",
parallel_info={"tp_rank": tp_rank, "tp_size": 2},
)
_dump_rank(
target_dir,
rank=tp_rank,
name="my_bad_dims_tensor",
tensor=shard,
dims="b s h[cp]",
parallel_info={"tp_rank": tp_rank, "tp_size": 2},
)
baseline_exp = baseline_dir / _EXP_NAME
target_exp = target_dir / _EXP_NAME
# Step 4: run normal, then verbose
for verbosity in ("normal", "verbose"):
report_path = tmp_path / f"report_{verbosity}.jsonl"
_run(
baseline_exp,
target_exp,
report_path=report_path,
output_format="text",
verbosity=verbosity,
)
_assert_summary(report_path, passed=2, failed=1, skipped=1, errored=1)
# Step 5: verify error record content
records = _read_report(tmp_path / "report_verbose.jsonl")
errors = [r for r in records if isinstance(r, ComparisonErrorRecord)]
assert len(errors) == 1
assert "tp" in errors[0].exception_message
assert "--override-dims" in errors[0].traceback_str
# ── Helpers ──────────────────────────────────────────────────────────
def _assert_summary(
report_path: Path, *, passed: int, failed: int, skipped: int, errored: int = 0
) -> None:
records = _read_report(report_path)
summary = next(r for r in records if isinstance(r, SummaryRecord))
assert summary.passed == passed
assert summary.failed == failed
assert summary.skipped == skipped
assert summary.errored == errored
def _dump_single(directory: Path, *, name: str, tensor: torch.Tensor) -> None:
_dump_rank(directory, rank=0, name=name, tensor=tensor)
def _dump_tp_sharded(
directory: Path,
*,
name: str,
full_tensor: torch.Tensor,
tp_size: int,
) -> None:
"""Dump TP-sharded tensor: dims="b s h[tp]", shard along last dim."""
shards = list(full_tensor.chunk(tp_size, dim=-1))
for tp_rank, shard in enumerate(shards):
_dump_rank(
directory,
rank=tp_rank,
name=name,
tensor=shard,
dims="b s h[tp]",
parallel_info={"tp_rank": tp_rank, "tp_size": tp_size},
)
def _dump_cp_zigzag_sp_sharded(
directory: Path,
*,
name: str,
full_tensor: torch.Tensor,
cp_size: int,
sp_size: int,
) -> None:
"""Dump CP-zigzag+SP sharded tensor: dims="b s[cp:zigzag,sp] h", shard seq dim."""
seq_dim = 1
num_chunks = cp_size * 2
natural_chunks = list(full_tensor.chunk(num_chunks, dim=seq_dim))
zigzag_order: List[int] = []
for i in range(cp_size):
zigzag_order.append(i)
zigzag_order.append(num_chunks - 1 - i)
zigzagged = torch.cat([natural_chunks[idx] for idx in zigzag_order], dim=seq_dim)
cp_chunks = list(zigzagged.chunk(cp_size, dim=seq_dim))
rank = 0
for cp_rank in range(cp_size):
sp_chunks = list(cp_chunks[cp_rank].chunk(sp_size, dim=seq_dim))
for sp_rank in range(sp_size):
_dump_rank(
directory,
rank=rank,
name=name,
tensor=sp_chunks[sp_rank],
dims="b s[cp:zigzag,sp] h",
parallel_info={
"cp_rank": cp_rank,
"cp_size": cp_size,
"sp_rank": sp_rank,
"sp_size": sp_size,
},
)
rank += 1
def _dump_rank(
directory: Path,
*,
rank: int,
name: str,
tensor: torch.Tensor,
dims: Optional[str] = None,
parallel_info: Optional[Dict[str, int]] = None,
) -> None:
with pytest.MonkeyPatch.context() as mp:
mp.setattr(_dumper_module, "_get_rank", lambda: rank)
dumper = _Dumper(
config=DumperConfig(enable=True, dir=str(directory), exp_name=_EXP_NAME)
)
static_meta: Dict[str, object] = {"world_rank": rank, "world_size": 1}
if parallel_info is not None:
static_meta["sglang_parallel_info"] = parallel_info
dumper.__dict__["_static_meta"] = static_meta
dumper.dump(name, tensor, dims=dims)
dumper.step()
def _run(
baseline_path: Path,
target_path: Path,
*,
report_path: Path,
output_format: str = "text",
verbosity: str = "normal",
) -> int:
argv = [
"--baseline-path",
str(baseline_path),
"--target-path",
str(target_path),
"--output-format",
output_format,
"--verbosity",
verbosity,
"--preset",
"sglang_dev",
"--report-path",
str(report_path),
]
print(
f"\n $ python -m sglang.srt.debug_utils.comparator {' '.join(argv)}\n",
flush=True,
)
return run(parse_args(argv))
def _read_report(report_path: Path) -> List[AnyRecord]:
return [
parse_record_json(line) for line in report_path.read_text().strip().splitlines()
]
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-s", "-v"]))

View File

@@ -77,7 +77,7 @@ patches:
hidden_states=hidden_states,
forward_batch=forward_batch,
)
append: "dumper.dump('attn_output', hidden_states, dims='t h[tp:partial]')"
append: "dumper.dump('attn_output', hidden_states, dims='t h[attn_tp:partial] # tp:replicated')"
- match: |
hidden_states, residual = self.layer_communicator.prepare_mlp(
hidden_states, residual, forward_batch
@@ -87,13 +87,13 @@ patches:
hidden_states = self.mlp(
hidden_states, forward_batch, should_allreduce_fusion, use_reduce_scatter
)
append: "dumper.dump('mlp_output', hidden_states, dims='t h[tp:partial]')"
append: "dumper.dump('mlp_output', hidden_states, dims='t h[moe_tp:partial] # tp:replicated')"
# --- attention internals ---
- target: sglang.srt.models.qwen3_moe.Qwen3MoeAttention.forward_core
edits:
- match: "output, _ = self.o_proj(attn_output)"
prepend: "dumper.dump('attn_pre_o_proj', attn_output, dims='t attn_h[tp]')"
prepend: "dumper.dump('attn_pre_o_proj', attn_output, dims='t attn_h[attn_tp] # tp:replicated')"
# --- moe internals ---
- target: sglang.srt.models.qwen3_moe.Qwen3MoeSparseMoeBlock.forward_normal
@@ -101,14 +101,19 @@ patches:
- match: "router_logits, _ = self.gate(hidden_states)"
append: "dumper.dump('moe_router_logits', router_logits, dims='t num_experts # tp:replicated')"
- match: "final_hidden_states = self.experts(hidden_states, topk_output)"
append: "dumper.dump('moe_expert_output', final_hidden_states, dims='t h[tp:partial]')"
append: "dumper.dump('moe_expert_output', final_hidden_states, dims='t h[moe_tp:partial] # tp:replicated')"
"""
PATCH_CONFIG_DP_ATTENTION_YAML: str = """\
patches:
# --- decoder layer level (aligned with miles test) ---
# In dp-attention mode: attn tensors are NOT TP-sharded (attn_tp_size=1),
# and mlp_output is already all-reduced inside forward_normal().
# dp-attention TP=2 DP=2 uses only 2 GPUs:
# GPU 0: tp=0, attn_tp=0 (attn_tp_size=1), attn_dp=0
# GPU 1: tp=1, attn_tp=0 (attn_tp_size=1), attn_dp=1
# All sub-axes (attn_tp, moe_tp, attn_dp) are uniquely determined by tp_rank,
# so only tp:replicated is needed — sub-axes are auto-resolved as implicitly replicated.
#
# Attn tensors are NOT TP-sharded, mlp_output is already all-reduced.
# layer_input is dumped after prepare_attn which DP-distributes tokens,
# so it needs dp:=attn_dp to filter to the non-empty DP rank.
- target: sglang.srt.models.qwen3_moe.Qwen3MoeDecoderLayer.forward
@@ -146,7 +151,7 @@ patches:
- target: sglang.srt.models.qwen3_moe.Qwen3MoeAttention.forward_core
edits:
- match: "output, _ = self.o_proj(attn_output)"
prepend: "dumper.dump('attn_pre_o_proj', attn_output, dims='t attn_h # tp:replicated')"
prepend: "dumper.dump('attn_pre_o_proj', attn_output, dims='t attn_h # tp:replicated dp:=attn_dp')"
# --- moe internals ---
- target: sglang.srt.models.qwen3_moe.Qwen3MoeSparseMoeBlock.forward_normal
@@ -154,7 +159,7 @@ patches:
- match: "router_logits, _ = self.gate(hidden_states)"
append: "dumper.dump('moe_router_logits', router_logits, dims='t num_experts # tp:replicated')"
- match: "final_hidden_states = self.experts(hidden_states, topk_output)"
append: "dumper.dump('moe_expert_output', final_hidden_states, dims='t h[tp:partial]')"
append: "dumper.dump('moe_expert_output', final_hidden_states, dims='t h[moe_tp:partial] # tp:replicated')"
"""
@@ -279,6 +284,7 @@ def _run_server_and_generate(
"--mem-fraction-static",
"0.5",
"--disable-cuda-graph",
"--disable-piecewise-cuda-graph",
"--disable-radix-cache",
]
if extra_server_args: