diff --git a/python/pyproject.toml b/python/pyproject.toml index b934e83351f8..1246b2b00420 100755 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -123,7 +123,7 @@ diffusion = [ ] ray = [ - "ray[default]>=2.55.1", + "ray[default]>=2.56.0", ] tracing = [ diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 25a2a6213032..7387b32bb12b 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -802,10 +802,7 @@ def _launch_subprocesses( # Start the engine info bootstrap server if per-rank info is needed. engine_info_bootstrap_server = None - if ( - server_args.remote_instance_weight_loader_start_seed_via_transfer_engine - and server_args.node_rank == 0 - ): + if server_args.needs_engine_info_bootstrap() and server_args.node_rank == 0: bootstrap_port = server_args.engine_info_bootstrap_port if not is_port_available(bootstrap_port): raise RuntimeError( diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 6b32cb645fca..45b0f821003b 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -610,6 +610,7 @@ def init_memory_saver_adapter(self): def maybe_init_remote_instance_transfer_engine(self): if self.server_args.remote_instance_weight_loader_use_transfer_engine(): self.remote_instance_weight_transporter.init_engine() + self.remote_instance_weight_transporter.maybe_init_parallelism_config() def maybe_init_expert_location_metadata(self): if self.is_draft_worker: diff --git a/python/sglang/srt/model_executor/model_runner_components/remote_instance_weight_transporter.py b/python/sglang/srt/model_executor/model_runner_components/remote_instance_weight_transporter.py index 476e8070c6d9..ba85474997e5 100644 --- a/python/sglang/srt/model_executor/model_runner_components/remote_instance_weight_transporter.py +++ b/python/sglang/srt/model_executor/model_runner_components/remote_instance_weight_transporter.py @@ -53,9 +53,12 @@ def init_engine(self): self.session_id = NetworkAddress( local_ip, self.engine.get_rpc_port() ).to_host_port_str() - self.parallelism_config = RankParallelismConfig.from_parallel_state( - self.tp_rank - ) + + def maybe_init_parallelism_config(self) -> None: + if self.server_args.registers_parallelism_config(): + self.parallelism_config = RankParallelismConfig.from_parallel_state( + self.tp_rank + ) def maybe_register_and_publish_weight_info(self) -> None: if ( @@ -75,7 +78,7 @@ def maybe_register_and_publish_weight_info(self) -> None: # The P2P weight-update client needs each rank's parallelism layout to # map training-side parameters onto this rank's shards. if ( - self.server_args.remote_instance_weight_loader_use_transfer_engine() + self.server_args.registers_parallelism_config() and self.parallelism_config is not None ): self._register_parallelism_config_to_bootstrap() diff --git a/python/sglang/srt/ray/__init__.py b/python/sglang/srt/ray/__init__.py index 5927c789f926..ba619e3bc7a5 100644 --- a/python/sglang/srt/ray/__init__.py +++ b/python/sglang/srt/ray/__init__.py @@ -1,3 +1,4 @@ from sglang.srt.ray.engine import RayEngine +from sglang.srt.ray.http_server import launch_engine, launch_server, serve_http -__all__ = ["RayEngine"] +__all__ = ["RayEngine", "launch_engine", "launch_server", "serve_http"] diff --git a/python/sglang/srt/ray/engine.py b/python/sglang/srt/ray/engine.py index 57c1d16b86df..3e3fc400ba0b 100644 --- a/python/sglang/srt/ray/engine.py +++ b/python/sglang/srt/ray/engine.py @@ -15,9 +15,11 @@ from __future__ import annotations +import contextvars import dataclasses import logging import threading +from contextlib import contextmanager from typing import Callable, List, Optional import ray @@ -36,6 +38,20 @@ logger = logging.getLogger(__name__) +_caller_placement_group: contextvars.ContextVar[Optional[PlacementGroup]] = ( + contextvars.ContextVar("sglang_ray_caller_placement_group", default=None) +) + + +@contextmanager +def _placement_group_context(placement_group: Optional[PlacementGroup]): + """Expose launch-only state while Engine synchronously starts schedulers.""" + token = _caller_placement_group.set(placement_group) + try: + yield + finally: + _caller_placement_group.reset(token) + @dataclasses.dataclass class RaySchedulerInitResult(SchedulerInitResult): @@ -175,6 +191,22 @@ def _validate_custom_placement_group(pg: PlacementGroup, world_size: int) -> Non ) +def get_scheduler_actor_name( + *, + rank0_node_ip: str, + dp_rank: int, + pp_rank: int, + tp_rank: int, + pg_id_hex: str, + bundle_idx: int, +) -> str: + return ( + f"sglang_scheduler_node{rank0_node_ip}" + f"_dp{dp_rank}_pp{pp_rank}_tp{tp_rank}" + f"_pg{pg_id_hex}_bundle{bundle_idx}" + ) + + def _create_scheduler_actor( pg: PlacementGroup, bundle_idx: int, @@ -200,14 +232,25 @@ def _create_scheduler_actor( server_args, tp_rank ) + rdt = server_args.enable_rdt_weight_sync + return SchedulerActor.options( num_cpus=0, num_gpus=1, - name=( - f"sglang_scheduler_node{rank0_node_ip}" - f"_dp{dp_rank}_pp{pp_rank}_tp{tp_rank}" - f"_pg{pg.id.hex()[:8]}_bundle{bundle_idx}" + # run_event_loop() blocks one thread for the actor's lifetime; leave a spare + # for pull_weights, which the trainer calls while generation is paused. + max_concurrency=2 if rdt else 1, + name=get_scheduler_actor_name( + rank0_node_ip=rank0_node_ip, + dp_rank=dp_rank, + pp_rank=pp_rank, + tp_rank=tp_rank, + pg_id_hex=pg.id.hex()[:8], + bundle_idx=bundle_idx, ), + # SchedulerActor calls set_device() with the absolute id from + # get_accelerator_ids(), which is only valid if Ray leaves the mask alone. + runtime_env={"env_vars": {"RAY_EXPERIMENTAL_NOSET_CUDA_VISIBLE_DEVICES": "1"}}, scheduling_strategy=PlacementGroupSchedulingStrategy( placement_group=pg, placement_group_bundle_index=bundle_idx, @@ -227,15 +270,15 @@ def _create_scheduler_actor( class RayEngine(Engine): - """Engine using Ray actors for scheduler processes.""" + """Engine using Ray actors for scheduler processes. + + Same constructor kwargs as :class:`Engine`, plus ``placement_group`` (a Ray + PlacementGroup handle, not a ServerArgs field). + """ - def __init__(self, **kwargs): - placement_group = kwargs.pop("placement_group", None) - if "log_level" not in kwargs: - kwargs["log_level"] = "error" - server_args = ServerArgs(**kwargs) - server_args.override("ray.placement_group", placement_group=placement_group) - super().__init__(server_args=server_args) + def __init__(self, *, placement_group: Optional[PlacementGroup] = None, **kwargs): + with _placement_group_context(placement_group): + super().__init__(**kwargs) def shutdown(self): """Shutdown the engine — kill Ray scheduler actors then local processes.""" @@ -259,7 +302,8 @@ def _launch_scheduler_processes( Tuple of (RaySchedulerInitResult, None). scheduler_procs is None since Ray uses actors instead of mp.Process. """ - pg = server_args.placement_group or ray.util.get_current_placement_group() + placement_group = _caller_placement_group.get() + pg = placement_group or ray.util.get_current_placement_group() if pg is None: from ray.util.placement_group import ( placement_group as create_placement_group, @@ -288,7 +332,7 @@ def _launch_scheduler_processes( ) ray.get(pg.ready()) - is_custom_pg = server_args.placement_group is not None + is_custom_pg = placement_group is not None nnodes = server_args.nnodes world_size = _compute_world_size(server_args) @@ -424,6 +468,7 @@ def wait_for_completion(): pg, bundle_for_node, rank0_node_ip, + is_custom_pg, ), None, ) @@ -436,6 +481,7 @@ def _launch_dp_scheduler_processes( pg, bundle_for_node: Optional[List[int]], rank0_node_ip: str, + is_custom_pg: bool = False, ) -> RaySchedulerInitResult: """Launch DP schedulers via RayDataParallelController.""" from sglang.srt.ray.data_parallel_controller import ( @@ -461,16 +507,16 @@ def _launch_dp_scheduler_processes( server_args, dist_init_addr=f"{rank0_node_ip}:{port_args.nccl_port}", ) - # dataclasses.replace only copies declared fields; placement_group is - # a dynamic attribute that must be manually appended after the rebuild. - dp_server_args.override( - "ray.placement_group", placement_group=server_args.placement_group - ) # Create the DP controller in-process. This blocks until all actors # are initialized and their event loops have started. controller = RayDataParallelController( - dp_server_args, port_args, pg, bundle_for_node, rank0_node_ip + dp_server_args, + port_args, + pg, + bundle_for_node, + rank0_node_ip, + is_custom_pg, ) # Start the DP controller's event loop in a daemon thread. diff --git a/python/sglang/srt/ray/http_server.py b/python/sglang/srt/ray/http_server.py index 7d7e23567aa2..f8c568956934 100644 --- a/python/sglang/srt/ray/http_server.py +++ b/python/sglang/srt/ray/http_server.py @@ -15,6 +15,8 @@ from typing import Callable, Optional +from ray.util.placement_group import PlacementGroup + from sglang.srt.entrypoints.engine import ( init_tokenizer_manager, run_detokenizer_process, @@ -23,41 +25,48 @@ from sglang.srt.server_args import ServerArgs -def launch_server( +def launch_engine( server_args: ServerArgs, init_tokenizer_manager_func: Callable = init_tokenizer_manager, run_scheduler_process_func: Callable = run_scheduler_process, run_detokenizer_process_func: Callable = run_detokenizer_process, + *, + placement_group: Optional[PlacementGroup] = None, +): + """Create RayEngine subprocesses / SchedulerActors.""" + from sglang.srt.ray.engine import RayEngine, _placement_group_context + + with _placement_group_context(placement_group): + return RayEngine._launch_subprocesses( + server_args, + init_tokenizer_manager_func=init_tokenizer_manager_func, + run_scheduler_process_func=run_scheduler_process_func, + run_detokenizer_process_func=run_detokenizer_process_func, + ) + + +def serve_http( + engine, + server_args: ServerArgs, execute_warmup_func: Optional[Callable] = None, launch_callback: Optional[Callable[[], None]] = None, ): - """Launch HTTP server with Ray-based scheduler actors. - - Mirrors http_server.launch_server() but uses RayEngine for scheduler launching. - """ + """Block in uvicorn. ``engine`` is the 5-tuple from ``launch_engine``.""" from sglang.srt.entrypoints.http_server import ( _execute_server_warmup, _setup_and_run_http_server, ) - from sglang.srt.ray.engine import RayEngine if execute_warmup_func is None: execute_warmup_func = _execute_server_warmup - server_args.override("ray.http_server.clear_placement_group", placement_group=None) - ( tokenizer_manager, template_manager, port_args, scheduler_init_result, subprocess_watchdog, - ) = RayEngine._launch_subprocesses( - server_args, - init_tokenizer_manager_func=init_tokenizer_manager_func, - run_scheduler_process_func=run_scheduler_process_func, - run_detokenizer_process_func=run_detokenizer_process_func, - ) + ) = engine _setup_and_run_http_server( server_args, @@ -69,3 +78,28 @@ def launch_server( execute_warmup_func=execute_warmup_func, launch_callback=launch_callback, ) + + +def launch_server( + server_args: ServerArgs, + init_tokenizer_manager_func: Callable = init_tokenizer_manager, + run_scheduler_process_func: Callable = run_scheduler_process, + run_detokenizer_process_func: Callable = run_detokenizer_process, + execute_warmup_func: Optional[Callable] = None, + launch_callback: Optional[Callable[[], None]] = None, +): + """Launch HTTP server with Ray-based scheduler actors. + + Mirrors http_server.launch_server() but uses RayEngine for scheduler launching. + """ + serve_http( + launch_engine( + server_args, + init_tokenizer_manager_func=init_tokenizer_manager_func, + run_scheduler_process_func=run_scheduler_process_func, + run_detokenizer_process_func=run_detokenizer_process_func, + ), + server_args, + execute_warmup_func=execute_warmup_func, + launch_callback=launch_callback, + ) diff --git a/python/sglang/srt/ray/scheduler_actor.py b/python/sglang/srt/ray/scheduler_actor.py index e9090ec9a576..e3a95b704a40 100644 --- a/python/sglang/srt/ray/scheduler_actor.py +++ b/python/sglang/srt/ray/scheduler_actor.py @@ -16,9 +16,10 @@ from __future__ import annotations import logging -from typing import TYPE_CHECKING, Any, Dict, Optional +from typing import TYPE_CHECKING, Any, Dict, List, Optional import ray +from ray import ObjectRef if TYPE_CHECKING: from sglang.srt.server_args import PortArgs, ServerArgs @@ -120,6 +121,38 @@ def get_info(self) -> Dict[str, Any]: """Return scheduler initialization info for handshake.""" return self.scheduler.get_init_info() + def register_weight_for_rdt(self) -> None: + """Pin model parameters with NIXL for repeated RDT pulls.""" + if self.scheduler.server_args.enable_memory_saver: + return + + import torch + from ray.experimental import register_nixl_memory + + torch.cuda.set_device(self.scheduler.ps.gpu_id) + model = self.scheduler.tp_worker.model_runner.model + for _, param in model.named_parameters(): + register_nixl_memory(param.data) + + def pull_weights( + self, weights_refs: List[ObjectRef], param_names: List[str] + ) -> None: + """Pull a pre-sharded weight bucket from the trainer via RDT zero-copy. + + ``weights_refs`` is a list so Ray does not resolve the ref on call. + """ + import torch + from ray.experimental import set_target_for_ref + + # Runs on a different thread than run_event_loop, which owns the device binding. + torch.cuda.set_device(self.scheduler.ps.gpu_id) + + model = self.scheduler.tp_worker.model_runner.model + params_dict = dict(model.named_parameters()) + target_buffers = [params_dict[name].data for name in param_names] + set_target_for_ref(weights_refs[0], target_buffers) + ray.get(weights_refs[0]) + def run_event_loop(self) -> None: """Run the scheduler's event loop. Blocks until shutdown.""" try: diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 04ba960b9ad2..8687e29a4d5c 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -2669,6 +2669,14 @@ class ServerArgs: bool, "Start seed server via transfer engine backend for remote instance weight loader.", ] = False + enable_engine_info_bootstrap: A[ + bool, + "Start the EngineInfoBootstrapServer and register per-rank parallelism config, without the mooncake/verbs P2P transfer-engine seeding.", + ] = False + enable_rdt_weight_sync: A[ + bool, + "Expose SchedulerActor.pull_weights for RDT (Ray Direct Transport / NIXL) weight sync. Requires --use-ray; implies --enable-engine-info-bootstrap.", + ] = False engine_info_bootstrap_port: A[ int, "Port for the engine info bootstrap server. Default is 6789. Must be set explicitly when running multiple instances on the same node.", @@ -2927,6 +2935,8 @@ def __post_init__(self): # _handle_model_specific_adjustments never runs. self._resolved_overrides = [] + self._handle_rdt_weight_sync() + if self.model_path.lower() in ["none", "dummy"]: return @@ -3160,6 +3170,14 @@ def _handle_model_source_paths(self): ): ObjectStorageModel.download_and_get_path(self.tokenizer_path) + def _handle_rdt_weight_sync(self): + if not self.enable_rdt_weight_sync: + return + assert ( + self.use_ray + ), "--enable-rdt-weight-sync requires --use-ray: the trainer pulls weights through SchedulerActors." + self.enable_engine_info_bootstrap = True + def _handle_pd_disaggregation(self): from sglang.srt.arg_groups.pd_disaggregation_hook import ( handle_pd_disaggregation, @@ -8091,6 +8109,20 @@ def remote_instance_weight_loader_use_transfer_engine(self): else: return False + def needs_engine_info_bootstrap(self) -> bool: + """Whether this node (rank 0) hosts the EngineInfoBootstrapServer.""" + return ( + self.remote_instance_weight_loader_start_seed_via_transfer_engine + or self.enable_engine_info_bootstrap + ) + + def registers_parallelism_config(self) -> bool: + """Whether this rank publishes its parallelism config to the bootstrap server.""" + return ( + self.remote_instance_weight_loader_use_transfer_engine() + or self.enable_engine_info_bootstrap + ) + def describe_kv_events_publisher(self) -> Optional[dict]: """Return a structured description of this server's KV-event publisher, or `None` if publishing is disabled / misconfigured.