初始化项目,由ModelHub XC社区提供模型

Model: ayh015/myLightningOPD
Source: Original Platform
This commit is contained in:
ModelHub XC
2026-08-27 23:50:14 +08:00
commit d4e0a1af66
368 changed files with 559583 additions and 0 deletions

View File

@@ -0,0 +1,3 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

View 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)

View 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",
]