From 638a19c1e92dc29a522784096ee7e4e3522095d8 Mon Sep 17 00:00:00 2001 From: William Arnold <7565007+Aphoh@users.noreply.github.com> Date: Sun, 26 Jul 2026 22:19:36 +0000 Subject: [PATCH 01/12] Support externally routed streaming rollouts --- .github/workflows/pr-test.yml | 4 + .github/workflows/pr-test.yml.j2 | 1 + slime/backends/sglang_utils/external.py | 36 +++- slime/rollout/sglang_rollout.py | 19 +- slime/rollout/sglang_streaming_rollout.py | 118 +++++------ slime/rollout/streaming_utils.py | 37 ++++ slime/utils/arguments.py | 23 ++- tests/test_external_sglang_engines.py | 55 ++++++ tests/test_streaming_rollout.py | 231 ++++++++++++++++++++++ 9 files changed, 457 insertions(+), 67 deletions(-) create mode 100644 slime/rollout/streaming_utils.py create mode 100644 tests/test_streaming_rollout.py diff --git a/.github/workflows/pr-test.yml b/.github/workflows/pr-test.yml index cb7432f974..1c6abeabf1 100644 --- a/.github/workflows/pr-test.yml +++ b/.github/workflows/pr-test.yml @@ -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" diff --git a/.github/workflows/pr-test.yml.j2 b/.github/workflows/pr-test.yml.j2 index fba22ea91c..f47d8a9f6d 100644 --- a/.github/workflows/pr-test.yml.j2 +++ b/.github/workflows/pr-test.yml.j2 @@ -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}, diff --git a/slime/backends/sglang_utils/external.py b/slime/backends/sglang_utils/external.py index 41efc86fb4..8cfc0cb71f 100644 --- a/slime/backends/sglang_utils/external.py +++ b/slime/backends/sglang_utils/external.py @@ -8,6 +8,8 @@ import requests +from slime.utils.misc import load_function + logger = logging.getLogger(__name__) @@ -122,13 +124,23 @@ def discover_external_engines(addrs: list[str], timeout: float = 30.0) -> list[E 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 - if not addrs: - raise ValueError("apply_external_engine_info_to_args requires --rollout-external-engine-addrs.") + discovery_path = getattr(args, "rollout_external_engine_discovery_path", None) + if discovery_path is not None: + discovered = load_function(discovery_path)(args) + infos = [ExternalEngineInfo(**info) if isinstance(info, dict) else info for info in discovered] + else: + addrs = args.rollout_external_engine_addrs + if not addrs: + raise ValueError( + "External rollout requires --rollout-external-engine-addrs or " + "--rollout-external-engine-discovery-path." + ) + infos = discover_external_engines(addrs) - infos = discover_external_engines(addrs) if not infos: - raise ValueError("--rollout-external-engine-addrs did not contain any engines.") + raise ValueError("External rollout engine discovery returned no engines.") + if not all(isinstance(info, ExternalEngineInfo) for info in infos): + raise TypeError("External rollout engine discovery must return ExternalEngineInfo objects or dictionaries.") args.rollout_external_engine_infos = [info.to_dict() for info in infos] args.rollout_num_engines = len(infos) @@ -192,10 +204,20 @@ def external_engine_infos_from_args(args) -> list[ExternalEngineInfo]: return [ExternalEngineInfo(**info) if isinstance(info, dict) else info for info in raw_infos] +def get_external_engine_class(args): + """Return the control actor class for externally managed engines.""" + engine_class_path = getattr(args, "rollout_external_engine_class_path", None) + if engine_class_path is not None: + return load_function(engine_class_path) + + from slime.backends.sglang_utils.sglang_engine import SGLangEngine + + return SGLangEngine + + 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) @@ -207,7 +229,7 @@ def start_external_rollout_servers(args, *, start_router) -> tuple[dict[str, Ext engine_gpu_counts = [] engine_gpu_offsets = [] init_handles = [] - RolloutRayActor = ray.remote(SGLangEngine) + RolloutRayActor = ray.remote(get_external_engine_class(args)) gpu_offset = 0 for rank, info in enumerate(infos): rollout_engine = RolloutRayActor.options( diff --git a/slime/rollout/sglang_rollout.py b/slime/rollout/sglang_rollout.py index b76013bfbe..ae36b9b14c 100644 --- a/slime/rollout/sglang_rollout.py +++ b/slime/rollout/sglang_rollout.py @@ -132,6 +132,8 @@ def reset(self) -> None: self.remaining_batch_size = 0 self.pendings = set() self.aborted = False + self.streaming_generation = False + self.streaming_tasks = set() def submit_generate_tasks(self, samples: list[list[Sample]]) -> None: for group in samples: @@ -341,10 +343,15 @@ async def abort(args: Namespace, rollout_id: int) -> list[list[Sample]]: assert not state.aborted state.aborted = True - response = await get(f"http://{args.sglang_router_ip}:{args.sglang_router_port}/workers") - urls = [worker["url"] for worker in response["workers"]] - - await abort_servers_until_idle(urls) + if state.streaming_generation: + streaming_tasks = list(state.streaming_tasks) + for task in streaming_tasks: + task.cancel() + await asyncio.gather(*streaming_tasks, return_exceptions=True) + else: + response = await get(f"http://{args.sglang_router_ip}:{args.sglang_router_port}/workers") + urls = [worker["url"] for worker in response["workers"]] + await abort_servers_until_idle(urls) # make sure all the pending tasks are finished count = 0 @@ -357,6 +364,10 @@ async def abort(args: Namespace, rollout_id: int) -> list[list[Sample]]: # for partial rollout, collect the partial samples into the data buffer for task in done: group = task.result() + if not all(sample.status == Sample.Status.ABORTED for sample in group): + continue + if not any(sample.response_length > 0 for sample in group): + continue for sample in group: if sample.response and "start_rollout_id" not in sample.metadata: sample.metadata["start_rollout_id"] = rollout_id diff --git a/slime/rollout/sglang_streaming_rollout.py b/slime/rollout/sglang_streaming_rollout.py index 12471148c9..03c81d4a65 100644 --- a/slime/rollout/sglang_streaming_rollout.py +++ b/slime/rollout/sglang_streaming_rollout.py @@ -17,19 +17,18 @@ partial-rollout buffer hand-off) is still owned by ``sglang_rollout``; this file only replaces the inner HTTP call. -sglang's default streaming output is cumulative — server-side -``state.output_token_logprobs`` accumulates and every chunk references the -full list-so-far (see ``tokenizer_manager.py``). If anyone ever flips -``--incremental-streaming-output`` on the sglang server, the text/output_ids -deltas will need different handling here. +Both cumulative and disjoint SGLang streams are accepted. The latter is used +when the server enables incremental streaming output. """ +import asyncio import json import logging from argparse import Namespace from typing import Any from slime.rollout.sglang_rollout import GenerateState, _prepare_prompt_ids +from slime.rollout.streaming_utils import merge_stream_chunk from slime.utils import http_utils from slime.utils.processing_utils import encode_image_for_rollout_engine from slime.utils.trace_utils import build_sglang_meta_trace_attrs, trace_span @@ -50,6 +49,7 @@ async def generate_streaming(args: Namespace, sample: Sample, sampling_params: d assert isinstance(sample.prompt, str) state = GenerateState(args) + state.streaming_generation = True url = f"http://{args.sglang_router_ip}:{args.sglang_router_port}/generate" assert sample.status in ( @@ -108,56 +108,64 @@ async def generate_streaming(args: Namespace, sample: Sample, sampling_params: d client = http_utils._http_client assert client is not None, "http client not initialized; call init_http_client first" - with trace_span( - sample, "sglang_generate_stream", attrs={"max_new_tokens": sampling_params["max_new_tokens"]} - ) as span: - async with client.stream("POST", url, json=payload, headers=headers) as response: - response.raise_for_status() - async for raw_line in response.aiter_lines(): - if not raw_line or not raw_line.startswith("data:"): - continue - data_str = raw_line[len("data:") :].strip() - if not data_str or data_str == "[DONE]": - continue - try: - chunk = json.loads(data_str) - except json.JSONDecodeError: - logger.warning("sglang_streaming: skipping non-JSON chunk: %r", data_str[:120]) - continue - - meta = chunk.get("meta_info") or {} - last_meta_info = meta - - call_text = chunk.get("text", call_text) - if "output_token_logprobs" in meta: - call_tokens = [item[1] for item in meta["output_token_logprobs"]] - call_log_probs = [item[0] for item in meta["output_token_logprobs"]] - - # Surface partial state on the sample immediately. If the - # outer abort path cuts us, whatever we've written so far is - # what survives — no /abort_request round-trip needed. - sample.tokens = list(base_tokens) - sample.response = base_response - sample.response_length = base_response_length - sample.rollout_log_probs = None if base_log_probs is None else list(base_log_probs) - sample.rollout_top_p_token_ids = base_top_p_token_ids - sample.rollout_top_p_token_offsets = base_top_p_token_offsets - sample.loss_mask = None if base_loss_mask is None else list(base_loss_mask) - sample.append_response_tokens( - args, - tokens=call_tokens, - log_probs=call_log_probs, - trainable=True, - meta_info=meta, - text=call_text, - update_terminal_info=bool(meta.get("finish_reason")), - ) - - if state.aborted: - break - - if last_meta_info.get("finish_reason"): - span.update(build_sglang_meta_trace_attrs(last_meta_info)) + current_task = asyncio.current_task() + assert current_task is not None + state.streaming_tasks.add(current_task) + try: + with trace_span( + sample, "sglang_generate_stream", attrs={"max_new_tokens": sampling_params["max_new_tokens"]} + ) as span: + async with client.stream("POST", url, json=payload, headers=headers) as response: + response.raise_for_status() + async for raw_line in response.aiter_lines(): + if not raw_line or not raw_line.startswith("data:"): + continue + data_str = raw_line[len("data:") :].strip() + if not data_str or data_str == "[DONE]": + continue + try: + chunk = json.loads(data_str) + except json.JSONDecodeError: + logger.warning("sglang_streaming: skipping non-JSON chunk: %r", data_str[:120]) + continue + + last_meta_info = chunk.get("meta_info") or {} + call_tokens, call_log_probs, call_text = merge_stream_chunk( + tokens=call_tokens, + log_probs=call_log_probs, + text=call_text, + chunk=chunk, + ) + if chunk.get("text") is None: + call_text = state.tokenizer.decode(call_tokens, skip_special_tokens=False) + + # Rebuild from the pre-call snapshot so the sample always + # exposes the coherent prefix represented by this chunk. + sample.tokens = list(base_tokens) + sample.response = base_response + sample.response_length = base_response_length + sample.rollout_log_probs = None if base_log_probs is None else list(base_log_probs) + sample.rollout_top_p_token_ids = base_top_p_token_ids + sample.rollout_top_p_token_offsets = base_top_p_token_offsets + sample.loss_mask = None if base_loss_mask is None else list(base_loss_mask) + sample.append_response_tokens( + args, + tokens=call_tokens, + log_probs=call_log_probs, + trainable=True, + meta_info=last_meta_info, + text=call_text, + update_terminal_info=bool(last_meta_info.get("finish_reason")), + ) + + if last_meta_info.get("finish_reason"): + span.update(build_sglang_meta_trace_attrs(last_meta_info)) + except asyncio.CancelledError: + if not last_meta_info.get("finish_reason"): + sample.status = Sample.Status.ABORTED + return sample + finally: + state.streaming_tasks.discard(current_task) if state.aborted and not last_meta_info.get("finish_reason"): sample.status = Sample.Status.ABORTED diff --git a/slime/rollout/streaming_utils.py b/slime/rollout/streaming_utils.py new file mode 100644 index 0000000000..5fbf4e193d --- /dev/null +++ b/slime/rollout/streaming_utils.py @@ -0,0 +1,37 @@ +"""Backend-neutral helpers for token/logprob streaming responses.""" + +from typing import Any + + +def merge_stream_chunk( + *, + tokens: list[int], + log_probs: list[float], + text: str, + chunk: dict[str, Any], +) -> tuple[list[int], list[float], str]: + """Merge one cumulative or disjoint SGLang-compatible stream chunk.""" + meta = chunk.get("meta_info") or {} + pairs = meta.get("output_token_logprobs") or [] + chunk_tokens = [item[1] for item in pairs] + chunk_log_probs = [item[0] for item in pairs] + output_length = meta.get("output_token_logprobs_length") + + if output_length is not None: + output_length = int(output_length) + if len(chunk_tokens) == output_length: + cumulative = True + elif len(tokens) + len(chunk_tokens) == output_length: + cumulative = False + else: + raise ValueError( + "Inconsistent streaming output_token_logprobs_length: " + f"received={len(chunk_tokens)}, accumulated={len(tokens)}, reported={output_length}." + ) + else: + cumulative = len(chunk_tokens) >= len(tokens) and chunk_tokens[: len(tokens)] == tokens + + chunk_text = chunk.get("text") + if cumulative: + return chunk_tokens, chunk_log_probs, text if chunk_text is None else chunk_text + return tokens + chunk_tokens, log_probs + chunk_log_probs, text + (chunk_text or "") diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index 7eda64eb70..f77950c84b 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -572,6 +572,25 @@ def add_rollout_arguments(parser): nargs="+", help="Address and ports of the external engines.", ) + parser.add_argument( + "--rollout-external-engine-discovery-path", + type=str, + default=None, + help=( + "Optional path to a synchronous function that discovers externally managed rollout engines. " + "The function receives args and returns ExternalEngineInfo objects or equivalent dictionaries." + ), + ) + parser.add_argument( + "--rollout-external-engine-class-path", + type=str, + default=None, + help=( + "Optional path to the control actor class used for externally managed rollout engines. " + "The class must implement the SGLangEngine control interface used by the selected " + "weight-update mode." + ), + ) return parser def add_fault_tolerance_arguments(parser): @@ -1891,7 +1910,9 @@ def slime_validate_args(args): ) args.debug_train_only = True - args.rollout_external = args.rollout_external_engine_addrs is not None + args.rollout_external = ( + args.rollout_external_engine_addrs is not None or args.rollout_external_engine_discovery_path is not None + ) if args.rollout_external and not args.debug_train_only: apply_external_engine_info_to_args(args, logger=logger) diff --git a/tests/test_external_sglang_engines.py b/tests/test_external_sglang_engines.py index 61b32c460d..46b291f007 100644 --- a/tests/test_external_sglang_engines.py +++ b/tests/test_external_sglang_engines.py @@ -9,11 +9,19 @@ if str(REPO_ROOT) not in sys.path: sys.path.insert(0, str(REPO_ROOT)) +<<<<<<< HEAD from slime.backends.sglang_utils.external import ( ExternalEngineInfo, apply_external_engine_info_to_args, discover_external_engines, start_external_rollout_servers, +======= +from slime.backends.sglang_utils import external +from slime.backends.sglang_utils.external import ( + apply_external_engine_info_to_args, + discover_external_engines, + get_external_engine_class, +>>>>>>> 0b31af6a (Support externally routed streaming rollouts) ) from slime.utils.http_utils import get_rollout_num_engines @@ -176,6 +184,53 @@ def fake_get(url, timeout): assert args.rollout_num_engines == 1 +def test_apply_external_engine_info_uses_discovery_hook(monkeypatch): + args = Namespace( + rollout_external_engine_addrs=None, + rollout_external_engine_discovery_path="deployment.discover_engines", + ) + + def discover(received_args): + assert received_args is args + return [ + { + "url": "http://worker:9000", + "host": "worker", + "port": 9000, + "worker_type": "regular", + "num_gpus": 2, + "server_info": {"tp_size": 2}, + } + ] + + def fake_load_function(path): + assert path == "deployment.discover_engines" + return discover + + monkeypatch.setattr(external, "load_function", fake_load_function) + + apply_external_engine_info_to_args(args) + + assert args.rollout_num_engines == 1 + assert args.rollout_num_gpus == 2 + assert args.rollout_external_engine_infos[0]["url"] == "http://worker:9000" + + +def test_get_external_engine_class_uses_control_actor_hook(monkeypatch): + class DeploymentControlActor: + pass + + def fake_load_function(path): + assert path == "deployment.ControlActor" + return DeploymentControlActor + + monkeypatch.setattr(external, "load_function", fake_load_function) + + actor_class = get_external_engine_class(Namespace(rollout_external_engine_class_path="deployment.ControlActor")) + + assert actor_class is DeploymentControlActor + + def test_apply_external_engine_info_requires_addrs(): args = Namespace(rollout_external_engine_addrs=None) diff --git a/tests/test_streaming_rollout.py b/tests/test_streaming_rollout.py new file mode 100644 index 0000000000..076e37ab2d --- /dev/null +++ b/tests/test_streaming_rollout.py @@ -0,0 +1,231 @@ +import asyncio +import json +import sys +from contextlib import nullcontext +from pathlib import Path +from types import SimpleNamespace + +import pytest + +REPO_ROOT = Path(__file__).resolve().parents[1] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from slime.rollout.streaming_utils import merge_stream_chunk +from slime.utils.types import Sample + +NUM_GPUS = 0 + + +def _chunk(text, pairs, output_length): + return { + "text": text, + "meta_info": { + "output_token_logprobs": pairs, + "output_token_logprobs_length": output_length, + }, + } + + +def test_merge_cumulative_stream_chunks(): + tokens, log_probs, text = merge_stream_chunk( + tokens=[], + log_probs=[], + text="", + chunk=_chunk("a", [[-0.1, 11, None]], 1), + ) + tokens, log_probs, text = merge_stream_chunk( + tokens=tokens, + log_probs=log_probs, + text=text, + chunk=_chunk("ab", [[-0.1, 11, None], [-0.2, 12, None]], 2), + ) + + assert tokens == [11, 12] + assert log_probs == [-0.1, -0.2] + assert text == "ab" + + +def test_merge_disjoint_stream_chunks(): + tokens, log_probs, text = merge_stream_chunk( + tokens=[], + log_probs=[], + text="", + chunk=_chunk("a", [[-0.1, 11, None]], 1), + ) + tokens, log_probs, text = merge_stream_chunk( + tokens=tokens, + log_probs=log_probs, + text=text, + chunk=_chunk("b", [[-0.2, 12, None]], 2), + ) + tokens, log_probs, text = merge_stream_chunk( + tokens=tokens, + log_probs=log_probs, + text=text, + chunk=_chunk("", [], 2), + ) + + assert tokens == [11, 12] + assert log_probs == [-0.1, -0.2] + assert text == "ab" + + +def test_merge_stream_rejects_inconsistent_length(): + with pytest.raises(ValueError, match="output_token_logprobs_length"): + merge_stream_chunk( + tokens=[11], + log_probs=[-0.1], + text="a", + chunk=_chunk("bc", [[-0.2, 12, None], [-0.3, 13, None]], 4), + ) + + +def test_merge_cumulative_stream_preserves_text_when_intermediate_text_is_null(): + tokens, log_probs, text = merge_stream_chunk( + tokens=[11], + log_probs=[-0.1], + text="a", + chunk=_chunk(None, [[-0.1, 11, None]], 1), + ) + + assert tokens == [11] + assert log_probs == [-0.1] + assert text == "a" + + +def test_stream_cancellation_closes_request_and_keeps_prefix(monkeypatch): + try: + import sglang_router # noqa: F401 + except ImportError: + monkeypatch.setitem(sys.modules, "sglang_router", SimpleNamespace(__version__="0.3.0")) + + try: + import transformers # noqa: F401 + except ImportError: + monkeypatch.setitem( + sys.modules, + "transformers", + SimpleNamespace( + AutoProcessor=object, + AutoTokenizer=object, + PreTrainedTokenizerBase=object, + ProcessorMixin=object, + ), + ) + + from slime.rollout import sglang_streaming_rollout as streaming + + first_chunk_seen = asyncio.Event() + request_closed = False + state = SimpleNamespace( + tokenizer=SimpleNamespace( + decode=lambda token_ids, skip_special_tokens=False: "".join(f"<{token_id}>" for token_id in token_ids) + ), + processor=None, + aborted=False, + streaming_generation=False, + streaming_tasks=set(), + ) + + class FakeResponse: + async def __aenter__(self): + return self + + async def __aexit__(self, *_exc): + nonlocal request_closed + request_closed = True + + def raise_for_status(self): + pass + + async def aiter_lines(self): + yield "data: " + json.dumps( + { + "text": None, + "meta_info": { + "output_token_logprobs": [[-0.1, 11, None]], + "output_token_logprobs_length": 1, + }, + } + ) + first_chunk_seen.set() + await asyncio.Event().wait() + + class FakeClient: + def stream(self, method, url, json, headers): + assert method == "POST" + assert url == "http://frontend:8000/generate" + assert json["stream"] is True + return FakeResponse() + + monkeypatch.setattr(streaming, "GenerateState", lambda _args: state) + monkeypatch.setattr(streaming, "_prepare_prompt_ids", lambda *_args: [1, 2]) + monkeypatch.setattr(streaming.http_utils, "_http_client", FakeClient()) + monkeypatch.setattr( + streaming, + "trace_span", + lambda *_args, **_kwargs: nullcontext(SimpleNamespace(update=lambda *_args, **_kwargs: None)), + ) + + args = SimpleNamespace( + ci_test=False, + sglang_router_ip="frontend", + sglang_router_port=8000, + use_rollout_routing_replay=False, + router_policy=None, + ) + sample = Sample(prompt="hello") + + async def exercise(): + task = asyncio.create_task(streaming.generate_streaming(args, sample, {"max_new_tokens": 8})) + await first_chunk_seen.wait() + state.aborted = True + task.cancel() + return await task + + result = asyncio.run(exercise()) + + assert result.status == Sample.Status.ABORTED + assert result.tokens == [1, 2, 11] + assert result.rollout_log_probs == [-0.1] + assert result.response == "<11>" + assert request_closed is True + assert state.streaming_generation is True + assert not state.streaming_tasks + + +def test_partial_abort_buffers_only_nonempty_fully_aborted_groups(monkeypatch): + from slime.rollout import sglang_rollout + + partial = Sample(response="x", response_length=1, status=Sample.Status.ABORTED) + terminal = Sample(response="done", response_length=2, status=Sample.Status.TRUNCATED) + empty = Sample(response="", response_length=0, status=Sample.Status.ABORTED) + + async def exercise(): + async def return_group(group): + return group + + state = SimpleNamespace( + aborted=False, + streaming_generation=True, + streaming_tasks=set(), + pendings={ + asyncio.create_task(return_group([partial])), + asyncio.create_task(return_group([terminal])), + asyncio.create_task(return_group([empty])), + }, + ) + monkeypatch.setattr(sglang_rollout, "GenerateState", lambda _args: state) + + return await sglang_rollout.abort( + SimpleNamespace(partial_rollout=True), + rollout_id=7, + ) + + assert asyncio.run(exercise()) == [[partial]] + assert partial.metadata["start_rollout_id"] == 7 + + +if __name__ == "__main__": + raise SystemExit(pytest.main([__file__])) From 2f2896ffe0e0e5620fc65c9c79a7187dd0f8785d Mon Sep 17 00:00:00 2001 From: William Arnold <7565007+Aphoh@users.noreply.github.com> Date: Sun, 26 Jul 2026 23:01:24 +0000 Subject: [PATCH 02/12] Fix external engine parallel metadata --- slime/backends/sglang_utils/external.py | 1 + tests/test_external_sglang_engines.py | 6 ++++++ 2 files changed, 7 insertions(+) diff --git a/slime/backends/sglang_utils/external.py b/slime/backends/sglang_utils/external.py index 8cfc0cb71f..f09179c075 100644 --- a/slime/backends/sglang_utils/external.py +++ b/slime/backends/sglang_utils/external.py @@ -173,6 +173,7 @@ class ExternalRolloutServer: update_weights: bool = True num_new_engines: int = 0 server_groups: list = dataclasses.field(default_factory=list) + engine_parallel_configs: list[dict] = dataclasses.field(default_factory=list) @property def all_engines(self): diff --git a/tests/test_external_sglang_engines.py b/tests/test_external_sglang_engines.py index 46b291f007..1c8f26f9a8 100644 --- a/tests/test_external_sglang_engines.py +++ b/tests/test_external_sglang_engines.py @@ -238,5 +238,11 @@ def test_apply_external_engine_info_requires_addrs(): apply_external_engine_info_to_args(args) +def test_external_rollout_server_has_neutral_parallel_config(): + server = external.ExternalRolloutServer(engines=[], engine_gpu_counts=[], engine_gpu_offsets=[]) + + assert server.engine_parallel_configs == [] + + if __name__ == "__main__": raise SystemExit(pytest.main([__file__])) From cd309dadf21e2d56e7ceb9a7ee22c04adc308433 Mon Sep 17 00:00:00 2001 From: William Arnold <7565007+Aphoh@users.noreply.github.com> Date: Mon, 27 Jul 2026 23:20:50 +0000 Subject: [PATCH 03/12] fix external debug rollout GPU allocation --- slime/ray/actor_group.py | 4 ++++ slime/utils/arguments.py | 12 +++++++---- tests/test_megatron_argument_validation.py | 24 ++++++++++++++++++++++ tests/test_placement_group.py | 14 +++++++++++++ 4 files changed, 50 insertions(+), 4 deletions(-) diff --git a/slime/ray/actor_group.py b/slime/ray/actor_group.py index 7662a0de69..f72f676c08 100644 --- a/slime/ray/actor_group.py +++ b/slime/ray/actor_group.py @@ -191,6 +191,10 @@ def create(self, rollout_manager=None): if rollout_manager is not None: self._rollout_manager = rollout_manager self.args.update_weight_start_version = self._disk_weight_version + if self._num_nodes * self._num_gpus_per_node == 0: + assert self.args.debug_rollout_only, "zero-sized train groups are only valid in rollout-only debug mode" + start_rollout_id = self.args.start_rollout_id + return [0 if start_rollout_id is None else start_rollout_id] self._allocate_gpus_for_actor(self._pg, self._num_gpus_per_actor) start_rollout_ids = ray.get( [ diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index f77950c84b..5f32cc1910 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -1918,9 +1918,6 @@ def slime_validate_args(args): apply_external_engine_info_to_args(args, logger=logger) args.use_critic = args.advantage_estimator == "ppo" - # Critic always uses the same GPU count as actor. - args.critic_num_gpus_per_node = args.actor_num_gpus_per_node - args.critic_num_nodes = args.actor_num_nodes if args.offload: args.offload_train = True @@ -1928,7 +1925,10 @@ def slime_validate_args(args): del args.offload if args.debug_rollout_only: - if args.colocate and args.rollout_num_gpus is None: + if args.rollout_external: + args.actor_num_gpus_per_node = 0 + args.actor_num_nodes = 0 + elif args.colocate and args.rollout_num_gpus is None: args.rollout_num_gpus = args.actor_num_gpus_per_node * args.actor_num_nodes elif args.rollout_num_gpus == 0: args.actor_num_gpus_per_node = 0 @@ -1939,6 +1939,10 @@ def slime_validate_args(args): args.colocate = False args.offload_train = args.offload_rollout = False + # Critic always uses the same GPU count as actor, including debug-mode overrides. + args.critic_num_gpus_per_node = args.actor_num_gpus_per_node + args.critic_num_nodes = args.actor_num_nodes + assert not (args.debug_rollout_only and args.debug_train_only), ( "debug_rollout_only and debug_train_only cannot be set at the same time, " "please set only one of them." ) diff --git a/tests/test_megatron_argument_validation.py b/tests/test_megatron_argument_validation.py index 9510b78ffd..fb3333c320 100644 --- a/tests/test_megatron_argument_validation.py +++ b/tests/test_megatron_argument_validation.py @@ -218,6 +218,7 @@ def make_slime_validate_args(**overrides): save_debug_train_data=None, load_debug_rollout_data=None, rollout_external_engine_addrs=None, + rollout_external_engine_discovery_path=None, debug_train_only=False, actor_num_gpus_per_node=8, actor_num_nodes=1, @@ -345,5 +346,28 @@ def test_update_weight_delta_requires_local_checkpoint_dir(monkeypatch): module.slime_validate_args(args) +@pytest.mark.unit +def test_external_debug_rollout_does_not_allocate_train_gpus(monkeypatch): + module = load_slime_arguments_module(monkeypatch) + + def discover_external_engine(args, logger=None): + args.rollout_num_gpus = 2 + + module.apply_external_engine_info_to_args = discover_external_engine + args = make_slime_validate_args( + rollout_external_engine_discovery_path="deployment.discover_engines", + debug_rollout_only=True, + ) + + module.slime_validate_args(args) + + assert args.rollout_external is True + assert args.rollout_num_gpus == 2 + assert args.actor_num_gpus_per_node == 0 + assert args.actor_num_nodes == 0 + assert args.critic_num_gpus_per_node == 0 + assert args.critic_num_nodes == 0 + + if __name__ == "__main__": raise SystemExit(pytest.main([__file__])) diff --git a/tests/test_placement_group.py b/tests/test_placement_group.py index c1ae8aedef..481dd605ec 100644 --- a/tests/test_placement_group.py +++ b/tests/test_placement_group.py @@ -8,6 +8,7 @@ if str(REPO_ROOT) not in sys.path: sys.path.insert(0, str(REPO_ROOT)) +from slime.ray.actor_group import RayTrainGroup from slime.ray.placement_group import _create_placement_group, _get_placement_group_layout NUM_GPUS = 0 @@ -50,5 +51,18 @@ def test_create_zero_gpu_placement_group_is_empty(): assert _create_placement_group(0) == (None, [], []) +@pytest.mark.parametrize(("start_rollout_id", "expected"), [(None, 0), (7, 7)]) +def test_zero_sized_debug_train_group_uses_configured_rollout_id(start_rollout_id, expected): + args = Namespace(debug_rollout_only=True, start_rollout_id=start_rollout_id) + group = RayTrainGroup( + args=args, + num_nodes=0, + num_gpus_per_node=0, + pg=(None, [], []), + ) + + assert group.create() == [expected] + + if __name__ == "__main__": raise SystemExit(pytest.main([__file__])) From 7cc2345afe6194a115caadc8c95b183c45c15d9f Mon Sep 17 00:00:00 2001 From: William Arnold <7565007+Aphoh@users.noreply.github.com> Date: Tue, 28 Jul 2026 00:05:04 +0000 Subject: [PATCH 04/12] test: streamline external streaming coverage --- tests/test_external_sglang_engines.py | 27 ++-- tests/test_streaming_rollout.py | 208 ++++++++++---------------- 2 files changed, 93 insertions(+), 142 deletions(-) diff --git a/tests/test_external_sglang_engines.py b/tests/test_external_sglang_engines.py index 1c8f26f9a8..7e525e22ec 100644 --- a/tests/test_external_sglang_engines.py +++ b/tests/test_external_sglang_engines.py @@ -41,6 +41,14 @@ def json(self): return self.payload +def _loader(expected_path, result): + def load(path): + assert path == expected_path + return result + + return load + + def test_discover_external_engines_reads_server_info(monkeypatch): def fake_get(url, timeout): assert timeout == 30.0 @@ -203,11 +211,7 @@ def discover(received_args): } ] - def fake_load_function(path): - assert path == "deployment.discover_engines" - return discover - - monkeypatch.setattr(external, "load_function", fake_load_function) + monkeypatch.setattr(external, "load_function", _loader("deployment.discover_engines", discover)) apply_external_engine_info_to_args(args) @@ -220,12 +224,7 @@ def test_get_external_engine_class_uses_control_actor_hook(monkeypatch): class DeploymentControlActor: pass - def fake_load_function(path): - assert path == "deployment.ControlActor" - return DeploymentControlActor - - monkeypatch.setattr(external, "load_function", fake_load_function) - + monkeypatch.setattr(external, "load_function", _loader("deployment.ControlActor", DeploymentControlActor)) actor_class = get_external_engine_class(Namespace(rollout_external_engine_class_path="deployment.ControlActor")) assert actor_class is DeploymentControlActor @@ -238,11 +237,5 @@ def test_apply_external_engine_info_requires_addrs(): apply_external_engine_info_to_args(args) -def test_external_rollout_server_has_neutral_parallel_config(): - server = external.ExternalRolloutServer(engines=[], engine_gpu_counts=[], engine_gpu_offsets=[]) - - assert server.engine_parallel_configs == [] - - if __name__ == "__main__": raise SystemExit(pytest.main([__file__])) diff --git a/tests/test_streaming_rollout.py b/tests/test_streaming_rollout.py index 076e37ab2d..9993359842 100644 --- a/tests/test_streaming_rollout.py +++ b/tests/test_streaming_rollout.py @@ -1,7 +1,7 @@ import asyncio import json import sys -from contextlib import nullcontext +from contextlib import asynccontextmanager, nullcontext from pathlib import Path from types import SimpleNamespace @@ -11,6 +11,22 @@ if str(REPO_ROOT) not in sys.path: sys.path.insert(0, str(REPO_ROOT)) +try: + import sglang_router # noqa: F401 +except ImportError: + sys.modules["sglang_router"] = SimpleNamespace(__version__="0.3.0") + +try: + import transformers # noqa: F401 +except ImportError: + sys.modules["transformers"] = SimpleNamespace( + AutoProcessor=object, + AutoTokenizer=object, + PreTrainedTokenizerBase=object, + ProcessorMixin=object, + ) + +from slime.rollout import sglang_rollout, sglang_streaming_rollout as streaming from slime.rollout.streaming_utils import merge_stream_chunk from slime.utils.types import Sample @@ -27,48 +43,45 @@ def _chunk(text, pairs, output_length): } -def test_merge_cumulative_stream_chunks(): - tokens, log_probs, text = merge_stream_chunk( - tokens=[], - log_probs=[], - text="", - chunk=_chunk("a", [[-0.1, 11, None]], 1), - ) - tokens, log_probs, text = merge_stream_chunk( - tokens=tokens, - log_probs=log_probs, - text=text, - chunk=_chunk("ab", [[-0.1, 11, None], [-0.2, 12, None]], 2), - ) - - assert tokens == [11, 12] - assert log_probs == [-0.1, -0.2] - assert text == "ab" - - -def test_merge_disjoint_stream_chunks(): - tokens, log_probs, text = merge_stream_chunk( - tokens=[], - log_probs=[], - text="", - chunk=_chunk("a", [[-0.1, 11, None]], 1), - ) - tokens, log_probs, text = merge_stream_chunk( - tokens=tokens, - log_probs=log_probs, - text=text, - chunk=_chunk("b", [[-0.2, 12, None]], 2), - ) - tokens, log_probs, text = merge_stream_chunk( - tokens=tokens, - log_probs=log_probs, - text=text, - chunk=_chunk("", [], 2), - ) +@pytest.mark.parametrize( + ("initial", "chunks", "expected"), + [ + ( + ([], [], ""), + [ + _chunk("a", [[-0.1, 11, None]], 1), + _chunk("ab", [[-0.1, 11, None], [-0.2, 12, None]], 2), + ], + ([11, 12], [-0.1, -0.2], "ab"), + ), + ( + ([], [], ""), + [ + _chunk("a", [[-0.1, 11, None]], 1), + _chunk("b", [[-0.2, 12, None]], 2), + _chunk("", [], 2), + ], + ([11, 12], [-0.1, -0.2], "ab"), + ), + ( + ([11], [-0.1], "a"), + [_chunk(None, [[-0.1, 11, None]], 1)], + ([11], [-0.1], "a"), + ), + ], + ids=["cumulative", "disjoint", "null-text"], +) +def test_merge_stream_chunks(initial, chunks, expected): + tokens, log_probs, text = initial + for chunk in chunks: + tokens, log_probs, text = merge_stream_chunk( + tokens=tokens, + log_probs=log_probs, + text=text, + chunk=chunk, + ) - assert tokens == [11, 12] - assert log_probs == [-0.1, -0.2] - assert text == "ab" + assert (tokens, log_probs, text) == expected def test_merge_stream_rejects_inconsistent_length(): @@ -81,41 +94,7 @@ def test_merge_stream_rejects_inconsistent_length(): ) -def test_merge_cumulative_stream_preserves_text_when_intermediate_text_is_null(): - tokens, log_probs, text = merge_stream_chunk( - tokens=[11], - log_probs=[-0.1], - text="a", - chunk=_chunk(None, [[-0.1, 11, None]], 1), - ) - - assert tokens == [11] - assert log_probs == [-0.1] - assert text == "a" - - def test_stream_cancellation_closes_request_and_keeps_prefix(monkeypatch): - try: - import sglang_router # noqa: F401 - except ImportError: - monkeypatch.setitem(sys.modules, "sglang_router", SimpleNamespace(__version__="0.3.0")) - - try: - import transformers # noqa: F401 - except ImportError: - monkeypatch.setitem( - sys.modules, - "transformers", - SimpleNamespace( - AutoProcessor=object, - AutoTokenizer=object, - PreTrainedTokenizerBase=object, - ProcessorMixin=object, - ), - ) - - from slime.rollout import sglang_streaming_rollout as streaming - first_chunk_seen = asyncio.Event() request_closed = False state = SimpleNamespace( @@ -128,46 +107,38 @@ def test_stream_cancellation_closes_request_and_keeps_prefix(monkeypatch): streaming_tasks=set(), ) - class FakeResponse: - async def __aenter__(self): - return self - - async def __aexit__(self, *_exc): - nonlocal request_closed + async def lines(): + yield "data: " + json.dumps( + { + "text": None, + "meta_info": { + "output_token_logprobs": [[-0.1, 11, None]], + "output_token_logprobs_length": 1, + }, + } + ) + first_chunk_seen.set() + await asyncio.Event().wait() + + response = SimpleNamespace(raise_for_status=lambda: None, aiter_lines=lines) + + @asynccontextmanager + async def stream(method, url, json, headers): + nonlocal request_closed + assert (method, url, json["stream"]) == ("POST", "http://frontend:8000/generate", True) + try: + yield response + finally: request_closed = True - def raise_for_status(self): - pass - - async def aiter_lines(self): - yield "data: " + json.dumps( - { - "text": None, - "meta_info": { - "output_token_logprobs": [[-0.1, 11, None]], - "output_token_logprobs_length": 1, - }, - } - ) - first_chunk_seen.set() - await asyncio.Event().wait() - - class FakeClient: - def stream(self, method, url, json, headers): - assert method == "POST" - assert url == "http://frontend:8000/generate" - assert json["stream"] is True - return FakeResponse() - monkeypatch.setattr(streaming, "GenerateState", lambda _args: state) monkeypatch.setattr(streaming, "_prepare_prompt_ids", lambda *_args: [1, 2]) - monkeypatch.setattr(streaming.http_utils, "_http_client", FakeClient()) + monkeypatch.setattr(streaming.http_utils, "_http_client", SimpleNamespace(stream=stream)) monkeypatch.setattr( streaming, "trace_span", lambda *_args, **_kwargs: nullcontext(SimpleNamespace(update=lambda *_args, **_kwargs: None)), ) - args = SimpleNamespace( ci_test=False, sglang_router_ip="frontend", @@ -175,10 +146,9 @@ def stream(self, method, url, json, headers): use_rollout_routing_replay=False, router_policy=None, ) - sample = Sample(prompt="hello") async def exercise(): - task = asyncio.create_task(streaming.generate_streaming(args, sample, {"max_new_tokens": 8})) + task = asyncio.create_task(streaming.generate_streaming(args, Sample(prompt="hello"), {"max_new_tokens": 8})) await first_chunk_seen.wait() state.aborted = True task.cancel() @@ -187,41 +157,29 @@ async def exercise(): result = asyncio.run(exercise()) assert result.status == Sample.Status.ABORTED - assert result.tokens == [1, 2, 11] - assert result.rollout_log_probs == [-0.1] - assert result.response == "<11>" + assert (result.tokens, result.rollout_log_probs, result.response) == ([1, 2, 11], [-0.1], "<11>") assert request_closed is True assert state.streaming_generation is True assert not state.streaming_tasks def test_partial_abort_buffers_only_nonempty_fully_aborted_groups(monkeypatch): - from slime.rollout import sglang_rollout - partial = Sample(response="x", response_length=1, status=Sample.Status.ABORTED) terminal = Sample(response="done", response_length=2, status=Sample.Status.TRUNCATED) empty = Sample(response="", response_length=0, status=Sample.Status.ABORTED) async def exercise(): - async def return_group(group): - return group - state = SimpleNamespace( aborted=False, streaming_generation=True, streaming_tasks=set(), pendings={ - asyncio.create_task(return_group([partial])), - asyncio.create_task(return_group([terminal])), - asyncio.create_task(return_group([empty])), + asyncio.create_task(asyncio.sleep(0, result=group)) + for group in ([partial], [terminal], [empty]) }, ) monkeypatch.setattr(sglang_rollout, "GenerateState", lambda _args: state) - - return await sglang_rollout.abort( - SimpleNamespace(partial_rollout=True), - rollout_id=7, - ) + return await sglang_rollout.abort(SimpleNamespace(partial_rollout=True), rollout_id=7) assert asyncio.run(exercise()) == [[partial]] assert partial.metadata["start_rollout_id"] == 7 From 4959a5e6d2f95a755aa1270b90f93cb3121aa061 Mon Sep 17 00:00:00 2001 From: William Arnold <7565007+Aphoh@users.noreply.github.com> Date: Tue, 28 Jul 2026 07:06:26 +0000 Subject: [PATCH 05/12] feat: support shared external rollout endpoints --- slime/backends/sglang_utils/external.py | 59 +++++++++++++++----- slime/backends/sglang_utils/sglang_engine.py | 26 ++++++--- slime/utils/arguments.py | 12 ++++ tests/test_external_sglang_engines.py | 31 ++++++++++ 4 files changed, 106 insertions(+), 22 deletions(-) diff --git a/slime/backends/sglang_utils/external.py b/slime/backends/sglang_utils/external.py index f09179c075..c2e4ce63a6 100644 --- a/slime/backends/sglang_utils/external.py +++ b/slime/backends/sglang_utils/external.py @@ -22,6 +22,8 @@ class ExternalEngineInfo: num_gpus: int disaggregation_bootstrap_port: int | None = None server_info: dict = dataclasses.field(default_factory=dict) + engine_api_prefix: str = "" + rollout_url: str | None = None @property def is_pd_worker(self) -> bool: @@ -61,6 +63,12 @@ def normalize_external_engine_addr(addr: str) -> str: return addr +def engine_control_url(url: str, engine_api_prefix: str, method: str) -> str: + """Build a per-worker SGLang control URL.""" + prefix = engine_api_prefix.strip("/") + return "{}/{}{}".format(url.rstrip("/"), prefix + "/" if prefix else "", method) + + def external_engine_init_kwargs(info: ExternalEngineInfo) -> dict: init_kwargs = { "dist_init_addr": f"{info.host}:{info.port}", @@ -70,14 +78,17 @@ def external_engine_init_kwargs(info: ExternalEngineInfo) -> dict: } if info.worker_type == "prefill": init_kwargs["disaggregation_bootstrap_port"] = info.disaggregation_bootstrap_port + if info.engine_api_prefix: + init_kwargs["engine_api_prefix"] = info.engine_api_prefix return init_kwargs -def get_server_info(url: str, timeout: float = 30.0) -> dict: +def get_server_info(url: str, engine_api_prefix: str = "", timeout: float = 30.0) -> dict: errors = [] - for endpoint in ("/server_info", "/get_server_info"): + for method in ("server_info", "get_server_info"): + endpoint = engine_control_url(url, engine_api_prefix, method) try: - response = requests.get(f"{url}{endpoint}", timeout=timeout) + response = requests.get(endpoint, timeout=timeout) response.raise_for_status() return response.json() except Exception as exc: @@ -94,13 +105,15 @@ def _infer_worker_type(server_info: dict) -> str: return "regular" -def discover_external_engines(addrs: list[str], timeout: float = 30.0) -> list[ExternalEngineInfo]: +def discover_external_engines( + addrs: list[str], engine_api_prefix: str = "", rollout_url: str | None = None, timeout: float = 30.0 +) -> list[ExternalEngineInfo]: infos = [] for addr in addrs: url = normalize_external_engine_addr(addr) parsed = urlparse(url) assert parsed.hostname is not None and parsed.port is not None - server_info = get_server_info(url, timeout=timeout) + server_info = get_server_info(url, engine_api_prefix=engine_api_prefix, timeout=timeout) pp_size = int(server_info.get("pp_size") or server_info.get("pipeline_parallel_size") or 1) tp_size = int(server_info.get("tp_size") or server_info.get("tensor_parallel_size") or 1) @@ -115,6 +128,8 @@ def discover_external_engines(addrs: list[str], timeout: float = 30.0) -> list[E port=parsed.port, worker_type=_infer_worker_type(server_info), num_gpus=num_gpus, + engine_api_prefix=engine_api_prefix, + rollout_url=normalize_external_engine_addr(rollout_url) if rollout_url else None, disaggregation_bootstrap_port=bootstrap_port, server_info=server_info, ) @@ -135,7 +150,11 @@ def apply_external_engine_info_to_args(args, logger=None) -> None: "External rollout requires --rollout-external-engine-addrs or " "--rollout-external-engine-discovery-path." ) - infos = discover_external_engines(addrs) + infos = discover_external_engines( + addrs, + engine_api_prefix=getattr(args, "rollout_external_engine_api_prefix", ""), + rollout_url=getattr(args, "rollout_external_rollout_url", None), + ) if not infos: raise ValueError("External rollout engine discovery returned no engines.") @@ -216,13 +235,28 @@ def get_external_engine_class(args): return SGLangEngine +def get_external_rollout_url(infos: list[ExternalEngineInfo]) -> str | None: + rollout_urls = {info.rollout_url for info in infos} + if rollout_urls == {None}: + return None + if None in rollout_urls or len(rollout_urls) != 1: + raise ValueError("External rollout engines must all use the same rollout_url, or none.") + return rollout_urls.pop() + + def start_external_rollout_servers(args, *, start_router) -> tuple[dict[str, ExternalRolloutServer], list]: import ray 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 = get_external_rollout_url(infos) + 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 @@ -248,13 +282,10 @@ def start_external_rollout_servers(args, *, start_router) -> tuple[dict[str, Ext 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, - ) - ) + init_kwargs = external_engine_init_kwargs(info) + if rollout_url is not None: + init_kwargs["register_to_router"] = False + init_handles.append(rollout_engine.init.remote(**init_kwargs, router_ip=router_ip, router_port=router_port)) args.sglang_model_routers = {"default": (router_ip, router_port)} servers = { diff --git a/slime/backends/sglang_utils/sglang_engine.py b/slime/backends/sglang_utils/sglang_engine.py index 3e562a7e38..64e69ef2be 100644 --- a/slime/backends/sglang_utils/sglang_engine.py +++ b/slime/backends/sglang_utils/sglang_engine.py @@ -10,7 +10,7 @@ from sglang.srt.utils import kill_process_tree from urllib3.exceptions import NewConnectionError -from slime.backends.sglang_utils.external import get_server_info +from slime.backends.sglang_utils.external import engine_control_url, get_server_info from slime.ray.ray_actor import RayActor from slime.utils.http_utils import get_host_info @@ -128,9 +128,13 @@ def init( disaggregation_bootstrap_port=None, router_ip=None, router_port=None, + engine_api_prefix="", + register_to_router=True, ): self.router_ip = router_ip if router_ip is not None else self.args.sglang_router_ip self.router_port = router_port if router_port is not None else self.args.sglang_router_port + self.engine_api_prefix = engine_api_prefix + self.register_to_router = register_to_router host = host or get_host_info()[1] @@ -182,9 +186,12 @@ def _sanity_check_server_args(actual_server_args, expect_server_args): actual_value == expect_value ), f"{name=} {expect_value=} {actual_value=} {expect_server_args=} {actual_server_args=}" - actual_server_args = get_server_info(f"http://{self.server_host}:{self.server_port}") + actual_server_args = get_server_info( + f"http://{self.server_host}:{self.server_port}", engine_api_prefix=self.engine_api_prefix + ) _sanity_check_server_args(actual_server_args, expect_server_args) - self._register_to_router(expect_server_args) + if self.register_to_router: + self._register_to_router(expect_server_args) def _init_normal(self, server_args_dict): logger.info(f"Launch HttpServerEngineAdapter at: {self.server_host}:{self.server_port}") @@ -228,7 +235,7 @@ def _make_request(self, endpoint: str, payload: dict | None = None): if self.node_rank != 0: return - url = f"http://{self.server_host}:{self.server_port}/{endpoint}" + url = self._engine_control_url(endpoint) response = requests.post(url, json=payload or {}) try: response.raise_for_status() @@ -237,6 +244,9 @@ def _make_request(self, endpoint: str, payload: dict | None = None): raise return response.json() + def _engine_control_url(self, method: str) -> str: + return engine_control_url(f"http://{self.server_host}:{self.server_port}", self.engine_api_prefix, method) + def health_generate(self, timeout: float = 5.0) -> bool: """Run /health_generate on the underlying SGLang HTTP server. @@ -253,7 +263,7 @@ def health_generate(self, timeout: float = 5.0) -> bool: return True response = requests.get( - f"http://{self.server_host}:{self.server_port}/health_generate", + self._engine_control_url("health_generate"), timeout=timeout, ) response.raise_for_status() @@ -291,7 +301,7 @@ def flush_cache(self): # 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") + response = requests.get(self._engine_control_url("flush_cache")) if response.status_code == 200: break logger.info(f"Error flushing cache: HTTP {response.status_code} {response.text!r}") @@ -440,14 +450,14 @@ def update_weights_from_distributed( def pause_generation(self): if self.node_rank != 0: return - response = requests.post(f"http://{self.server_host}:{self.server_port}/pause_generation", json={}) + response = requests.post(self._engine_control_url("pause_generation"), json={}) response.raise_for_status() return response def continue_generation(self): if self.node_rank != 0: return - response = requests.post(f"http://{self.server_host}:{self.server_port}/continue_generation", json={}) + response = requests.post(self._engine_control_url("continue_generation"), json={}) response.raise_for_status() return response diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index 5f32cc1910..2251829792 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -572,6 +572,18 @@ def add_rollout_arguments(parser): nargs="+", help="Address and ports of the external engines.", ) + parser.add_argument( + "--rollout-external-engine-api-prefix", + type=str, + default="", + help="Optional per-worker control API prefix, for example /engine.", + ) + parser.add_argument( + "--rollout-external-rollout-url", + type=str, + default=None, + help="Optional shared rollout endpoint. When set, Slime sends /generate there instead of starting a router.", + ) parser.add_argument( "--rollout-external-engine-discovery-path", type=str, diff --git a/tests/test_external_sglang_engines.py b/tests/test_external_sglang_engines.py index 7e525e22ec..218d49431f 100644 --- a/tests/test_external_sglang_engines.py +++ b/tests/test_external_sglang_engines.py @@ -20,6 +20,8 @@ from slime.backends.sglang_utils.external import ( apply_external_engine_info_to_args, discover_external_engines, + engine_control_url, + get_external_rollout_url, get_external_engine_class, >>>>>>> 0b31af6a (Support externally routed streaming rollouts) ) @@ -118,6 +120,35 @@ def remote(self, **kwargs): assert len(init_handles) == 1 +def test_discover_external_engines_uses_control_prefix_and_shared_rollout_url(monkeypatch): + def fake_get(url, timeout): + assert timeout == 30.0 + assert url == "http://worker:9090/engine/server_info" + return _Response({"tp_size": 2}) + + monkeypatch.setattr("slime.backends.sglang_utils.external.requests.get", fake_get) + + info = discover_external_engines( + ["worker:9090"], engine_api_prefix="/engine", rollout_url="http://frontend:8000" + )[0] + + assert info.engine_api_prefix == "/engine" + assert info.rollout_url == "http://frontend:8000" + assert engine_control_url(info.url, info.engine_api_prefix, "flush_cache") == "http://worker:9090/engine/flush_cache" + + +def test_get_external_rollout_url_requires_one_shared_frontend(): + engine = external.ExternalEngineInfo("http://worker:9090", "worker", 9090, "regular", 1) + shared = external.ExternalEngineInfo( + "http://worker:9091", "worker", 9091, "regular", 1, rollout_url="http://frontend:8000" + ) + + assert get_external_rollout_url([engine]) is None + assert get_external_rollout_url([shared]) == "http://frontend:8000" + with pytest.raises(ValueError, match="same rollout_url"): + get_external_rollout_url([engine, shared]) + + def test_apply_external_engine_info_handles_pd(monkeypatch): payloads = { "http://prefill:10090/server_info": { From 31552a6a9062dca84b9a959eef5a34187b1033fa Mon Sep 17 00:00:00 2001 From: William Arnold <7565007+Aphoh@users.noreply.github.com> Date: Tue, 28 Jul 2026 23:04:09 +0000 Subject: [PATCH 06/12] refactor: simplify external engine setup --- slime/backends/sglang_utils/external.py | 53 +++++------------- slime/utils/arguments.py | 26 +-------- tests/test_external_sglang_engines.py | 65 ++-------------------- tests/test_megatron_argument_validation.py | 24 -------- 4 files changed, 18 insertions(+), 150 deletions(-) diff --git a/slime/backends/sglang_utils/external.py b/slime/backends/sglang_utils/external.py index c2e4ce63a6..a3efa99162 100644 --- a/slime/backends/sglang_utils/external.py +++ b/slime/backends/sglang_utils/external.py @@ -8,7 +8,6 @@ import requests -from slime.utils.misc import load_function logger = logging.getLogger(__name__) @@ -56,10 +55,7 @@ def normalize_external_engine_addr(addr: str) -> str: addr = addr.rstrip("/") 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)." - ) + raise ValueError(f"Invalid external SGLang engine address {addr!r}. Use host:port or http://host:port (IPv6 must be bracketed).") return addr @@ -105,9 +101,7 @@ def _infer_worker_type(server_info: dict) -> str: return "regular" -def discover_external_engines( - addrs: list[str], engine_api_prefix: str = "", rollout_url: str | None = None, timeout: float = 30.0 -) -> list[ExternalEngineInfo]: +def discover_external_engines(addrs: list[str], engine_api_prefix: str = "", rollout_url: str | None = None, timeout: float = 30.0) -> list[ExternalEngineInfo]: infos = [] for addr in addrs: url = normalize_external_engine_addr(addr) @@ -139,22 +133,14 @@ def discover_external_engines( def apply_external_engine_info_to_args(args, logger=None) -> None: """Detect external engines and store the derived topology on ``args``.""" - discovery_path = getattr(args, "rollout_external_engine_discovery_path", None) - if discovery_path is not None: - discovered = load_function(discovery_path)(args) - infos = [ExternalEngineInfo(**info) if isinstance(info, dict) else info for info in discovered] - else: - addrs = args.rollout_external_engine_addrs - if not addrs: - raise ValueError( - "External rollout requires --rollout-external-engine-addrs or " - "--rollout-external-engine-discovery-path." - ) - infos = discover_external_engines( - addrs, - engine_api_prefix=getattr(args, "rollout_external_engine_api_prefix", ""), - rollout_url=getattr(args, "rollout_external_rollout_url", None), - ) + addrs = args.rollout_external_engine_addrs + if not addrs: + raise ValueError("External rollout requires --rollout-external-engine-addrs.") + infos = discover_external_engines( + addrs, + engine_api_prefix=getattr(args, "rollout_external_engine_api_prefix", ""), + rollout_url=getattr(args, "rollout_external_rollout_url", None), + ) if not infos: raise ValueError("External rollout engine discovery returned no engines.") @@ -217,24 +203,10 @@ def onload_kv(self): 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." - ) + raise RuntimeError("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 get_external_engine_class(args): - """Return the control actor class for externally managed engines.""" - engine_class_path = getattr(args, "rollout_external_engine_class_path", None) - if engine_class_path is not None: - return load_function(engine_class_path) - - from slime.backends.sglang_utils.sglang_engine import SGLangEngine - - return SGLangEngine - - def get_external_rollout_url(infos: list[ExternalEngineInfo]) -> str | None: rollout_urls = {info.rollout_url for info in infos} if rollout_urls == {None}: @@ -248,6 +220,7 @@ def start_external_rollout_servers(args, *, start_router) -> tuple[dict[str, Ext import ray from slime.ray.utils import add_default_ray_env_vars + from slime.backends.sglang_utils.sglang_engine import SGLangEngine infos = external_engine_infos_from_args(args) rollout_url = get_external_rollout_url(infos) @@ -264,7 +237,7 @@ def start_external_rollout_servers(args, *, start_router) -> tuple[dict[str, Ext engine_gpu_counts = [] engine_gpu_offsets = [] init_handles = [] - RolloutRayActor = ray.remote(get_external_engine_class(args)) + RolloutRayActor = ray.remote(SGLangEngine) gpu_offset = 0 for rank, info in enumerate(infos): rollout_engine = RolloutRayActor.options( diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index 2251829792..732632c071 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -584,25 +584,6 @@ def add_rollout_arguments(parser): default=None, help="Optional shared rollout endpoint. When set, Slime sends /generate there instead of starting a router.", ) - parser.add_argument( - "--rollout-external-engine-discovery-path", - type=str, - default=None, - help=( - "Optional path to a synchronous function that discovers externally managed rollout engines. " - "The function receives args and returns ExternalEngineInfo objects or equivalent dictionaries." - ), - ) - parser.add_argument( - "--rollout-external-engine-class-path", - type=str, - default=None, - help=( - "Optional path to the control actor class used for externally managed rollout engines. " - "The class must implement the SGLangEngine control interface used by the selected " - "weight-update mode." - ), - ) return parser def add_fault_tolerance_arguments(parser): @@ -1922,9 +1903,7 @@ def slime_validate_args(args): ) args.debug_train_only = True - args.rollout_external = ( - args.rollout_external_engine_addrs is not None or args.rollout_external_engine_discovery_path is not None - ) + args.rollout_external = args.rollout_external_engine_addrs is not None if args.rollout_external and not args.debug_train_only: apply_external_engine_info_to_args(args, logger=logger) @@ -1951,9 +1930,6 @@ def slime_validate_args(args): args.colocate = False args.offload_train = args.offload_rollout = False - # Critic always uses the same GPU count as actor, including debug-mode overrides. - args.critic_num_gpus_per_node = args.actor_num_gpus_per_node - args.critic_num_nodes = args.actor_num_nodes assert not (args.debug_rollout_only and args.debug_train_only), ( "debug_rollout_only and debug_train_only cannot be set at the same time, " "please set only one of them." diff --git a/tests/test_external_sglang_engines.py b/tests/test_external_sglang_engines.py index 218d49431f..cfdc90dfad 100644 --- a/tests/test_external_sglang_engines.py +++ b/tests/test_external_sglang_engines.py @@ -9,21 +9,14 @@ if str(REPO_ROOT) not in sys.path: sys.path.insert(0, str(REPO_ROOT)) -<<<<<<< HEAD -from slime.backends.sglang_utils.external import ( - ExternalEngineInfo, - apply_external_engine_info_to_args, - discover_external_engines, - start_external_rollout_servers, -======= from slime.backends.sglang_utils import external from slime.backends.sglang_utils.external import ( + ExternalEngineInfo, apply_external_engine_info_to_args, discover_external_engines, engine_control_url, get_external_rollout_url, - get_external_engine_class, ->>>>>>> 0b31af6a (Support externally routed streaming rollouts) + start_external_rollout_servers, ) from slime.utils.http_utils import get_rollout_num_engines @@ -43,14 +36,6 @@ def json(self): return self.payload -def _loader(expected_path, result): - def load(path): - assert path == expected_path - return result - - return load - - def test_discover_external_engines_reads_server_info(monkeypatch): def fake_get(url, timeout): assert timeout == 30.0 @@ -128,9 +113,7 @@ def fake_get(url, timeout): monkeypatch.setattr("slime.backends.sglang_utils.external.requests.get", fake_get) - info = discover_external_engines( - ["worker:9090"], engine_api_prefix="/engine", rollout_url="http://frontend:8000" - )[0] + info = discover_external_engines(["worker:9090"], engine_api_prefix="/engine", rollout_url="http://frontend:8000")[0] assert info.engine_api_prefix == "/engine" assert info.rollout_url == "http://frontend:8000" @@ -139,9 +122,7 @@ def fake_get(url, timeout): def test_get_external_rollout_url_requires_one_shared_frontend(): engine = external.ExternalEngineInfo("http://worker:9090", "worker", 9090, "regular", 1) - shared = external.ExternalEngineInfo( - "http://worker:9091", "worker", 9091, "regular", 1, rollout_url="http://frontend:8000" - ) + shared = external.ExternalEngineInfo("http://worker:9091", "worker", 9091, "regular", 1, rollout_url="http://frontend:8000") assert get_external_rollout_url([engine]) is None assert get_external_rollout_url([shared]) == "http://frontend:8000" @@ -223,44 +204,6 @@ def fake_get(url, timeout): assert args.rollout_num_engines == 1 -def test_apply_external_engine_info_uses_discovery_hook(monkeypatch): - args = Namespace( - rollout_external_engine_addrs=None, - rollout_external_engine_discovery_path="deployment.discover_engines", - ) - - def discover(received_args): - assert received_args is args - return [ - { - "url": "http://worker:9000", - "host": "worker", - "port": 9000, - "worker_type": "regular", - "num_gpus": 2, - "server_info": {"tp_size": 2}, - } - ] - - monkeypatch.setattr(external, "load_function", _loader("deployment.discover_engines", discover)) - - apply_external_engine_info_to_args(args) - - assert args.rollout_num_engines == 1 - assert args.rollout_num_gpus == 2 - assert args.rollout_external_engine_infos[0]["url"] == "http://worker:9000" - - -def test_get_external_engine_class_uses_control_actor_hook(monkeypatch): - class DeploymentControlActor: - pass - - monkeypatch.setattr(external, "load_function", _loader("deployment.ControlActor", DeploymentControlActor)) - actor_class = get_external_engine_class(Namespace(rollout_external_engine_class_path="deployment.ControlActor")) - - assert actor_class is DeploymentControlActor - - def test_apply_external_engine_info_requires_addrs(): args = Namespace(rollout_external_engine_addrs=None) diff --git a/tests/test_megatron_argument_validation.py b/tests/test_megatron_argument_validation.py index fb3333c320..9510b78ffd 100644 --- a/tests/test_megatron_argument_validation.py +++ b/tests/test_megatron_argument_validation.py @@ -218,7 +218,6 @@ def make_slime_validate_args(**overrides): save_debug_train_data=None, load_debug_rollout_data=None, rollout_external_engine_addrs=None, - rollout_external_engine_discovery_path=None, debug_train_only=False, actor_num_gpus_per_node=8, actor_num_nodes=1, @@ -346,28 +345,5 @@ def test_update_weight_delta_requires_local_checkpoint_dir(monkeypatch): module.slime_validate_args(args) -@pytest.mark.unit -def test_external_debug_rollout_does_not_allocate_train_gpus(monkeypatch): - module = load_slime_arguments_module(monkeypatch) - - def discover_external_engine(args, logger=None): - args.rollout_num_gpus = 2 - - module.apply_external_engine_info_to_args = discover_external_engine - args = make_slime_validate_args( - rollout_external_engine_discovery_path="deployment.discover_engines", - debug_rollout_only=True, - ) - - module.slime_validate_args(args) - - assert args.rollout_external is True - assert args.rollout_num_gpus == 2 - assert args.actor_num_gpus_per_node == 0 - assert args.actor_num_nodes == 0 - assert args.critic_num_gpus_per_node == 0 - assert args.critic_num_nodes == 0 - - if __name__ == "__main__": raise SystemExit(pytest.main([__file__])) From 29a62610a0441551825a53d3e7b8393125f6c9c6 Mon Sep 17 00:00:00 2001 From: William Arnold <7565007+Aphoh@users.noreply.github.com> Date: Tue, 28 Jul 2026 23:42:53 +0000 Subject: [PATCH 07/12] style: run external rollout checks --- slime/backends/sglang_utils/external.py | 19 ++++++++++++------- slime/utils/arguments.py | 1 - tests/test_external_sglang_engines.py | 16 +++++++++++++--- tests/test_streaming_rollout.py | 6 +++--- 4 files changed, 28 insertions(+), 14 deletions(-) diff --git a/slime/backends/sglang_utils/external.py b/slime/backends/sglang_utils/external.py index a3efa99162..313fc51a48 100644 --- a/slime/backends/sglang_utils/external.py +++ b/slime/backends/sglang_utils/external.py @@ -8,7 +8,6 @@ import requests - logger = logging.getLogger(__name__) @@ -55,7 +54,9 @@ def normalize_external_engine_addr(addr: str) -> str: addr = addr.rstrip("/") 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).") + raise ValueError( + f"Invalid external SGLang engine address {addr!r}. Use host:port or http://host:port (IPv6 must be bracketed)." + ) return addr @@ -101,7 +102,9 @@ def _infer_worker_type(server_info: dict) -> str: return "regular" -def discover_external_engines(addrs: list[str], engine_api_prefix: str = "", rollout_url: str | None = None, timeout: float = 30.0) -> list[ExternalEngineInfo]: +def discover_external_engines( + addrs: list[str], engine_api_prefix: str = "", rollout_url: str | None = None, timeout: float = 30.0 +) -> list[ExternalEngineInfo]: infos = [] for addr in addrs: url = normalize_external_engine_addr(addr) @@ -138,8 +141,8 @@ def apply_external_engine_info_to_args(args, logger=None) -> None: raise ValueError("External rollout requires --rollout-external-engine-addrs.") infos = discover_external_engines( addrs, - engine_api_prefix=getattr(args, "rollout_external_engine_api_prefix", ""), - rollout_url=getattr(args, "rollout_external_rollout_url", None), + engine_api_prefix=args.rollout_external_engine_api_prefix, + rollout_url=args.rollout_external_rollout_url, ) if not infos: @@ -203,7 +206,9 @@ def onload_kv(self): 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.") + raise RuntimeError( + "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] @@ -219,8 +224,8 @@ def get_external_rollout_url(infos: list[ExternalEngineInfo]) -> str | None: def start_external_rollout_servers(args, *, start_router) -> tuple[dict[str, ExternalRolloutServer], list]: import ray - from slime.ray.utils import add_default_ray_env_vars 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) rollout_url = get_external_rollout_url(infos) diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index 732632c071..93f27f8443 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -1930,7 +1930,6 @@ def slime_validate_args(args): args.colocate = False args.offload_train = args.offload_rollout = False - assert not (args.debug_rollout_only and args.debug_train_only), ( "debug_rollout_only and debug_train_only cannot be set at the same time, " "please set only one of them." ) diff --git a/tests/test_external_sglang_engines.py b/tests/test_external_sglang_engines.py index cfdc90dfad..c84648bd71 100644 --- a/tests/test_external_sglang_engines.py +++ b/tests/test_external_sglang_engines.py @@ -113,16 +113,22 @@ def fake_get(url, timeout): monkeypatch.setattr("slime.backends.sglang_utils.external.requests.get", fake_get) - info = discover_external_engines(["worker:9090"], engine_api_prefix="/engine", rollout_url="http://frontend:8000")[0] + info = discover_external_engines(["worker:9090"], engine_api_prefix="/engine", rollout_url="http://frontend:8000")[ + 0 + ] assert info.engine_api_prefix == "/engine" assert info.rollout_url == "http://frontend:8000" - assert engine_control_url(info.url, info.engine_api_prefix, "flush_cache") == "http://worker:9090/engine/flush_cache" + assert ( + engine_control_url(info.url, info.engine_api_prefix, "flush_cache") == "http://worker:9090/engine/flush_cache" + ) def test_get_external_rollout_url_requires_one_shared_frontend(): engine = external.ExternalEngineInfo("http://worker:9090", "worker", 9090, "regular", 1) - shared = external.ExternalEngineInfo("http://worker:9091", "worker", 9091, "regular", 1, rollout_url="http://frontend:8000") + shared = external.ExternalEngineInfo( + "http://worker:9091", "worker", 9091, "regular", 1, rollout_url="http://frontend:8000" + ) assert get_external_rollout_url([engine]) is None assert get_external_rollout_url([shared]) == "http://frontend:8000" @@ -156,6 +162,8 @@ def fake_get(url, timeout): args = Namespace( rollout_external=True, rollout_external_engine_addrs=["prefill:10090", "decode:10091"], + rollout_external_engine_api_prefix="", + rollout_external_rollout_url=None, rollout_num_gpus=None, rollout_num_gpus_per_engine=1, sglang_pipeline_parallel_size=1, @@ -193,6 +201,8 @@ def fake_get(url, timeout): args = Namespace( rollout_external=True, rollout_external_engine_addrs=["regular:10090"], + rollout_external_engine_api_prefix="", + rollout_external_rollout_url=None, router_pd_disaggregation=True, ) diff --git a/tests/test_streaming_rollout.py b/tests/test_streaming_rollout.py index 9993359842..9a397f305f 100644 --- a/tests/test_streaming_rollout.py +++ b/tests/test_streaming_rollout.py @@ -26,7 +26,8 @@ ProcessorMixin=object, ) -from slime.rollout import sglang_rollout, sglang_streaming_rollout as streaming +from slime.rollout import sglang_rollout +from slime.rollout import sglang_streaming_rollout as streaming from slime.rollout.streaming_utils import merge_stream_chunk from slime.utils.types import Sample @@ -174,8 +175,7 @@ async def exercise(): streaming_generation=True, streaming_tasks=set(), pendings={ - asyncio.create_task(asyncio.sleep(0, result=group)) - for group in ([partial], [terminal], [empty]) + asyncio.create_task(asyncio.sleep(0, result=group)) for group in ([partial], [terminal], [empty]) }, ) monkeypatch.setattr(sglang_rollout, "GenerateState", lambda _args: state) From 0f6753543cd3e2ed0aacba667a586d835cad85a3 Mon Sep 17 00:00:00 2001 From: William Arnold <7565007+Aphoh@users.noreply.github.com> Date: Fri, 31 Jul 2026 00:31:12 +0000 Subject: [PATCH 08/12] refactor: use external control base URLs --- slime/backends/sglang_utils/external.py | 42 +++++++++++-------- slime/backends/sglang_utils/sglang_engine.py | 10 ++--- slime/utils/arguments.py | 15 ++++--- tests/test_external_sglang_engines.py | 43 ++++++++++++++------ 4 files changed, 71 insertions(+), 39 deletions(-) diff --git a/slime/backends/sglang_utils/external.py b/slime/backends/sglang_utils/external.py index 313fc51a48..45cfa42ca1 100644 --- a/slime/backends/sglang_utils/external.py +++ b/slime/backends/sglang_utils/external.py @@ -8,6 +8,8 @@ import requests +from slime.utils.misc import load_function + logger = logging.getLogger(__name__) @@ -20,7 +22,6 @@ class ExternalEngineInfo: num_gpus: int disaggregation_bootstrap_port: int | None = None server_info: dict = dataclasses.field(default_factory=dict) - engine_api_prefix: str = "" rollout_url: str | None = None @property @@ -60,10 +61,9 @@ def normalize_external_engine_addr(addr: str) -> str: return addr -def engine_control_url(url: str, engine_api_prefix: str, method: str) -> str: +def engine_control_url(url: str, method: str) -> str: """Build a per-worker SGLang control URL.""" - prefix = engine_api_prefix.strip("/") - return "{}/{}{}".format(url.rstrip("/"), prefix + "/" if prefix else "", method) + return f"{url.rstrip('/')}/{method}" def external_engine_init_kwargs(info: ExternalEngineInfo) -> dict: @@ -72,18 +72,17 @@ 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 - if info.engine_api_prefix: - init_kwargs["engine_api_prefix"] = info.engine_api_prefix return init_kwargs -def get_server_info(url: str, engine_api_prefix: str = "", timeout: float = 30.0) -> dict: +def get_server_info(url: str, timeout: float = 30.0) -> dict: errors = [] for method in ("server_info", "get_server_info"): - endpoint = engine_control_url(url, engine_api_prefix, method) + endpoint = engine_control_url(url, method) try: response = requests.get(endpoint, timeout=timeout) response.raise_for_status() @@ -103,14 +102,14 @@ def _infer_worker_type(server_info: dict) -> str: def discover_external_engines( - addrs: list[str], engine_api_prefix: str = "", rollout_url: str | None = None, timeout: float = 30.0 + addrs: list[str], rollout_url: str | None = None, timeout: float = 30.0 ) -> list[ExternalEngineInfo]: infos = [] for addr in addrs: url = normalize_external_engine_addr(addr) parsed = urlparse(url) assert parsed.hostname is not None and parsed.port is not None - server_info = get_server_info(url, engine_api_prefix=engine_api_prefix, timeout=timeout) + server_info = get_server_info(url, timeout=timeout) pp_size = int(server_info.get("pp_size") or server_info.get("pipeline_parallel_size") or 1) tp_size = int(server_info.get("tp_size") or server_info.get("tensor_parallel_size") or 1) @@ -125,7 +124,6 @@ def discover_external_engines( port=parsed.port, worker_type=_infer_worker_type(server_info), num_gpus=num_gpus, - engine_api_prefix=engine_api_prefix, rollout_url=normalize_external_engine_addr(rollout_url) if rollout_url else None, disaggregation_bootstrap_port=bootstrap_port, server_info=server_info, @@ -134,14 +132,26 @@ def discover_external_engines( return infos +def external_engine_addrs_from_args(args) -> list[str]: + """Return external-engine control base URLs from the CLI or a discovery plugin.""" + if args.rollout_external_engine_discovery_path: + addrs = load_function(args.rollout_external_engine_discovery_path)(args) + else: + addrs = args.rollout_external_engine_addrs + + if not addrs: + raise ValueError( + "External rollout requires --rollout-external-engine-addrs or " "--rollout-external-engine-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 + + 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 - if not addrs: - raise ValueError("External rollout requires --rollout-external-engine-addrs.") infos = discover_external_engines( - addrs, - engine_api_prefix=args.rollout_external_engine_api_prefix, + external_engine_addrs_from_args(args), rollout_url=args.rollout_external_rollout_url, ) diff --git a/slime/backends/sglang_utils/sglang_engine.py b/slime/backends/sglang_utils/sglang_engine.py index 64e69ef2be..acb9430436 100644 --- a/slime/backends/sglang_utils/sglang_engine.py +++ b/slime/backends/sglang_utils/sglang_engine.py @@ -128,12 +128,11 @@ def init( disaggregation_bootstrap_port=None, router_ip=None, router_port=None, - engine_api_prefix="", + control_url=None, register_to_router=True, ): self.router_ip = router_ip if router_ip is not None else self.args.sglang_router_ip self.router_port = router_port if router_port is not None else self.args.sglang_router_port - self.engine_api_prefix = engine_api_prefix self.register_to_router = register_to_router host = host or get_host_info()[1] @@ -169,6 +168,7 @@ def _format_v6_uri(addr): self.node_rank = server_args_dict["node_rank"] self.server_host = server_args_dict["host"] # with [] if ipv6 self.server_port = server_args_dict["port"] + self.control_url = control_url.rstrip("/") if control_url else f"http://{self.server_host}:{self.server_port}" if self.args.rollout_external: self._init_external(server_args_dict, external_engine_need_check_fields=external_engine_need_check_fields) @@ -186,9 +186,7 @@ def _sanity_check_server_args(actual_server_args, expect_server_args): actual_value == expect_value ), f"{name=} {expect_value=} {actual_value=} {expect_server_args=} {actual_server_args=}" - actual_server_args = get_server_info( - f"http://{self.server_host}:{self.server_port}", engine_api_prefix=self.engine_api_prefix - ) + actual_server_args = get_server_info(self.control_url) _sanity_check_server_args(actual_server_args, expect_server_args) if self.register_to_router: self._register_to_router(expect_server_args) @@ -245,7 +243,7 @@ def _make_request(self, endpoint: str, payload: dict | None = None): return response.json() def _engine_control_url(self, method: str) -> str: - return engine_control_url(f"http://{self.server_host}:{self.server_port}", self.engine_api_prefix, method) + return engine_control_url(self.control_url, method) def health_generate(self, timeout: float = 5.0) -> bool: """Run /health_generate on the underlying SGLang HTTP server. diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index 93f27f8443..3509a091ce 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -570,13 +570,16 @@ def add_rollout_arguments(parser): type=str, default=None, nargs="+", - help="Address and ports of the external engines.", + help="Control base URLs of the external engines, optionally including a path prefix.", ) parser.add_argument( - "--rollout-external-engine-api-prefix", + "--rollout-external-engine-discovery-path", type=str, - default="", - help="Optional per-worker control API prefix, for example /engine.", + default=None, + help=( + "Optional path to a synchronous function that receives args and returns a list of external " + "engine control base URLs." + ), ) parser.add_argument( "--rollout-external-rollout-url", @@ -1903,7 +1906,9 @@ def slime_validate_args(args): ) args.debug_train_only = True - args.rollout_external = args.rollout_external_engine_addrs is not None + args.rollout_external = ( + args.rollout_external_engine_addrs is not None or args.rollout_external_engine_discovery_path is not None + ) if args.rollout_external and not args.debug_train_only: apply_external_engine_info_to_args(args, logger=logger) diff --git a/tests/test_external_sglang_engines.py b/tests/test_external_sglang_engines.py index c84648bd71..4b5d34b26d 100644 --- a/tests/test_external_sglang_engines.py +++ b/tests/test_external_sglang_engines.py @@ -15,6 +15,7 @@ apply_external_engine_info_to_args, discover_external_engines, engine_control_url, + external_engine_init_kwargs, get_external_rollout_url, start_external_rollout_servers, ) @@ -105,7 +106,7 @@ def remote(self, **kwargs): assert len(init_handles) == 1 -def test_discover_external_engines_uses_control_prefix_and_shared_rollout_url(monkeypatch): +def test_discover_external_engines_uses_control_base_url_and_shared_rollout_url(monkeypatch): def fake_get(url, timeout): assert timeout == 30.0 assert url == "http://worker:9090/engine/server_info" @@ -113,16 +114,34 @@ def fake_get(url, timeout): monkeypatch.setattr("slime.backends.sglang_utils.external.requests.get", fake_get) - info = discover_external_engines(["worker:9090"], engine_api_prefix="/engine", rollout_url="http://frontend:8000")[ - 0 - ] + info = discover_external_engines(["worker:9090/engine"], rollout_url="http://frontend:8000")[0] - assert info.engine_api_prefix == "/engine" + assert info.url == "http://worker:9090/engine" assert info.rollout_url == "http://frontend:8000" - assert ( - engine_control_url(info.url, info.engine_api_prefix, "flush_cache") == "http://worker:9090/engine/flush_cache" + assert external_engine_init_kwargs(info)["control_url"] == "http://worker:9090/engine" + assert engine_control_url(info.url, "flush_cache") == "http://worker:9090/engine/flush_cache" + + +def test_apply_external_engine_info_uses_discovery_control_base_urls(monkeypatch): + monkeypatch.setattr(external, "load_function", lambda path: lambda args: ["worker:9090/engine"]) + + def fake_get(url, timeout): + assert url == "http://worker:9090/engine/server_info" + return _Response({"tp_size": 2}) + + monkeypatch.setattr("slime.backends.sglang_utils.external.requests.get", fake_get) + args = Namespace( + rollout_external_engine_addrs=None, + rollout_external_engine_discovery_path="example.discover", + rollout_external_rollout_url="http://frontend:8000", ) + apply_external_engine_info_to_args(args) + + assert args.rollout_external_engine_infos[0]["url"] == "http://worker:9090/engine" + assert args.rollout_external_engine_infos[0]["host"] == "worker" + assert args.rollout_external_engine_infos[0]["port"] == 9090 + def test_get_external_rollout_url_requires_one_shared_frontend(): engine = external.ExternalEngineInfo("http://worker:9090", "worker", 9090, "regular", 1) @@ -162,7 +181,7 @@ def fake_get(url, timeout): args = Namespace( rollout_external=True, rollout_external_engine_addrs=["prefill:10090", "decode:10091"], - rollout_external_engine_api_prefix="", + rollout_external_engine_discovery_path=None, rollout_external_rollout_url=None, rollout_num_gpus=None, rollout_num_gpus_per_engine=1, @@ -201,7 +220,7 @@ def fake_get(url, timeout): args = Namespace( rollout_external=True, rollout_external_engine_addrs=["regular:10090"], - rollout_external_engine_api_prefix="", + rollout_external_engine_discovery_path=None, rollout_external_rollout_url=None, router_pd_disaggregation=True, ) @@ -214,10 +233,10 @@ def fake_get(url, timeout): assert args.rollout_num_engines == 1 -def test_apply_external_engine_info_requires_addrs(): - args = Namespace(rollout_external_engine_addrs=None) +def test_apply_external_engine_info_requires_addrs_or_discovery(): + args = Namespace(rollout_external_engine_addrs=None, rollout_external_engine_discovery_path=None) - with pytest.raises(ValueError, match="rollout-external-engine-addrs"): + with pytest.raises(ValueError, match="rollout-external-engine-addrs or"): apply_external_engine_info_to_args(args) From 84dcdf02e5a96ed31edf66fc38f3649a9f3f982c Mon Sep 17 00:00:00 2001 From: William Arnold <7565007+Aphoh@users.noreply.github.com> Date: Fri, 31 Jul 2026 00:56:54 +0000 Subject: [PATCH 09/12] feat: refresh dynamic external engines before updates --- slime/backends/megatron_utils/actor.py | 8 +- .../update_weight_from_distributed.py | 5 +- slime/backends/sglang_utils/external.py | 161 +++++++++++++----- slime/ray/rollout.py | 7 +- slime/utils/arguments.py | 9 +- tests/test_external_sglang_engines.py | 74 +++++++- 6 files changed, 210 insertions(+), 54 deletions(-) diff --git a/slime/backends/megatron_utils/actor.py b/slime/backends/megatron_utils/actor.py index 9bd653363a..c91d3df2e1 100644 --- a/slime/backends/megatron_utils/actor.py +++ b/slime/backends/megatron_utils/actor.py @@ -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 = getattr(self.args, "rollout_external_dynamic_discovery_path", None) + 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()) ( diff --git a/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py b/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py index 1ba987f06b..4565223630 100644 --- a/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py +++ b/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py @@ -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 @@ -82,7 +82,7 @@ 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, @@ -90,6 +90,7 @@ def connect_rollout_engines( 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: diff --git a/slime/backends/sglang_utils/external.py b/slime/backends/sglang_utils/external.py index 45cfa42ca1..55af9596fc 100644 --- a/slime/backends/sglang_utils/external.py +++ b/slime/backends/sglang_utils/external.py @@ -4,6 +4,7 @@ import dataclasses import logging +from itertools import accumulate from urllib.parse import urlparse import requests @@ -133,28 +134,34 @@ def discover_external_engines( def external_engine_addrs_from_args(args) -> list[str]: - """Return external-engine control base URLs from the CLI or a discovery plugin.""" - if args.rollout_external_engine_discovery_path: - addrs = load_function(args.rollout_external_engine_discovery_path)(args) + """Return external-engine control base URLs from the CLI or dynamic discovery.""" + dynamic_discovery_path = getattr(args, "rollout_external_dynamic_discovery_path", None) + if dynamic_discovery_path: + addrs = load_function(dynamic_discovery_path)(args) else: addrs = args.rollout_external_engine_addrs if not addrs: raise ValueError( - "External rollout requires --rollout-external-engine-addrs or " "--rollout-external-engine-discovery-path." + "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 -def apply_external_engine_info_to_args(args, logger=None) -> None: - """Detect external engines and store the derived topology on ``args``.""" - infos = discover_external_engines( +def discover_external_engine_infos(args) -> list[ExternalEngineInfo]: + return discover_external_engines( external_engine_addrs_from_args(args), rollout_url=args.rollout_external_rollout_url, ) + +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) + if not infos: raise ValueError("External rollout engine discovery returned no engines.") if not all(isinstance(info, ExternalEngineInfo) for info in infos): @@ -177,6 +184,53 @@ 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.""" @@ -184,12 +238,15 @@ class ExternalRolloutServer: 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) @@ -197,6 +254,48 @@ class ExternalRolloutServer: def all_engines(self): return self.engines + def refresh(self) -> bool: + """Refresh dynamic external-engine membership before a weight update.""" + if not getattr(self.args, "rollout_external_dynamic_discovery_path", None): + 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.") @@ -232,13 +331,10 @@ def get_external_rollout_url(infos: list[ExternalEngineInfo]) -> str | None: 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) rollout_url = get_external_rollout_url(infos) + if getattr(args, "rollout_external_dynamic_discovery_path", None) 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: @@ -248,33 +344,15 @@ def start_external_rollout_servers(args, *, start_router) -> tuple[dict[str, Ext 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_kwargs = external_engine_init_kwargs(info) - if rollout_url is not None: - init_kwargs["register_to_router"] = False - init_handles.append(rollout_engine.init.remote(**init_kwargs, 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( @@ -282,6 +360,9 @@ def start_external_rollout_servers(args, *, start_router) -> tuple[dict[str, Ext 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", diff --git a/slime/ray/rollout.py b/slime/ray/rollout.py index abd8e35d70..8b270fd1ba 100644 --- a/slime/ray/rollout.py +++ b/slime/ray/rollout.py @@ -638,6 +638,11 @@ def onload_kv(self): for srv in self.servers.values(): srv.onload_kv() + def refresh_updatable_engines(self): + """Refresh dynamic external-engine membership before the next weight update.""" + srv = self._get_updatable_server() + return srv.refresh() if srv else False + def recover_updatable_engines(self): """Restart dead updatable rollout engines before the next weight update. @@ -655,7 +660,7 @@ def clear_updatable_num_new_engines(self): # when fault tolerance is not enabled, we need to manually clear num_new_engines after update_weights srv = self._get_updatable_server() if srv: - srv.num_new_engines = 0 + srv.clear_num_new_engines() def health_monitoring_pause(self) -> None: for monitor in self._health_monitors: diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index 3509a091ce..bea0f34030 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -573,12 +573,12 @@ def add_rollout_arguments(parser): help="Control base URLs of the external engines, optionally including a path prefix.", ) parser.add_argument( - "--rollout-external-engine-discovery-path", + "--rollout-external-dynamic-discovery-path", type=str, default=None, help=( - "Optional path to a synchronous function that receives args and returns a list of external " - "engine control base URLs." + "Optional path to a synchronous function called before every weight update. It receives args " + "and returns the current external engine control base URLs." ), ) parser.add_argument( @@ -1907,7 +1907,8 @@ def slime_validate_args(args): args.debug_train_only = True args.rollout_external = ( - args.rollout_external_engine_addrs is not None or args.rollout_external_engine_discovery_path is not None + args.rollout_external_engine_addrs is not None + or getattr(args, "rollout_external_dynamic_discovery_path", None) is not None ) if args.rollout_external and not args.debug_train_only: diff --git a/tests/test_external_sglang_engines.py b/tests/test_external_sglang_engines.py index 4b5d34b26d..aedc86e732 100644 --- a/tests/test_external_sglang_engines.py +++ b/tests/test_external_sglang_engines.py @@ -122,7 +122,7 @@ def fake_get(url, timeout): assert engine_control_url(info.url, "flush_cache") == "http://worker:9090/engine/flush_cache" -def test_apply_external_engine_info_uses_discovery_control_base_urls(monkeypatch): +def test_apply_external_engine_info_uses_dynamic_discovery_control_base_urls(monkeypatch): monkeypatch.setattr(external, "load_function", lambda path: lambda args: ["worker:9090/engine"]) def fake_get(url, timeout): @@ -132,7 +132,7 @@ def fake_get(url, timeout): monkeypatch.setattr("slime.backends.sglang_utils.external.requests.get", fake_get) args = Namespace( rollout_external_engine_addrs=None, - rollout_external_engine_discovery_path="example.discover", + rollout_external_dynamic_discovery_path="example.discover", rollout_external_rollout_url="http://frontend:8000", ) @@ -143,6 +143,70 @@ def fake_get(url, timeout): assert args.rollout_external_engine_infos[0]["port"] == 9090 +def test_dynamic_discovery_refreshes_changed_external_engine_membership(monkeypatch): + args = Namespace( + rollout_external_engine_addrs=None, + rollout_external_dynamic_discovery_path="example.discover", + rollout_external_rollout_url="http://frontend:8000", + ) + old_info = external.ExternalEngineInfo( + "http://old:9090/engine", "old", 9090, "regular", 2, server_info={"tp_size": 2} + ) + new_info = external.ExternalEngineInfo( + "http://new:9090/engine", "new", 9090, "regular", 4, server_info={"tp_size": 4} + ) + server = external.ExternalRolloutServer( + engines=["old-engine"], + engine_gpu_counts=[2], + engine_gpu_offsets=[0], + args=args, + engine_infos=[old_info], + register_to_router=False, + ) + monkeypatch.setattr(external, "discover_external_engine_infos", lambda _args: [new_info]) + monkeypatch.setattr( + external, + "_start_external_engine_actors", + lambda *args, **kwargs: (["new-engine"], []), + ) + + assert server.refresh() is True + assert server.engines == ["new-engine"] + assert server.engine_gpu_counts == [4] + assert server.engine_gpu_offsets == [0] + assert server.num_new_engines == 1 + assert server.retired_engines == ["old-engine"] + assert args.rollout_external_engine_infos == [new_info.to_dict()] + + +def test_dynamic_discovery_skips_unchanged_external_engine_membership(monkeypatch): + args = Namespace( + rollout_external_engine_addrs=None, + rollout_external_dynamic_discovery_path="example.discover", + rollout_external_rollout_url="http://frontend:8000", + ) + info = external.ExternalEngineInfo( + "http://worker:9090/engine", "worker", 9090, "regular", 2, server_info={"tp_size": 2} + ) + server = external.ExternalRolloutServer( + engines=["engine"], + engine_gpu_counts=[2], + engine_gpu_offsets=[0], + args=args, + engine_infos=[info], + register_to_router=False, + ) + monkeypatch.setattr(external, "discover_external_engine_infos", lambda _args: [info]) + monkeypatch.setattr( + external, + "_start_external_engine_actors", + lambda *args, **kwargs: pytest.fail("unchanged membership must not recreate control actors"), + ) + + assert server.refresh() is False + assert server.engines == ["engine"] + + def test_get_external_rollout_url_requires_one_shared_frontend(): engine = external.ExternalEngineInfo("http://worker:9090", "worker", 9090, "regular", 1) shared = external.ExternalEngineInfo( @@ -181,7 +245,7 @@ def fake_get(url, timeout): args = Namespace( rollout_external=True, rollout_external_engine_addrs=["prefill:10090", "decode:10091"], - rollout_external_engine_discovery_path=None, + rollout_external_dynamic_discovery_path=None, rollout_external_rollout_url=None, rollout_num_gpus=None, rollout_num_gpus_per_engine=1, @@ -220,7 +284,7 @@ def fake_get(url, timeout): args = Namespace( rollout_external=True, rollout_external_engine_addrs=["regular:10090"], - rollout_external_engine_discovery_path=None, + rollout_external_dynamic_discovery_path=None, rollout_external_rollout_url=None, router_pd_disaggregation=True, ) @@ -234,7 +298,7 @@ def fake_get(url, timeout): def test_apply_external_engine_info_requires_addrs_or_discovery(): - args = Namespace(rollout_external_engine_addrs=None, rollout_external_engine_discovery_path=None) + args = Namespace(rollout_external_engine_addrs=None, rollout_external_dynamic_discovery_path=None) with pytest.raises(ValueError, match="rollout-external-engine-addrs or"): apply_external_engine_info_to_args(args) From 1b0ec068074695aeea0423b7edfb9debd5b18178 Mon Sep 17 00:00:00 2001 From: William Arnold <7565007+Aphoh@users.noreply.github.com> Date: Wed, 12 Aug 2026 01:26:48 +0000 Subject: [PATCH 10/12] refactor: streamline external rollout URL handling --- slime/backends/megatron_utils/actor.py | 2 +- slime/backends/sglang_utils/external.py | 34 ++++++++----------------- slime/utils/arguments.py | 3 +-- tests/test_external_sglang_engines.py | 24 +++++------------ 4 files changed, 19 insertions(+), 44 deletions(-) diff --git a/slime/backends/megatron_utils/actor.py b/slime/backends/megatron_utils/actor.py index c91d3df2e1..7c31b0979f 100644 --- a/slime/backends/megatron_utils/actor.py +++ b/slime/backends/megatron_utils/actor.py @@ -572,7 +572,7 @@ def update_weights(self) -> None: if self.args.debug_train_only or self.args.debug_rollout_only: return - dynamic_discovery_path = getattr(self.args, "rollout_external_dynamic_discovery_path", None) + 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: if dynamic_discovery_path: diff --git a/slime/backends/sglang_utils/external.py b/slime/backends/sglang_utils/external.py index 55af9596fc..9c13bb1953 100644 --- a/slime/backends/sglang_utils/external.py +++ b/slime/backends/sglang_utils/external.py @@ -23,7 +23,6 @@ class ExternalEngineInfo: num_gpus: int disaggregation_bootstrap_port: int | None = None server_info: dict = dataclasses.field(default_factory=dict) - rollout_url: str | None = None @property def is_pd_worker(self) -> bool: @@ -102,9 +101,7 @@ def _infer_worker_type(server_info: dict) -> str: return "regular" -def discover_external_engines( - addrs: list[str], rollout_url: str | None = None, timeout: float = 30.0 -) -> list[ExternalEngineInfo]: +def discover_external_engines(addrs: list[str], timeout: float = 30.0) -> list[ExternalEngineInfo]: infos = [] for addr in addrs: url = normalize_external_engine_addr(addr) @@ -125,7 +122,6 @@ def discover_external_engines( port=parsed.port, worker_type=_infer_worker_type(server_info), num_gpus=num_gpus, - rollout_url=normalize_external_engine_addr(rollout_url) if rollout_url else None, disaggregation_bootstrap_port=bootstrap_port, server_info=server_info, ) @@ -135,7 +131,7 @@ def discover_external_engines( def external_engine_addrs_from_args(args) -> list[str]: """Return external-engine control base URLs from the CLI or dynamic discovery.""" - dynamic_discovery_path = getattr(args, "rollout_external_dynamic_discovery_path", None) + dynamic_discovery_path = args.rollout_external_dynamic_discovery_path if dynamic_discovery_path: addrs = load_function(dynamic_discovery_path)(args) else: @@ -152,10 +148,7 @@ def external_engine_addrs_from_args(args) -> list[str]: def discover_external_engine_infos(args) -> list[ExternalEngineInfo]: - return discover_external_engines( - external_engine_addrs_from_args(args), - rollout_url=args.rollout_external_rollout_url, - ) + return discover_external_engines(external_engine_addrs_from_args(args)) def apply_external_engine_info_to_args(args, logger=None) -> None: @@ -164,8 +157,6 @@ def apply_external_engine_info_to_args(args, logger=None) -> None: if not infos: raise ValueError("External rollout engine discovery returned no engines.") - if not all(isinstance(info, ExternalEngineInfo) for info in infos): - raise TypeError("External rollout engine discovery must return ExternalEngineInfo objects or dictionaries.") args.rollout_external_engine_infos = [info.to_dict() for info in infos] args.rollout_num_engines = len(infos) @@ -256,7 +247,7 @@ def all_engines(self): def refresh(self) -> bool: """Refresh dynamic external-engine membership before a weight update.""" - if not getattr(self.args, "rollout_external_dynamic_discovery_path", None): + if not self.args.rollout_external_dynamic_discovery_path: return False infos = discover_external_engine_infos(self.args) @@ -321,19 +312,14 @@ def external_engine_infos_from_args(args) -> list[ExternalEngineInfo]: return [ExternalEngineInfo(**info) if isinstance(info, dict) else info for info in raw_infos] -def get_external_rollout_url(infos: list[ExternalEngineInfo]) -> str | None: - rollout_urls = {info.rollout_url for info in infos} - if rollout_urls == {None}: - return None - if None in rollout_urls or len(rollout_urls) != 1: - raise ValueError("External rollout engines must all use the same rollout_url, or none.") - return rollout_urls.pop() - - def start_external_rollout_servers(args, *, start_router) -> tuple[dict[str, ExternalRolloutServer], list]: infos = external_engine_infos_from_args(args) - rollout_url = get_external_rollout_url(infos) - if getattr(args, "rollout_external_dynamic_discovery_path", None) and rollout_url is None: + 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)) diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index bea0f34030..706abb3022 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -1907,8 +1907,7 @@ def slime_validate_args(args): args.debug_train_only = True args.rollout_external = ( - args.rollout_external_engine_addrs is not None - or getattr(args, "rollout_external_dynamic_discovery_path", None) is not None + args.rollout_external_engine_addrs is not None or args.rollout_external_dynamic_discovery_path is not None ) if args.rollout_external and not args.debug_train_only: diff --git a/tests/test_external_sglang_engines.py b/tests/test_external_sglang_engines.py index aedc86e732..c2f2fce7a7 100644 --- a/tests/test_external_sglang_engines.py +++ b/tests/test_external_sglang_engines.py @@ -16,7 +16,6 @@ discover_external_engines, engine_control_url, external_engine_init_kwargs, - get_external_rollout_url, start_external_rollout_servers, ) from slime.utils.http_utils import get_rollout_num_engines @@ -98,7 +97,11 @@ def remote(self, **kwargs): num_gpus=8, server_info={"tp_size": 4, "pp_size": 2, "ep_size": 4, "moe_dp_size": 2}, ) - args = Namespace(rollout_external_engine_infos=[info.to_dict()]) + args = Namespace( + rollout_external_engine_infos=[info.to_dict()], + rollout_external_dynamic_discovery_path=None, + rollout_external_rollout_url=None, + ) servers, init_handles = start_external_rollout_servers(args, start_router=lambda *args, **kwargs: ("host1", 30000)) @@ -106,7 +109,7 @@ def remote(self, **kwargs): assert len(init_handles) == 1 -def test_discover_external_engines_uses_control_base_url_and_shared_rollout_url(monkeypatch): +def test_discover_external_engines_uses_control_base_url(monkeypatch): def fake_get(url, timeout): assert timeout == 30.0 assert url == "http://worker:9090/engine/server_info" @@ -114,10 +117,9 @@ def fake_get(url, timeout): monkeypatch.setattr("slime.backends.sglang_utils.external.requests.get", fake_get) - info = discover_external_engines(["worker:9090/engine"], rollout_url="http://frontend:8000")[0] + info = discover_external_engines(["worker:9090/engine"])[0] assert info.url == "http://worker:9090/engine" - assert info.rollout_url == "http://frontend:8000" assert external_engine_init_kwargs(info)["control_url"] == "http://worker:9090/engine" assert engine_control_url(info.url, "flush_cache") == "http://worker:9090/engine/flush_cache" @@ -207,18 +209,6 @@ def test_dynamic_discovery_skips_unchanged_external_engine_membership(monkeypatc assert server.engines == ["engine"] -def test_get_external_rollout_url_requires_one_shared_frontend(): - engine = external.ExternalEngineInfo("http://worker:9090", "worker", 9090, "regular", 1) - shared = external.ExternalEngineInfo( - "http://worker:9091", "worker", 9091, "regular", 1, rollout_url="http://frontend:8000" - ) - - assert get_external_rollout_url([engine]) is None - assert get_external_rollout_url([shared]) == "http://frontend:8000" - with pytest.raises(ValueError, match="same rollout_url"): - get_external_rollout_url([engine, shared]) - - def test_apply_external_engine_info_handles_pd(monkeypatch): payloads = { "http://prefill:10090/server_info": { From bf0b668cdf8de77881b997010a803fe23707320a Mon Sep 17 00:00:00 2001 From: William Arnold <7565007+Aphoh@users.noreply.github.com> Date: Wed, 12 Aug 2026 05:38:45 +0000 Subject: [PATCH 11/12] refactor: inline external engine control URLs --- slime/backends/sglang_utils/external.py | 12 +++--------- slime/backends/sglang_utils/sglang_engine.py | 16 ++++++---------- tests/test_external_sglang_engines.py | 2 -- 3 files changed, 9 insertions(+), 21 deletions(-) diff --git a/slime/backends/sglang_utils/external.py b/slime/backends/sglang_utils/external.py index 9c13bb1953..81102ffa29 100644 --- a/slime/backends/sglang_utils/external.py +++ b/slime/backends/sglang_utils/external.py @@ -61,11 +61,6 @@ def normalize_external_engine_addr(addr: str) -> str: return addr -def engine_control_url(url: str, method: str) -> str: - """Build a per-worker SGLang control URL.""" - return f"{url.rstrip('/')}/{method}" - - def external_engine_init_kwargs(info: ExternalEngineInfo) -> dict: init_kwargs = { "dist_init_addr": f"{info.host}:{info.port}", @@ -81,14 +76,13 @@ def external_engine_init_kwargs(info: ExternalEngineInfo) -> dict: def get_server_info(url: str, timeout: float = 30.0) -> dict: errors = [] - for method in ("server_info", "get_server_info"): - endpoint = engine_control_url(url, method) + for endpoint in ("/server_info", "/get_server_info"): try: - response = requests.get(endpoint, timeout=timeout) + response = requests.get(f"{url}{endpoint}", timeout=timeout) 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)}") diff --git a/slime/backends/sglang_utils/sglang_engine.py b/slime/backends/sglang_utils/sglang_engine.py index acb9430436..751d1309d9 100644 --- a/slime/backends/sglang_utils/sglang_engine.py +++ b/slime/backends/sglang_utils/sglang_engine.py @@ -10,7 +10,7 @@ from sglang.srt.utils import kill_process_tree from urllib3.exceptions import NewConnectionError -from slime.backends.sglang_utils.external import engine_control_url, get_server_info +from slime.backends.sglang_utils.external import get_server_info from slime.ray.ray_actor import RayActor from slime.utils.http_utils import get_host_info @@ -233,8 +233,7 @@ def _make_request(self, endpoint: str, payload: dict | None = None): if self.node_rank != 0: return - url = self._engine_control_url(endpoint) - response = requests.post(url, json=payload or {}) + response = requests.post(f"{self.control_url}/{endpoint}", json=payload or {}) try: response.raise_for_status() except requests.exceptions.HTTPError as e: @@ -242,9 +241,6 @@ def _make_request(self, endpoint: str, payload: dict | None = None): raise return response.json() - def _engine_control_url(self, method: str) -> str: - return engine_control_url(self.control_url, method) - def health_generate(self, timeout: float = 5.0) -> bool: """Run /health_generate on the underlying SGLang HTTP server. @@ -261,7 +257,7 @@ def health_generate(self, timeout: float = 5.0) -> bool: return True response = requests.get( - self._engine_control_url("health_generate"), + f"{self.control_url}/health_generate", timeout=timeout, ) response.raise_for_status() @@ -299,7 +295,7 @@ def flush_cache(self): # flush cache will not return status_code 200 when there are pending requests for _ in range(60): try: - response = requests.get(self._engine_control_url("flush_cache")) + response = requests.get(f"{self.control_url}/flush_cache") if response.status_code == 200: break logger.info(f"Error flushing cache: HTTP {response.status_code} {response.text!r}") @@ -448,14 +444,14 @@ def update_weights_from_distributed( def pause_generation(self): if self.node_rank != 0: return - response = requests.post(self._engine_control_url("pause_generation"), json={}) + response = requests.post(f"{self.control_url}/pause_generation", json={}) response.raise_for_status() return response def continue_generation(self): if self.node_rank != 0: return - response = requests.post(self._engine_control_url("continue_generation"), json={}) + response = requests.post(f"{self.control_url}/continue_generation", json={}) response.raise_for_status() return response diff --git a/tests/test_external_sglang_engines.py b/tests/test_external_sglang_engines.py index c2f2fce7a7..595f0269ce 100644 --- a/tests/test_external_sglang_engines.py +++ b/tests/test_external_sglang_engines.py @@ -14,7 +14,6 @@ ExternalEngineInfo, apply_external_engine_info_to_args, discover_external_engines, - engine_control_url, external_engine_init_kwargs, start_external_rollout_servers, ) @@ -121,7 +120,6 @@ def fake_get(url, timeout): assert info.url == "http://worker:9090/engine" assert external_engine_init_kwargs(info)["control_url"] == "http://worker:9090/engine" - assert engine_control_url(info.url, "flush_cache") == "http://worker:9090/engine/flush_cache" def test_apply_external_engine_info_uses_dynamic_discovery_control_base_urls(monkeypatch): From 971bf61320aeaa69dccec710740a9ec718c9ce9c Mon Sep 17 00:00:00 2001 From: William Arnold <7565007+Aphoh@users.noreply.github.com> Date: Wed, 12 Aug 2026 06:15:54 +0000 Subject: [PATCH 12/12] refactor: tighten external engine integration --- slime/backends/sglang_utils/external.py | 3 --- slime/backends/sglang_utils/sglang_engine.py | 4 ++-- slime/ray/actor_group.py | 4 ---- slime/utils/arguments.py | 8 ++++---- tests/test_placement_group.py | 14 -------------- 5 files changed, 6 insertions(+), 27 deletions(-) diff --git a/slime/backends/sglang_utils/external.py b/slime/backends/sglang_utils/external.py index 81102ffa29..2ad861545d 100644 --- a/slime/backends/sglang_utils/external.py +++ b/slime/backends/sglang_utils/external.py @@ -149,9 +149,6 @@ 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) - if not infos: - raise ValueError("External rollout engine discovery returned no engines.") - args.rollout_external_engine_infos = [info.to_dict() for info in infos] args.rollout_num_engines = len(infos) args.rollout_num_gpus = sum(info.num_gpus for info in infos) diff --git a/slime/backends/sglang_utils/sglang_engine.py b/slime/backends/sglang_utils/sglang_engine.py index 751d1309d9..4feadca81d 100644 --- a/slime/backends/sglang_utils/sglang_engine.py +++ b/slime/backends/sglang_utils/sglang_engine.py @@ -133,7 +133,7 @@ def init( ): self.router_ip = router_ip if router_ip is not None else self.args.sglang_router_ip self.router_port = router_port if router_port is not None else self.args.sglang_router_port - self.register_to_router = register_to_router + self.should_register_to_router = register_to_router host = host or get_host_info()[1] @@ -188,7 +188,7 @@ def _sanity_check_server_args(actual_server_args, expect_server_args): actual_server_args = get_server_info(self.control_url) _sanity_check_server_args(actual_server_args, expect_server_args) - if self.register_to_router: + if self.should_register_to_router: self._register_to_router(expect_server_args) def _init_normal(self, server_args_dict): diff --git a/slime/ray/actor_group.py b/slime/ray/actor_group.py index f72f676c08..7662a0de69 100644 --- a/slime/ray/actor_group.py +++ b/slime/ray/actor_group.py @@ -191,10 +191,6 @@ def create(self, rollout_manager=None): if rollout_manager is not None: self._rollout_manager = rollout_manager self.args.update_weight_start_version = self._disk_weight_version - if self._num_nodes * self._num_gpus_per_node == 0: - assert self.args.debug_rollout_only, "zero-sized train groups are only valid in rollout-only debug mode" - start_rollout_id = self.args.start_rollout_id - return [0 if start_rollout_id is None else start_rollout_id] self._allocate_gpus_for_actor(self._pg, self._num_gpus_per_actor) start_rollout_ids = ray.get( [ diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index 706abb3022..f0a3184c22 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -1914,6 +1914,9 @@ def slime_validate_args(args): apply_external_engine_info_to_args(args, logger=logger) args.use_critic = args.advantage_estimator == "ppo" + # Critic always uses the same GPU count as actor. + args.critic_num_gpus_per_node = args.actor_num_gpus_per_node + args.critic_num_nodes = args.actor_num_nodes if args.offload: args.offload_train = True @@ -1921,10 +1924,7 @@ def slime_validate_args(args): del args.offload if args.debug_rollout_only: - if args.rollout_external: - args.actor_num_gpus_per_node = 0 - args.actor_num_nodes = 0 - elif args.colocate and args.rollout_num_gpus is None: + if args.colocate and args.rollout_num_gpus is None: args.rollout_num_gpus = args.actor_num_gpus_per_node * args.actor_num_nodes elif args.rollout_num_gpus == 0: args.actor_num_gpus_per_node = 0 diff --git a/tests/test_placement_group.py b/tests/test_placement_group.py index 481dd605ec..c1ae8aedef 100644 --- a/tests/test_placement_group.py +++ b/tests/test_placement_group.py @@ -8,7 +8,6 @@ if str(REPO_ROOT) not in sys.path: sys.path.insert(0, str(REPO_ROOT)) -from slime.ray.actor_group import RayTrainGroup from slime.ray.placement_group import _create_placement_group, _get_placement_group_layout NUM_GPUS = 0 @@ -51,18 +50,5 @@ def test_create_zero_gpu_placement_group_is_empty(): assert _create_placement_group(0) == (None, [], []) -@pytest.mark.parametrize(("start_rollout_id", "expected"), [(None, 0), (7, 7)]) -def test_zero_sized_debug_train_group_uses_configured_rollout_id(start_rollout_id, expected): - args = Namespace(debug_rollout_only=True, start_rollout_id=start_rollout_id) - group = RayTrainGroup( - args=args, - num_nodes=0, - num_gpus_per_node=0, - pg=(None, [], []), - ) - - assert group.create() == [expected] - - if __name__ == "__main__": raise SystemExit(pytest.main([__file__]))