Support non-intrusive arbitrary dumping in dumper and add e2e tests (#19563)
This commit is contained in:
@@ -138,6 +138,7 @@ class DumperConfig(_BaseConfig):
|
|||||||
collective_timeout: int = 60
|
collective_timeout: int = 60
|
||||||
server_port: str = "-1"
|
server_port: str = "-1"
|
||||||
non_intrusive_mode: str = "core"
|
non_intrusive_mode: str = "core"
|
||||||
|
source_patcher_config: Optional[str] = None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _env_prefix(cls) -> str:
|
def _env_prefix(cls) -> str:
|
||||||
@@ -325,6 +326,25 @@ class _Dumper:
|
|||||||
|
|
||||||
return decorator
|
return decorator
|
||||||
|
|
||||||
|
def apply_source_patches(self) -> None:
|
||||||
|
"""Apply source patches from DUMPER_SOURCE_PATCHER_CONFIG if set.
|
||||||
|
|
||||||
|
Automatically injects ``from sglang.srt.debug_utils.dumper import dumper``
|
||||||
|
into every replacement block so users don't need to write it in YAML.
|
||||||
|
"""
|
||||||
|
config_path = self._config.source_patcher_config
|
||||||
|
if not config_path:
|
||||||
|
return
|
||||||
|
|
||||||
|
from sglang.srt.debug_utils.source_patcher import apply_patches_from_config
|
||||||
|
|
||||||
|
yaml_content: str = Path(config_path).read_text()
|
||||||
|
print(f"[source_patcher] loading config from {config_path}")
|
||||||
|
apply_patches_from_config(
|
||||||
|
yaml_content,
|
||||||
|
extra_imports=["from sglang.srt.debug_utils.dumper import dumper"],
|
||||||
|
)
|
||||||
|
|
||||||
def register_non_intrusive_dumper(
|
def register_non_intrusive_dumper(
|
||||||
self,
|
self,
|
||||||
model: "torch.nn.Module",
|
model: "torch.nn.Module",
|
||||||
@@ -588,11 +608,13 @@ class _NonIntrusiveDumper:
|
|||||||
def _detect_module_ctx(
|
def _detect_module_ctx(
|
||||||
cls, module_name: str, module: "torch.nn.Module"
|
cls, module_name: str, module: "torch.nn.Module"
|
||||||
) -> Optional[dict]:
|
) -> Optional[dict]:
|
||||||
if cls._LAYER_NAME_RE.fullmatch(module_name):
|
match = cls._LAYER_NAME_RE.fullmatch(module_name)
|
||||||
|
if match:
|
||||||
for plugin in _plugins:
|
for plugin in _plugins:
|
||||||
layer_id = plugin.detect_layer_id(module)
|
layer_id = plugin.detect_layer_id(module)
|
||||||
if layer_id is not None:
|
if layer_id is not None:
|
||||||
return {"layer_id": layer_id}
|
return {"layer_id": layer_id}
|
||||||
|
return {"layer_id": int(match.group(1))}
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _register_ctx_hooks(self, module: "torch.nn.Module", *, ctx: dict) -> None:
|
def _register_ctx_hooks(self, module: "torch.nn.Module", *, ctx: dict) -> None:
|
||||||
|
|||||||
@@ -1061,6 +1061,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if dumper.may_enable:
|
if dumper.may_enable:
|
||||||
|
dumper.apply_source_patches()
|
||||||
dumper.register_non_intrusive_dumper(self.model)
|
dumper.register_non_intrusive_dumper(self.model)
|
||||||
|
|
||||||
# Pre-expand RoPE cache before CUDA Graph capture
|
# Pre-expand RoPE cache before CUDA Graph capture
|
||||||
@@ -1763,9 +1764,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
init_new_workspace=init_new_workspace,
|
init_new_workspace=init_new_workspace,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.prefill_attention_backend_str, self.decode_attention_backend_str = (
|
(
|
||||||
self.server_args.get_attention_backends()
|
self.prefill_attention_backend_str,
|
||||||
)
|
self.decode_attention_backend_str,
|
||||||
|
) = self.server_args.get_attention_backends()
|
||||||
|
|
||||||
if self.decode_attention_backend_str != self.prefill_attention_backend_str:
|
if self.decode_attention_backend_str != self.prefill_attention_backend_str:
|
||||||
from sglang.srt.layers.attention.hybrid_attn_backend import (
|
from sglang.srt.layers.attention.hybrid_attn_backend import (
|
||||||
|
|||||||
@@ -0,0 +1,54 @@
|
|||||||
|
"""Test dumper.apply_source_patches() integration with source_patcher."""
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
from types import ModuleType
|
||||||
|
|
||||||
|
import yaml
|
||||||
|
|
||||||
|
from sglang.srt.debug_utils.dumper import DumperConfig, _Dumper
|
||||||
|
|
||||||
|
SAMPLE_MODULE_NAME = "_source_patcher_test_fixtures.sample_module"
|
||||||
|
|
||||||
|
|
||||||
|
class TestDumperApplySourcePatches:
|
||||||
|
def test_no_config_is_noop(self) -> None:
|
||||||
|
config = DumperConfig(source_patcher_config=None)
|
||||||
|
d = _Dumper(config=config)
|
||||||
|
d.apply_source_patches()
|
||||||
|
|
||||||
|
def test_patches_applied_from_yaml(
|
||||||
|
self, sample_module: ModuleType, tmp_path: Path
|
||||||
|
) -> None:
|
||||||
|
cls = sample_module.SampleClass
|
||||||
|
obj = cls()
|
||||||
|
assert obj.greet("world") == "hello world"
|
||||||
|
|
||||||
|
original_code = cls.greet.__code__
|
||||||
|
|
||||||
|
patch_config = {
|
||||||
|
"patches": [
|
||||||
|
{
|
||||||
|
"target": f"{SAMPLE_MODULE_NAME}.SampleClass.greet",
|
||||||
|
"edits": [
|
||||||
|
{
|
||||||
|
"match": 'greeting = f"hello {name}"',
|
||||||
|
"replacement": 'greeting = f"dumper_patched {name}"',
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
config_path = tmp_path / "patch_config.yaml"
|
||||||
|
config_path.write_text(yaml.dump(patch_config))
|
||||||
|
|
||||||
|
config = DumperConfig(source_patcher_config=str(config_path))
|
||||||
|
d = _Dumper(config=config)
|
||||||
|
|
||||||
|
try:
|
||||||
|
d.apply_source_patches()
|
||||||
|
assert obj.greet("world") == "dumper_patched world"
|
||||||
|
finally:
|
||||||
|
cls.greet.__code__ = original_code
|
||||||
|
|
||||||
|
assert obj.greet("world") == "hello world"
|
||||||
@@ -1755,7 +1755,6 @@ def _make_forward_batch():
|
|||||||
|
|
||||||
|
|
||||||
class TestNonIntrusiveDumperConfigMode(_NonIntrusiveTestBase):
|
class TestNonIntrusiveDumperConfigMode(_NonIntrusiveTestBase):
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _build_model() -> torch.nn.Module:
|
def _build_model() -> torch.nn.Module:
|
||||||
class SubLayer(torch.nn.Module):
|
class SubLayer(torch.nn.Module):
|
||||||
@@ -1904,8 +1903,8 @@ class TestNonIntrusiveLayerIdCtx(_NonIntrusiveTestBase):
|
|||||||
assert layer_key in captured
|
assert layer_key in captured
|
||||||
assert captured[layer_key]["meta"]["layer_id"] == 5
|
assert captured[layer_key]["meta"]["layer_id"] == 5
|
||||||
|
|
||||||
def test_no_layer_id_when_no_attr(self, tmp_path):
|
def test_layer_id_fallback_from_module_name(self, tmp_path):
|
||||||
"""layers.N modules without layer_number/layer_id -> no layer_id injected."""
|
"""layers.N modules without layer_number/layer_id -> layer_id from module name."""
|
||||||
|
|
||||||
class Inner(torch.nn.Module):
|
class Inner(torch.nn.Module):
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
@@ -1922,8 +1921,17 @@ class TestNonIntrusiveLayerIdCtx(_NonIntrusiveTestBase):
|
|||||||
captured, x, output = self._run(tmp_path, Inner)
|
captured, x, output = self._run(tmp_path, Inner)
|
||||||
|
|
||||||
assert len(captured) > 0
|
assert len(captured) > 0
|
||||||
for key, entry in captured.items():
|
input_keys: list[str] = [
|
||||||
assert "layer_id" not in entry["meta"], f"{key} has unexpected layer_id"
|
k for k in captured if "model.layers." in k and "inputs" in k
|
||||||
|
]
|
||||||
|
assert len(input_keys) > 0
|
||||||
|
for key in input_keys:
|
||||||
|
meta = captured[key]["meta"]
|
||||||
|
assert "layer_id" in meta, f"{key} missing layer_id"
|
||||||
|
if "layers.0" in key:
|
||||||
|
assert meta["layer_id"] == 0
|
||||||
|
elif "layers.1" in key:
|
||||||
|
assert meta["layer_id"] == 1
|
||||||
|
|
||||||
def test_filter_by_layer_id(self, tmp_path):
|
def test_filter_by_layer_id(self, tmp_path):
|
||||||
"""filter='layer_id=0' keeps only layer 0 dumps."""
|
"""filter='layer_id=0' keeps only layer 0 dumps."""
|
||||||
@@ -2171,7 +2179,6 @@ class TestMegatronConvertValue:
|
|||||||
|
|
||||||
|
|
||||||
class TestNonIntrusiveKwargsModel(_NonIntrusiveTestBase):
|
class TestNonIntrusiveKwargsModel(_NonIntrusiveTestBase):
|
||||||
|
|
||||||
def test_kwargs_core_fields(self, tmp_path):
|
def test_kwargs_core_fields(self, tmp_path):
|
||||||
class KwargsModel(torch.nn.Module):
|
class KwargsModel(torch.nn.Module):
|
||||||
def forward(self, *, input_ids, position_ids):
|
def forward(self, *, input_ids, position_ids):
|
||||||
|
|||||||
@@ -0,0 +1,216 @@
|
|||||||
|
"""E2E test: source patcher + dumper + comparator on SGLang server.
|
||||||
|
|
||||||
|
Patches Qwen3DecoderLayer.forward to insert dumper.dump() calls,
|
||||||
|
launches 1-GPU baseline and 2-GPU TP=2 target servers, runs inference,
|
||||||
|
verifies patched dump fields exist, then runs comparator to verify
|
||||||
|
numerical consistency.
|
||||||
|
|
||||||
|
The dumper.apply_source_patches() auto-injects ``from ... import dumper``
|
||||||
|
so the YAML only needs ``dumper.dump(...)`` calls.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
|
import subprocess
|
||||||
|
import tempfile
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import requests
|
||||||
|
|
||||||
|
from sglang.srt.debug_utils.comparator.output_types import (
|
||||||
|
AnyRecord,
|
||||||
|
SummaryRecord,
|
||||||
|
parse_record_json,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils import kill_process_tree
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
DEFAULT_URL_FOR_TEST,
|
||||||
|
popen_launch_server,
|
||||||
|
)
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=120, suite="nightly-2-gpu", nightly=True)
|
||||||
|
|
||||||
|
MODEL = "Qwen/Qwen3-0.6B"
|
||||||
|
EXP_NAME = "e2e_source_patcher"
|
||||||
|
DUMPER_FILTER = r"layer_id=[012]"
|
||||||
|
|
||||||
|
PATCH_CONFIG_YAML: str = """\
|
||||||
|
patches:
|
||||||
|
- target: sglang.srt.models.qwen3.Qwen3DecoderLayer.forward
|
||||||
|
edits:
|
||||||
|
- match: "hidden_states = self.mlp(hidden_states)"
|
||||||
|
prepend: "dumper.dump('patched_attn_output', hidden_states, dims='t h')"
|
||||||
|
- match: "return hidden_states, residual"
|
||||||
|
prepend: "dumper.dump('patched_mlp_output', hidden_states, dims='t h')"
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
class TestSourcePatcherE2ESGLang:
|
||||||
|
"""E2E: patch Qwen3 forward -> dump -> compare 1gpu vs 2gpu-tp2."""
|
||||||
|
|
||||||
|
@pytest.mark.timeout(300)
|
||||||
|
def test_patch_dump_and_compare(self, tmp_path: Path) -> None:
|
||||||
|
patched_fields: list[str] = ["patched_attn_output", "patched_mlp_output"]
|
||||||
|
base_url: str = DEFAULT_URL_FOR_TEST
|
||||||
|
|
||||||
|
config_path: Path = tmp_path / "patch_config.yaml"
|
||||||
|
config_path.write_text(PATCH_CONFIG_YAML)
|
||||||
|
|
||||||
|
# Run 1: baseline (1 GPU)
|
||||||
|
baseline_dir: Path = tmp_path / "baseline"
|
||||||
|
_run_server_and_generate(
|
||||||
|
dump_dir=baseline_dir,
|
||||||
|
config_path=config_path,
|
||||||
|
tp=1,
|
||||||
|
base_url=base_url,
|
||||||
|
)
|
||||||
|
_verify_patched_fields(dump_dir=baseline_dir, field_names=patched_fields)
|
||||||
|
|
||||||
|
# Run 2: target (2 GPU TP=2)
|
||||||
|
target_dir: Path = tmp_path / "target"
|
||||||
|
_run_server_and_generate(
|
||||||
|
dump_dir=target_dir,
|
||||||
|
config_path=config_path,
|
||||||
|
tp=2,
|
||||||
|
base_url=base_url,
|
||||||
|
)
|
||||||
|
_verify_patched_fields(dump_dir=target_dir, field_names=patched_fields)
|
||||||
|
|
||||||
|
# Compare baseline vs target
|
||||||
|
baseline_exp: Path = baseline_dir / EXP_NAME
|
||||||
|
target_exp: Path = target_dir / EXP_NAME
|
||||||
|
|
||||||
|
result: subprocess.CompletedProcess[str] = subprocess.run(
|
||||||
|
[
|
||||||
|
"python",
|
||||||
|
"-m",
|
||||||
|
"sglang.srt.debug_utils.comparator",
|
||||||
|
"--baseline-path",
|
||||||
|
str(baseline_exp),
|
||||||
|
"--target-path",
|
||||||
|
str(target_exp),
|
||||||
|
"--output-format",
|
||||||
|
"json",
|
||||||
|
"--grouping",
|
||||||
|
"logical",
|
||||||
|
],
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
debug_file: Path = _save_comparator_output(
|
||||||
|
stdout=result.stdout, stderr=result.stderr
|
||||||
|
)
|
||||||
|
print(f"Comparator debug output: {debug_file}")
|
||||||
|
|
||||||
|
assert result.returncode == 0, (
|
||||||
|
f"Comparator failed (rc={result.returncode}). "
|
||||||
|
f"Debug output: {debug_file}"
|
||||||
|
)
|
||||||
|
|
||||||
|
records: list[AnyRecord] = [
|
||||||
|
parse_record_json(line)
|
||||||
|
for line in result.stdout.strip().splitlines()
|
||||||
|
if line.strip()
|
||||||
|
]
|
||||||
|
assert (
|
||||||
|
len(records) > 0
|
||||||
|
), f"Comparator produced no output records. Debug: {debug_file}"
|
||||||
|
|
||||||
|
summary: SummaryRecord = _find_summary(records=records, debug_file=debug_file)
|
||||||
|
assert (
|
||||||
|
summary.passed > 0
|
||||||
|
), f"No comparisons passed (total={summary.total}). Debug: {debug_file}"
|
||||||
|
assert summary.failed == 0, (
|
||||||
|
f"{summary.failed} comparisons failed "
|
||||||
|
f"(passed={summary.passed}, skipped={summary.skipped}). "
|
||||||
|
f"Debug: {debug_file}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# --------------------------------- helpers ---------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _run_server_and_generate(
|
||||||
|
*,
|
||||||
|
dump_dir: Path,
|
||||||
|
config_path: Path,
|
||||||
|
tp: int,
|
||||||
|
base_url: str,
|
||||||
|
) -> None:
|
||||||
|
"""Launch SGLang server with source patcher + dumper, send a generate request."""
|
||||||
|
env: dict[str, str] = {
|
||||||
|
**os.environ,
|
||||||
|
"DUMPER_SOURCE_PATCHER_CONFIG": str(config_path),
|
||||||
|
"DUMPER_DIR": str(dump_dir),
|
||||||
|
"DUMPER_EXP_NAME": EXP_NAME,
|
||||||
|
"DUMPER_SERVER_PORT": "reuse",
|
||||||
|
}
|
||||||
|
|
||||||
|
proc = popen_launch_server(
|
||||||
|
MODEL,
|
||||||
|
base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=["--tp", str(tp), "--max-total-tokens", "128"],
|
||||||
|
env=env,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
requests.post(
|
||||||
|
f"{base_url}/dumper/configure",
|
||||||
|
json={
|
||||||
|
"enable": True,
|
||||||
|
"filter": DUMPER_FILTER,
|
||||||
|
"cleanup_previous": True,
|
||||||
|
},
|
||||||
|
).raise_for_status()
|
||||||
|
|
||||||
|
resp = requests.post(
|
||||||
|
f"{base_url}/generate",
|
||||||
|
json={
|
||||||
|
"text": "The capital of France is",
|
||||||
|
"sampling_params": {"max_new_tokens": 8},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert resp.status_code == 200, f"Generate failed: {resp.text}"
|
||||||
|
finally:
|
||||||
|
kill_process_tree(proc.pid)
|
||||||
|
|
||||||
|
|
||||||
|
def _verify_patched_fields(*, dump_dir: Path, field_names: list[str]) -> None:
|
||||||
|
"""Verify that patched dump fields exist as .pt files."""
|
||||||
|
for field in field_names:
|
||||||
|
matches: list[Path] = list(dump_dir.rglob(f"*name={field}*.pt"))
|
||||||
|
assert len(matches) > 0, (
|
||||||
|
f"Expected patched field '{field}' not found under {dump_dir}. "
|
||||||
|
f"Available files: {sorted(f.name for f in dump_dir.rglob('*.pt'))[:20]}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _find_summary(*, records: list[AnyRecord], debug_file: Path) -> SummaryRecord:
|
||||||
|
"""Extract the SummaryRecord from comparator output."""
|
||||||
|
summaries: list[SummaryRecord] = [
|
||||||
|
r for r in records if isinstance(r, SummaryRecord)
|
||||||
|
]
|
||||||
|
assert len(summaries) == 1, (
|
||||||
|
f"Expected 1 summary record, got {len(summaries)}. "
|
||||||
|
f"Record types: {[type(r).__name__ for r in records]}. "
|
||||||
|
f"Debug: {debug_file}"
|
||||||
|
)
|
||||||
|
return summaries[0]
|
||||||
|
|
||||||
|
|
||||||
|
def _save_comparator_output(*, stdout: str, stderr: str) -> Path:
|
||||||
|
"""Save comparator stdout+stderr to a temp file that persists for debugging."""
|
||||||
|
fd, path_str = tempfile.mkstemp(prefix="comparator_e2e_", suffix=".log", dir="/tmp")
|
||||||
|
with os.fdopen(fd, "w") as f:
|
||||||
|
f.write("=== STDOUT ===\n")
|
||||||
|
f.write(stdout)
|
||||||
|
f.write("\n=== STDERR ===\n")
|
||||||
|
f.write(stderr)
|
||||||
|
return Path(path_str)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
pytest.main([__file__, "-v"])
|
||||||
Reference in New Issue
Block a user