初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD Source: Original Platform
This commit is contained in:
491
slime/backends/sglang_utils/sglang_engine.py
Normal file
491
slime/backends/sglang_utils/sglang_engine.py
Normal file
@@ -0,0 +1,491 @@
|
||||
# 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",
|
||||
]
|
||||
Reference in New Issue
Block a user