Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion python/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -123,7 +123,7 @@ diffusion = [
]

ray = [
"ray[default]>=2.55.1",
"ray[default]>=2.56.0",
]

tracing = [
Expand Down
5 changes: 1 addition & 4 deletions python/sglang/srt/entrypoints/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
1 change: 1 addition & 0 deletions python/sglang/srt/model_executor/model_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand All @@ -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()
Expand Down
3 changes: 2 additions & 1 deletion python/sglang/srt/ray/__init__.py
Original file line number Diff line number Diff line change
@@ -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"]
86 changes: 66 additions & 20 deletions python/sglang/srt/ray/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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):
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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."""
Expand All @@ -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,
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -424,6 +468,7 @@ def wait_for_completion():
pg,
bundle_for_node,
rank0_node_ip,
is_custom_pg,
),
None,
)
Expand All @@ -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 (
Expand All @@ -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.
Expand Down
62 changes: 48 additions & 14 deletions python/sglang/srt/ray/http_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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,
)
35 changes: 34 additions & 1 deletion python/sglang/srt/ray/scheduler_actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down
Loading
Loading