Skip to content
Closed
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
4 changes: 4 additions & 0 deletions .github/workflows/pr-test.yml
Original file line number Diff line number Diff line change
Expand Up @@ -669,6 +669,10 @@ jobs:
"num_gpus": 0,
"test_file": "test_external_sglang_engines.py"
},
{
"num_gpus": 0,
"test_file": "test_streaming_rollout.py"
},
{
"num_gpus": 0,
"test_file": "test_empty_colocated_weight_bucket.py"
Expand Down
1 change: 1 addition & 0 deletions .github/workflows/pr-test.yml.j2
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,7 @@
{'test_file': 'test_reloadable_process_group_world.py', 'num_gpus': 0},
{'test_file': 'test_placement_group.py', 'num_gpus': 0},
{'test_file': 'test_external_sglang_engines.py', 'num_gpus': 0},
{'test_file': 'test_streaming_rollout.py', 'num_gpus': 0},
{'test_file': 'test_empty_colocated_weight_bucket.py', 'num_gpus': 0},
{'test_file': 'test_expert_routing.py', 'num_gpus': 0},
{'test_file': 'test_glm5_indexer_short_context.py', 'num_gpus': 0},
Expand Down
8 changes: 6 additions & 2 deletions slime/backends/megatron_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -572,9 +572,13 @@ def update_weights(self) -> None:
if self.args.debug_train_only or self.args.debug_rollout_only:
return

if self.args.use_fault_tolerance:
dynamic_discovery_path = self.args.rollout_external_dynamic_discovery_path
if dynamic_discovery_path or self.args.use_fault_tolerance:
if dist.get_rank() == 0:
ray.get(self.rollout_manager.recover_updatable_engines.remote())
if dynamic_discovery_path:
ray.get(self.rollout_manager.refresh_updatable_engines.remote())
if self.args.use_fault_tolerance:
ray.get(self.rollout_manager.recover_updatable_engines.remote())
dist.barrier(group=get_gloo_group())

(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -65,7 +65,7 @@ def connect_rollout_engines(
"""
Create NCCL "slime-pp_{pp_rank}" if PP source (DP=TP=0). Lock prevents concurrent broadcasts.
"""
self.rollout_engines = rollout_engines
previous_rollout_engines = getattr(self, "rollout_engines", [])
self.rollout_engine_lock = rollout_engine_lock
self._engine_gpu_counts = engine_gpu_counts

Expand All @@ -82,14 +82,15 @@ def connect_rollout_engines(
if self._is_pp_src_rank:
if self._model_update_groups is not None:
disconnect_rollout_engines_from_distributed(
self.args, self._group_name, self._model_update_groups, self.rollout_engines
self.args, self._group_name, self._model_update_groups, previous_rollout_engines
)
self._model_update_groups = connect_rollout_engines_from_distributed(
self.args,
self._group_name,
rollout_engines,
engine_gpu_counts=engine_gpu_counts,
)
self.rollout_engines = rollout_engines

def disconnect_rollout_engines(self) -> None:
if not getattr(self, "_is_pp_src_rank", False) or self._model_update_groups is None:
Expand Down
198 changes: 149 additions & 49 deletions slime/backends/sglang_utils/external.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,10 +4,13 @@

import dataclasses
import logging
from itertools import accumulate
from urllib.parse import urlparse

import requests

from slime.utils.misc import load_function

logger = logging.getLogger(__name__)


Expand Down Expand Up @@ -53,8 +56,7 @@ def normalize_external_engine_addr(addr: str) -> str:
parsed = urlparse(addr)
if parsed.scheme != "http" or parsed.hostname is None or parsed.port is None:
raise ValueError(
f"Invalid external SGLang engine address {addr!r}. "
"Use host:port or http://host:port (IPv6 must be bracketed)."
f"Invalid external SGLang engine address {addr!r}. Use host:port or http://host:port (IPv6 must be bracketed)."
)
return addr

Expand All @@ -65,6 +67,7 @@ def external_engine_init_kwargs(info: ExternalEngineInfo) -> dict:
"nccl_port": None,
"host": info.host,
"port": info.port,
"control_url": info.url,
}
if info.worker_type == "prefill":
init_kwargs["disaggregation_bootstrap_port"] = info.disaggregation_bootstrap_port
Expand All @@ -79,7 +82,7 @@ def get_server_info(url: str, timeout: float = 30.0) -> dict:
response.raise_for_status()
return response.json()
except Exception as exc:
errors.append(f"{endpoint}: {exc}")
errors.append(f"{url}{endpoint}: {exc}")
raise RuntimeError(f"Failed to fetch SGLang server info from {url}: {'; '.join(errors)}")


Expand Down Expand Up @@ -120,15 +123,31 @@ def discover_external_engines(addrs: list[str], timeout: float = 30.0) -> list[E
return infos


def apply_external_engine_info_to_args(args, logger=None) -> None:
"""Detect external engines and store the derived topology on ``args``."""
addrs = args.rollout_external_engine_addrs
def external_engine_addrs_from_args(args) -> list[str]:
"""Return external-engine control base URLs from the CLI or dynamic discovery."""
dynamic_discovery_path = args.rollout_external_dynamic_discovery_path
if dynamic_discovery_path:
addrs = load_function(dynamic_discovery_path)(args)
else:
addrs = args.rollout_external_engine_addrs

if not addrs:
raise ValueError("apply_external_engine_info_to_args requires --rollout-external-engine-addrs.")
raise ValueError(
"External rollout requires --rollout-external-engine-addrs or "
"--rollout-external-dynamic-discovery-path."
)
if not all(isinstance(addr, str) for addr in addrs):
raise TypeError("External engine discovery must return a list of control base URLs.")
return addrs


infos = discover_external_engines(addrs)
if not infos:
raise ValueError("--rollout-external-engine-addrs did not contain any engines.")
def discover_external_engine_infos(args) -> list[ExternalEngineInfo]:
return discover_external_engines(external_engine_addrs_from_args(args))


def apply_external_engine_info_to_args(args, logger=None) -> None:
"""Detect external engines and store the derived topology on ``args``."""
infos = discover_external_engine_infos(args)

args.rollout_external_engine_infos = [info.to_dict() for info in infos]
args.rollout_num_engines = len(infos)
Expand All @@ -147,25 +166,118 @@ def apply_external_engine_info_to_args(args, logger=None) -> None:
logger.info(f"Detected external SGLang engines: {summary}")


def _topology_signature(infos: list[ExternalEngineInfo]) -> tuple:
"""Return fields that require a new control/NCCL membership when changed."""
return tuple(
sorted(
(
info.url,
info.worker_type,
info.num_gpus,
info.disaggregation_bootstrap_port,
info.server_info.get("tp_size"),
info.server_info.get("pp_size"),
info.server_info.get("dp_size"),
info.server_info.get("ep_size"),
)
for info in infos
)
)


def _start_external_engine_actors(args, infos, router_ip, router_port, *, register_to_router):
import ray

from slime.backends.sglang_utils.sglang_engine import SGLangEngine
from slime.ray.utils import add_default_ray_env_vars

engines = []
init_handles = []
RolloutRayActor = ray.remote(SGLangEngine)
for rank, info in enumerate(infos):
rollout_engine = RolloutRayActor.options(
num_cpus=0.2,
num_gpus=0,
runtime_env={"env_vars": add_default_ray_env_vars()},
).remote(
args=args,
rank=rank,
worker_type=info.worker_type,
base_gpu_id=0,
num_gpus_per_engine=info.num_gpus,
)
init_kwargs = external_engine_init_kwargs(info)
init_kwargs["register_to_router"] = register_to_router
init_handles.append(rollout_engine.init.remote(**init_kwargs, router_ip=router_ip, router_port=router_port))
engines.append(rollout_engine)
return engines, init_handles


@dataclasses.dataclass
class ExternalRolloutServer:
"""Rollout server backed by pre-launched external SGLang engines."""

engines: list
engine_gpu_counts: list[int]
engine_gpu_offsets: list[int]
engine_parallel_configs: list[dict[str, int]]
args: object
engine_infos: list[ExternalEngineInfo]
register_to_router: bool
router_ip: str | None = None
router_port: int | None = None
model_name: str = "default"
update_weights: bool = True
num_new_engines: int = 0
retired_engines: list = dataclasses.field(default_factory=list)
server_groups: list = dataclasses.field(default_factory=list)
engine_parallel_configs: list[dict] = dataclasses.field(default_factory=list)

@property
def all_engines(self):
return self.engines

def refresh(self) -> bool:
"""Refresh dynamic external-engine membership before a weight update."""
if not self.args.rollout_external_dynamic_discovery_path:
return False

infos = discover_external_engine_infos(self.args)
if _topology_signature(infos) == _topology_signature(self.engine_infos):
return False

engines, init_handles = _start_external_engine_actors(
self.args,
infos,
self.router_ip,
self.router_port,
register_to_router=self.register_to_router,
)
if init_handles:
import ray

ray.get(init_handles)
self.retired_engines.extend(self.engines)
self.engines = engines
self.engine_gpu_counts = [info.num_gpus for info in infos]
self.engine_gpu_offsets = list(accumulate([0, *self.engine_gpu_counts[:-1]]))
self.engine_parallel_configs = [info.parallel_config for info in infos]
self.engine_infos = infos
self.num_new_engines = len(engines)
self.args.rollout_external_engine_infos = [info.to_dict() for info in infos]
self.args.rollout_num_engines = len(infos)
self.args.rollout_num_gpus = sum(self.engine_gpu_counts)
logger.info("Refreshed external rollout engines: %s", [info.url for info in infos])
return True

def clear_num_new_engines(self) -> None:
self.num_new_engines = 0
if self.retired_engines:
import ray

for engine in self.retired_engines:
ray.kill(engine, no_restart=True)
self.retired_engines.clear()

def recover(self):
logger.warning("Fault tolerance is not supported for external rollout engines; skip recover.")

Expand All @@ -186,60 +298,48 @@ def external_engine_infos_from_args(args) -> list[ExternalEngineInfo]:
raw_infos = getattr(args, "rollout_external_engine_infos", None)
if raw_infos is None:
raise RuntimeError(
"External rollout engine info is missing. "
"apply_external_engine_info_to_args must run before starting external rollout servers."
"External rollout engine info is missing. apply_external_engine_info_to_args must run before starting external rollout servers."
)
return [ExternalEngineInfo(**info) if isinstance(info, dict) else info for info in raw_infos]


def start_external_rollout_servers(args, *, start_router) -> tuple[dict[str, ExternalRolloutServer], list]:
import ray

from slime.backends.sglang_utils.sglang_engine import SGLangEngine
from slime.ray.utils import add_default_ray_env_vars

infos = external_engine_infos_from_args(args)
router_ip, router_port = start_router(args, has_pd_disaggregation=any(info.is_pd_worker for info in infos))
rollout_url = (
normalize_external_engine_addr(args.rollout_external_rollout_url)
if args.rollout_external_rollout_url
else None
)
if args.rollout_external_dynamic_discovery_path and rollout_url is None:
raise ValueError("--rollout-external-dynamic-discovery-path requires --rollout-external-rollout-url.")
if rollout_url is None:
router_ip, router_port = start_router(args, has_pd_disaggregation=any(info.is_pd_worker for info in infos))
else:
parsed = urlparse(rollout_url)
assert parsed.hostname is not None and parsed.port is not None
router_ip, router_port = parsed.hostname, parsed.port
args.sglang_router_ip = router_ip
args.sglang_router_port = router_port

engines = []
engine_gpu_counts = []
engine_gpu_offsets = []
init_handles = []
RolloutRayActor = ray.remote(SGLangEngine)
gpu_offset = 0
for rank, info in enumerate(infos):
rollout_engine = RolloutRayActor.options(
num_cpus=0.2,
num_gpus=0,
runtime_env={"env_vars": add_default_ray_env_vars()},
).remote(
args=args,
rank=rank,
worker_type=info.worker_type,
base_gpu_id=0,
num_gpus_per_engine=info.num_gpus,
)
engines.append(rollout_engine)
engine_gpu_counts.append(info.num_gpus)
engine_gpu_offsets.append(gpu_offset)
gpu_offset += info.num_gpus
init_handles.append(
rollout_engine.init.remote(
**external_engine_init_kwargs(info),
router_ip=router_ip,
router_port=router_port,
)
)

engines, init_handles = _start_external_engine_actors(
args,
infos,
router_ip,
router_port,
register_to_router=rollout_url is None,
)
engine_gpu_counts = [info.num_gpus for info in infos]
engine_gpu_offsets = list(accumulate([0, *engine_gpu_counts[:-1]]))
args.sglang_model_routers = {"default": (router_ip, router_port)}
servers = {
"default": ExternalRolloutServer(
engines=engines,
engine_gpu_counts=engine_gpu_counts,
engine_gpu_offsets=engine_gpu_offsets,
engine_parallel_configs=[info.parallel_config for info in infos],
args=args,
engine_infos=infos,
register_to_router=rollout_url is None,
router_ip=router_ip,
router_port=router_port,
model_name="default",
Expand Down
Loading
Loading