Configure and call dumper in main SGLang logic (#19093)
This commit is contained in:
@@ -20,9 +20,7 @@ def main(args):
|
|||||||
)
|
)
|
||||||
if args.filter:
|
if args.filter:
|
||||||
df_target = df_target.filter(pl.col("filename").str.contains(args.filter))
|
df_target = df_target.filter(pl.col("filename").str.contains(args.filter))
|
||||||
assert all(
|
assert all(c in df_target.columns for c in ["rank", "step", "dump_index", "name"])
|
||||||
c in df_target.columns for c in ["rank", "step", "dump_index", "name"]
|
|
||||||
)
|
|
||||||
|
|
||||||
df_baseline = read_meta(args.baseline_path)
|
df_baseline = read_meta(args.baseline_path)
|
||||||
print("df_target", df_target)
|
print("df_target", df_target)
|
||||||
@@ -41,9 +39,7 @@ def main(args):
|
|||||||
baseline_step = location_info.baseline_step
|
baseline_step = location_info.baseline_step
|
||||||
baseline_token_slice = location_info.baseline_token_slice
|
baseline_token_slice = location_info.baseline_token_slice
|
||||||
else:
|
else:
|
||||||
baseline_step = (
|
baseline_step = row["step"] - args.start_id + args.baseline_start_id
|
||||||
row["step"] - args.start_id + args.baseline_start_id
|
|
||||||
)
|
|
||||||
baseline_token_slice = None
|
baseline_token_slice = None
|
||||||
|
|
||||||
tensor_dim_desc = None
|
tensor_dim_desc = None
|
||||||
|
|||||||
@@ -168,6 +168,10 @@ class _Dumper:
|
|||||||
|
|
||||||
# ------------------------------- public :: core ---------------------------------
|
# ------------------------------- public :: core ---------------------------------
|
||||||
|
|
||||||
|
@property
|
||||||
|
def may_enable(self) -> bool:
|
||||||
|
return self._config.enable or self._config.server_port_parsed is not None
|
||||||
|
|
||||||
def step(self):
|
def step(self):
|
||||||
"""This should be called on all ranks at the end of each iteration."""
|
"""This should be called on all ranks at the end of each iteration."""
|
||||||
|
|
||||||
@@ -237,6 +241,7 @@ class _Dumper:
|
|||||||
self,
|
self,
|
||||||
model: "torch.nn.Module",
|
model: "torch.nn.Module",
|
||||||
) -> Optional["_NonIntrusiveDumper"]:
|
) -> Optional["_NonIntrusiveDumper"]:
|
||||||
|
self._ensure_http_server()
|
||||||
mode = self._config.non_intrusive_mode
|
mode = self._config.non_intrusive_mode
|
||||||
if mode == "off":
|
if mode == "off":
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -50,6 +50,7 @@ from sglang.srt.configs.load_config import LoadConfig, LoadFormat
|
|||||||
from sglang.srt.configs.model_config import AttentionArch, ModelConfig, ModelImpl
|
from sglang.srt.configs.model_config import AttentionArch, ModelConfig, ModelImpl
|
||||||
from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp
|
from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp
|
||||||
from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS
|
from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS
|
||||||
|
from sglang.srt.debug_utils.dumper import dumper
|
||||||
from sglang.srt.debug_utils.tensor_dump_forward_hook import (
|
from sglang.srt.debug_utils.tensor_dump_forward_hook import (
|
||||||
register_forward_hook_for_model,
|
register_forward_hook_for_model,
|
||||||
)
|
)
|
||||||
@@ -1055,6 +1056,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
self.pp_rank,
|
self.pp_rank,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if dumper.may_enable:
|
||||||
|
dumper.register_non_intrusive_dumper(self.model)
|
||||||
|
|
||||||
# Pre-expand RoPE cache before CUDA Graph capture
|
# Pre-expand RoPE cache before CUDA Graph capture
|
||||||
reserve_rope_cache_for_long_sequences(
|
reserve_rope_cache_for_long_sequences(
|
||||||
self.model,
|
self.model,
|
||||||
@@ -2438,6 +2442,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
if self.eplb_manager is not None:
|
if self.eplb_manager is not None:
|
||||||
self.eplb_manager.on_forward_pass_end()
|
self.eplb_manager.on_forward_pass_end()
|
||||||
|
|
||||||
|
if dumper.may_enable:
|
||||||
|
dumper.step()
|
||||||
|
|
||||||
return output
|
return output
|
||||||
|
|
||||||
def _forward_raw(
|
def _forward_raw(
|
||||||
|
|||||||
@@ -102,6 +102,21 @@ class TestDumperConfig:
|
|||||||
d.configure_default(filter="from_code")
|
d.configure_default(filter="from_code")
|
||||||
assert d._config.filter == "from_code"
|
assert d._config.filter == "from_code"
|
||||||
|
|
||||||
|
def test_may_enable_default_false(self):
|
||||||
|
d = _Dumper(config=_DumperConfig())
|
||||||
|
assert d.may_enable is False
|
||||||
|
|
||||||
|
def test_may_enable_true_when_enabled(self):
|
||||||
|
d = _Dumper(config=_DumperConfig(enable=True))
|
||||||
|
assert d.may_enable is True
|
||||||
|
|
||||||
|
def test_may_enable_true_when_server_port_set(self):
|
||||||
|
d = _Dumper(config=_DumperConfig(server_port="40000"))
|
||||||
|
assert d.may_enable is True
|
||||||
|
|
||||||
|
d2 = _Dumper(config=_DumperConfig(server_port="reuse"))
|
||||||
|
assert d2.may_enable is True
|
||||||
|
|
||||||
|
|
||||||
class TestDumperPureFunctions:
|
class TestDumperPureFunctions:
|
||||||
def test_get_truncated_value(self):
|
def test_get_truncated_value(self):
|
||||||
@@ -1548,5 +1563,81 @@ class TestNonIntrusiveLayerIdCtx(_NonIntrusiveTestBase):
|
|||||||
assert len(layer1_keys) == 0, f"layer 1 dumps should be filtered: {layer1_keys}"
|
assert len(layer1_keys) == 0, f"layer 1 dumps should be filtered: {layer1_keys}"
|
||||||
|
|
||||||
|
|
||||||
|
class TestDumperE2E:
|
||||||
|
def test_step_and_non_intrusive_hooks(self, tmp_path):
|
||||||
|
base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
dump_dir = str(tmp_path)
|
||||||
|
env = {
|
||||||
|
**os.environ,
|
||||||
|
"DUMPER_SERVER_PORT": "reuse",
|
||||||
|
}
|
||||||
|
proc = popen_launch_server(
|
||||||
|
"Qwen/Qwen3-0.6B",
|
||||||
|
base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=["--tp", "2", "--max-total-tokens", "128"],
|
||||||
|
env=env,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
states = requests.post(f"{base_url}/dumper/get_state", json={}).json()
|
||||||
|
assert len(states) == 2, f"Expected 2 ranks (tp=2), got {len(states)}"
|
||||||
|
for state in states:
|
||||||
|
assert state["config"]["enable"] is False
|
||||||
|
assert state["step"] == 0
|
||||||
|
|
||||||
|
requests.post(
|
||||||
|
f"{base_url}/dumper/configure",
|
||||||
|
json={"enable": True, "dir": dump_dir},
|
||||||
|
).raise_for_status()
|
||||||
|
|
||||||
|
states = requests.post(f"{base_url}/dumper/get_state", json={}).json()
|
||||||
|
assert len(states) == 2
|
||||||
|
for rank, state in enumerate(states):
|
||||||
|
assert (
|
||||||
|
state["config"]["enable"] is True
|
||||||
|
), f"rank {rank}: enable should be True after configure"
|
||||||
|
assert state["config"]["dir"] == dump_dir
|
||||||
|
|
||||||
|
resp = requests.post(
|
||||||
|
f"{base_url}/generate",
|
||||||
|
json={"text": "Hello", "sampling_params": {"max_new_tokens": 8}},
|
||||||
|
)
|
||||||
|
assert resp.status_code == 200, f"Generate failed: {resp.text}"
|
||||||
|
|
||||||
|
states = requests.post(f"{base_url}/dumper/get_state", json={}).json()
|
||||||
|
assert len(states) == 2
|
||||||
|
steps = [s["step"] for s in states]
|
||||||
|
for rank, step in enumerate(steps):
|
||||||
|
assert step > 0, f"rank {rank}: step should be > 0, got {step}"
|
||||||
|
assert steps[0] == steps[1], f"step mismatch across ranks: {steps}"
|
||||||
|
|
||||||
|
dump_files = list(Path(dump_dir).glob("dump_*/*.pt"))
|
||||||
|
assert len(dump_files) > 0, f"No dump files in {dump_dir}"
|
||||||
|
filenames = {f.name for f in dump_files}
|
||||||
|
|
||||||
|
for field in ("input_ids", "positions"):
|
||||||
|
assert any(f"name={field}" in f for f in filenames), (
|
||||||
|
f"Missing {field} dump from non-intrusive hooks, "
|
||||||
|
f"got: {sorted(filenames)[:10]}"
|
||||||
|
)
|
||||||
|
|
||||||
|
for rank in range(2):
|
||||||
|
assert any(
|
||||||
|
f"rank={rank}" in f for f in filenames
|
||||||
|
), f"No dump files for rank {rank}"
|
||||||
|
|
||||||
|
sample_file = dump_files[0]
|
||||||
|
loaded = torch.load(sample_file, map_location="cpu", weights_only=False)
|
||||||
|
assert isinstance(loaded, dict), f"Expected dict, got {type(loaded)}"
|
||||||
|
assert (
|
||||||
|
"value" in loaded and "meta" in loaded
|
||||||
|
), f"Missing value/meta keys: {loaded.keys()}"
|
||||||
|
assert "name" in loaded["meta"]
|
||||||
|
assert "rank" in loaded["meta"]
|
||||||
|
assert "step" in loaded["meta"]
|
||||||
|
finally:
|
||||||
|
kill_process_tree(proc.pid)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__]))
|
sys.exit(pytest.main([__file__]))
|
||||||
|
|||||||
Reference in New Issue
Block a user