Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
39 commits
Select commit Hold shift + click to select a range
aac55ac
feat(proto): carry the engine-reported weight_version on TokenSpeed g…
key4ng Sep 22, 2026
a0395b6
feat(grpc): stamp meta_info.weight_version from the engine-reported v…
key4ng Sep 22, 2026
cf24909
feat(servicer): TokenSpeed reports live weight version and pause stat…
key4ng Sep 22, 2026
e7d5535
feat(discovery): read the TokenSpeed RL control endpoint and capabili…
key4ng Sep 22, 2026
dd13c4a
feat(rl): model a worker's control endpoint separately from its data …
key4ng Sep 22, 2026
e39c60c
feat(rl): proxy control calls to the worker's control endpoint regard…
key4ng Sep 22, 2026
a4cb1d3
feat(rl): resolve control endpoints for gRPC and ZMQ workers from the…
key4ng Sep 22, 2026
1b9b806
feat(rl): report control_url in discovery, add the TokenSpeed capabil…
key4ng Sep 22, 2026
cff21aa
style: apply rustfmt to the control-endpoint changes
key4ng Sep 22, 2026
226b62f
test(mock_worker): advertise custom server_args and stamp a weight ve…
key4ng Sep 22, 2026
a2c187c
chore(mock_worker): record the prost-types dependency in Cargo.lock
key4ng Sep 22, 2026
7e6062a
test(rl): drive a TokenSpeed gRPC worker through /v1/rl via its contr…
key4ng Sep 22, 2026
bd80947
feat(serve): wire the TokenSpeed control app for ZMQ workers and stam…
key4ng Sep 22, 2026
d6744dc
docs(rl): control endpoints, TokenSpeed drift rows, and the TokenSpee…
key4ng Sep 22, 2026
f1b9f69
test(e2e): launch TokenSpeed gRPC workers with their in-engine RL con…
key4ng Sep 22, 2026
fff0d2d
test(e2e): run the RL control-plane lane on TokenSpeed gRPC and add a…
key4ng Sep 22, 2026
3ded419
fix(servicer): accept a synchronous is_scheduler_paused
key4ng Sep 22, 2026
20cbd50
fix(rl): treat a blank rl.control_url label as absent; document base_url
key4ng Sep 22, 2026
fd6376b
refactor(grpc): route SSE weight_version through effective_weight_ver…
key4ng Sep 22, 2026
fca4315
fix(serve): stop retrying 4xx when stamping rl.control_url; require zmq
key4ng Sep 22, 2026
4f1ae8d
docs(rl): HTTP workers ignore rl.control_url; narrow recorder doc
key4ng Sep 22, 2026
21b527d
style(rl): drop redundant serde_json qualifications in the discovery …
key4ng Sep 22, 2026
7099c5a
fix(rl): TokenSpeed refits are distributed-only; drop disk from the s…
key4ng Sep 22, 2026
9a4c7a4
test(e2e): TokenSpeed lane asserts disk refits are refused per worker…
key4ng Sep 22, 2026
e2ef759
feat(rl): trainer-side NCCL refit example for TokenSpeed
key4ng Sep 22, 2026
fdae590
chore(rl): make refit_from_trainer.py executable like the disk example
key4ng Sep 22, 2026
2a4b8d6
fix(rl): let the trainer refit override a missing tp_size
key4ng Sep 22, 2026
4b0329c
docs(rl): TokenSpeed control-endpoint acceptance run on the H200 node
key4ng Sep 22, 2026
7f568d5
docs(rl): scope the acceptance cleanup to this lane's own processes
key4ng Sep 22, 2026
897f147
docs(rl): record the confirmed cause of the TokenSpeed refit no-op
key4ng Sep 22, 2026
99ac741
fix(rl): review round 1 on the trainer-side refit example
key4ng Sep 22, 2026
681188b
fix(rl): fall back to attn_tp_size when an engine reports no tp_size
key4ng Sep 22, 2026
4b0f04c
fix(rl): bound the refit verification to one --timeout and check flus…
key4ng Sep 22, 2026
31f2ee1
docs(rl): TokenSpeed control-endpoint acceptance passes on the H200 node
key4ng Sep 22, 2026
a5e45fa
fix(router): never forward the wildcard model placeholder on /generate
key4ng Sep 23, 2026
123074b
fix(grpc): serve slime's model-less single-prompt /generate like SGLang
key4ng Sep 23, 2026
b156402
fix(grpc): tighten the single-prompt unwrap and the model default aft…
key4ng Sep 23, 2026
9763f71
docs(rl): slime A/B, Task 19/20 decision and the wildcard-model fix i…
key4ng Sep 23, 2026
16f353e
docs(rl): list the gateway-side files that carry RL data and name bot…
key4ng Sep 23, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

8 changes: 5 additions & 3 deletions bindings/python/src/smg/rl.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,9 +8,10 @@
SGLang requires a JSON body on `pause_generation` and `continue_generation`
(a bodyless POST is a 400), so bodyless routes are sent as `{}`.

Only HTTP workers can be proxied. A gRPC or ZMQ worker matched by a selector is
reported in `failed[]` as `unsupported_connection_mode`, which makes `fanout`
raise `FanoutError` unless `allow_partial=True`.
A worker with no control endpoint (a gRPC or ZMQ worker without an
`rl.control_url` label) matched by a selector is reported in `failed[]` as
`no_control_endpoint`, which makes `fanout` raise `FanoutError` unless
`allow_partial=True`.
"""

from __future__ import annotations
Expand Down Expand Up @@ -52,6 +53,7 @@ class Worker:
role: str | None
health: str
weight_version: str | None
control_url: str | None = None
labels: dict[str, str] = field(default_factory=dict)
capabilities: dict[str, Any] = field(default_factory=dict)

Expand Down
104 changes: 102 additions & 2 deletions bindings/python/src/smg/serve.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,14 +8,18 @@

import argparse
import atexit
import json
import logging
import os
import random
import signal
import socket
import subprocess
import sys
import threading
import time
import urllib.error
import urllib.request
from abc import ABC, abstractmethod

from smg.launch_router import launch_router
Expand Down Expand Up @@ -83,6 +87,19 @@ def _zmq_handshake_port(ipc_url: str) -> int:
return _ZMQ_HANDSHAKE_PORT_BASE + (h % _ZMQ_HANDSHAKE_PORT_SPAN)


def _rl_control_port(port: int) -> int:
"""Port for a TokenSpeed worker's in-engine RL control app.

Offset from the worker port so co-located workers (consecutive ports) and
the engine's own +233 distributed store never collide, reflected below the
u16 ceiling and hopped past SMG's ZMQ handshake band like ``dist_port``.
"""
p = port + 400 if port + 400 <= 65535 else port - 400
if _ZMQ_HANDSHAKE_PORT_BASE <= p < _ZMQ_HANDSHAKE_PORT_BASE + _ZMQ_HANDSHAKE_PORT_SPAN:
p += _ZMQ_HANDSHAKE_PORT_SPAN
return p


def _reject_handshake_port_collisions(ports: list[int]) -> None:
"""Fail before launch if two workers derive the same ZMQ handshake port.

Expand Down Expand Up @@ -354,6 +371,10 @@ class TokenspeedWorkerLauncher(WorkerLauncher):
def _get_tp_size(self, args: argparse.Namespace) -> int:
return getattr(args, "tensor_parallel_size", 1) or 1

def control_url(self, port: int) -> str:
"""URL of the in-engine RL control app this launcher started for ``port``."""
return f"http://127.0.0.1:{_rl_control_port(port)}"

def build_command(
self, args: argparse.Namespace, backend_args: list[str], host: str, port: int
) -> list[str]:
Expand Down Expand Up @@ -420,6 +441,10 @@ def _build_zmq_command(
str(rpc_port),
"--zmq-engine-index",
"0",
"--rl-control-host",
"127.0.0.1",
"--rl-control-port",
str(_rl_control_port(port)),
]
cmd.extend(
self._backend_arg_defaults(
Expand Down Expand Up @@ -449,6 +474,8 @@ def _build_zmq_command(
"--data-parallel-address",
"--data-parallel-rpc-port",
"--zmq-engine-index",
"--rl-control-host",
"--rl-control-port",
],
)
)
Expand Down Expand Up @@ -547,8 +574,6 @@ def build_command(
def _http_health_check(url: str, timeout: float) -> bool:
"""GET the URL and return True on HTTP 200."""
try:
import urllib.request

req = urllib.request.Request(url, method="GET")
with urllib.request.urlopen(req, timeout=timeout) as resp:
return resp.status == 200
Expand Down Expand Up @@ -868,6 +893,60 @@ def parse_serve_args(
_WORKER_SHUTDOWN_TIMEOUT = 30


def _stamp_rl_control_labels(
gateway_url: str,
api_key: str | None,
targets: list[tuple[str, str]],
deadline_s: float,
) -> None:
"""Label each ZMQ worker with its control endpoint once the gateway lists it.

ZMQ discovery yields no labels, so the launcher, which owns both ends,
stamps ``rl.control_url`` through the worker update route (labels merge).
"""
headers = {"Content-Type": "application/json"}
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
pending = dict(targets)
stop_at = time.monotonic() + deadline_s
while pending and time.monotonic() < stop_at:
try:
req = urllib.request.Request(f"{gateway_url}/workers", headers=headers, method="GET")
with urllib.request.urlopen(req, timeout=5) as resp:
workers = json.loads(resp.read()).get("workers", [])
except Exception as e: # noqa: BLE001 — the gateway may not be up yet
logger.debug("rl.control_url stamping: gateway not ready: %s", e)
time.sleep(1)
continue
for w in workers:
url = str(w.get("url", ""))
if url not in pending:
continue
body = json.dumps({"labels": {"rl.control_url": pending[url]}}).encode()
req = urllib.request.Request(
f"{gateway_url}/workers/{w['id']}", data=body, headers=headers, method="PATCH"
)
try:
with urllib.request.urlopen(req, timeout=5):
pass
logger.info("stamped rl.control_url=%s on %s", pending[url], url)
del pending[url]
except urllib.error.HTTPError as e:
if 400 <= e.code < 500:
logger.warning(
"rl.control_url stamping got HTTP %s for %s; not retrying", e.code, url
)
del pending[url]
else:
logger.warning("rl.control_url stamping failed for %s: %s", url, e)
except Exception as e: # noqa: BLE001
logger.warning("rl.control_url stamping failed for %s: %s", url, e)
if pending:
time.sleep(1)
for url in pending:
logger.warning("rl.control_url never stamped on %s (gateway did not list it)", url)


class ServeOrchestrator:
"""Coordinate worker launch, health checking, router startup, and shutdown."""

Expand All @@ -890,6 +969,27 @@ def run(self) -> None:
self._launch_workers()
self._wait_healthy()
router_args = self._build_router_args()
if (
getattr(router_args, "enable_rl", False)
and self.backend == "tokenspeed"
and getattr(self.args, "connection_mode", "grpc") == "zmq"
):
control = getattr(self.launcher, "control_url", None)
if callable(control):
targets = [
(
self.launcher.worker_url(self.args, self.args.worker_host, port),
control(port),
)
for _, port in self.workers
]
gateway_url = f"http://127.0.0.1:{router_args.port}"
threading.Thread(
target=_stamp_rl_control_labels,
args=(gateway_url, getattr(router_args, "api_key", None), targets, 300.0),
name="smg-rl-control-labels",
daemon=True,
).start()
launch_router(router_args)
finally:
self._cleanup_workers()
Expand Down
124 changes: 124 additions & 0 deletions bindings/python/tests/test_refit_from_trainer_layout.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,124 @@
"""The trainer-side refit example's rank layout (no gateway, no GPU, no engine).

`rank_layout` decides which ranks of the trainer's NCCL group each engine
occupies, which is the one piece of `examples/rl/refit_from_trainer.py` that is
pure arithmetic and the one an off-by-one silently deadlocks.
"""

from __future__ import annotations

import importlib.util
import sys
from collections.abc import Iterator
from contextlib import contextmanager
from pathlib import Path
from typing import Any

import pytest

REPO = Path(__file__).resolve().parents[3]
SCRIPT = REPO / "examples" / "rl" / "refit_from_trainer.py"
SRC = str(REPO / "bindings" / "python" / "src")


def _smg_modules() -> list[str]:
return [name for name in sys.modules if name == "smg" or name.startswith("smg.")]


@contextmanager
def _loaded_script() -> Iterator[Any]:
"""Import the example by path, with this checkout's `smg.rl` importable.

The example imports `smg.rl`, torch and transformers at the top, the way a
trainer would, and `smg` is usually already imported from wherever the
package was installed -- a checkout that need not be this one, and whose
compiled `smg_rs` extension is not in this source tree. So point `smg` at
this checkout's sources for the import and put the interpreter back
afterwards: leaving either `sys.path` or `sys.modules` shifted breaks every
later test in the session that imports `smg.smg_rs`.
"""
saved_path = list(sys.path)
saved_modules = {name: sys.modules[name] for name in _smg_modules()}
try:
sys.path.insert(0, SRC)
for name in list(saved_modules):
del sys.modules[name]
spec = importlib.util.spec_from_file_location("refit_from_trainer", SCRIPT)
assert spec and spec.loader
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
yield module
finally:
sys.path[:] = saved_path
for name in _smg_modules():
del sys.modules[name]
sys.modules.update(saved_modules)


@pytest.fixture(scope="module")
def script() -> Iterator[Any]:
pytest.importorskip("torch")
pytest.importorskip("transformers")
with _loaded_script() as module:
yield module


@pytest.mark.unit
@pytest.mark.parametrize(
("tp_sizes", "expected"),
[
([], (1, [])),
([1], (2, [1])),
([1, 1], (3, [1, 2])),
([2, 4], (7, [1, 3])),
],
)
def test_rank_layout(script, tp_sizes, expected):
assert script.rank_layout(tp_sizes) == expected


@pytest.mark.unit
def test_rank_layout_rejects_a_nonsense_tp_size(script):
with pytest.raises(ValueError, match="tp_size must be >= 1"):
script.rank_layout([1, 0])


@pytest.mark.unit
@pytest.mark.parametrize(
("host", "loopback"),
[
("127.0.0.1", True),
("127.1.2.3", True),
("localhost", True),
("::1", True),
("[::1]", True),
("0.0.0.0", True),
("10.0.1.52", False),
("trainer.internal", False),
(None, False),
("", False),
],
)
def test_is_loopback(script, host, loopback):
assert script._is_loopback(host) is loopback


@pytest.mark.unit
def test_loading_the_example_restores_the_interpreter():
"""Importing the example must not leave `smg` pointed at this source tree.

Every other test file in this suite imports `smg.smg_rs` from inside a test
body, and this source tree has no compiled extension, so a leaked
`sys.path` entry or a purged `sys.modules` entry fails all of them.
"""
pytest.importorskip("torch")
pytest.importorskip("transformers")
before_path = list(sys.path)
before_modules = {name: sys.modules[name] for name in _smg_modules()}

with _loaded_script() as module:
assert module.rank_layout([1]) == (2, [1])
assert sys.path[0] == SRC

assert sys.path == before_path
assert {name: sys.modules[name] for name in _smg_modules()} == before_modules
4 changes: 4 additions & 0 deletions bindings/python/tests/test_rl_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
"model_id": "m",
"worker_type": "regular",
"connection_mode": "http",
"control_url": "http://e:1",
"tp_size": 1,
"dp_size": 1,
"pp_size": 1,
Expand Down Expand Up @@ -89,6 +90,7 @@ def test_workers_and_worker(stub):
ws = rl.workers()
assert len(ws) == 1 and ws[0].id == "w1" and ws[0].engine == "sglang"
assert ws[0].capabilities["pause_modes"] == ["abort"]
assert ws[0].control_url == "http://e:1"
assert rl.worker("w1").weight_version == "7"
assert _Stub.seen[0]["auth"] == "Bearer k"

Expand Down Expand Up @@ -155,11 +157,13 @@ def test_worker_from_json_defaults_missing_dicts():
del d["labels"]
del d["capabilities"]
del d["role"]
del d["control_url"]
d["future_field"] = 1
w = Worker.from_json(d)
assert w.labels == {}
assert w.capabilities == {}
assert w.role is None
assert w.control_url is None


def test_call_raises_on_smg_error_envelope(stub):
Expand Down
Loading
Loading