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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 5 additions & 5 deletions nemo_rl/distributed/virtual_cluster.py
Original file line number Diff line number Diff line change
Expand Up @@ -113,12 +113,12 @@ class PY_EXECUTABLES:
DEFAULT_SGLANG_ROUTER_PORT_RANGE_HIGH = 8799
DEFAULT_SGLANG_PROMETHEUS_PORT_RANGE_LOW = 8800
DEFAULT_SGLANG_PROMETHEUS_PORT_RANGE_HIGH = 8999
# vLLM Router control-plane ports occupy the reserved gap between Ray's
# Inference Router control-plane ports occupy the reserved gap between Ray's
# management ports and the master / TCPStore range.
DEFAULT_VLLM_ROUTER_PORT_RANGE_LOW = 1320
DEFAULT_VLLM_ROUTER_PORT_RANGE_HIGH = 1360
DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_LOW = 1360
DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_HIGH = 1400
DEFAULT_INFERENCE_ROUTER_PORT_RANGE_LOW = 1320
DEFAULT_INFERENCE_ROUTER_PORT_RANGE_HIGH = 1360
DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_LOW = 1360
DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_HIGH = 1400
# Master address / TCPStore range, tucked below the Ray worker-gRPC band (2000+).
DEFAULT_MASTER_PORT_RANGE_LOW = 1400
DEFAULT_MASTER_PORT_RANGE_HIGH = 1999
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -15,25 +15,76 @@
import subprocess
import sys
import time
from typing import Literal
from urllib.error import URLError
from urllib.request import urlopen

from pydantic import BaseModel


class VllmRouterConfig(BaseModel, extra="forbid"):
from pydantic import BaseModel, model_validator

RouterBackend = Literal["vllm_router", "smg"]
RouterPolicy = Literal[
"random",
"round_robin",
"cache_aware",
"power_of_two",
"consistent_hash",
"rendezvous_hash",
"passthrough",
"least_load",
"manual",
"prefix_hash",
]

_BACKEND_POLICIES = {
"vllm_router": frozenset(
{
"random",
"round_robin",
"cache_aware",
"power_of_two",
"consistent_hash",
}
),
"smg": frozenset(
{
"random",
"round_robin",
"cache_aware",
"power_of_two",
"consistent_hash",
"passthrough",
"least_load",
"manual",
"prefix_hash",
}
),
}


class InferenceRouterConfig(BaseModel, extra="forbid"):
enabled: bool = False
policy: str = "consistent_hash"
backend: RouterBackend = "vllm_router"
policy: RouterPolicy = "consistent_hash"

@model_validator(mode="after")
def validate_backend_policy(self) -> "InferenceRouterConfig":
if self.policy not in _BACKEND_POLICIES[self.backend]:
supported = ", ".join(sorted(_BACKEND_POLICIES[self.backend]))
raise ValueError(
f"Router backend {self.backend!r} does not support policy "
f"{self.policy!r}; supported policies: {supported}"
)
return self

class VllmRouterProcess:

class InferenceRouterProcess:
def __init__(
self,
worker_base_urls: list[str],
host: str,
port: int,
prometheus_port: int,
config: VllmRouterConfig,
config: InferenceRouterConfig,
) -> None:
self.worker_base_urls = [
base_url.rstrip("/").removesuffix("/v1") for base_url in worker_base_urls
Expand All @@ -44,16 +95,37 @@ def __init__(
self.config = config
self._process: subprocess.Popen | None = None

@property
def name(self) -> str:
return "vLLM Router" if self.config.backend == "vllm_router" else "SMG"

@property
def session_affinity_header(self) -> str:
if self.config.backend == "vllm_router":
return "X-Session-ID"
return "X-SMG-Routing-Key"

@property
def command(self) -> list[str]:
if self.config.backend == "vllm_router":
module = "vllm_router.launch_router"
policy = self.config.policy
else:
module = "smg.launch_router"
policy = (
"consistent_hashing"
if self.config.policy == "consistent_hash"
else self.config.policy
)

return [
sys.executable,
"-m",
"vllm_router.launch_router",
module,
"--worker-urls",
*self.worker_base_urls,
"--policy",
self.config.policy,
policy,
"--host",
self.host,
"--port",
Expand All @@ -72,7 +144,7 @@ def readiness_url(self) -> str:

def start(self) -> None:
if self._process is not None:
raise RuntimeError("vLLM Router process has already been started")
raise RuntimeError(f"{self.name} process has already been started")
self._process = subprocess.Popen(self.command)

def wait_until_ready(
Expand All @@ -82,15 +154,15 @@ def wait_until_ready(
) -> None:
process = self._process
if process is None:
raise RuntimeError("vLLM Router process has not been started")
raise RuntimeError(f"{self.name} process has not been started")

deadline = time.monotonic() + timeout
while True:
return_code = process.poll()
if return_code is not None:
self._process = None
raise RuntimeError(
f"vLLM Router process exited with code {return_code} "
f"{self.name} process exited with code {return_code} "
"before becoming ready"
)

Expand All @@ -106,12 +178,12 @@ def wait_until_ready(

if time.monotonic() >= deadline:
raise TimeoutError(
f"vLLM Router did not become ready within {timeout} seconds"
f"{self.name} did not become ready within {timeout} seconds"
)

time.sleep(poll_interval)

def stop(self, timeout: float = 5.0) -> None:
def stop(self, timeout: float = 10.0) -> None:
process = self._process
if process is None:
return
Expand Down
65 changes: 34 additions & 31 deletions nemo_rl/environments/nemo_gym.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,15 +38,18 @@
from nemo_rl.distributed.virtual_cluster import (
DEFAULT_GYM_PORT_RANGE_HIGH,
DEFAULT_GYM_PORT_RANGE_LOW,
DEFAULT_VLLM_ROUTER_PORT_RANGE_HIGH,
DEFAULT_VLLM_ROUTER_PORT_RANGE_LOW,
DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_HIGH,
DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_LOW,
DEFAULT_INFERENCE_ROUTER_PORT_RANGE_HIGH,
DEFAULT_INFERENCE_ROUTER_PORT_RANGE_LOW,
DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_HIGH,
DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_LOW,
_get_free_port_local,
_get_node_ip_local,
)
from nemo_rl.environments.inference_router import (
InferenceRouterConfig,
InferenceRouterProcess,
)
from nemo_rl.environments.interfaces import EnvironmentInterface
from nemo_rl.environments.vllm_router import VllmRouterConfig, VllmRouterProcess
from nemo_rl.models.policy import TokenizerConfig
from nemo_rl.utils.routed_experts_codec import decode_routed_experts
from nemo_rl.utils.timer import Timer
Expand Down Expand Up @@ -107,7 +110,7 @@ class NemoGymConfig(TypedDict):
model_name: str
base_urls: List[str]
initial_global_config_dict: Dict[str, Any]
vllm_router: NotRequired[VllmRouterConfig]
router: NotRequired[InferenceRouterConfig]
# Port range for Gym HTTP servers (head server + subprocess servers).
# Defaults to DEFAULT_GYM_PORT_RANGE_LOW/HIGH (5000-5999) from
# nemo_rl.distributed.virtual_cluster. See the port layout there.
Expand Down Expand Up @@ -361,7 +364,7 @@ class NemoGym(EnvironmentInterface):

def __init__(self, cfg: NemoGymConfig):
self.cfg = cfg
self._vllm_router: VllmRouterProcess | None = None
self._router: InferenceRouterProcess | None = None
# Reconstruct the processor inside the actor (rather than serializing it
# per rollout call) for full-trajectory multimodal postprocessing.
self._processor: Optional[Any] = None
Expand Down Expand Up @@ -394,24 +397,24 @@ def _spinup(self) -> None:
self.head_server_port = _get_free_port_local(_gym_port_low, _gym_port_high)

policy_base_urls = self.cfg["base_urls"]
vllm_router_config = self.cfg.get("vllm_router")
if vllm_router_config is not None and vllm_router_config.enabled:
router_config = self.cfg.get("router")
if router_config is not None and router_config.enabled:
router_port = _get_free_port_local(
DEFAULT_VLLM_ROUTER_PORT_RANGE_LOW,
DEFAULT_VLLM_ROUTER_PORT_RANGE_HIGH,
DEFAULT_INFERENCE_ROUTER_PORT_RANGE_LOW,
DEFAULT_INFERENCE_ROUTER_PORT_RANGE_HIGH,
)
prometheus_port = _get_free_port_local(
DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_LOW,
DEFAULT_VLLM_ROUTER_PROMETHEUS_PORT_RANGE_HIGH,
DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_LOW,
DEFAULT_INFERENCE_ROUTER_PROMETHEUS_PORT_RANGE_HIGH,
)
self._vllm_router = VllmRouterProcess(
self._router = InferenceRouterProcess(
worker_base_urls=self.cfg["base_urls"],
host=self.node_ip,
port=router_port,
prometheus_port=prometheus_port,
config=vllm_router_config,
config=router_config,
)
policy_base_urls = [self._vllm_router.openai_base_url]
policy_base_urls = [self._router.openai_base_url]

from nemo_gym.cli import GlobalConfigDictParserConfig, RunHelper
from nemo_gym.rollout_collection import RolloutCollectionHelper
Expand All @@ -426,7 +429,7 @@ def _spinup(self) -> None:
initial_global_config_dict = DictConfig(
self.cfg.get("initial_global_config_dict") or {}
)
if self._vllm_router is not None:
if self._router is not None:
initial_global_config_dict = cast(
DictConfig,
OmegaConf.merge(
Expand All @@ -435,7 +438,7 @@ def _spinup(self) -> None:
"policy_model": {
"responses_api_models": {
"vllm_model": {
"session_affinity_header": "X-Session-ID"
"session_affinity_header": self._router.session_affinity_header
}
}
}
Expand Down Expand Up @@ -497,9 +500,9 @@ def _spinup(self) -> None:

self.rh = RunHelper()
try:
if self._vllm_router is not None:
self._vllm_router.start()
self._vllm_router.wait_until_ready()
if self._router is not None:
self._router.start()
self._router.wait_until_ready()
self.rh.start(
global_config_dict_parser_config=GlobalConfigDictParserConfig(
dotenv_path=Path(__file__.removesuffix(RELATIVE_PATH)).absolute()
Expand All @@ -509,8 +512,8 @@ def _spinup(self) -> None:
)
)
except Exception:
router = self._vllm_router
self._vllm_router = None
router = self._router
self._router = None
if router is not None:
router.stop()
raise
Expand Down Expand Up @@ -871,8 +874,8 @@ def shutdown(self) -> None:
try:
self.rh.shutdown()
finally:
router = self._vllm_router
self._vllm_router = None
router = self._router
self._router = None
if router is not None:
router.stop()

Expand Down Expand Up @@ -1010,11 +1013,11 @@ def spinup_nemo_gym_actor(
The spun-up NemoGym Ray actor handle (_spinup already awaited).
"""
nemo_gym_dict = dict(env_configs["nemo_gym"])
vllm_router_dict = nemo_gym_dict.pop("vllm_router", None)
vllm_router_config = (
VllmRouterConfig.model_validate(vllm_router_dict)
if vllm_router_dict is not None
else VllmRouterConfig()
router_dict = nemo_gym_dict.pop("router", None)
router_config = (
InferenceRouterConfig.model_validate(router_dict)
if router_dict is not None
else InferenceRouterConfig()
)

# NeMo-RL-side detection knobs are top-level NemoGymConfig fields
Expand All @@ -1035,7 +1038,7 @@ def spinup_nemo_gym_actor(
nemo_gym_cfg = NemoGymConfig(
model_name=model_name,
base_urls=base_urls,
vllm_router=vllm_router_config,
router=router_config,
invalid_tool_call_patterns=invalid_tool_call_patterns,
thinking_tags=thinking_tags,
tokenizer_config=tokenizer_config,
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -213,7 +213,7 @@ nvrx = [
nemo_gym = ["nemo_gym"]

[dependency-groups]
nemo_gym_router = ["vllm-router==0.1.15"]
nemo_gym_router = ["smg==1.9.0", "vllm-router==0.1.15"]

# This is a default group so that we install these even with bare `uv sync`
build = [
Expand Down
2 changes: 1 addition & 1 deletion pyrefly.toml
Original file line number Diff line number Diff line change
Expand Up @@ -139,7 +139,7 @@ project-includes = [
"nemo_rl/environments/metrics.py",
"nemo_rl/environments/rewards.py",
"nemo_rl/environments/utils.py",
"nemo_rl/environments/vllm_router.py",
"nemo_rl/environments/inference_router.py",
"nemo_rl/environments/vlm_environment.py",
"nemo_rl/evals/__init__.py",
"nemo_rl/evals/answer_parsing.py",
Expand Down
Loading
Loading