Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 16 additions & 2 deletions python/sglang/test/chunked_prefill_test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,10 @@
from types import SimpleNamespace
from typing import ClassVar, List, Optional

import torch

from sglang.kernels.ops.kv_canary._dispatch import use_torch_reference
from sglang.srt.utils import get_device
from sglang.test.run_eval import run_eval
from sglang.test.server_fixtures.disaggregation_fixture import (
PDDisaggregationServerBase,
Expand Down Expand Up @@ -38,6 +42,16 @@
]


def _canary_args(enabled: bool) -> List[str]:
if not enabled:
return []
args = list(KV_CANARY_ARGS)
if use_torch_reference(torch.device(get_device())):
# install_canary refuses a captured decode over the torch reference.
args.append("--cuda-graph-backend-decode=disabled")
return args


class ChunkedGsm8kMixin:
__test__ = False
use_kv_canary: ClassVar[bool] = True
Expand All @@ -52,7 +66,7 @@ class ChunkedGsm8kMixin:
gsm8k_threshold: ClassVar[float]

def build_prefill_side_args(self) -> List[str]:
canary = list(KV_CANARY_ARGS) if self.use_kv_canary else []
canary = _canary_args(self.use_kv_canary)
return (
["--chunked-prefill-size", str(self.chunked_prefill_size)]
+ list(self.feature_args)
Expand Down Expand Up @@ -118,7 +132,7 @@ def setUpClass(cls):
cls.extra_prefill_args = cls(
"test_mixed_prefix_gsm8k_chunked"
).build_prefill_side_args()
canary = list(KV_CANARY_ARGS) if cls.use_kv_canary else []
canary = _canary_args(cls.use_kv_canary)
cls.extra_decode_args = canary + list(cls.decode_feature_args)
PDDisaggregationServerBase.setUpClass()
cls.model = try_cached_model(cls.model)
Expand Down
49 changes: 42 additions & 7 deletions python/sglang/test/scripted_runtime/http_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,12 +8,15 @@
from typing import Any, Callable, Dict, Optional, Tuple

import requests
import torch
import zmq

from sglang.kernels.ops.kv_canary._dispatch import use_torch_reference
from sglang.srt.entrypoints.http_server import launch_server
from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.common import get_device
from sglang.srt.utils.network import get_free_port, get_zmq_socket_on_host
from sglang.test.scripted_runtime.io_struct import (
HookReady,
Expand All @@ -26,12 +29,20 @@
from sglang.test.scripted_runtime.utils import close_zmq_socket

DEFAULT_RUN_TIMEOUT_S: float = 120.0

CANARY_REFERENCE_TIMEOUT_SCALE: float = 4.0
SHUTDOWN_JOIN_TIMEOUT_S: float = 60.0
LISTENER_ACCEPT_TIMEOUT_S: float = 300.0
HTTP_READY_TIMEOUT_S: float = 300.0
HTTP_READY_POLL_INTERVAL_S: float = 0.5
SERVER_HOST: str = "127.0.0.1"

CANARY_LAUNCH_DEFAULTS: Dict[str, Any] = dict(
kv_canary="raise",
kv_canary_real_data="partial",
kv_canary_sweep_interval=100,
)


class ScriptedHttpServer:
def __init__(
Expand All @@ -42,6 +53,7 @@ def __init__(
server_process: mp.process.BaseProcess,
out_of_band_error_path: Path,
http_port: int,
run_timeout_s: float = DEFAULT_RUN_TIMEOUT_S,
) -> None:
self._ctx = ctx
self._socket = socket
Expand All @@ -50,14 +62,15 @@ def __init__(
self._base_url = f"http://{SERVER_HOST}:{http_port}"
self._shutdown_done = False
self._dirty: Optional[str] = None
self._run_timeout_s = run_timeout_s

@classmethod
def start(cls, **engine_kwargs: Any) -> ScriptedHttpServer:
out_of_band_error_path = _create_oob_error_file()

ctx = zmq.Context()
dispatch_port, socket = get_zmq_socket_on_host(ctx, zmq.PAIR, host=SERVER_HOST)
server_process, http_port = _spawn_server_process(
server_process, http_port, run_timeout_s = _spawn_server_process(
endpoint=f"tcp://{SERVER_HOST}:{dispatch_port}",
out_of_band_error_path=out_of_band_error_path,
engine_kwargs=engine_kwargs,
Expand All @@ -69,6 +82,7 @@ def start(cls, **engine_kwargs: Any) -> ScriptedHttpServer:
server_process=server_process,
out_of_band_error_path=out_of_band_error_path,
http_port=http_port,
run_timeout_s=run_timeout_s,
)
try:
self._await_handshake()
Expand All @@ -83,8 +97,10 @@ def execute_script(
script_fn: Callable,
*,
args: Tuple[Any, ...] = (),
timeout_s: float = DEFAULT_RUN_TIMEOUT_S,
timeout_s: Optional[float] = None,
) -> None:
if timeout_s is None:
timeout_s = self._run_timeout_s
if self._dirty:
raise RuntimeError(f"ScriptedHttpServer is dirty: {self._dirty}")

Expand Down Expand Up @@ -217,23 +233,42 @@ def _create_oob_error_file() -> Path:
return Path(err_path)


def _default_run_timeout_s(*, kv_canary: str, device: str) -> float:
if _canary_decode_graph_override(kv_canary=kv_canary, device=device):
return DEFAULT_RUN_TIMEOUT_S * CANARY_REFERENCE_TIMEOUT_SCALE
return DEFAULT_RUN_TIMEOUT_S


def _canary_decode_graph_override(*, kv_canary: str, device: str) -> Dict[str, Any]:
if kv_canary != "none" and use_torch_reference(torch.device(device)):
# install_canary refuses a captured decode over the torch reference.
return {"cuda_graph_backend_decode": "disabled"}
return {}


def _spawn_server_process(
*,
endpoint: str,
out_of_band_error_path: Path,
engine_kwargs: Dict[str, Any],
) -> Tuple[mp.process.BaseProcess, int]:
) -> Tuple[mp.process.BaseProcess, int, float]:
mp_ctx = mp.get_context("spawn")
device = engine_kwargs.get("device") or get_device()
launch_kwargs: Dict[str, Any] = dict(
host=SERVER_HOST,
port=get_free_port(),
kv_canary="raise",
kv_canary_real_data="partial",
kv_canary_sweep_interval=100,
disable_prefill_cuda_graph=True,
**CANARY_LAUNCH_DEFAULTS,
)
launch_kwargs.update(engine_kwargs)
for key, value in _canary_decode_graph_override(
kv_canary=launch_kwargs["kv_canary"], device=device.split(":")[0]
).items():
launch_kwargs.setdefault(key, value)
http_port = launch_kwargs["port"]
run_timeout_s = _default_run_timeout_s(
kv_canary=launch_kwargs["kv_canary"], device=device.split(":")[0]
)
server_process = mp_ctx.Process(
target=_launch_scripted_http_server,
kwargs=launch_kwargs,
Expand All @@ -252,7 +287,7 @@ def _spawn_server_process(
):
server_process.start()

return server_process, http_port
return server_process, http_port, run_timeout_s


def _launch_scripted_http_server(**engine_kwargs: Any) -> None:
Expand Down
22 changes: 22 additions & 0 deletions python/sglang/test/scripted_runtime_chunked_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,10 @@

SMALL_MODEL: str = "Qwen/Qwen3-0.6B"

RID_RELEASE_SETTLE_STEPS: int = 40

DRAIN_RELEASE_STEPS: int = 12


def base_engine_kwargs(
*,
Expand Down Expand Up @@ -40,6 +44,24 @@ def run_until_finished(handle, *, max_steps: int = DEFAULT_MAX_STEPS):
yield from run_until(handle, lambda h: h.finished, max_steps=max_steps)


def run_until_finished_then_settle(handle, *, max_steps: int = DEFAULT_MAX_STEPS):
yield from run_until_finished(handle, max_steps=max_steps)
for _ in range(RID_RELEASE_SETTLE_STEPS):
yield


def drain_until_released(t, *handles, max_steps: int = DRAIN_RELEASE_STEPS):
for _ in range(max_steps):
if all(
h.kv_pages == 0
and h.lock_refs == 0
and (h.req is None or h.req.kv.req_pool_idx is None)
for h in handles
):
return
yield


def run_until_all_finished(handles: List[Any], *, max_steps: int = DEFAULT_MAX_STEPS):
done = [False] * len(handles)
for _ in range(max_steps):
Expand Down
Loading
Loading