Support using SGLang port in dumper (#19038)
This commit is contained in:
@@ -7,11 +7,11 @@ import time
|
|||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
from dataclasses import dataclass, fields, replace
|
from dataclasses import asdict, dataclass, fields, replace
|
||||||
from functools import cached_property
|
from functools import cached_property
|
||||||
from http.server import BaseHTTPRequestHandler, HTTPServer
|
from http.server import BaseHTTPRequestHandler, HTTPServer
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import List, Optional, Self, get_args, get_type_hints
|
from typing import Any, List, Literal, Optional, Self, Union, get_args, get_type_hints
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
@@ -73,7 +73,10 @@ class _FrozenConfig(ABC):
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _parse_env_field(env_name: str, default):
|
def _parse_env_field(env_name: str, default):
|
||||||
raw = os.getenv(env_name)
|
return _FrozenConfig._parse_env_value(os.getenv(env_name), default)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _parse_env_value(raw, default):
|
||||||
if raw is None or not raw.strip():
|
if raw is None or not raw.strip():
|
||||||
return default
|
return default
|
||||||
if isinstance(default, bool):
|
if isinstance(default, bool):
|
||||||
@@ -98,11 +101,22 @@ class _DumperConfig(_FrozenConfig):
|
|||||||
enable_http_server: bool = True
|
enable_http_server: bool = True
|
||||||
cleanup_previous: bool = False
|
cleanup_previous: bool = False
|
||||||
collective_timeout: int = 60
|
collective_timeout: int = 60
|
||||||
|
server_port: str = "-1"
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _env_prefix(cls) -> str:
|
def _env_prefix(cls) -> str:
|
||||||
return "SGLANG_DUMPER_"
|
return "SGLANG_DUMPER_"
|
||||||
|
|
||||||
|
@property
|
||||||
|
def server_port_parsed(self) -> Optional[Union[int, Literal["reuse"]]]:
|
||||||
|
raw = self.server_port
|
||||||
|
if raw == "reuse":
|
||||||
|
return "reuse"
|
||||||
|
port = int(raw)
|
||||||
|
if port <= 0:
|
||||||
|
return None
|
||||||
|
return port
|
||||||
|
|
||||||
|
|
||||||
# -------------------------------------- dumper core ------------------------------------------
|
# -------------------------------------- dumper core ------------------------------------------
|
||||||
|
|
||||||
@@ -129,9 +143,10 @@ class _Dumper:
|
|||||||
|
|
||||||
Alternatively, disable at startup and configure via HTTP:
|
Alternatively, disable at startup and configure via HTTP:
|
||||||
1. `python ...`
|
1. `python ...`
|
||||||
2. `curl -X POST http://localhost:40000/dumper/configure -d '{"enable": true}'`
|
2. sglang mode: `curl -X POST http://localhost:30000/dumper/configure -d '{"enable": true}'`
|
||||||
3. `curl -X POST http://localhost:40000/dumper/configure -d '{"enable": true, "filter": "layer_id=[0-3]"}'`
|
standalone: `curl -X POST http://localhost:40000/dumper/configure -d '{"enable": true}'`
|
||||||
4. `curl -X POST http://localhost:40000/dumper/reset`
|
3. `curl -X POST http://localhost:30000/dumper/configure -d '{"enable": true, "filter": "layer_id=[0-3]"}'`
|
||||||
|
4. `curl -X POST http://localhost:30000/dumper/reset`
|
||||||
|
|
||||||
Related: `sglang.srt.debug_utils.dump_comparator` for dump comparison
|
Related: `sglang.srt.debug_utils.dump_comparator` for dump comparison
|
||||||
"""
|
"""
|
||||||
@@ -146,6 +161,7 @@ class _Dumper:
|
|||||||
self._forward_pass_id = 0
|
self._forward_pass_id = 0
|
||||||
self._global_ctx: dict = {}
|
self._global_ctx: dict = {}
|
||||||
self._captured_output_data: Optional[dict] = None
|
self._captured_output_data: Optional[dict] = None
|
||||||
|
self._rpc_broadcast: "_RpcBroadcastBase" = _LocalOnlyBroadcast(self)
|
||||||
|
|
||||||
def on_forward_pass_start(self):
|
def on_forward_pass_start(self):
|
||||||
"""This should be called on all ranks."""
|
"""This should be called on all ranks."""
|
||||||
@@ -168,7 +184,28 @@ class _Dumper:
|
|||||||
if self._http_server_handled:
|
if self._http_server_handled:
|
||||||
return
|
return
|
||||||
self._http_server_handled = True
|
self._http_server_handled = True
|
||||||
_start_maybe_http_server(self, timeout_seconds=self._config.collective_timeout)
|
|
||||||
|
http_port = self._config.server_port_parsed
|
||||||
|
if http_port is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
rpc_broadcast = _create_zmq_rpc_broadcast(
|
||||||
|
self,
|
||||||
|
base_port=get_int_env_var("SGLANG_DUMPER_ZMQ_BASE_PORT", 16800),
|
||||||
|
timeout_seconds=self._config.collective_timeout,
|
||||||
|
)
|
||||||
|
|
||||||
|
if _get_rank() == 0:
|
||||||
|
assert rpc_broadcast is not None
|
||||||
|
self._rpc_broadcast = rpc_broadcast
|
||||||
|
|
||||||
|
if http_port == "reuse":
|
||||||
|
print(
|
||||||
|
"[Dumper] Standalone HTTP server disabled, reusing existing ports"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
_start_http_server(prefix="/dumper/", target=self, http_port=http_port)
|
||||||
|
print(f"[Dumper] HTTP server started on port {http_port}")
|
||||||
|
|
||||||
def _ensure_partial_name(self):
|
def _ensure_partial_name(self):
|
||||||
if self._config.partial_name is None:
|
if self._config.partial_name is None:
|
||||||
@@ -203,6 +240,34 @@ class _Dumper:
|
|||||||
finally:
|
finally:
|
||||||
self._captured_output_data = None
|
self._captured_output_data = None
|
||||||
|
|
||||||
|
def get_state(self) -> dict:
|
||||||
|
return {
|
||||||
|
"config": asdict(self._config),
|
||||||
|
"dump_index": self._dump_index,
|
||||||
|
"forward_pass_id": self._forward_pass_id,
|
||||||
|
}
|
||||||
|
|
||||||
|
def _handle_http_control_request(
|
||||||
|
self, *, method: str, body: dict[str, Any]
|
||||||
|
) -> list[dict]:
|
||||||
|
return self._rpc_broadcast._handle_http_control_request_inner(
|
||||||
|
method=method, body=body
|
||||||
|
)
|
||||||
|
|
||||||
|
def _handle_http_control_request_inner(
|
||||||
|
self, *, method: str, body: dict[str, Any]
|
||||||
|
) -> dict:
|
||||||
|
if method == "get_state":
|
||||||
|
return self.get_state()
|
||||||
|
elif method == "configure":
|
||||||
|
self.configure(**body)
|
||||||
|
return {}
|
||||||
|
elif method == "reset":
|
||||||
|
self.reset()
|
||||||
|
return {}
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown dumper control method: {method!r}")
|
||||||
|
|
||||||
def configure(self, **kwargs) -> None:
|
def configure(self, **kwargs) -> None:
|
||||||
self._config = replace(self._config, **kwargs)
|
self._config = replace(self._config, **kwargs)
|
||||||
|
|
||||||
@@ -612,22 +677,11 @@ def _collect_megatron_parallel_info():
|
|||||||
# -------------------------------------- http control server ------------------------------------------
|
# -------------------------------------- http control server ------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def _start_maybe_http_server(dumper, timeout_seconds: int = 60):
|
def _start_http_server(*, prefix: str, target: object, http_port: int):
|
||||||
http_port = get_int_env_var("SGLANG_DUMPER_SERVER_PORT", 40000)
|
handler_class = _make_http_handler(prefix=prefix, target=target)
|
||||||
zmq_base_port = get_int_env_var("SGLANG_DUMPER_ZMQ_BASE_PORT", 16800)
|
server = HTTPServer(("0.0.0.0", http_port), handler_class)
|
||||||
if http_port <= 0:
|
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||||
return
|
thread.start()
|
||||||
|
|
||||||
rpc_broadcast = _create_zmq_rpc_broadcast(
|
|
||||||
dumper, base_port=zmq_base_port, timeout_seconds=timeout_seconds
|
|
||||||
)
|
|
||||||
|
|
||||||
if _get_rank() == 0:
|
|
||||||
handler_class = _make_http_handler(prefix="/dumper/", target=rpc_broadcast)
|
|
||||||
server = HTTPServer(("0.0.0.0", http_port), handler_class)
|
|
||||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
|
||||||
thread.start()
|
|
||||||
print(f"[Dumper] HTTP server started on port {http_port}")
|
|
||||||
|
|
||||||
|
|
||||||
def _make_http_handler(*, prefix: str, target):
|
def _make_http_handler(*, prefix: str, target):
|
||||||
@@ -638,11 +692,17 @@ def _make_http_handler(*, prefix: str, target):
|
|||||||
return
|
return
|
||||||
method = self.path[len(prefix) :]
|
method = self.path[len(prefix) :]
|
||||||
try:
|
try:
|
||||||
kwargs = self._get_request_body()
|
req_body = self._get_request_body()
|
||||||
print(f"[Dumper#{_get_rank()}] HTTP {self.path} {kwargs=}")
|
print(f"[Dumper#{_get_rank()}] HTTP {self.path} {req_body=}")
|
||||||
getattr(target, method)(**kwargs)
|
result = target._handle_http_control_request(
|
||||||
|
method=method, body=req_body
|
||||||
|
)
|
||||||
|
resp_body = json.dumps(result).encode()
|
||||||
self.send_response(200)
|
self.send_response(200)
|
||||||
|
self.send_header("Content-Type", "application/json")
|
||||||
|
self.send_header("Content-Length", str(len(resp_body)))
|
||||||
self.end_headers()
|
self.end_headers()
|
||||||
|
self.wfile.write(resp_body)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.send_error(400, str(e))
|
self.send_error(400, str(e))
|
||||||
|
|
||||||
@@ -736,16 +796,43 @@ class _ZmqRpcHandle:
|
|||||||
return call
|
return call
|
||||||
|
|
||||||
|
|
||||||
class _ZmqRpcBroadcast:
|
class _RpcBroadcastBase:
|
||||||
"""Broadcasts method calls to all ZMQ RPC handles."""
|
"""Base for broadcasting method calls to dumper instance(s)."""
|
||||||
|
|
||||||
|
def __getattr__(self, method_name: str):
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def __init__(self, handles: List[_ZmqRpcHandle]):
|
||||||
|
self._handles = handles
|
||||||
|
|
||||||
|
class _LocalOnlyBroadcast(_RpcBroadcastBase):
|
||||||
|
"""Calls methods directly on the local dumper, wrapping the result in a list."""
|
||||||
|
|
||||||
|
def __init__(self, dumper: "_Dumper"):
|
||||||
|
self._dumper = dumper
|
||||||
|
|
||||||
|
def __getattr__(self, method_name: str):
|
||||||
|
def call(*args, **kwargs):
|
||||||
|
return [getattr(self._dumper, method_name)(*args, **kwargs)]
|
||||||
|
|
||||||
|
return call
|
||||||
|
|
||||||
|
|
||||||
|
class _ZmqRpcBroadcast(_RpcBroadcastBase):
|
||||||
|
"""Broadcasts method calls to all ZMQ RPC handles.
|
||||||
|
|
||||||
|
Returns a list of results, one per rank (ordered by rank).
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(self, handles: List[_ZmqRpcHandle]):
|
def __init__(self, handles: List[_ZmqRpcHandle]):
|
||||||
self._handles = handles
|
self._handles = handles
|
||||||
|
|
||||||
def __getattr__(self, method_name: str):
|
def __getattr__(self, method_name: str):
|
||||||
def call(*args, **kwargs):
|
def call(*args, **kwargs):
|
||||||
for handle in self._handles:
|
return [
|
||||||
getattr(handle, method_name)(*args, **kwargs)
|
getattr(handle, method_name)(*args, **kwargs)
|
||||||
|
for handle in self._handles
|
||||||
|
]
|
||||||
|
|
||||||
return call
|
return call
|
||||||
|
|
||||||
|
|||||||
@@ -99,6 +99,7 @@ from sglang.srt.managers.io_struct import (
|
|||||||
ConfigureLoggingReq,
|
ConfigureLoggingReq,
|
||||||
ContinueGenerationReqInput,
|
ContinueGenerationReqInput,
|
||||||
DestroyWeightsUpdateGroupReqInput,
|
DestroyWeightsUpdateGroupReqInput,
|
||||||
|
DumperControlReqInput,
|
||||||
EmbeddingReqInput,
|
EmbeddingReqInput,
|
||||||
GenerateReqInput,
|
GenerateReqInput,
|
||||||
GetWeightsByNameReqInput,
|
GetWeightsByNameReqInput,
|
||||||
@@ -618,6 +619,22 @@ async def set_internal_state(obj: SetInternalStateReq, request: Request):
|
|||||||
return res
|
return res
|
||||||
|
|
||||||
|
|
||||||
|
# Do not import `dumper.py` to avoid dependency
|
||||||
|
if os.environ.get("SGLANG_DUMPER_SERVER_PORT") == "reuse":
|
||||||
|
|
||||||
|
@app.api_route("/dumper/{method}", methods=["POST"])
|
||||||
|
@auth_level(AuthLevel.ADMIN_OPTIONAL)
|
||||||
|
async def _dumper_control_handler(method: str, request: Request):
|
||||||
|
body_bytes = await request.body()
|
||||||
|
body = await request.json() if body_bytes else {}
|
||||||
|
obj = DumperControlReqInput(method=method, body=body)
|
||||||
|
results = await _global_state.tokenizer_manager.dumper_control(obj)
|
||||||
|
if any(not r.success for r in results):
|
||||||
|
errors = [r.error for r in results if not r.success]
|
||||||
|
return ORJSONResponse(status_code=400, content={"error": errors})
|
||||||
|
return [x for result in results for x in result.response]
|
||||||
|
|
||||||
|
|
||||||
# fastapi implicitly converts json in the request to obj (dataclass)
|
# fastapi implicitly converts json in the request to obj (dataclass)
|
||||||
@app.api_route("/generate", methods=["POST", "PUT"])
|
@app.api_route("/generate", methods=["POST", "PUT"])
|
||||||
async def generate_request(obj: GenerateReqInput, request: Request):
|
async def generate_request(obj: GenerateReqInput, request: Request):
|
||||||
|
|||||||
@@ -1970,6 +1970,19 @@ class LazyDumpTensorsReqOutput(BaseReq):
|
|||||||
success: bool
|
success: bool
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class DumperControlReqInput(BaseReq):
|
||||||
|
method: str
|
||||||
|
body: Dict[str, Any]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class DumperControlReqOutput(BaseReq):
|
||||||
|
success: bool
|
||||||
|
response: List[Dict[str, Any]]
|
||||||
|
error: str = ""
|
||||||
|
|
||||||
|
|
||||||
def _check_all_req_types():
|
def _check_all_req_types():
|
||||||
"""A helper function to check all request types are defined in this file."""
|
"""A helper function to check all request types are defined in this file."""
|
||||||
import inspect
|
import inspect
|
||||||
|
|||||||
@@ -35,6 +35,7 @@ from torch.distributed import barrier
|
|||||||
|
|
||||||
from sglang.srt.configs.model_config import ModelConfig
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
from sglang.srt.constrained.grammar_manager import GrammarManager
|
from sglang.srt.constrained.grammar_manager import GrammarManager
|
||||||
|
from sglang.srt.debug_utils.dumper import dumper
|
||||||
from sglang.srt.disaggregation.decode import (
|
from sglang.srt.disaggregation.decode import (
|
||||||
DecodePreallocQueue,
|
DecodePreallocQueue,
|
||||||
DecodeTransferQueue,
|
DecodeTransferQueue,
|
||||||
@@ -90,6 +91,8 @@ from sglang.srt.managers.io_struct import (
|
|||||||
DestroyWeightsUpdateGroupReqInput,
|
DestroyWeightsUpdateGroupReqInput,
|
||||||
DetachHiCacheStorageReqInput,
|
DetachHiCacheStorageReqInput,
|
||||||
DetachHiCacheStorageReqOutput,
|
DetachHiCacheStorageReqOutput,
|
||||||
|
DumperControlReqInput,
|
||||||
|
DumperControlReqOutput,
|
||||||
ExpertDistributionReq,
|
ExpertDistributionReq,
|
||||||
ExpertDistributionReqOutput,
|
ExpertDistributionReqOutput,
|
||||||
ExpertDistributionReqType,
|
ExpertDistributionReqType,
|
||||||
@@ -1082,6 +1085,7 @@ class Scheduler(
|
|||||||
(GetLoadsReqInput, self.get_loads),
|
(GetLoadsReqInput, self.get_loads),
|
||||||
(PauseGenerationReqInput, self.pause_generation),
|
(PauseGenerationReqInput, self.pause_generation),
|
||||||
(ContinueGenerationReqInput, self.continue_generation),
|
(ContinueGenerationReqInput, self.continue_generation),
|
||||||
|
(DumperControlReqInput, self.handle_dumper_control),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -2960,6 +2964,26 @@ class Scheduler(
|
|||||||
self.send_to_detokenizer.send_output(recv_req, recv_req)
|
self.send_to_detokenizer.send_output(recv_req, recv_req)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def handle_dumper_control(self, recv_req: DumperControlReqInput):
|
||||||
|
try:
|
||||||
|
response: list = []
|
||||||
|
if (
|
||||||
|
not torch.distributed.is_initialized()
|
||||||
|
or torch.distributed.get_rank() == 0
|
||||||
|
):
|
||||||
|
response = dumper._handle_http_control_request(
|
||||||
|
method=recv_req.method, body=recv_req.body
|
||||||
|
)
|
||||||
|
self.send_to_tokenizer.send_output(
|
||||||
|
DumperControlReqOutput(success=True, response=response), recv_req
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
print(f"[Scheduler] handle_dumper_control error: {e}", flush=True)
|
||||||
|
self.send_to_tokenizer.send_output(
|
||||||
|
DumperControlReqOutput(success=False, response=[], error=str(e)),
|
||||||
|
recv_req,
|
||||||
|
)
|
||||||
|
|
||||||
# placeholder for override
|
# placeholder for override
|
||||||
def update_cache_from_scheduler(
|
def update_cache_from_scheduler(
|
||||||
self, schedule_batch: ScheduleBatch, batch_result: GenerationBatchResult
|
self, schedule_batch: ScheduleBatch, batch_result: GenerationBatchResult
|
||||||
|
|||||||
@@ -34,6 +34,8 @@ from sglang.srt.managers.io_struct import (
|
|||||||
DestroyWeightsUpdateGroupReqOutput,
|
DestroyWeightsUpdateGroupReqOutput,
|
||||||
DetachHiCacheStorageReqInput,
|
DetachHiCacheStorageReqInput,
|
||||||
DetachHiCacheStorageReqOutput,
|
DetachHiCacheStorageReqOutput,
|
||||||
|
DumperControlReqInput,
|
||||||
|
DumperControlReqOutput,
|
||||||
ExpertDistributionReq,
|
ExpertDistributionReq,
|
||||||
ExpertDistributionReqOutput,
|
ExpertDistributionReqOutput,
|
||||||
ExpertDistributionReqType,
|
ExpertDistributionReqType,
|
||||||
@@ -233,6 +235,9 @@ class TokenizerCommunicatorMixin:
|
|||||||
self.get_loads_communicator = _Communicator(
|
self.get_loads_communicator = _Communicator(
|
||||||
self.send_to_scheduler, server_args.dp_size
|
self.send_to_scheduler, server_args.dp_size
|
||||||
)
|
)
|
||||||
|
self.dumper_control_communicator = _Communicator(
|
||||||
|
self.send_to_scheduler, server_args.dp_size
|
||||||
|
)
|
||||||
|
|
||||||
self._result_dispatcher += self._get_communicator_dispatcher()
|
self._result_dispatcher += self._get_communicator_dispatcher()
|
||||||
|
|
||||||
@@ -331,6 +336,10 @@ class TokenizerCommunicatorMixin:
|
|||||||
GetLoadsReqOutput,
|
GetLoadsReqOutput,
|
||||||
self.get_loads_communicator.handle_recv,
|
self.get_loads_communicator.handle_recv,
|
||||||
),
|
),
|
||||||
|
(
|
||||||
|
DumperControlReqOutput,
|
||||||
|
self.dumper_control_communicator.handle_recv,
|
||||||
|
),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -861,6 +870,11 @@ class TokenizerCommunicatorMixin:
|
|||||||
)
|
)
|
||||||
return [res.updated for res in responses]
|
return [res.updated for res in responses]
|
||||||
|
|
||||||
|
async def dumper_control(
|
||||||
|
self: TokenizerManager, obj: DumperControlReqInput
|
||||||
|
) -> List[DumperControlReqOutput]:
|
||||||
|
return await self.dumper_control_communicator(obj)
|
||||||
|
|
||||||
async def get_load(self: TokenizerManager) -> List[GetLoadReqOutput]:
|
async def get_load(self: TokenizerManager) -> List[GetLoadReqOutput]:
|
||||||
req = GetLoadReqInput()
|
req = GetLoadReqInput()
|
||||||
return await self.get_load_communicator(req)
|
return await self.get_load_communicator(req)
|
||||||
|
|||||||
@@ -1,4 +1,6 @@
|
|||||||
import io
|
import io
|
||||||
|
import multiprocessing
|
||||||
|
import os
|
||||||
import sys
|
import sys
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
@@ -20,12 +22,20 @@ from sglang.srt.debug_utils.dumper import (
|
|||||||
_materialize_value,
|
_materialize_value,
|
||||||
_obj_to_dict,
|
_obj_to_dict,
|
||||||
_torch_save,
|
_torch_save,
|
||||||
|
dumper,
|
||||||
get_tensor_info,
|
get_tensor_info,
|
||||||
get_truncated_value,
|
get_truncated_value,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import temp_set_env
|
from sglang.srt.environ import temp_set_env
|
||||||
|
from sglang.srt.utils import kill_process_tree
|
||||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
from sglang.test.test_utils import run_distributed_test
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
DEFAULT_URL_FOR_TEST,
|
||||||
|
find_available_port,
|
||||||
|
popen_launch_server,
|
||||||
|
run_distributed_test,
|
||||||
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=30, suite="nightly-2-gpu", nightly=True)
|
register_cuda_ci(est_time=30, suite="nightly-2-gpu", nightly=True)
|
||||||
register_amd_ci(est_time=60, suite="nightly-amd", nightly=True)
|
register_amd_ci(est_time=60, suite="nightly-amd", nightly=True)
|
||||||
@@ -193,8 +203,6 @@ class TestDumperDistributed:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _test_basic_func(rank, tmpdir):
|
def _test_basic_func(rank, tmpdir):
|
||||||
from sglang.srt.debug_utils.dumper import dumper
|
|
||||||
|
|
||||||
tensor = torch.randn(10, 10, device=f"cuda:{rank}")
|
tensor = torch.randn(10, 10, device=f"cuda:{rank}")
|
||||||
|
|
||||||
dumper.on_forward_pass_start()
|
dumper.on_forward_pass_start()
|
||||||
@@ -247,85 +255,6 @@ 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_configure(self):
|
|
||||||
with temp_set_env(allow_sglang=True, SGLANG_DUMPER_ENABLE="0"):
|
|
||||||
run_distributed_test(self._test_http_configure_func)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _test_http_configure_func(rank):
|
|
||||||
from sglang.srt.debug_utils.dumper import dumper
|
|
||||||
|
|
||||||
assert not dumper._config.enable
|
|
||||||
dumper.on_forward_pass_start()
|
|
||||||
|
|
||||||
base_url = "http://localhost:40000"
|
|
||||||
|
|
||||||
# (1) enable toggle
|
|
||||||
for enable in [True, False]:
|
|
||||||
dist.barrier()
|
|
||||||
if rank == 0:
|
|
||||||
time.sleep(0.1)
|
|
||||||
requests.post(
|
|
||||||
f"{base_url}/dumper/configure", json={"enable": enable}
|
|
||||||
).raise_for_status()
|
|
||||||
dist.barrier()
|
|
||||||
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,
|
||||||
@@ -336,8 +265,6 @@ class TestDumperDistributed:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _test_file_content_func(rank, tmpdir):
|
def _test_file_content_func(rank, tmpdir):
|
||||||
from sglang.srt.debug_utils.dumper import dumper
|
|
||||||
|
|
||||||
tensor = torch.arange(12, device=f"cuda:{rank}").reshape(3, 4).float()
|
tensor = torch.arange(12, device=f"cuda:{rank}").reshape(3, 4).float()
|
||||||
|
|
||||||
dumper.on_forward_pass_start()
|
dumper.on_forward_pass_start()
|
||||||
@@ -365,8 +292,6 @@ class TestDumperFileWriteControl:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _test_filter_func(rank, tmpdir):
|
def _test_filter_func(rank, tmpdir):
|
||||||
from sglang.srt.debug_utils.dumper import dumper
|
|
||||||
|
|
||||||
dumper.on_forward_pass_start()
|
dumper.on_forward_pass_start()
|
||||||
dumper.dump("keep_this", torch.randn(5, device=f"cuda:{rank}"))
|
dumper.dump("keep_this", torch.randn(5, device=f"cuda:{rank}"))
|
||||||
dumper.dump("skip_this", torch.randn(5, device=f"cuda:{rank}"))
|
dumper.dump("skip_this", torch.randn(5, device=f"cuda:{rank}"))
|
||||||
@@ -390,8 +315,6 @@ class TestDumperFileWriteControl:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _test_save_false_func(rank, tmpdir):
|
def _test_save_false_func(rank, tmpdir):
|
||||||
from sglang.srt.debug_utils.dumper import dumper
|
|
||||||
|
|
||||||
dumper.on_forward_pass_start()
|
dumper.on_forward_pass_start()
|
||||||
dumper.dump("no_save_tensor", torch.randn(5, device=f"cuda:{rank}"), save=False)
|
dumper.dump("no_save_tensor", torch.randn(5, device=f"cuda:{rank}"), save=False)
|
||||||
|
|
||||||
@@ -901,5 +824,141 @@ class TestReset:
|
|||||||
assert "dump_index=1" in post_file.name
|
assert "dump_index=1" in post_file.name
|
||||||
|
|
||||||
|
|
||||||
|
class TestDumperHttp:
|
||||||
|
"""Test /dumper/* HTTP control — parametrized over standalone vs sglang server."""
|
||||||
|
|
||||||
|
@pytest.fixture(scope="class", params=["standalone", "sglang"])
|
||||||
|
def dumper_http_url(self, request):
|
||||||
|
if request.param == "standalone":
|
||||||
|
http_port = find_available_port(40000)
|
||||||
|
base_url = f"http://127.0.0.1:{http_port}"
|
||||||
|
stop_event = multiprocessing.get_context("spawn").Event()
|
||||||
|
thread = threading.Thread(
|
||||||
|
target=run_distributed_test,
|
||||||
|
args=(TestDumperHttp._standalone_mode_worker,),
|
||||||
|
kwargs={"http_port": http_port, "stop_event": stop_event},
|
||||||
|
)
|
||||||
|
thread.start()
|
||||||
|
try:
|
||||||
|
TestDumperHttp._wait_for_http(base_url)
|
||||||
|
yield base_url
|
||||||
|
finally:
|
||||||
|
stop_event.set()
|
||||||
|
thread.join(timeout=10)
|
||||||
|
else:
|
||||||
|
base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
env = {**os.environ, "SGLANG_DUMPER_SERVER_PORT": "reuse"}
|
||||||
|
proc = popen_launch_server(
|
||||||
|
"Qwen/Qwen3-0.6B",
|
||||||
|
base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=["--max-total-tokens", "128"],
|
||||||
|
env=env,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
yield base_url
|
||||||
|
finally:
|
||||||
|
kill_process_tree(proc.pid)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _standalone_mode_worker(rank, http_port: int, stop_event):
|
||||||
|
dumper.configure(enable=False, server_port=str(http_port))
|
||||||
|
dumper.on_forward_pass_start()
|
||||||
|
stop_event.wait()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _wait_for_http(url: str, timeout: float = 30) -> None:
|
||||||
|
deadline = time.time() + timeout
|
||||||
|
while time.time() < deadline:
|
||||||
|
try:
|
||||||
|
requests.post(f"{url}/dumper/configure", json={}, timeout=2)
|
||||||
|
return
|
||||||
|
except requests.ConnectionError:
|
||||||
|
time.sleep(0.5)
|
||||||
|
raise TimeoutError(f"Standalone dumper HTTP server not reachable at {url}")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _post(base_url: str, method: str, **kwargs) -> list[dict]:
|
||||||
|
resp = requests.post(f"{base_url}/dumper/{method}", json=kwargs or None)
|
||||||
|
resp.raise_for_status()
|
||||||
|
states = resp.json()
|
||||||
|
assert isinstance(states, list) and len(states) >= 1
|
||||||
|
return states
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _assert_all_ranks(states: list[dict], path: str, expected):
|
||||||
|
"""Assert that ``state[path]`` equals ``expected`` on every rank."""
|
||||||
|
keys = path.split(".")
|
||||||
|
for rank, state in enumerate(states):
|
||||||
|
val = state
|
||||||
|
for k in keys:
|
||||||
|
val = val[k]
|
||||||
|
assert (
|
||||||
|
val == expected
|
||||||
|
), f"rank {rank}: {path}={val!r}, expected {expected!r}"
|
||||||
|
|
||||||
|
def test_configure_enable_toggle(self, dumper_http_url: str):
|
||||||
|
for enable in [True, False]:
|
||||||
|
self._post(dumper_http_url, "configure", enable=enable)
|
||||||
|
states = self._post(dumper_http_url, "get_state")
|
||||||
|
self._assert_all_ranks(states, "config.enable", enable)
|
||||||
|
|
||||||
|
def test_configure_multi_field(self, dumper_http_url: str):
|
||||||
|
self._post(
|
||||||
|
dumper_http_url,
|
||||||
|
"configure",
|
||||||
|
enable=True,
|
||||||
|
filter="layer_id=0",
|
||||||
|
dir="/tmp/test_http",
|
||||||
|
)
|
||||||
|
states = self._post(dumper_http_url, "get_state")
|
||||||
|
self._assert_all_ranks(states, "config.enable", True)
|
||||||
|
self._assert_all_ranks(states, "config.filter", "layer_id=0")
|
||||||
|
self._assert_all_ranks(states, "config.dir", "/tmp/test_http")
|
||||||
|
|
||||||
|
def test_configure_clear_optional(self, dumper_http_url: str):
|
||||||
|
self._post(dumper_http_url, "configure", filter="layer_id=0")
|
||||||
|
self._post(dumper_http_url, "configure", filter=None)
|
||||||
|
states = self._post(dumper_http_url, "get_state")
|
||||||
|
self._assert_all_ranks(states, "config.filter", None)
|
||||||
|
|
||||||
|
def test_reset(self, dumper_http_url: str):
|
||||||
|
self._post(dumper_http_url, "configure", enable=True)
|
||||||
|
self._post(dumper_http_url, "reset")
|
||||||
|
states = self._post(dumper_http_url, "get_state")
|
||||||
|
self._assert_all_ranks(states, "dump_index", 0)
|
||||||
|
self._assert_all_ranks(states, "forward_pass_id", 0)
|
||||||
|
|
||||||
|
def test_get_state(self, dumper_http_url: str):
|
||||||
|
self._post(dumper_http_url, "configure", enable=True, filter="layer_id=[0-3]")
|
||||||
|
states = self._post(dumper_http_url, "get_state")
|
||||||
|
self._assert_all_ranks(states, "config.enable", True)
|
||||||
|
self._assert_all_ranks(states, "config.filter", "layer_id=[0-3]")
|
||||||
|
for state in states:
|
||||||
|
assert "dump_index" in state
|
||||||
|
assert "forward_pass_id" in state
|
||||||
|
|
||||||
|
def test_all_ranks_consistent(self, dumper_http_url: str):
|
||||||
|
self._post(dumper_http_url, "configure", enable=True, dir="/tmp/multi")
|
||||||
|
states = self._post(dumper_http_url, "get_state")
|
||||||
|
configs = [s["config"] for s in states]
|
||||||
|
for rank_config in configs[1:]:
|
||||||
|
assert rank_config == configs[0], f"rank configs diverged: {configs}"
|
||||||
|
|
||||||
|
def test_error_unknown_field(self, dumper_http_url: str):
|
||||||
|
resp = requests.post(
|
||||||
|
f"{dumper_http_url}/dumper/configure",
|
||||||
|
json={"nonexistent_field": 123},
|
||||||
|
)
|
||||||
|
assert resp.status_code == 400
|
||||||
|
|
||||||
|
def test_error_wrong_type(self, dumper_http_url: str):
|
||||||
|
resp = requests.post(
|
||||||
|
f"{dumper_http_url}/dumper/configure",
|
||||||
|
json={"enable": "not_a_bool"},
|
||||||
|
)
|
||||||
|
assert resp.status_code == 400
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__]))
|
sys.exit(pytest.main([__file__]))
|
||||||
|
|||||||
Reference in New Issue
Block a user