492 lines
18 KiB
Python
492 lines
18 KiB
Python
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
|
# SPDX-License-Identifier: Apache-2.0
|
|
|
|
import dataclasses
|
|
import logging
|
|
import multiprocessing
|
|
import time
|
|
from urllib.parse import quote
|
|
|
|
import requests
|
|
import sglang_router
|
|
from packaging.version import parse
|
|
from sglang.srt.server_args import ServerArgs
|
|
from sglang.srt.utils import kill_process_tree
|
|
from urllib3.exceptions import NewConnectionError
|
|
|
|
from slime.ray.ray_actor import RayActor
|
|
from slime.utils.http_utils import get_host_info
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def get_base_gpu_id(args, rank):
|
|
num_gpus = min(args.num_gpus_per_node, args.rollout_num_gpus_per_engine)
|
|
if args.colocate:
|
|
start_index = (rank * num_gpus) % args.num_gpus_per_node
|
|
else:
|
|
num_actor_gpus = 0 if args.debug_rollout_only else args.actor_num_gpus_per_node * args.actor_num_nodes
|
|
start_index = (num_actor_gpus + rank * num_gpus) % args.num_gpus_per_node
|
|
if args.use_critic:
|
|
num_critic_gpus = args.critic_num_gpus_per_node * args.critic_num_nodes
|
|
start_index = (num_actor_gpus + num_critic_gpus + rank * num_gpus) % args.num_gpus_per_node
|
|
return start_index
|
|
|
|
|
|
def launch_server_process(server_args: ServerArgs) -> multiprocessing.Process:
|
|
from sglang.srt.entrypoints.http_server import launch_server
|
|
|
|
multiprocessing.set_start_method("spawn", force=True)
|
|
server_args.host = server_args.host.strip("[]")
|
|
p = multiprocessing.Process(target=launch_server, args=(server_args,))
|
|
p.start()
|
|
|
|
if server_args.node_rank != 0:
|
|
return
|
|
|
|
_wait_server_healthy(
|
|
base_url=server_args.url(),
|
|
api_key=server_args.api_key,
|
|
is_process_alive=lambda: p.is_alive(),
|
|
)
|
|
|
|
return p
|
|
|
|
|
|
def _wait_server_healthy(base_url, api_key, is_process_alive):
|
|
headers = {
|
|
"Content-Type": "application/json; charset=utf-8",
|
|
"Authorization": f"Bearer {api_key}",
|
|
}
|
|
|
|
with requests.Session() as session:
|
|
while True:
|
|
try:
|
|
response = session.get(f"{base_url}/health_generate", headers=headers)
|
|
if response.status_code == 200:
|
|
break
|
|
except requests.RequestException:
|
|
pass
|
|
|
|
if not is_process_alive():
|
|
raise Exception("Server process terminated unexpectedly.")
|
|
|
|
time.sleep(2)
|
|
|
|
# use flush_cache to make sure the working queue is empty, so that we can do offload
|
|
while True:
|
|
try:
|
|
response = session.get(f"{base_url}/flush_cache", headers=headers)
|
|
if response.status_code == 200:
|
|
break
|
|
|
|
except requests.RequestException:
|
|
pass
|
|
|
|
if not is_process_alive():
|
|
raise Exception("Server process terminated unexpectedly.")
|
|
|
|
time.sleep(2)
|
|
|
|
|
|
class SGLangEngine(RayActor):
|
|
def __init__(self, args, rank: int, worker_type: str = "regular"):
|
|
self.args = args
|
|
self.rank = rank
|
|
self.worker_type = worker_type
|
|
|
|
def init(self, dist_init_addr, port, nccl_port, host=None, disaggregation_bootstrap_port=None):
|
|
self.router_ip = self.args.sglang_router_ip
|
|
self.router_port = self.args.sglang_router_port
|
|
|
|
host = host or get_host_info()[1]
|
|
|
|
# support ipv6 address
|
|
if ":" in host and not host.startswith("["):
|
|
host = f"[{host}]"
|
|
|
|
# dist_init_addr may be 2605:...:10163, should split port
|
|
*addr_parts, port_str = dist_init_addr.split(":")
|
|
ipv6_addr = ":".join(addr_parts)
|
|
if ":" in ipv6_addr and not ipv6_addr.startswith("["):
|
|
dist_init_addr = f"[{ipv6_addr}]:{port_str}"
|
|
|
|
server_args_dict, external_engine_need_check_fields = _compute_server_args(
|
|
self.args,
|
|
self.rank,
|
|
dist_init_addr,
|
|
nccl_port,
|
|
host,
|
|
port,
|
|
self.worker_type,
|
|
disaggregation_bootstrap_port,
|
|
)
|
|
|
|
self.node_rank = server_args_dict["node_rank"]
|
|
self.server_host = server_args_dict["host"]
|
|
self.server_port = server_args_dict["port"]
|
|
|
|
if self.args.rollout_external:
|
|
self._init_external(server_args_dict, external_engine_need_check_fields=external_engine_need_check_fields)
|
|
else:
|
|
self._init_normal(server_args_dict)
|
|
|
|
def _init_external(self, expect_server_args, external_engine_need_check_fields):
|
|
logger.info(f"Use external SGLang engine (rank={self.rank}, expect_server_args={expect_server_args})")
|
|
|
|
def _get_actual_server_args():
|
|
response = requests.get(f"http://{self.server_host}:{self.server_port}/get_server_info")
|
|
response.raise_for_status()
|
|
return response.json()
|
|
|
|
def _sanity_check_server_args(actual_server_args, expect_server_args):
|
|
for name in external_engine_need_check_fields:
|
|
expect_value = expect_server_args.get(name)
|
|
actual_value = actual_server_args.get(name)
|
|
assert (
|
|
actual_value == expect_value
|
|
), f"{name=} {expect_value=} {actual_value=} {expect_server_args=} {actual_server_args=}"
|
|
|
|
_wait_server_healthy(
|
|
base_url=f"http://{self.server_host}:{self.server_port}",
|
|
api_key=None,
|
|
is_process_alive=lambda: True,
|
|
)
|
|
actual_server_args = _get_actual_server_args()
|
|
_sanity_check_server_args(actual_server_args, expect_server_args)
|
|
|
|
def _init_normal(self, server_args_dict):
|
|
logger.info(f"Launch HttpServerEngineAdapter at: {self.server_host}:{self.server_port}")
|
|
self.process = launch_server_process(ServerArgs(**server_args_dict))
|
|
|
|
if self.node_rank == 0 and self.router_ip and self.router_port:
|
|
if parse(sglang_router.__version__) <= parse("0.2.1") or self.args.use_slime_router:
|
|
assert (
|
|
self.worker_type == "regular"
|
|
), "pd disaggregation is not supported in old router or slime router."
|
|
response = requests.post(
|
|
f"http://{self.router_ip}:{self.router_port}/add_worker?url=http://{self.server_host}:{self.server_port}"
|
|
)
|
|
else:
|
|
payload = {
|
|
"url": f"http://{self.server_host}:{self.server_port}",
|
|
"worker_type": self.worker_type,
|
|
}
|
|
if self.worker_type == "prefill":
|
|
payload["bootstrap_port"] = server_args_dict["disaggregation_bootstrap_port"]
|
|
response = requests.post(
|
|
f"http://{self.router_ip}:{self.router_port}/workers",
|
|
json=payload,
|
|
)
|
|
response.raise_for_status()
|
|
|
|
def _make_request(self, endpoint: str, payload: dict | None = None):
|
|
"""Make a POST request to the specified endpoint with the given payload.
|
|
|
|
Args:
|
|
endpoint: The API endpoint to call
|
|
payload: The JSON payload to send (default: empty dict)
|
|
|
|
Returns:
|
|
The JSON response from the server
|
|
"""
|
|
if self.node_rank != 0:
|
|
return
|
|
|
|
url = f"http://{self.server_host}:{self.server_port}/{endpoint}"
|
|
response = requests.post(url, json=payload or {})
|
|
try:
|
|
response.raise_for_status()
|
|
except requests.exceptions.HTTPError as e:
|
|
e.add_note(f"{response.text=}")
|
|
raise
|
|
return response.json()
|
|
|
|
def health_generate(self, timeout: float = 5.0) -> bool:
|
|
"""Run /health_generate on the underlying SGLang HTTP server.
|
|
|
|
Args:
|
|
timeout: Timeout for the health request in seconds.
|
|
|
|
Returns:
|
|
True if the server responds with HTTP 200.
|
|
|
|
Raises:
|
|
requests.RequestException: If the request fails for any reason, including timeout.
|
|
"""
|
|
if self.node_rank != 0:
|
|
return True
|
|
|
|
response = requests.get(
|
|
f"http://{self.server_host}:{self.server_port}/health_generate",
|
|
timeout=timeout,
|
|
)
|
|
response.raise_for_status()
|
|
return True
|
|
|
|
def update_weights_from_tensor(
|
|
self,
|
|
serialized_named_tensors: list[str],
|
|
load_format: str | None = None,
|
|
flush_cache: bool = False,
|
|
weight_version: str | None = None,
|
|
):
|
|
"""
|
|
Update model weights from tensor data. The HTTP server will only post meta data, and the real weights will be copied directly from GPUs.
|
|
|
|
Note: The model should be on GPUs rather than CPU for this functionality to work properly.
|
|
If you encounter issues, ensure your model is loaded on GPU devices rather than CPU.
|
|
"""
|
|
payload = {
|
|
"serialized_named_tensors": serialized_named_tensors,
|
|
"load_format": load_format,
|
|
"flush_cache": flush_cache,
|
|
}
|
|
if weight_version is not None:
|
|
payload["weight_version"] = weight_version
|
|
return self._make_request(
|
|
"update_weights_from_tensor",
|
|
payload,
|
|
)
|
|
|
|
def flush_cache(self):
|
|
"""Flush the cache of the server."""
|
|
if self.node_rank != 0:
|
|
return
|
|
# flush cache will not return status_code 200 when there are pending requests
|
|
for _ in range(60):
|
|
try:
|
|
response = requests.get(f"http://{self.server_host}:{self.server_port}/flush_cache")
|
|
if response.status_code == 200:
|
|
break
|
|
except NewConnectionError as e:
|
|
raise e
|
|
except Exception as e:
|
|
logger.info(f"Error flushing cache: {e}")
|
|
time.sleep(1)
|
|
continue
|
|
else:
|
|
raise TimeoutError("Timeout while flushing cache.")
|
|
|
|
def shutdown(self):
|
|
if self.args.rollout_external:
|
|
return
|
|
|
|
logger.info(f"Shutdown engine {self.server_host}:{self.server_port}...")
|
|
if self.node_rank == 0:
|
|
worker_url = f"http://{self.server_host}:{self.server_port}"
|
|
response = None
|
|
if parse(sglang_router.__version__) <= parse("0.2.1") or self.args.use_slime_router:
|
|
response = requests.post(
|
|
f"http://{self.router_ip}:{self.router_port}/remove_worker?url=http://{self.server_host}:{self.server_port}"
|
|
)
|
|
elif parse(sglang_router.__version__) < parse("0.3.0"):
|
|
worker_url = quote(worker_url, safe="")
|
|
response = requests.delete(f"http://{self.router_ip}:{self.router_port}/workers/{worker_url}")
|
|
else:
|
|
try:
|
|
all_workers = requests.get(f"http://{self.router_ip}:{self.router_port}/workers").json()["workers"]
|
|
for worker in all_workers:
|
|
if worker["url"] == worker_url:
|
|
worker_id = worker["id"]
|
|
response = requests.delete(
|
|
f"http://{self.router_ip}:{self.router_port}/workers/{worker_id}"
|
|
)
|
|
break
|
|
else:
|
|
logger.warning(f"Worker {worker_url} not found in router during shutdown.")
|
|
except Exception as e:
|
|
logger.warning(f"Failed to fetch workers list or remove worker: {e}")
|
|
|
|
if response is not None:
|
|
response.raise_for_status()
|
|
kill_process_tree(self.process.pid)
|
|
|
|
def get_weight_version(self):
|
|
if self.node_rank != 0:
|
|
return
|
|
url = f"http://{self.server_host}:{self.server_port}/get_weight_version"
|
|
response = requests.get(url)
|
|
response.raise_for_status()
|
|
return response.json()["weight_version"]
|
|
|
|
def release_memory_occupation(self):
|
|
self.flush_cache()
|
|
return self._make_request("release_memory_occupation")
|
|
|
|
def resume_memory_occupation(self, tags: list[str] = None):
|
|
"""
|
|
Available tags for multi-stage resume: weights, kv_cache
|
|
"""
|
|
return self._make_request(
|
|
"resume_memory_occupation",
|
|
{"tags": tags},
|
|
)
|
|
|
|
def check_weights(self, action: str):
|
|
return self._make_request("weights_checker", {"action": action})
|
|
|
|
def init_weights_update_group(self, master_address, master_port, rank_offset, world_size, group_name, backend):
|
|
return self._make_request(
|
|
"init_weights_update_group",
|
|
{
|
|
"master_address": master_address,
|
|
"master_port": master_port,
|
|
"rank_offset": rank_offset,
|
|
"world_size": world_size,
|
|
"group_name": group_name,
|
|
"backend": backend,
|
|
},
|
|
)
|
|
|
|
def destroy_weights_update_group(self, group_name):
|
|
try:
|
|
return self._make_request(
|
|
"destroy_weights_update_group",
|
|
{
|
|
"group_name": group_name,
|
|
},
|
|
)
|
|
except requests.exceptions.RequestException:
|
|
# catch the case there the engine is just created and does not have the group.
|
|
pass
|
|
|
|
def update_weights_from_distributed(
|
|
self, names, dtypes, shapes, group_name, flush_cache=False, weight_version: str | None = None
|
|
):
|
|
payload = {
|
|
"names": names,
|
|
"dtypes": [str(dtype).replace("torch.", "") for dtype in dtypes],
|
|
"shapes": shapes,
|
|
"group_name": group_name,
|
|
"flush_cache": flush_cache,
|
|
}
|
|
if weight_version is not None:
|
|
payload["weight_version"] = weight_version
|
|
return self._make_request(
|
|
"update_weights_from_distributed",
|
|
payload,
|
|
)
|
|
|
|
def pause_generation(self):
|
|
response = requests.post(f"http://{self.server_host}:{self.server_port}/pause_generation", json={})
|
|
response.raise_for_status()
|
|
return response
|
|
|
|
def continue_generation(self):
|
|
response = requests.post(f"http://{self.server_host}:{self.server_port}/continue_generation", json={})
|
|
response.raise_for_status()
|
|
return response
|
|
|
|
def start_profile(
|
|
self,
|
|
# The output directory
|
|
output_dir: str | None = None,
|
|
# If set, it profile as many as this number of steps.
|
|
# If it is set, profiling is automatically stopped after this step, and
|
|
# the caller doesn't need to run stop_profile.
|
|
start_step: int | None = None,
|
|
num_steps: int | None = None,
|
|
activities: list[str] | None = None,
|
|
profile_by_stage: bool = False,
|
|
with_stack: bool | None = None,
|
|
record_shapes: bool | None = None,
|
|
):
|
|
response = requests.post(
|
|
f"http://{self.server_host}:{self.server_port}/start_profile",
|
|
json={
|
|
"output_dir": output_dir,
|
|
"start_step": start_step,
|
|
"num_steps": num_steps,
|
|
"activities": activities,
|
|
"profile_by_stage": profile_by_stage,
|
|
"with_stack": with_stack,
|
|
"record_shapes": record_shapes,
|
|
},
|
|
)
|
|
response.raise_for_status()
|
|
return response
|
|
|
|
def stop_profile(self):
|
|
response = requests.post(f"http://{self.server_host}:{self.server_port}/stop_profile", json={})
|
|
response.raise_for_status()
|
|
return response
|
|
|
|
|
|
def _compute_server_args(
|
|
args,
|
|
rank,
|
|
dist_init_addr,
|
|
nccl_port,
|
|
host,
|
|
port,
|
|
worker_type: str = "regular",
|
|
disaggregation_bootstrap_port: int | None = None,
|
|
):
|
|
nnodes = max(1, args.rollout_num_gpus_per_engine // args.num_gpus_per_node)
|
|
node_rank = rank % nnodes
|
|
kwargs = {
|
|
"model_path": args.hf_checkpoint,
|
|
"trust_remote_code": True,
|
|
"random_seed": args.seed + rank,
|
|
# memory
|
|
"enable_memory_saver": args.offload_rollout,
|
|
# distributed
|
|
"host": host,
|
|
"port": port,
|
|
"nccl_port": nccl_port,
|
|
"nnodes": nnodes,
|
|
"node_rank": node_rank,
|
|
"dist_init_addr": dist_init_addr,
|
|
"gpu_id_step": 1,
|
|
"base_gpu_id": get_base_gpu_id(args, rank),
|
|
# parallel
|
|
"tp_size": args.rollout_num_gpus_per_engine,
|
|
"dp_size": args.sglang_dp_size,
|
|
"pp_size": args.sglang_pp_size,
|
|
"ep_size": args.sglang_ep_size,
|
|
# always skip warmup to prevent warmup timeout.
|
|
"skip_server_warmup": True,
|
|
}
|
|
|
|
if worker_type == "prefill":
|
|
kwargs["disaggregation_mode"] = "prefill"
|
|
kwargs["load_balance_method"] = "round_robin"
|
|
assert (
|
|
disaggregation_bootstrap_port is not None
|
|
), "disaggregation_bootstrap_port must be set for prefill worker"
|
|
kwargs["disaggregation_bootstrap_port"] = disaggregation_bootstrap_port
|
|
elif worker_type == "decode":
|
|
kwargs["disaggregation_mode"] = "decode"
|
|
kwargs["prefill_round_robin_balance"] = True
|
|
|
|
if args.use_rollout_routing_replay:
|
|
kwargs["enable_return_routed_experts"] = True
|
|
if args.fp16:
|
|
kwargs["dtype"] = "float16"
|
|
external_engine_need_check_fields = [k for k in kwargs.keys() if k not in _EXTERNAL_ENGINE_SKIP_CHECK_FIELDS]
|
|
|
|
unused_keys = set(kwargs.keys())
|
|
for attr in dataclasses.fields(ServerArgs):
|
|
if hasattr(args, f"sglang_{attr.name}") and attr.name not in kwargs:
|
|
kwargs[attr.name] = getattr(args, f"sglang_{attr.name}")
|
|
unused_keys.discard(attr.name)
|
|
|
|
# for compatibility with old args
|
|
if len(unused_keys) > 0:
|
|
logger.info(f"Warning: The following arguments is not supported in the current sglang: {unused_keys}.")
|
|
for key in unused_keys:
|
|
kwargs.pop(key)
|
|
|
|
return kwargs, external_engine_need_check_fields
|
|
|
|
|
|
_EXTERNAL_ENGINE_SKIP_CHECK_FIELDS = [
|
|
"model_path",
|
|
"trust_remote_code",
|
|
"random_seed",
|
|
"nccl_port",
|
|
"dist_init_addr",
|
|
"skip_server_warmup",
|
|
]
|