初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD Source: Original Platform
This commit is contained in:
3
slime/backends/sglang_utils/__init__.py
Normal file
3
slime/backends/sglang_utils/__init__.py
Normal file
@@ -0,0 +1,3 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
BIN
slime/backends/sglang_utils/__pycache__/__init__.cpython-312.pyc
Normal file
BIN
slime/backends/sglang_utils/__pycache__/__init__.cpython-312.pyc
Normal file
Binary file not shown.
Binary file not shown.
Binary file not shown.
133
slime/backends/sglang_utils/arguments.py
Normal file
133
slime/backends/sglang_utils/arguments.py
Normal file
@@ -0,0 +1,133 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import sglang
|
||||
from packaging.version import parse
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from slime.utils.http_utils import _wrap_ipv6
|
||||
|
||||
|
||||
# TODO: use all sglang router arguments with `--sglang-router` prefix
|
||||
def add_sglang_router_arguments(parser):
|
||||
"""
|
||||
Add arguments to the parser for the SGLang router.
|
||||
"""
|
||||
parser.add_argument(
|
||||
"--sglang-router-ip",
|
||||
type=str,
|
||||
default=None,
|
||||
help="IP address of the SGLang router",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sglang-router-port",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Port of the SGLang router",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sglang-router-request-timeout-secs",
|
||||
type=int,
|
||||
default=14400,
|
||||
help="Timeout for requests to the SGLang router in seconds",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
def add_sglang_arguments(parser):
|
||||
"""
|
||||
Add arguments to the parser for the SGLang server.
|
||||
"""
|
||||
parser = add_sglang_router_arguments(parser)
|
||||
parser.add_argument("--sglang-server-concurrency", type=int, default=512)
|
||||
|
||||
old_add_argument = parser.add_argument
|
||||
|
||||
skipped_args = [
|
||||
"model_path",
|
||||
"dtype",
|
||||
"trust_remote_code",
|
||||
"random_seed",
|
||||
# memory
|
||||
"enable_memory_saver",
|
||||
# distributed
|
||||
"tp_size",
|
||||
"port",
|
||||
"nnodes",
|
||||
"node_rank",
|
||||
"dist_init_addr",
|
||||
"gpu_id_step",
|
||||
"base_gpu_id",
|
||||
"nccl_port",
|
||||
"skip_server_warmup",
|
||||
"enable_return_routed_experts",
|
||||
]
|
||||
|
||||
def new_add_argument_wrapper(*name_or_flags, **kwargs):
|
||||
"""
|
||||
Add arguments to the parser, ensuring that the server arguments are prefixed and skippable.
|
||||
"""
|
||||
# Determine the canonical name for skip check (e.g., "model_path")
|
||||
canonical_name_for_skip_check = None
|
||||
if "dest" in kwargs:
|
||||
canonical_name_for_skip_check = kwargs["dest"]
|
||||
else:
|
||||
for flag_name_candidate in name_or_flags:
|
||||
if isinstance(flag_name_candidate, str) and flag_name_candidate.startswith("--"):
|
||||
# Derive from first long flag: --foo-bar -> foo_bar
|
||||
stem = flag_name_candidate[2:]
|
||||
canonical_name_for_skip_check = stem.replace("-", "_")
|
||||
break
|
||||
# If no long flag and no dest, skip logic might not catch it unless short flags imply a dest.
|
||||
|
||||
if canonical_name_for_skip_check and canonical_name_for_skip_check in skipped_args:
|
||||
return # Skip this entire argument definition
|
||||
|
||||
# If not skipped, proceed to prefix flags and dest
|
||||
new_name_or_flags_list = []
|
||||
for item_flag in name_or_flags:
|
||||
if isinstance(item_flag, str) and item_flag.startswith("-"):
|
||||
original_flag_stem = item_flag.lstrip("-") # "foo-bar" from "--foo-bar", or "f" from "-f"
|
||||
prefixed_item = f"--sglang-{original_flag_stem}"
|
||||
new_name_or_flags_list.append(prefixed_item)
|
||||
else:
|
||||
# Positional arguments or non-string items
|
||||
new_name_or_flags_list.append(item_flag)
|
||||
|
||||
# Prepare kwargs for the actual add_argument call.
|
||||
# Make a copy to avoid modifying the original kwargs dict.
|
||||
final_kwargs = kwargs.copy()
|
||||
|
||||
# If 'dest' is explicitly provided and is a string, prefix it.
|
||||
# This ensures the attribute on the args namespace becomes, e.g., args.sglang_dest_name.
|
||||
if "dest" in final_kwargs and isinstance(final_kwargs["dest"], str):
|
||||
original_dest = final_kwargs["dest"]
|
||||
# Avoid double prefixing if dest somehow already starts with sglang_
|
||||
if not original_dest.startswith("sglang_"):
|
||||
final_kwargs["dest"] = f"sglang_{original_dest}"
|
||||
# If 'dest' is not explicitly provided (or is None/not a string),
|
||||
# argparse will derive 'dest' from the (now prefixed) flag names.
|
||||
# E.g., if the first flag is "--sglang-foo-bar", argparse sets dest to "sglang_foo_bar".
|
||||
|
||||
old_add_argument(*new_name_or_flags_list, **final_kwargs)
|
||||
|
||||
parser.add_argument = new_add_argument_wrapper
|
||||
ServerArgs.add_cli_args(parser)
|
||||
parser.add_argument = old_add_argument
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
def validate_args(args):
|
||||
if parse(sglang.__version__) == parse("0.4.10") and getattr(args, "sglang_enable_ep_moe", False):
|
||||
args.sglang_expert_parallel_size = args.rollout_num_gpus_per_engine
|
||||
|
||||
args.sglang_tp_size = args.rollout_num_gpus_per_engine
|
||||
args.sglang_dp_size = args.sglang_data_parallel_size
|
||||
args.sglang_pp_size = args.sglang_pipeline_parallel_size
|
||||
args.sglang_ep_size = args.sglang_expert_parallel_size
|
||||
|
||||
if args.sglang_dp_size > 1:
|
||||
assert args.sglang_enable_dp_attention
|
||||
|
||||
if getattr(args, "sglang_router_ip", None):
|
||||
args.sglang_router_ip = _wrap_ipv6(args.sglang_router_ip)
|
||||
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