Support resetting and enhance HTTP endpoints for dumper (#19046)
Co-authored-by: Yueming Yuan <112649537+yueming-yuan@users.noreply.github.com>
This commit is contained in:
@@ -127,9 +127,11 @@ class _Dumper:
|
|||||||
Auto-cleanup old dumps before first write:
|
Auto-cleanup old dumps before first write:
|
||||||
`SGLANG_DUMPER_CLEANUP_PREVIOUS=1 python ...`
|
`SGLANG_DUMPER_CLEANUP_PREVIOUS=1 python ...`
|
||||||
|
|
||||||
Alternatively, disable at startup and enable via HTTP:
|
Alternatively, disable at startup and configure via HTTP:
|
||||||
1. `python ...`
|
1. `python ...`
|
||||||
2. `curl -X POST http://localhost:40000/dumper -d '{"enable": true}'`
|
2. `curl -X POST http://localhost:40000/dumper/configure -d '{"enable": true}'`
|
||||||
|
3. `curl -X POST http://localhost:40000/dumper/configure -d '{"enable": true, "filter": "layer_id=[0-3]"}'`
|
||||||
|
4. `curl -X POST http://localhost:40000/dumper/reset`
|
||||||
|
|
||||||
Related: `sglang.srt.debug_utils.dump_comparator` for dump comparison
|
Related: `sglang.srt.debug_utils.dump_comparator` for dump comparison
|
||||||
"""
|
"""
|
||||||
@@ -187,6 +189,11 @@ class _Dumper:
|
|||||||
k: v for k, v in (self._global_ctx | kwargs).items() if v is not None
|
k: v for k, v in (self._global_ctx | kwargs).items() if v is not None
|
||||||
}
|
}
|
||||||
|
|
||||||
|
def reset(self) -> None:
|
||||||
|
self._dump_index = 0
|
||||||
|
self._forward_pass_id = 0
|
||||||
|
self._global_ctx = {}
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def capture_output(self):
|
def capture_output(self):
|
||||||
assert self._captured_output_data is None
|
assert self._captured_output_data is None
|
||||||
@@ -611,60 +618,50 @@ def _start_maybe_http_server(dumper, timeout_seconds: int = 60):
|
|||||||
if http_port <= 0:
|
if http_port <= 0:
|
||||||
return
|
return
|
||||||
|
|
||||||
local_handler = _DumperRpcHandler(dumper)
|
rpc_broadcast = _create_zmq_rpc_broadcast(
|
||||||
rpc_handles = _create_zmq_rpc_handles(
|
dumper, base_port=zmq_base_port, timeout_seconds=timeout_seconds
|
||||||
local_handler, base_port=zmq_base_port, timeout_seconds=timeout_seconds
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if _get_rank() == 0:
|
if _get_rank() == 0:
|
||||||
handler_class = _make_dumper_http_handler(rpc_handles=rpc_handles)
|
handler_class = _make_http_handler(prefix="/dumper/", target=rpc_broadcast)
|
||||||
server = HTTPServer(("0.0.0.0", http_port), handler_class)
|
server = HTTPServer(("0.0.0.0", http_port), handler_class)
|
||||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||||
thread.start()
|
thread.start()
|
||||||
print(f"[Dumper] HTTP server started on port {http_port}")
|
print(f"[Dumper] HTTP server started on port {http_port}")
|
||||||
|
|
||||||
|
|
||||||
def _make_dumper_http_handler(rpc_handles):
|
def _make_http_handler(*, prefix: str, target):
|
||||||
class _DumperHTTPHandler(BaseHTTPRequestHandler):
|
class _HTTPHandler(BaseHTTPRequestHandler):
|
||||||
def do_POST(self):
|
def do_POST(self):
|
||||||
if self.path == "/dumper":
|
if not self.path.startswith(prefix):
|
||||||
try:
|
|
||||||
self._handle_endpoint_dumper()
|
|
||||||
self.send_response(200)
|
|
||||||
self.end_headers()
|
|
||||||
except Exception as e:
|
|
||||||
self.send_error(400, str(e))
|
|
||||||
else:
|
|
||||||
self.send_error(404)
|
self.send_error(404)
|
||||||
|
return
|
||||||
|
method = self.path[len(prefix) :]
|
||||||
|
try:
|
||||||
|
kwargs = self._get_request_body()
|
||||||
|
print(f"[Dumper#{_get_rank()}] HTTP {self.path} {kwargs=}")
|
||||||
|
getattr(target, method)(**kwargs)
|
||||||
|
self.send_response(200)
|
||||||
|
self.end_headers()
|
||||||
|
except Exception as e:
|
||||||
|
self.send_error(400, str(e))
|
||||||
|
|
||||||
def _get_request_body(self):
|
def _get_request_body(self) -> dict:
|
||||||
content_length = int(self.headers.get("Content-Length", 0))
|
content_length = int(self.headers.get("Content-Length", 0))
|
||||||
|
if content_length == 0:
|
||||||
|
return {}
|
||||||
return json.loads(self.rfile.read(content_length))
|
return json.loads(self.rfile.read(content_length))
|
||||||
|
|
||||||
def _handle_endpoint_dumper(self):
|
return _HTTPHandler
|
||||||
data = self._get_request_body()
|
|
||||||
print(f"[Dumper#{_get_rank()}] Handle HTTP endpoint {data=}")
|
|
||||||
for rpc_handle in rpc_handles:
|
|
||||||
rpc_handle.set_enable(data["enable"])
|
|
||||||
|
|
||||||
return _DumperHTTPHandler
|
|
||||||
|
|
||||||
|
|
||||||
class _DumperRpcHandler:
|
|
||||||
def __init__(self, dumper):
|
|
||||||
self._dumper = dumper
|
|
||||||
|
|
||||||
def set_enable(self, enable: bool):
|
|
||||||
print(f"[DumperRpcHandler] set_enable {enable=}")
|
|
||||||
self._dumper.configure(enable=enable)
|
|
||||||
|
|
||||||
|
|
||||||
# -------------------------------------- zmq rpc ------------------------------------------
|
# -------------------------------------- zmq rpc ------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def _create_zmq_rpc_handles(
|
def _create_zmq_rpc_broadcast(
|
||||||
handler, base_port: int, timeout_seconds: int = 60
|
handler, base_port: int, timeout_seconds: int = 60
|
||||||
) -> Optional[List["_ZmqRpcHandle"]]:
|
) -> Optional["_ZmqRpcBroadcast"]:
|
||||||
|
"""A general-purpose minimal RPC to support broadcasting executions to multi processes"""
|
||||||
import zmq
|
import zmq
|
||||||
|
|
||||||
rank = _get_rank()
|
rank = _get_rank()
|
||||||
@@ -695,7 +692,7 @@ def _create_zmq_rpc_handles(
|
|||||||
all_addresses = [None] * world_size
|
all_addresses = [None] * world_size
|
||||||
_collective_with_timeout(
|
_collective_with_timeout(
|
||||||
lambda: dist.all_gather_object(all_addresses, local_addr),
|
lambda: dist.all_gather_object(all_addresses, local_addr),
|
||||||
operation_name="all_gather_object in _create_zmq_rpc_handles",
|
operation_name="all_gather_object in _create_zmq_rpc_broadcast",
|
||||||
timeout_seconds=timeout_seconds,
|
timeout_seconds=timeout_seconds,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -708,7 +705,7 @@ def _create_zmq_rpc_handles(
|
|||||||
req_socket = ctx.socket(zmq.REQ)
|
req_socket = ctx.socket(zmq.REQ)
|
||||||
req_socket.connect(addr)
|
req_socket.connect(addr)
|
||||||
handles.append(_ZmqRpcHandle(req_socket, debug_name=f"rank-{i}"))
|
handles.append(_ZmqRpcHandle(req_socket, debug_name=f"rank-{i}"))
|
||||||
return handles
|
return _ZmqRpcBroadcast(handles)
|
||||||
else:
|
else:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -716,7 +713,7 @@ def _create_zmq_rpc_handles(
|
|||||||
class _ZmqRpcHandle:
|
class _ZmqRpcHandle:
|
||||||
"""Proxy object to call remote handler methods via ZMQ."""
|
"""Proxy object to call remote handler methods via ZMQ."""
|
||||||
|
|
||||||
def __init__(self, socket, debug_name):
|
def __init__(self, socket, debug_name: str):
|
||||||
self._socket = socket
|
self._socket = socket
|
||||||
self._debug_name = debug_name
|
self._debug_name = debug_name
|
||||||
|
|
||||||
@@ -739,6 +736,20 @@ class _ZmqRpcHandle:
|
|||||||
return call
|
return call
|
||||||
|
|
||||||
|
|
||||||
|
class _ZmqRpcBroadcast:
|
||||||
|
"""Broadcasts method calls to all ZMQ RPC handles."""
|
||||||
|
|
||||||
|
def __init__(self, handles: List[_ZmqRpcHandle]):
|
||||||
|
self._handles = handles
|
||||||
|
|
||||||
|
def __getattr__(self, method_name: str):
|
||||||
|
def call(*args, **kwargs):
|
||||||
|
for handle in self._handles:
|
||||||
|
getattr(handle, method_name)(*args, **kwargs)
|
||||||
|
|
||||||
|
return call
|
||||||
|
|
||||||
|
|
||||||
# --------------------------------- copied code (avoid dependency) --------------------------------------
|
# --------------------------------- copied code (avoid dependency) --------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -247,27 +247,85 @@ class TestDumperDistributed:
|
|||||||
assert "WARNING" in output, f"Expected WARNING in rank 0 output: {output}"
|
assert "WARNING" in output, f"Expected WARNING in rank 0 output: {output}"
|
||||||
assert "has not completed after 3s" in output
|
assert "has not completed after 3s" in output
|
||||||
|
|
||||||
def test_http_enable(self):
|
def test_http_configure(self):
|
||||||
with temp_set_env(allow_sglang=True, SGLANG_DUMPER_ENABLE="0"):
|
with temp_set_env(allow_sglang=True, SGLANG_DUMPER_ENABLE="0"):
|
||||||
run_distributed_test(self._test_http_func)
|
run_distributed_test(self._test_http_configure_func)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _test_http_func(rank):
|
def _test_http_configure_func(rank):
|
||||||
from sglang.srt.debug_utils.dumper import dumper
|
from sglang.srt.debug_utils.dumper import dumper
|
||||||
|
|
||||||
assert not dumper._config.enable
|
assert not dumper._config.enable
|
||||||
dumper.on_forward_pass_start()
|
dumper.on_forward_pass_start()
|
||||||
|
|
||||||
|
base_url = "http://localhost:40000"
|
||||||
|
|
||||||
|
# (1) enable toggle
|
||||||
for enable in [True, False]:
|
for enable in [True, False]:
|
||||||
dist.barrier()
|
dist.barrier()
|
||||||
if rank == 0:
|
if rank == 0:
|
||||||
time.sleep(0.1)
|
time.sleep(0.1)
|
||||||
requests.post(
|
requests.post(
|
||||||
"http://localhost:40000/dumper", json={"enable": enable}
|
f"{base_url}/dumper/configure", json={"enable": enable}
|
||||||
).raise_for_status()
|
).raise_for_status()
|
||||||
dist.barrier()
|
dist.barrier()
|
||||||
assert dumper._config.enable == enable
|
assert dumper._config.enable == enable
|
||||||
|
|
||||||
|
# (2) multi-field configure
|
||||||
|
dist.barrier()
|
||||||
|
if rank == 0:
|
||||||
|
time.sleep(0.1)
|
||||||
|
requests.post(
|
||||||
|
f"{base_url}/dumper/configure",
|
||||||
|
json={"enable": True, "filter": "layer_id=0", "dir": "/tmp/test_http"},
|
||||||
|
).raise_for_status()
|
||||||
|
dist.barrier()
|
||||||
|
assert dumper._config.enable is True
|
||||||
|
assert dumper._config.filter == "layer_id=0"
|
||||||
|
assert dumper._config.dir == "/tmp/test_http"
|
||||||
|
|
||||||
|
# (3) clear optional field
|
||||||
|
dist.barrier()
|
||||||
|
if rank == 0:
|
||||||
|
time.sleep(0.1)
|
||||||
|
requests.post(
|
||||||
|
f"{base_url}/dumper/configure",
|
||||||
|
json={"filter": None},
|
||||||
|
).raise_for_status()
|
||||||
|
dist.barrier()
|
||||||
|
assert dumper._config.filter is None
|
||||||
|
|
||||||
|
# (4) reset
|
||||||
|
dumper._dump_index = 42
|
||||||
|
dumper._forward_pass_id = 99
|
||||||
|
dist.barrier()
|
||||||
|
if rank == 0:
|
||||||
|
time.sleep(0.1)
|
||||||
|
requests.post(f"{base_url}/dumper/reset").raise_for_status()
|
||||||
|
dist.barrier()
|
||||||
|
assert dumper._dump_index == 0
|
||||||
|
assert dumper._forward_pass_id == 0
|
||||||
|
|
||||||
|
# (5) error: unknown field -> 400
|
||||||
|
dist.barrier()
|
||||||
|
if rank == 0:
|
||||||
|
time.sleep(0.1)
|
||||||
|
resp = requests.post(
|
||||||
|
f"{base_url}/dumper/configure",
|
||||||
|
json={"nonexistent_field": 123},
|
||||||
|
)
|
||||||
|
assert resp.status_code == 400
|
||||||
|
|
||||||
|
# (6) error: wrong type -> 400
|
||||||
|
dist.barrier()
|
||||||
|
if rank == 0:
|
||||||
|
time.sleep(0.1)
|
||||||
|
resp = requests.post(
|
||||||
|
f"{base_url}/dumper/configure",
|
||||||
|
json={"enable": "not_a_bool"},
|
||||||
|
)
|
||||||
|
assert resp.status_code == 400
|
||||||
|
|
||||||
def test_file_content_correctness(self, tmp_path):
|
def test_file_content_correctness(self, tmp_path):
|
||||||
with temp_set_env(
|
with temp_set_env(
|
||||||
allow_sglang=True,
|
allow_sglang=True,
|
||||||
@@ -817,5 +875,31 @@ class TestCleanup:
|
|||||||
_assert_files(_get_filenames(tmp_path), exist=["new_tensor"])
|
_assert_files(_get_filenames(tmp_path), exist=["new_tensor"])
|
||||||
|
|
||||||
|
|
||||||
|
class TestReset:
|
||||||
|
def test_reset_clears_state(self, tmp_path):
|
||||||
|
d = _make_test_dumper(tmp_path)
|
||||||
|
d.set_ctx(layer_id=1)
|
||||||
|
d.dump("before_reset", torch.randn(3, 3))
|
||||||
|
|
||||||
|
d.reset()
|
||||||
|
|
||||||
|
assert d._dump_index == 0
|
||||||
|
assert d._forward_pass_id == 0
|
||||||
|
assert d._global_ctx == {}
|
||||||
|
|
||||||
|
def test_dump_works_after_reset(self, tmp_path):
|
||||||
|
d = _make_test_dumper(tmp_path)
|
||||||
|
d.dump("pre", torch.randn(3, 3))
|
||||||
|
|
||||||
|
d.reset()
|
||||||
|
d.on_forward_pass_start()
|
||||||
|
d.dump("post", torch.randn(3, 3))
|
||||||
|
|
||||||
|
filenames = _get_filenames(tmp_path)
|
||||||
|
_assert_files(filenames, exist=["pre", "post"])
|
||||||
|
post_file = _find_dump_file(tmp_path, name="post")
|
||||||
|
assert "dump_index=1" in post_file.name
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__]))
|
sys.exit(pytest.main([__file__]))
|
||||||
|
|||||||
Reference in New Issue
Block a user