From 26d94deec4b3f9ed8b8ca08dcb937273d3914267 Mon Sep 17 00:00:00 2001 From: Zilin Zhu Date: Thu, 11 Jun 2026 08:48:42 +0000 Subject: [PATCH] Use /v1/loads to re-abort server --- slime/backends/sglang_utils/server_control.py | 67 +++++++++++++++++++ slime/rollout/sglang_rollout.py | 8 +-- 2 files changed, 69 insertions(+), 6 deletions(-) create mode 100644 slime/backends/sglang_utils/server_control.py diff --git a/slime/backends/sglang_utils/server_control.py b/slime/backends/sglang_utils/server_control.py new file mode 100644 index 0000000000..56da8e4c16 --- /dev/null +++ b/slime/backends/sglang_utils/server_control.py @@ -0,0 +1,67 @@ +import asyncio +import logging +from typing import Any + +from slime.utils.http_utils import get, post + +logger = logging.getLogger(__name__) + +ABORT_RETRY_INTERVAL_SECONDS = 3 + + +def num_requests_from_load(load: Any) -> int: + if isinstance(load, list): + return sum(num_requests_from_load(item) for item in load) + + if not isinstance(load, dict): + return 0 + + if "loads" in load: + return num_requests_from_load(load["loads"]) + + for key in ("num_reqs", "num_total_reqs", "total_reqs"): + value = load.get(key) + if isinstance(value, int): + return value + + running = load.get("num_running_reqs", load.get("total_running_reqs")) + waiting = load.get("num_waiting_reqs", load.get("total_waiting_reqs")) + return (running if isinstance(running, int) else 0) + (waiting if isinstance(waiting, int) else 0) + + +async def _abort_server_once(url: str) -> None: + try: + await post(f"{url}/abort_request", {"abort_all": True}) + except Exception as e: + logger.warning(f"Failed to abort SGLang server at {url}: {e}") + + +async def _get_server_num_requests(url: str) -> int: + return num_requests_from_load(await get(f"{url}/v1/loads?include=core")) + + +async def abort_server_until_idle(url: str, retry_interval: int = ABORT_RETRY_INTERVAL_SECONDS) -> None: + attempt = 1 + while True: + logger.info(f"Abort request for SGLang server {url}") + await _abort_server_once(url) + + try: + num_requests = await _get_server_num_requests(url) + except Exception as e: + logger.warning(f"Failed to get SGLang server load from {url}: {e}") + return + + if num_requests <= 0: + return + + logger.info( + f"SGLang server {url} still has {num_requests} requests after abort attempt {attempt}; " + f"retrying in {retry_interval} seconds." + ) + await asyncio.sleep(retry_interval) + attempt += 1 + + +async def abort_servers_until_idle(urls: list[str]) -> None: + await asyncio.gather(*(abort_server_until_idle(url) for url in urls)) diff --git a/slime/rollout/sglang_rollout.py b/slime/rollout/sglang_rollout.py index eee3a13680..bb87360639 100644 --- a/slime/rollout/sglang_rollout.py +++ b/slime/rollout/sglang_rollout.py @@ -14,6 +14,7 @@ from packaging.version import parse from tqdm import tqdm +from slime.backends.sglang_utils.server_control import abort_servers_until_idle from slime.rollout.base_types import RolloutFnEvalOutput, RolloutFnTrainOutput from slime.rollout.filter_hub.base_types import MetricGatherer, call_dynamic_filter from slime.utils.async_utils import run @@ -361,12 +362,7 @@ async def abort(args: Namespace, rollout_id: int) -> list[list[Sample]]: response = await get(f"http://{args.sglang_router_ip}:{args.sglang_router_port}/workers") urls = [worker["url"] for worker in response["workers"]] - logger.info(f"Abort request for {urls}") - abort_tasks = [post(f"{url}/abort_request", {"abort_all": True}) for url in urls] - abort_results = await asyncio.gather(*abort_tasks, return_exceptions=True) - for url, result in zip(urls, abort_results, strict=False): - if isinstance(result, Exception): - logger.warning(f"Failed to abort worker at {url}: {result}") + await abort_servers_until_idle(urls) # make sure all the pending tasks are finished count = 0