Generator
Generator API — GeneratorInterface, InferenceEngineInterface.
Core APIs
class GeneratorInterface
Bases: ABC
Functions:
| Name | Description |
|---|---|
generate | Generate trajectories for the input batch. |
Source code in skyrl/train/generators/base.py:71-83
class GeneratorInterface(ABC):
@abstractmethod
async def generate(self, input_batch: GeneratorInput) -> GeneratorOutput:
"""Generate trajectories for the input batch.
Returns outputs in the same order as the input batch.
Args:
input_batch (GeneratorInput): Input batch
Returns:
GeneratorOutput: Generated trajectories
"""
raise NotImplementedErrormethod async generate
generate(input_batch: GeneratorInput) -> GeneratorOutputGenerate trajectories for the input batch.
Returns outputs in the same order as the input batch.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
input_batch | GeneratorInput | Input batch | required |
Returns: GeneratorOutput: Generated trajectories
Source code in skyrl/train/generators/base.py:72-83
@abstractmethod
async def generate(self, input_batch: GeneratorInput) -> GeneratorOutput:
"""Generate trajectories for the input batch.
Returns outputs in the same order as the input batch.
Args:
input_batch (GeneratorInput): Input batch
Returns:
GeneratorOutput: Generated trajectories
"""
raise NotImplementedErrorclass InferenceEngineInterface
Bases: ABC
Functions:
| Name | Description |
|---|---|
generate | |
get_endpoint_url | Return the base URL of the data-plane (OpenAI-compatible) endpoint. |
chat_completion | Handles OpenAI-compatible HTTP endpoint. |
render_chat_completion | Apply the chat template and tokenize without generating. |
completion | Handles OpenAI-compatible HTTP endpoint. |
wake_up | |
sleep | |
init_weight_update_communicator | Initialize weight update communicator from init info. |
update_named_weights | |
teardown | |
reset_prefix_cache | |
pause_generation | Pause generation, freezing in-flight requests so they can be resumed later. |
resume_generation | Resume generation after a pause, continuing any frozen in-flight requests. |
finish_session | Notify the inference server that a session (trajectory) is complete. |
get_world_size | Return (total_world_size, world_size_per_server) across all inference workers. |
Attributes:
| Name | Type | Description |
|---|---|---|
model_name | str | The base model identifier the inference server was started with. |
Source code in skyrl/backends/skyrl_train/inference_servers/base.py:53-169
class InferenceEngineInterface(ABC):
@property
@abstractmethod
def model_name(self) -> str:
"""The base model identifier the inference server was started with.
Generators pass this as the ``model`` field on requests (e.g. when
rendering a chat completion) when they don't otherwise specify one.
"""
raise NotImplementedError
@abstractmethod
async def generate(
self,
input_batch: InferenceEngineInput,
model: Optional[str] = None,
) -> InferenceEngineOutput:
raise NotImplementedError
@abstractmethod
def get_endpoint_url(self) -> str:
"""Return the base URL of the data-plane (OpenAI-compatible) endpoint.
Generators point external clients at the inference server/router through
this URL (e.g. LiteLLM via ``OPENAI_BASE_URL``) without depending on the
concrete client type or how the URL is derived.
"""
raise NotImplementedError
@abstractmethod
async def chat_completion(self, request_payload: Dict[str, Any]) -> Dict[str, Any]:
"""Handles OpenAI-compatible HTTP endpoint.
Accepts a JSON payload: {"json": <request-body>, "headers": <headers-dict>}.
The request body will be used to construct a ChatCompletionRequest.
Returns a plain dict, either a ChatCompletionResponse or an ErrorResponse.
The specific fields of the response/request depend on the engine's backend (e.g. for vllm
these are defined in vllm.entrypoints.openai.protocol).
"""
raise NotImplementedError
@abstractmethod
async def render_chat_completion(self, request_payload: Dict[str, Any]) -> Dict[str, Any]:
"""Apply the chat template and tokenize without generating.
Accepts the same ``{"json": <request-body>}`` payload as
``chat_completion`` and returns the rendered prompt / token IDs. Used by
generators that need token-in/token-out rendering (e.g. multi-modal).
"""
raise NotImplementedError
@abstractmethod
async def completion(self, request_payload: Dict[str, Any]) -> Dict[str, Any]:
"""Handles OpenAI-compatible HTTP endpoint.
Accepts a JSON payload: {"json": <request-body>, "headers": <headers-dict>}.
The request body will be used to construct a CompletionRequest.
Returns a plain dict, either a CompletionResponse or an ErrorResponse.
The specific fields of the response/request depend on the engine's backend (e.g. for vllm
these are defined in vllm.entrypoints.openai.protocol).
"""
raise NotImplementedError
@abstractmethod
async def wake_up(self, *args: Any, **kwargs: Any):
raise NotImplementedError
@abstractmethod
async def sleep(self, *args: Any, **kwargs: Any):
raise NotImplementedError
@abstractmethod
async def init_weight_update_communicator(self, init_info: "WeightSyncInitInfo"):
"""Initialize weight update communicator from init info.
Args:
init_info: WeightSyncInitInfo from the sender containing all info needed
to create the appropriate receiver.
"""
raise NotImplementedError()
@abstractmethod
async def update_named_weights(self, request: "WeightUpdateRequest"):
raise NotImplementedError()
@abstractmethod
async def teardown(self):
raise NotImplementedError
@abstractmethod
async def reset_prefix_cache(self, reset_running_requests: bool = False):
raise NotImplementedError
@abstractmethod
async def pause_generation(self) -> None:
"""Pause generation, freezing in-flight requests so they can be resumed later."""
raise NotImplementedError
@abstractmethod
async def resume_generation(self) -> None:
"""Resume generation after a pause, continuing any frozen in-flight requests."""
raise NotImplementedError
@abstractmethod
async def finish_session(self, session_id: str) -> None:
"""Notify the inference server that a session (trajectory) is complete.
Best-effort: lets session-aware routing release the replica capacity the
session held. Generators call this in trajectory-cleanup paths.
"""
raise NotImplementedError
@abstractmethod
async def get_world_size(self) -> Tuple[int, int]:
"""Return ``(total_world_size, world_size_per_server)`` across all inference workers."""
raise NotImplementedErrorattr abstractmethod property model_name
model_name: strThe base model identifier the inference server was started with.
Generators pass this as the model field on requests (e.g. when
rendering a chat completion) when they don't otherwise specify one.
method async generate
generate(input_batch: InferenceEngineInput, model: Optional[str] = None) -> InferenceEngineOutputSource code in skyrl/backends/skyrl_train/inference_servers/base.py:65-71
@abstractmethod
async def generate(
self,
input_batch: InferenceEngineInput,
model: Optional[str] = None,
) -> InferenceEngineOutput:
raise NotImplementedErrormethod abstractmethod get_endpoint_url
get_endpoint_url() -> strReturn the base URL of the data-plane (OpenAI-compatible) endpoint.
Generators point external clients at the inference server/router through
this URL (e.g. LiteLLM via OPENAI_BASE_URL) without depending on the
concrete client type or how the URL is derived.
Source code in skyrl/backends/skyrl_train/inference_servers/base.py:73-81
@abstractmethod
def get_endpoint_url(self) -> str:
"""Return the base URL of the data-plane (OpenAI-compatible) endpoint.
Generators point external clients at the inference server/router through
this URL (e.g. LiteLLM via ``OPENAI_BASE_URL``) without depending on the
concrete client type or how the URL is derived.
"""
raise NotImplementedErrormethod abstractmethod async chat_completion
chat_completion(request_payload: Dict[str, Any]) -> Dict[str, Any]Handles OpenAI-compatible HTTP endpoint.
Accepts a JSON payload: {"json": <request-body>, "headers": <headers-dict>}. The request body will be used to construct a ChatCompletionRequest. Returns a plain dict, either a ChatCompletionResponse or an ErrorResponse. The specific fields of the response/request depend on the engine's backend (e.g. for vllm these are defined in vllm.entrypoints.openai.protocol).
Source code in skyrl/backends/skyrl_train/inference_servers/base.py:83-93
@abstractmethod
async def chat_completion(self, request_payload: Dict[str, Any]) -> Dict[str, Any]:
"""Handles OpenAI-compatible HTTP endpoint.
Accepts a JSON payload: {"json": <request-body>, "headers": <headers-dict>}.
The request body will be used to construct a ChatCompletionRequest.
Returns a plain dict, either a ChatCompletionResponse or an ErrorResponse.
The specific fields of the response/request depend on the engine's backend (e.g. for vllm
these are defined in vllm.entrypoints.openai.protocol).
"""
raise NotImplementedErrormethod abstractmethod async render_chat_completion
render_chat_completion(request_payload: Dict[str, Any]) -> Dict[str, Any]Apply the chat template and tokenize without generating.
Accepts the same {"json": <request-body>} payload as
chat_completion and returns the rendered prompt / token IDs. Used by
generators that need token-in/token-out rendering (e.g. multi-modal).
Source code in skyrl/backends/skyrl_train/inference_servers/base.py:95-103
@abstractmethod
async def render_chat_completion(self, request_payload: Dict[str, Any]) -> Dict[str, Any]:
"""Apply the chat template and tokenize without generating.
Accepts the same ``{"json": <request-body>}`` payload as
``chat_completion`` and returns the rendered prompt / token IDs. Used by
generators that need token-in/token-out rendering (e.g. multi-modal).
"""
raise NotImplementedErrormethod abstractmethod async completion
completion(request_payload: Dict[str, Any]) -> Dict[str, Any]Handles OpenAI-compatible HTTP endpoint.
Accepts a JSON payload: {"json": <request-body>, "headers": <headers-dict>}. The request body will be used to construct a CompletionRequest. Returns a plain dict, either a CompletionResponse or an ErrorResponse. The specific fields of the response/request depend on the engine's backend (e.g. for vllm these are defined in vllm.entrypoints.openai.protocol).
Source code in skyrl/backends/skyrl_train/inference_servers/base.py:105-115
@abstractmethod
async def completion(self, request_payload: Dict[str, Any]) -> Dict[str, Any]:
"""Handles OpenAI-compatible HTTP endpoint.
Accepts a JSON payload: {"json": <request-body>, "headers": <headers-dict>}.
The request body will be used to construct a CompletionRequest.
Returns a plain dict, either a CompletionResponse or an ErrorResponse.
The specific fields of the response/request depend on the engine's backend (e.g. for vllm
these are defined in vllm.entrypoints.openai.protocol).
"""
raise NotImplementedErrormethod abstractmethod async wake_up
wake_up(*args: Any, **kwargs: Any)Source code in skyrl/backends/skyrl_train/inference_servers/base.py:117-119
@abstractmethod
async def wake_up(self, *args: Any, **kwargs: Any):
raise NotImplementedErrormethod abstractmethod async sleep
sleep(*args: Any, **kwargs: Any)Source code in skyrl/backends/skyrl_train/inference_servers/base.py:121-123
@abstractmethod
async def sleep(self, *args: Any, **kwargs: Any):
raise NotImplementedErrormethod abstractmethod async init_weight_update_communicator
init_weight_update_communicator(init_info: WeightSyncInitInfo)Initialize weight update communicator from init info.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
init_info | WeightSyncInitInfo | WeightSyncInitInfo from the sender containing all info needed to create the appropriate receiver. | required |
Source code in skyrl/backends/skyrl_train/inference_servers/base.py:125-133
@abstractmethod
async def init_weight_update_communicator(self, init_info: "WeightSyncInitInfo"):
"""Initialize weight update communicator from init info.
Args:
init_info: WeightSyncInitInfo from the sender containing all info needed
to create the appropriate receiver.
"""
raise NotImplementedError()method abstractmethod async update_named_weights
update_named_weights(request: WeightUpdateRequest)Source code in skyrl/backends/skyrl_train/inference_servers/base.py:135-137
@abstractmethod
async def update_named_weights(self, request: "WeightUpdateRequest"):
raise NotImplementedError()method abstractmethod async teardown
teardown()Source code in skyrl/backends/skyrl_train/inference_servers/base.py:139-141
@abstractmethod
async def teardown(self):
raise NotImplementedErrormethod abstractmethod async reset_prefix_cache
reset_prefix_cache(reset_running_requests: bool = False)Source code in skyrl/backends/skyrl_train/inference_servers/base.py:143-145
@abstractmethod
async def reset_prefix_cache(self, reset_running_requests: bool = False):
raise NotImplementedErrormethod abstractmethod async pause_generation
pause_generation() -> NonePause generation, freezing in-flight requests so they can be resumed later.
Source code in skyrl/backends/skyrl_train/inference_servers/base.py:147-150
@abstractmethod
async def pause_generation(self) -> None:
"""Pause generation, freezing in-flight requests so they can be resumed later."""
raise NotImplementedErrormethod abstractmethod async resume_generation
resume_generation() -> NoneResume generation after a pause, continuing any frozen in-flight requests.
Source code in skyrl/backends/skyrl_train/inference_servers/base.py:152-155
@abstractmethod
async def resume_generation(self) -> None:
"""Resume generation after a pause, continuing any frozen in-flight requests."""
raise NotImplementedErrormethod abstractmethod async finish_session
finish_session(session_id: str) -> NoneNotify the inference server that a session (trajectory) is complete.
Best-effort: lets session-aware routing release the replica capacity the session held. Generators call this in trajectory-cleanup paths.
Source code in skyrl/backends/skyrl_train/inference_servers/base.py:157-164
@abstractmethod
async def finish_session(self, session_id: str) -> None:
"""Notify the inference server that a session (trajectory) is complete.
Best-effort: lets session-aware routing release the replica capacity the
session held. Generators call this in trajectory-cleanup paths.
"""
raise NotImplementedErrormethod abstractmethod async get_world_size
get_world_size() -> Tuple[int, int]Return (total_world_size, world_size_per_server) across all inference workers.
Source code in skyrl/backends/skyrl_train/inference_servers/base.py:166-169
@abstractmethod
async def get_world_size(self) -> Tuple[int, int]:
"""Return ``(total_world_size, world_size_per_server)`` across all inference workers."""
raise NotImplementedErrorclass RemoteInferenceClient
RemoteInferenceClient(proxy_url: str, server_urls: List[str], data_parallel_size: int, model_name: str = 'default', enable_return_routed_experts: bool = False, uses_lora_weight_sync: bool = False, tokenizer: Optional[Any] = None, _session: Optional[aiohttp.ClientSession] = None, _world_size: Optional[Tuple[int, int]] = None, _gen_sem: Optional[asyncio.Semaphore] = None, _detok_sem: Optional[asyncio.Semaphore] = None, _sem_loop: Optional[asyncio.AbstractEventLoop] = None, _weight_version: int = 0) -> NoneBases: InferenceEngineInterface
Serializable HTTP client for inference. The concrete InferenceEngineInterface.
This class maintains two URL types:
- proxy_url: Single URL for data plane operations (routed requests)
- server_urls: List of backend URLs for control plane operations (fan-out)
The router (proxy_url) is expected to be a data-plane-only router (like VLLMRouter or an external router). Control plane operations are always fanned out to all backends directly by this client.
Usage:
client = RemoteInferenceClient( proxy_url="http://router:8080", # Data plane (router) server_urls=["http://backend1:8000", "http://backend2:8000"], # Control plane data_parallel_size=1, # data parallel size for deployments )
Functions:
| Name | Description |
|---|---|
increment_weight_version | Advance the weight version. Called once per completed weight sync to the engines. |
get_endpoint_url | Data-plane endpoint base URL (the router/proxy that load-balances requests). |
generate | Generate completions via /v1/completions. |
sample | Sample completions via /inference/v1/generate (Tinker API). |
chat_completion | Chat completion via /v1/chat/completions. |
render_chat_completion | Render a chat completion (apply chat template + tokenize) via /v1/chat/completions/render. |
completion | Completion via /v1/completions. |
tokenize | Tokenize texts. |
detokenize | Detokenize token IDs. |
finish_session | Notify the router that a session (trajectory) is complete. |
pause | Pause generation on all backends. |
resume | Resume generation on all backends. |
pause_generation | Pause using keep mode. |
resume_generation | Resume after pause. |
sleep | Put all backends to sleep (offload weights to CPU). |
wake_up | Wake up all backends (load weights back to GPU). |
sleep_for_weight_sync | Free GPU memory for weight sync while keeping in-flight requests frozen. |
wake_for_weight_sync | Restore allocator pools by tag (see :meth:sleep_for_weight_sync). |
reset_prefix_cache | Reset KV cache on all backends. |
init_weight_update_communicator | Initialize weight sync via vLLM native /init_weight_transfer_engine. |
update_named_weights | Update model weights via vLLM native /update_weights. Used for full parameter fine-tuning. |
start_weight_update | Start a new chunked weight update via /collective_rpc. |
update_weights_ipc | Send a single weight chunk via /collective_rpc. |
update_weights_nccl | Send batched weight update via /collective_rpc to the broadcast receiver. |
finish_weight_update | Finish the current chunked weight update via /collective_rpc. |
load_lora_adapter | Load (or reload) a LoRA adapter on all backend servers via the SkyRL |
unload_lora_adapter | Unload a previously-loaded LoRA adapter on all backend servers via /v1/unload_lora_adapter. |
get_world_size | Get total and per-server world size across all inference workers. |
teardown | Close HTTP session. |
aclose |
Attributes:
| Name | Type | Description |
|---|---|---|
proxy_url | str | Data plane URL (single endpoint - router or direct server). |
server_urls | List[str] | Control plane URLs (list of backend servers for fan-out). |
data_parallel_size | int | Data parallel size. Used to compute total inference world size correctly: |
model_name | str | The model identifier accepted by the inference server for the base model. |
enable_return_routed_experts | bool | Whether to return routed expert indices (R3 / rollout router replay). |
uses_lora_weight_sync | bool | True when the trainer syncs LoRA adapters (rather than full/merged weights). When True, |
tokenizer | Optional[Any] | Optional HF tokenizer for local tokenize/detokenize (avoids HTTP round-trips). |
weight_version | int | Number of weight syncs to the engines so far (0 before the first sync); the policy version. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:165-1436
@dataclass
class RemoteInferenceClient(InferenceEngineInterface):
"""
Serializable HTTP client for inference. The concrete InferenceEngineInterface.
This class maintains two URL types:
- proxy_url: Single URL for data plane operations (routed requests)
- server_urls: List of backend URLs for control plane operations (fan-out)
The router (proxy_url) is expected to be a data-plane-only router (like
VLLMRouter or an external router). Control plane operations
are always fanned out to all backends directly by this client.
Usage:
client = RemoteInferenceClient(
proxy_url="http://router:8080", # Data plane (router)
server_urls=["http://backend1:8000", "http://backend2:8000"], # Control plane
data_parallel_size=1, # data parallel size for deployments
)
"""
proxy_url: str
"""Data plane URL (single endpoint - router or direct server)."""
server_urls: List[str]
"""Control plane URLs (list of backend servers for fan-out)."""
data_parallel_size: int
"""Data parallel size. Used to compute total inference world size correctly:
server_urls contains num_engines * data_parallel_size entries, but vLLM already
reports the full DP world size per server, so we divide by num_deployments."""
model_name: str = "default"
"""The model identifier accepted by the inference server for the base model.
This is usually the model path, but may be ``served_model_name`` when vLLM
is started with an alias. It is never a LoRA adapter name. LoRA adapters are
addressed by the names callers register them under via
``load_lora_adapter(name, path)``, and per-call routing is done by
passing that name as ``model`` on the data-plane methods.
Used internally only by ``tokenize``/``detokenize``, which are LoRA-
agnostic but still require a ``model`` field per the OpenAI schema.
"""
enable_return_routed_experts: bool = False
"""Whether to return routed expert indices (R3 / rollout router replay)."""
uses_lora_weight_sync: bool = False
"""True when the trainer syncs LoRA adapters (rather than full/merged weights). When True,
`sleep()` is forced to level=1: level=2 discards the base model from VRAM with no CPU backup,
and LoRA-only broadcasts cannot repopulate it. Must be kept in sync with the same gate vLLM
uses for `enable_lora` (see `_uses_lora_weight_sync` in inference_servers/utils.py)."""
tokenizer: Optional[Any] = None
"""Optional HF tokenizer for local tokenize/detokenize (avoids HTTP round-trips)."""
# Private fields excluded from repr for cleaner output
_session: Optional[aiohttp.ClientSession] = field(default=None, repr=False)
_world_size: Optional[Tuple[int, int]] = field(default=None, repr=False)
_gen_sem: Optional[asyncio.Semaphore] = field(default=None, repr=False)
_detok_sem: Optional[asyncio.Semaphore] = field(default=None, repr=False)
_sem_loop: Optional[asyncio.AbstractEventLoop] = field(default=None, repr=False)
# Monotonic counter of weight syncs (see `increment_weight_version`); source of the prefix-cache salt.
_weight_version: int = field(default=0, repr=False)
@property
def weight_version(self) -> int:
"""Number of weight syncs to the engines so far (0 before the first sync); the policy version."""
return self._weight_version
def increment_weight_version(self) -> None:
"""Advance the weight version. Called once per completed weight sync to the engines."""
self._weight_version += 1
def __post_init__(self):
if self.data_parallel_size <= 0:
raise ValueError(f"Expected `data_parallel_size` >0, got {self.data_parallel_size}")
if len(self.server_urls) % self.data_parallel_size != 0:
raise ValueError(
f"Expected number of servers to be divisible by data parallel size, got {self.server_urls} and {self.data_parallel_size}"
)
def get_endpoint_url(self) -> str:
"""Data-plane endpoint base URL (the router/proxy that load-balances requests)."""
return self.proxy_url
# ---------------------------
# Session Management
# ---------------------------
def _get_semaphores(self) -> Tuple[Optional[asyncio.Semaphore], Optional[asyncio.Semaphore]]:
"""Get or create the shared generate/detokenize semaphores for this client.
Semaphores are event-loop-bound (Python 3.10+). If the running loop has
changed since they were created, recreate them.
All concurrent generate() calls on the same client instance share these
semaphores, capping total in-flight requests at
SKYRL_GENERATE_CONCURRENCY_PER_ENGINE × num_engines.
"""
current_loop = asyncio.get_running_loop()
if self._sem_loop is not current_loop:
if SKYRL_GENERATE_CONCURRENCY_PER_ENGINE > 0:
concurrency = SKYRL_GENERATE_CONCURRENCY_PER_ENGINE * len(self.server_urls)
logger.info(f"Capping concurrency for generation to a maximum of {concurrency} requests")
self._gen_sem = asyncio.Semaphore(concurrency)
self._detok_sem = asyncio.Semaphore(concurrency)
else:
self._gen_sem = None
self._detok_sem = None
self._sem_loop = current_loop
return self._gen_sem, self._detok_sem
async def _get_session(self) -> aiohttp.ClientSession:
"""Get or create the aiohttp session."""
# Re-use the existing session object if it is not closed.
# Note that we also create a new session object if the event loop has changed, since
# aiohttp.ClientSession is tied to the event loop.
current_loop = asyncio.get_running_loop()
if self._session is not None and not self._session.closed and self._session.loop != current_loop:
# Event loop changed - the old session is unusable (bound to a dead loop).
self._session = None
if self._session is None or self._session.closed:
# keepalive_timeout must be shorter than the server's timeout_keep_alive
# (uvicorn default: 5s). Otherwise aiohttp reuses connections the server
# has already closed, causing ECONNRESET under high concurrency.
connector = aiohttp.TCPConnector(
limit=SKYRL_HTTP_CONNECTION_LIMIT,
keepalive_timeout=2,
)
self._session = aiohttp.ClientSession(connector=connector, timeout=aiohttp.ClientTimeout(total=None))
return self._session
async def _post(self, url: str, json: Dict[str, Any], headers: Optional[Dict[str, str]] = None) -> Any:
"""POST with retry + backoff on transient connection errors.
Between generate bursts the pool's keep-alive connections go stale
(server closes them after ``timeout_keep_alive``). An immediate
retry would grab another stale connection from the same pool, so we
sleep briefly to let the connector detect and purge dead sockets
before the next attempt.
"""
session = await self._get_session()
last_exc: Optional[Exception] = None
for attempt in range(_DATA_PLANE_RETRIES):
try:
async with session.post(url, json=json, headers=headers) as resp:
try:
body = await resp.json(content_type=None)
except Exception as e:
if 400 <= resp.status < 500:
# Non-JSON client error (e.g. plain text 422 from vllm-router).
# Raise immediately — client errors won't succeed on retry.
text = await resp.text()
raise aiohttp.ClientResponseError(
resp.request_info,
resp.history,
status=resp.status,
message=text or resp.reason,
headers=resp.headers,
)
last_exc = e
logger.debug(f"retry {attempt + 1}/{_DATA_PLANE_RETRIES} for {url=}: {e}")
await asyncio.sleep(1)
continue
raise_for_status(resp, body)
return body
except (aiohttp.ServerDisconnectedError, aiohttp.ClientOSError) as e:
last_exc = e
logger.debug(f"POST retry {attempt + 1}/{_DATA_PLANE_RETRIES} for {url=}: {e}")
await asyncio.sleep(1)
continue
raise last_exc # type: ignore[misc]
# ---------------------------
# Data Plane
# ---------------------------
def _resolve_model(self, model: Optional[str], method_name: str) -> str:
"""Pick the target model name for a data-plane call.
- If ``model`` is non-empty, use it as-is.
- Otherwise, when LoRA is in use (``uses_lora_weight_sync=True``) raise
``ValueError`` — the caller must name the adapter explicitly because
falling back to the base model would silently bypass LoRA.
- Otherwise return ``self.model_name`` (the base model the server was
started with).
"""
if model:
return model
if self.uses_lora_weight_sync:
raise ValueError(
f"RemoteInferenceClient.{method_name}: `model` is required when LoRA "
f"is enabled (uses_lora_weight_sync=True). Pass the LoRA adapter name "
f"explicitly so the request doesn't silently target the base model."
)
return self.model_name
async def generate(
self,
input_batch: InferenceEngineInput,
model: Optional[str] = None,
) -> InferenceEngineOutput:
"""
Generate completions via /v1/completions.
This is the interface for token-in-token-out workflows. Input will have
token ids, and the output is token ids as well.
Each prompt is sent as a separate request to allow the router to route
based on session_id. All requests are made in parallel.
With keep-mode pause, in-flight requests are frozen and resume
transparently after /resume -- no client-side retry needed.
Args:
input_batch: Contains prompt_token_ids, sampling_params, and optional session_ids.
model: Optional model identifier — the base model name or a loaded
LoRA adapter name. When omitted, defaults to ``self.model_name``
if LoRA is not in use; raises ``ValueError`` if it is.
Returns:
InferenceEngineOutput with responses, response_ids, and stop_reasons.
"""
model = self._resolve_model(model, "generate")
prompt_token_ids = input_batch.get("prompt_token_ids")
if prompt_token_ids is None:
raise ValueError("RemoteInferenceClient only accepts `prompt_token_ids`, not `prompts`.")
sampling_params = input_batch.get("sampling_params") or {}
if sampling_params.get("n", 1) > 1:
raise ValueError("n > 1 is not supported. Use `config.generator.n_samples_per_prompt` instead.")
session_ids = input_batch.get("session_ids")
mm_features = input_batch.get("mm_features")
cache_salt = input_batch.get("cache_salt")
get_logprobs = sampling_params.get("logprobs") is not None
# Two semaphores decouple the generate and detokenize stages:
# gen_sem: limits concurrent in-flight generate requests so we don't
# overwhelm the router/vLLM scheduler. Released as soon as
# generation finishes, so the GPU slot is freed immediately.
# detok_sem: limits concurrent detokenize calls independently. Uses the
# same concurrency limit so detokenize never starves generate.
# Semaphores are shared across all concurrent generate() calls on this client
# instance, so total in-flight requests are capped at
# SKYRL_GENERATE_CONCURRENCY_PER_ENGINE × num_engines regardless of how many
# callers invoke generate() simultaneously.
# TODO (sumanthrh) (RemoteInferenceClient data-plane-deprecation): We should move this outside of the client to a runner abstraction that will also parallelize client requests across processes.
gen_sem, detok_sem = self._get_semaphores()
batch_size = len(prompt_token_ids)
async def _throttled_generate(idx: int) -> Dict[str, Any]:
if gen_sem is None:
return await self._generate_single(
prompt_token_ids=prompt_token_ids[idx],
sampling_params=sampling_params,
session_id=session_ids[idx] if session_ids and idx < len(session_ids) else None,
mm_features=mm_features[idx] if mm_features and idx < len(mm_features) else None,
model=model,
cache_salt=cache_salt,
)
async with gen_sem:
return await self._generate_single(
prompt_token_ids=prompt_token_ids[idx],
sampling_params=sampling_params,
session_id=session_ids[idx] if session_ids and idx < len(session_ids) else None,
mm_features=mm_features[idx] if mm_features and idx < len(mm_features) else None,
model=model,
cache_salt=cache_salt,
)
async def _throttled_detokenize(token_ids: List[int]) -> str:
if detok_sem is None:
return (await self.detokenize([token_ids]))[0]
async with detok_sem:
return (await self.detokenize([token_ids]))[0]
raw_results = await asyncio.gather(*[_throttled_generate(idx) for idx in range(batch_size)])
responses = await asyncio.gather(*[_throttled_detokenize(r["response_ids"]) for r in raw_results])
rollout_expert_indices = [r.get("routed_experts") for r in raw_results]
has_routed_experts = any(x is not None for x in rollout_expert_indices)
return InferenceEngineOutput(
responses=responses,
stop_reasons=[r["stop_reason"] for r in raw_results],
response_ids=[r["response_ids"] for r in raw_results],
response_logprobs=[r["response_logprobs"] for r in raw_results] if get_logprobs else None,
rollout_expert_indices=rollout_expert_indices if has_routed_experts else None,
)
async def _generate_single(
self,
prompt_token_ids: List[int],
sampling_params: Dict[str, Any],
session_id: Optional[Any],
model: str,
mm_features: Optional[MultiModalFeatures] = None,
cache_salt: Optional[str] = None,
) -> Dict[str, Any]:
"""
Generate completion for a single prompt.
With keep-mode pause, in-flight requests are frozen by the vLLM
scheduler and resume where they left off after /resume. No retry
logic is needed.
Returns:
Dict with keys: stop_reason, response_ids, response_logprobs
"""
url = (
f"{self.proxy_url}/skyrl/v1/generate"
if self.enable_return_routed_experts
else f"{self.proxy_url}/inference/v1/generate"
)
payload: dict[str, Any] = {
"sampling_params": sampling_params,
"model": model,
"token_ids": prompt_token_ids,
}
if mm_features:
payload["features"] = mm_features
# `cache_salt` is a top-level request field (forwarded to vLLM's TokensPrompt), not a sampling
# param.
if cache_salt is not None:
payload["cache_salt"] = cache_salt
headers = {"Content-Type": "application/json"}
if session_id:
headers["X-Session-ID"] = str(session_id)
response = await self._post(url, json=payload, headers=headers)
choice = response["choices"][0]
token_ids = choice["token_ids"]
stop_reason = choice["finish_reason"]
response_logprobs: Optional[List[float]] = None
logprobs = choice.get("logprobs")
if logprobs is not None:
logprobs_content = logprobs.get("content", [])
if logprobs_content:
response_logprobs = [logprob_info["logprob"] for logprob_info in logprobs_content]
routed_experts = choice.get("routed_experts")
return {
"stop_reason": stop_reason,
"response_ids": token_ids,
"response_logprobs": response_logprobs,
"routed_experts": routed_experts,
}
async def _render_for_sample(
self,
prompt: Dict[str, Any],
session_id: Optional[str],
model: str,
) -> Tuple[List[int], Optional[MultiModalFeatures]]:
"""Build token_ids and optional multi-modal features from a Tinker prompt.
For text-only prompts this simply flattens chunk tokens (no HTTP call).
When image chunks are present, calls /v1/chat/completions/render to
process images, then splices the resulting placeholder tokens into the
pre-tokenized text stream and adjusts placeholder offsets.
Returns:
(token_ids, features) where features is None for text-only prompts.
"""
chunks = prompt.get("chunks", [])
# No images → flatten text tokens directly.
image_chunks = [c for c in chunks if c.get("type") in ("image", "image_asset_pointer")]
if not image_chunks:
token_ids = [tok for c in chunks for tok in c.get("tokens", [])]
return token_ids, None
# Build OpenAI chat template with only image_urls
content_parts: List[Dict[str, Any]] = []
for c in image_chunks:
if c["type"] == "image":
# model_dump() on Base64Bytes produces bytes with the b64 string.
raw = c["data"]
b64_str = raw.decode("ascii") if isinstance(raw, bytes) else raw
url = f"data:image/{c.get('format', 'jpeg')};base64,{b64_str}"
else: # image_asset_pointer
url = c["location"]
content_parts.append({"type": "image_url", "image_url": {"url": url}})
render_payload: Dict[str, Any] = {
"json": {
"model": model,
"messages": [{"role": "user", "content": content_parts}],
}
}
if session_id:
render_payload["json"]["session_id"] = session_id
render_resp = await self.render_chat_completion(render_payload)
# Extract per-image placeholder token slices from the render output.
features = render_resp.get("features") or {}
render_token_ids = render_resp.get("token_ids", [])
render_placeholders = features.get("mm_placeholders", {}).get("image", [])
placeholder_token_slices: List[List[int]] = []
for ph in render_placeholders:
offset, length = ph["offset"], ph["length"]
placeholder_token_slices.append(render_token_ids[offset : offset + length])
if len(placeholder_token_slices) != len(image_chunks):
raise ValueError(
f"Expected {len(image_chunks)} placeholder token slices, got {len(placeholder_token_slices)}"
)
# Splice: walk chunks in order, substituting image placeholder tokens.
final_token_ids: List[int] = []
new_placeholders: List[MMPlaceholderRangeInfo] = []
img_idx = 0
for c in chunks:
ctype = c.get("type", "encoded_text")
if ctype == "encoded_text":
final_token_ids.extend(c.get("tokens", []))
elif ctype in ("image", "image_asset_pointer"):
ph_tokens = placeholder_token_slices[img_idx]
new_placeholders.append({"offset": len(final_token_ids), "length": len(ph_tokens)})
final_token_ids.extend(ph_tokens)
img_idx += 1
# No need to decode, vllm handles decoding
adjusted_features: MultiModalFeatures = {
"mm_hashes": features.get("mm_hashes", {}),
"mm_placeholders": {"image": new_placeholders},
"kwargs_data": features.get("kwargs_data"),
}
return final_token_ids, adjusted_features
async def sample(
self,
request_payload: SampleRequestPayload,
) -> SampleResponse:
"""
Sample completions via /inference/v1/generate (Tinker API).
Maps Tinker-style sample requests to the vLLM generate endpoint.
Uses self._post() for automatic retry + backoff on transient errors.
Args:
request_payload: SampleRequestPayload with {"json": <request-body>}.
Expected keys in json: prompt, num_samples, sampling_params,
session_id, include_prompt_logprobs (bool), topk_prompt_logprobs (int).
``model`` is optional and resolved via ``_resolve_model``.
Returns:
SampleResponse with type="sample", sequences list, prompt_logprobs, and topk_prompt_logprobs.
"""
session_id, body = _extract_session_id_and_body(request_payload)
model = self._resolve_model(body.get("model"), "sample")
body["model"] = model
prompt = body.get("prompt", {})
num_samples = body.get("num_samples", 1)
tinker_params = body.get("sampling_params", {})
# Note: Tinker SampleRequest uses "prompt_logprobs" (bool), while
# SamplingClient.sample() uses "include_prompt_logprobs".
include_prompt_logprobs = body.get("include_prompt_logprobs", body.get("prompt_logprobs", False))
topk_prompt_logprobs_k = body.get("topk_prompt_logprobs", 0)
# vLLM prompt logprob mapping
prompt_logprobs_sp = None
if include_prompt_logprobs:
prompt_logprobs_sp = topk_prompt_logprobs_k if topk_prompt_logprobs_k > 0 else 0
# Render prompt: flatten text tokens and, if images are present,
# call the render endpoint to get placeholder tokens + features.
token_ids, mm_features = await self._render_for_sample(prompt, session_id, model=model)
# Map Tinker SamplingParams → vLLM format
sampling_params: Dict[str, Any] = {
"n": num_samples,
"logprobs": 0,
"output_kind": 2,
"prompt_logprobs": prompt_logprobs_sp,
}
for tinker_key, vllm_key in _TINKER_SAMPLE_TO_VLLM_PARAM_MAP.items():
val = tinker_params.get(tinker_key)
if val is not None:
sampling_params[vllm_key] = val
payload: Dict[str, Any] = {
"sampling_params": sampling_params,
"model": model,
"token_ids": token_ids,
}
if mm_features is not None:
payload["features"] = mm_features
headers = {"Content-Type": "application/json"}
if session_id:
headers["X-Session-ID"] = str(session_id)
url = f"{self.proxy_url}/inference/v1/generate"
gen_sem, _ = self._get_semaphores()
if gen_sem is None:
response = await self._post(url, json=payload, headers=headers)
else:
async with gen_sem:
response = await self._post(url, json=payload, headers=headers)
# vLLM returns: list[dict[str(token_id) → {"logprob": float, ...}] | None]
result_prompt_logprobs: Optional[List[Optional[float]]] = None
result_topk_prompt_logprobs: Optional[List[Optional[List[Tuple[int, float]]]]] = None
raw_prompt_logprobs = response.get("prompt_logprobs")
if raw_prompt_logprobs is not None and include_prompt_logprobs:
result_prompt_logprobs = [
(pos_dict.get(str(tid)) or {}).get("logprob") if pos_dict is not None else None
for tid, pos_dict in zip(token_ids, raw_prompt_logprobs)
]
if topk_prompt_logprobs_k > 0:
# vLLM returns k or k+1 logprobs per position (the extra entry is the
# prompt token when it falls outside the top-k). Tinker always returns
# exactly top-k, so we sort and truncate below.
result_topk_prompt_logprobs = [
(
sorted(
[(int(tid), entry["logprob"]) for tid, entry in pos_dict.items()],
key=lambda x: x[1],
reverse=True,
)[:topk_prompt_logprobs_k]
if pos_dict is not None
else None
)
for _, pos_dict in zip(token_ids, raw_prompt_logprobs)
]
# Transform response choices → sequences
sequences = []
for choice in response.get("choices", []):
seq_logprobs: Optional[List[float]] = None
logprobs_data = choice.get("logprobs")
if logprobs_data is not None:
logprobs_content = logprobs_data.get("content", [])
if logprobs_content:
seq_logprobs = [lp["logprob"] for lp in logprobs_content]
sequences.append(
{
"tokens": choice["token_ids"],
"logprobs": seq_logprobs,
"stop_reason": choice.get("finish_reason"),
}
)
return {
"type": "sample",
"sequences": sequences,
"prompt_logprobs": result_prompt_logprobs,
"topk_prompt_logprobs": result_topk_prompt_logprobs,
}
async def chat_completion(
self,
request_payload: Dict[str, Any],
) -> Dict[str, Any]:
"""
Chat completion via /v1/chat/completions.
Args:
request_payload: Dict with {"json": <request-body>, "headers": <headers-dict>}.
The request body must be an OpenAI-compatible chat completion
request. ``model`` is optional and resolved via
``_resolve_model``; if omitted the body is mutated to inject the
resolved value before forwarding to vLLM. ``session_id`` can be
included in the body for consistent routing.
Returns:
OpenAI-compatible chat completion response.
"""
session_id, body = _extract_session_id_and_body(request_payload)
body["model"] = self._resolve_model(body.get("model"), "chat_completion")
headers = {"Content-Type": "application/json"}
if session_id:
headers["X-Session-ID"] = str(session_id)
url = f"{self.proxy_url}/v1/chat/completions"
gen_sem, _ = self._get_semaphores()
if gen_sem is None:
return await self._post(url, json=body, headers=headers)
else:
async with gen_sem:
return await self._post(url, json=body, headers=headers)
async def render_chat_completion(
self,
request_payload: Dict[str, Any],
) -> Dict[str, Any]:
"""
Render a chat completion (apply chat template + tokenize) via /v1/chat/completions/render.
Args:
request_payload: Dict with {"json": <request-body>}.
The request body should be OpenAI-compatible chat completion
request. ``model`` is optional and resolved via
``_resolve_model``. session_id can be included in json for
consistent routing.
Returns:
Rendered chat completion response (template-applied prompt and token IDs).
"""
session_id, body = _extract_session_id_and_body(request_payload)
body["model"] = self._resolve_model(body.get("model"), "render_chat_completion")
headers = {"Content-Type": "application/json"}
if session_id:
headers["X-Session-ID"] = str(session_id)
url = f"{self.proxy_url}/v1/chat/completions/render"
gen_sem, _ = self._get_semaphores()
if gen_sem is None:
return await self._post(url, json=body, headers=headers)
else:
async with gen_sem:
return await self._post(url, json=body, headers=headers)
async def completion(
self,
request_payload: Dict[str, Any],
) -> Dict[str, Any]:
"""
Completion via /v1/completions.
Args:
request_payload: Dict with {"json": <request-body>, "headers": <headers-dict>}.
The request body should be OpenAI-compatible completion
request. ``model`` is optional and resolved via
``_resolve_model``. session_id can be included in json for
consistent routing.
Returns:
OpenAI-compatible completion response.
"""
session_id, body = _extract_session_id_and_body(request_payload)
body["model"] = self._resolve_model(body.get("model"), "completion")
headers = {"Content-Type": "application/json"}
if session_id:
headers["X-Session-ID"] = str(session_id)
url = f"{self.proxy_url}/v1/completions"
gen_sem, _ = self._get_semaphores()
if gen_sem is None:
return await self._post(url, json=body, headers=headers)
else:
async with gen_sem:
return await self._post(url, json=body, headers=headers)
async def tokenize(
self,
texts: List[str],
add_special_tokens: bool = True,
) -> List[List[int]]:
"""
Tokenize texts.
Uses the local tokenizer if available, otherwise falls back to HTTP /tokenize.
Args:
texts: List of texts to tokenize.
add_special_tokens: Whether to add special tokens.
Returns:
List of token ID lists.
"""
if self.tokenizer is not None:
return self.tokenizer(texts, add_special_tokens=add_special_tokens)["input_ids"]
url = f"{self.proxy_url}/tokenize"
# vLLM /tokenize expects individual requests, batch them
results = []
for text in texts:
payload = {
"model": self.model_name,
"prompt": text,
"add_special_tokens": add_special_tokens,
}
result = await self._post(url, json=payload)
results.append(result.get("tokens", []))
return results
async def detokenize(
self,
token_ids: List[List[int]],
) -> List[str]:
"""
Detokenize token IDs.
Uses the local tokenizer if available, otherwise falls back to HTTP /detokenize.
Args:
token_ids: List of token ID lists.
Returns:
List of decoded texts.
"""
if self.tokenizer is not None:
return self.tokenizer.batch_decode(token_ids)
url = f"{self.proxy_url}/detokenize"
# vLLM /detokenize expects individual requests, batch them
results = []
for ids in token_ids:
payload = {
"model": self.model_name,
"tokens": ids,
}
result = await self._post(url, json=payload)
results.append(result.get("prompt", ""))
return results
async def finish_session(self, session_id: str) -> None:
"""Notify the router that a session (trajectory) is complete.
Best-effort data-plane call to the router's ``/finish_session`` endpoint.
Session-aware routing policies (e.g. ``sticky_least_loaded``)
use this to release the replica capacity held by the session so that new
trajectories are balanced onto less-busy engines.
Failures are logged but never raised: this runs in trajectory cleanup
paths (``finally`` blocks, cancellation handlers) and must not mask the
original outcome. Routers/policies that don't track sessions treat this
as a no-op, and unknown session ids are ignored server-side.
"""
if not session_id:
return
url = f"{self.proxy_url}/finish_session"
try:
session = await self._get_session()
# Bound this best-effort cleanup call: the shared session has no
# timeout (total=None), so an unresponsive router would otherwise
# hang the trajectory's finally block forever and wedge the loop.
async with session.post(
url,
params={"session_id": str(session_id)},
timeout=aiohttp.ClientTimeout(total=10.0),
) as resp:
# Drain the body so the keep-alive connection can be reused.
await resp.read()
if resp.status >= 400:
logger.warning(f"finish_session for session_id={session_id!r} returned HTTP {resp.status}")
except asyncio.TimeoutError:
logger.warning(f"finish_session for session_id={session_id!r} timed out after 10s (router unresponsive)")
except Exception as e:
logger.warning(f"finish_session for session_id={session_id!r} failed: {e}")
# ---------------------------
# Control Plane (fan-out to all server_urls)
# ---------------------------
async def _call_server(
self,
server_url: str,
endpoint: str,
json: Optional[Dict[str, Any]] = None,
method: str = "POST",
params: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
"""
Call endpoint on a single server.
Args:
server_url: Base URL of the server.
endpoint: Endpoint path (e.g., "/pause").
json: JSON payload to send as request body.
method: HTTP method (default: POST).
params: URL query parameters (e.g., for FastAPI Query() params).
Returns:
Tuple of (server_url, {"status": <int>, "body": <response>}).
"""
session = await self._get_session()
url = f"{server_url}{endpoint}"
async with session.request(method, url, json=json, params=params) as resp:
body = await resp.json() if resp.content_length else None
raise_for_status(resp, body)
return server_url, {"status": resp.status, "body": body}
async def _call_all_servers(
self,
endpoint: str,
json: Optional[Dict[str, Any]] = None,
method: str = "POST",
params: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
"""
Call endpoint on all server_urls concurrently.
Args:
endpoint: Endpoint path (e.g., "/pause").
json: JSON payload to send as request body.
method: HTTP method (default: POST).
params: URL query parameters (e.g., for FastAPI Query() params).
Returns:
Dict mapping server_url to response.
"""
results = await asyncio.gather(
*[self._call_server(url, endpoint, json, method, params) for url in self.server_urls]
)
return {url: resp for url, resp in results}
async def pause(self, mode: Union[PauseMode, str] = PauseMode.KEEP, clear_cache: bool = False) -> Dict[str, Any]:
"""
Pause generation on all backends.
Args:
mode: Pause mode determining how in-flight requests are handled.
Can be a PauseMode enum or string ("abort", "keep", "wait").
- KEEP / "keep": Freeze in-flight requests in the scheduler.
They resume where they left off on /resume. KV cache is
preserved. No retry needed. (default)
- ABORT / "abort": Abort in-flight requests immediately. Clients
receive partial tokens and must retry with accumulated context.
- WAIT / "wait": Wait for in-flight requests to complete before
pausing. New requests are blocked. No retry needed.
clear_cache: Whether to clear the KV cache on pause. Defaults to False.
Returns:
Dict mapping server_url to response.
"""
if isinstance(mode, str):
mode = PauseMode(mode.lower())
params: Dict[str, Any] = {"mode": mode.value, "clear_cache": str(clear_cache).lower()}
return await self._call_all_servers("/pause", params=params)
async def resume(self) -> Dict[str, Any]:
"""Resume generation on all backends."""
return await self._call_all_servers("/resume")
async def pause_generation(self, clear_cache: bool = False) -> Dict[str, Any]:
"""Pause using keep mode."""
return await self.pause(mode=PauseMode.KEEP, clear_cache=clear_cache)
async def resume_generation(self) -> Dict[str, Any]:
"""Resume after pause."""
return await self.resume()
async def sleep(self, level: int = 2, tags: Optional[List[str]] = None) -> Dict[str, Any]:
"""
Put all backends to sleep (offload weights to CPU).
Args:
level: Sleep level (1 or 2). Level 2 offloads more aggressively.
tags: Optional list of tags to sleep specific resources.
Common tags: ["weights"], ["kv_cache"], or None for all.
Returns:
Dict mapping server_url to response.
"""
# Mirror BaseVLLMInferenceEngine.sleep: when the trainer syncs LoRA adapters
# only, force level=1 so the base model survives via CPU backup. level=2
# discards weights with no source to restore from on wake_up(["weights"]).
if self.uses_lora_weight_sync and level != 1:
logger.info(
"Forcing sleep level=1 (uses_lora_weight_sync=True); requested level=%d would discard the base model.",
level,
)
level = 1
params: Dict[str, Any] = {"level": str(level)}
if tags:
params["tags"] = tags
return await self._call_all_servers("/sleep", params=params)
async def wake_up(self, tags: Optional[List[str]] = None) -> Dict[str, Any]:
"""
Wake up all backends (load weights back to GPU).
Args:
tags: Optional list of tags to wake up specific resources.
Common tags: ["weights"], ["kv_cache"], or None for all.
"""
params = {"tags": tags} if tags else {}
return await self._call_all_servers("/wake_up", params=params)
async def sleep_for_weight_sync(self, offload_kv: bool = True) -> Dict[str, Any]:
"""Free GPU memory for weight sync while keeping in-flight requests frozen.
Sleeps the allocator via ``/collective_rpc`` without touching the scheduler
(unlike :meth:`sleep`, which preempts running requests and clears the prefix
cache). ``offload_kv`` offloads the KV cache to CPU so frozen requests resume
from it; set False to discard it (e.g. when the sync will reset it anyway).
Pair with :meth:`wake_for_weight_sync`; the caller must KEEP-pause before and
``resume`` after.
"""
return await self._call_all_servers(
"/collective_rpc",
{"method": "skyrl_sleep_for_weight_sync", "kwargs": {"offload_kv": offload_kv}},
)
async def wake_for_weight_sync(self, tags: List[str]) -> Dict[str, Any]:
"""Restore allocator pools by tag (see :meth:`sleep_for_weight_sync`).
Wake ``["weights"]`` before the broadcast and ``["kv_cache"]`` after. Does
not resume generation -- call :meth:`resume_generation` once KV is back.
"""
return await self._call_all_servers(
"/collective_rpc",
{"method": "skyrl_wake_for_weight_sync", "kwargs": {"tags": tags}},
)
async def reset_prefix_cache(
self,
reset_running_requests: bool = False,
) -> Dict[str, Any]:
"""
Reset KV cache on all backends.
Args:
reset_running_requests: Whether to reset running requests.
Returns:
Dict mapping server_url to response.
"""
return await self._call_all_servers("/reset_prefix_cache", {"reset_running_requests": reset_running_requests})
# ---------------------------
# Weight Sync (control plane - fan-out)
# ---------------------------
async def init_weight_update_communicator(
self,
init_info: "WeightSyncInitInfo",
) -> Dict[str, Any]:
"""
Initialize weight sync via vLLM native /init_weight_transfer_engine.
Fetches per-server world sizes, expands init_info into per-server
payloads (with correct NCCL rank offsets), and fans out to all servers.
Args:
init_info: A WeightSyncInitInfo (e.g. BroadcastInitInfo) that supports
for_servers() and to_api_payload().
Returns:
Dict mapping server_url to response.
"""
_, world_size_per_server = await self.get_world_size()
num_servers = len(self.server_urls)
server_infos = init_info.for_servers(world_size_per_server, num_servers, dp_size=self.data_parallel_size)
payloads = [{"init_info": x.to_api_payload()} for x in server_infos]
results = await asyncio.gather(
*[
self._call_server(url, "/init_weight_transfer_engine", payload)
for url, payload in zip(self.server_urls, payloads)
]
)
return {url: resp for url, resp in results}
async def update_named_weights(
self,
update_info: Dict[str, Any],
) -> Dict[str, Any]:
"""
Update model weights via vLLM native /update_weights. Used for full parameter fine-tuning.
For LoRA weight sync, use load_lora_adapter() instead.
Args:
update_info: Dict with keys expected by vLLM (names, dtype_names, shapes, packed, etc.)
Returns:
Dict mapping server_url to response.
"""
return await self._call_all_servers(
"/update_weights",
{"update_info": update_info},
)
# TODO: Once https://github.com/vllm-project/vllm/pull/39212 lands, switch
# these three methods from /collective_rpc to the native vLLM endpoints
# (/start_weight_update, /update_weights, /finish_weight_update) and remove
# the NewInferenceWorkerWrap worker extension.
async def start_weight_update(
self,
is_checkpoint_format: bool = True,
) -> Dict[str, Any]:
"""
Start a new chunked weight update via /collective_rpc.
Calls the NewInferenceWorkerWrap.skyrl_start_weight_update method on all
workers. For checkpoint-format weights this initializes layerwise
reload. Must be called before any update_weights_ipc calls.
Args:
is_checkpoint_format: True if weights are in checkpoint format
(need layerwise processing), False for kernel format.
Returns:
Dict mapping server_url to response.
"""
return await self._call_all_servers(
"/collective_rpc",
{
"method": "skyrl_start_weight_update",
"kwargs": {"is_checkpoint_format": is_checkpoint_format},
},
)
async def update_weights_ipc(
self,
update_info: Dict[str, Any],
) -> Dict[str, Any]:
"""
Send a single weight chunk via /collective_rpc.
Calls NewInferenceWorkerWrap.update_weights_ipc on all workers.
Can be called multiple times between skyrl_start_weight_update and
skyrl_finish_weight_update.
Args:
update_info: Dict with backend-specific update info (names,
dtype_names, shapes, ipc_handles_pickled or packed flag).
Returns:
Dict mapping server_url to response.
"""
return await self._call_all_servers(
"/collective_rpc",
{
"method": "update_weights_ipc",
"kwargs": {"update_info": update_info},
},
)
async def update_weights_nccl(
self,
update_info: Dict[str, Any],
) -> Dict[str, Any]:
"""
Send batched weight update via /collective_rpc to the broadcast receiver.
Calls NewInferenceWorkerWrap.update_weights_nccl on all workers,
which routes weight_transfer_engine.receive_weights through the
set_current_vllm_config wrap. Used by the broadcast (NCCL) sender as
a temporary substitute for vLLM's native /update_weights endpoint
until the upstream patch (vllm-project/vllm weight-sync-fix) lands.
Args:
update_info: Dict with backend-specific update info (names,
dtype_names, shapes, packed flag, etc.) — same shape vLLM's
native /update_weights expects.
Returns:
Dict mapping server_url to response.
"""
return await self._call_all_servers(
"/collective_rpc",
{
"method": "update_weights_nccl",
"kwargs": {"update_info": update_info},
},
)
async def finish_weight_update(self) -> Dict[str, Any]:
"""
Finish the current chunked weight update via /collective_rpc.
Calls NewInferenceWorkerWrap.skyrl_finish_weight_update on all workers.
For checkpoint-format weights, runs layerwise postprocessing.
Returns:
Dict mapping server_url to response.
"""
return await self._call_all_servers(
"/collective_rpc",
{"method": "skyrl_finish_weight_update"},
)
async def load_lora_adapter(
self,
lora_name: str,
lora_path: str,
) -> Dict[str, Any]:
"""
Load (or reload) a LoRA adapter on all backend servers via the SkyRL
custom /skyrl/v1/load_lora_adapter endpoint.
After loading, generation/chat/completion requests can target this LoRA
by passing ``model=lora_name``.
TODO(aaron): switch back to vLLM's /v1/load_lora_adapter once the
upstream fix in https://github.com/vllm-project/vllm/pull/41482 lands
in a vLLM release we depend on.
The custom endpoint (defined in vllm_server_actor.py) wraps add_lora
with load_inplace=True (so the engine reloads the freshly-written
safetensors) and then resets the cached LoRARequest's load_inplace=False
(so subsequent generates don't reload from disk on every step). This
avoids two vLLM 0.19.0 bugs that surface under colocate_all + tp=1 +
num_engines>=2 — see vllm_server_actor.py:_skyrl_load_lora_adapter for
the detailed explanation.
Args:
lora_name: Name to register the adapter under on each server.
lora_path: Path to the LoRA adapter on disk (must be accessible from servers).
Returns:
Dict mapping server_url to response.
"""
session = await self._get_session()
async def _load_on_server(server_url: str):
url = f"{server_url}/skyrl/v1/load_lora_adapter"
payload = {"lora_name": lora_name, "lora_path": lora_path}
async with session.post(url, json=payload) as resp:
if resp.status >= 400:
body = await resp.json()
raise_for_status(resp, body)
return server_url, {"status": resp.status, "body": await resp.text()}
results = await asyncio.gather(*[_load_on_server(url) for url in self.server_urls])
logger.info(f"Loaded LoRA adapter '{lora_name}' from {lora_path}")
return {url: resp for url, resp in results}
async def unload_lora_adapter(self, lora_name: str) -> Dict[str, Any]:
"""
Unload a previously-loaded LoRA adapter on all backend servers via /v1/unload_lora_adapter.
After unloading, ``lora_name`` is no longer accepted as a ``model``
target on any server. The underlying CPU/GPU LRU entries on vLLM age
out naturally as new adapters are loaded.
Args:
lora_name: Name of the adapter to unload.
Returns:
Dict mapping server_url to response.
"""
payload = {"lora_name": lora_name}
# Mirror load_lora_adapter: vLLM returns plain text on success and JSON
# ErrorResponse (e.g. 404) on failure.
session = await self._get_session()
async def _unload_on_server(server_url: str):
url = f"{server_url}/v1/unload_lora_adapter"
async with session.post(url, json=payload) as resp:
if resp.status >= 400:
body = await resp.json()
raise_for_status(resp, body)
return server_url, {"status": resp.status, "body": await resp.text()}
results = await asyncio.gather(*[_unload_on_server(url) for url in self.server_urls])
logger.info(f"Unloaded LoRA adapter '{lora_name}'")
return {url: resp for url, resp in results}
# ---------------------------
# Info
# ---------------------------
async def get_world_size(self) -> Tuple[int, int]:
"""
Get total and per-server world size across all inference workers.
Fetches from vLLM's /get_world_size endpoint on each server.
All servers are expected to have the same world size.
Result is cached after first call.
When data_parallel_size > 1, server_urls contains num_engines * dp_size entries.
vLLM reports the full DP * TP world size per server, which already
covers all DP ranks in one deployment. To avoid double-counting,
total_world_size = per_server_ws * num_deployments (not num_servers).
Returns:
Tuple of (total_world_size, world_size_per_server).
"""
if self._world_size is not None:
return self._world_size
results = await self._call_all_servers("/get_world_size", {}, method="GET")
per_server = []
for server_url in self.server_urls:
resp = results.get(server_url)
if resp is None:
raise RuntimeError(f"No response for server {server_url}")
body = resp.get("body", {})
world_size = body.get("world_size")
if world_size is None:
raise RuntimeError(f"Missing world_size in response from {server_url}")
per_server.append(world_size)
assert all(
ws == per_server[0] for ws in per_server
), f"All servers must have the same world_size, got {per_server}"
# Each server is one DP rank. vLLM reports world_size = dp_size * tp_size * pp_size,
# which is the worker count across ALL DP ranks in one deployment.
# num_deployments = num_servers / dp_size (each deployment has dp_size servers).
# Total unique workers = per_server_ws * num_deployments.
num_deployments = len(self.server_urls) // self.data_parallel_size
self._world_size = (per_server[0] * num_deployments, per_server[0])
return self._world_size
# ---------------------------
# Lifecycle
# ---------------------------
async def teardown(self) -> None:
"""Close HTTP session."""
if self._session and not self._session.closed:
await self._session.close()
self._session = None
async def __aenter__(self) -> "RemoteInferenceClient":
"""Async context manager entry."""
return self
async def __aexit__(self, exc_type, exc_val, exc_tb) -> None:
"""Async context manager exit."""
await self.teardown()
# ---------------------------
# Serialization
# ---------------------------
def __getstate__(self) -> dict:
"""Exclude non-serializable fields from pickle."""
state = self.__dict__.copy()
state["_session"] = None
state["_gen_sem"] = None
state["_detok_sem"] = None
state["_sem_loop"] = None
return state
def __setstate__(self, state: dict) -> None:
"""Restore state after unpickling."""
self.__dict__.update(state)
self._session = None
self._gen_sem = None
self._detok_sem = None
self._sem_loop = None
async def aclose(self):
if self._session is not None:
try:
await self._session.close()
except Exception as e:
logger.warning(f"Encountered exception {e} while closing client session")
pass
self._session = Noneattr proxy_url
proxy_url: strData plane URL (single endpoint - router or direct server).
attr server_urls
server_urls: List[str]Control plane URLs (list of backend servers for fan-out).
attr data_parallel_size
data_parallel_size: intData parallel size. Used to compute total inference world size correctly: server_urls contains num_engines * data_parallel_size entries, but vLLM already reports the full DP world size per server, so we divide by num_deployments.
attr abstractmethod property model_name
model_name: str = 'default'The model identifier accepted by the inference server for the base model.
This is usually the model path, but may be served_model_name when vLLM
is started with an alias. It is never a LoRA adapter name. LoRA adapters are
addressed by the names callers register them under via
load_lora_adapter(name, path), and per-call routing is done by
passing that name as model on the data-plane methods.
Used internally only by tokenize/detokenize, which are LoRA-
agnostic but still require a model field per the OpenAI schema.
attr enable_return_routed_experts
enable_return_routed_experts: bool = FalseWhether to return routed expert indices (R3 / rollout router replay).
attr uses_lora_weight_sync
uses_lora_weight_sync: bool = FalseTrue when the trainer syncs LoRA adapters (rather than full/merged weights). When True,
sleep() is forced to level=1: level=2 discards the base model from VRAM with no CPU backup,
and LoRA-only broadcasts cannot repopulate it. Must be kept in sync with the same gate vLLM
uses for enable_lora (see _uses_lora_weight_sync in inference_servers/utils.py).
attr tokenizer
tokenizer: Optional[Any] = NoneOptional HF tokenizer for local tokenize/detokenize (avoids HTTP round-trips).
attr property weight_version
weight_version: intNumber of weight syncs to the engines so far (0 before the first sync); the policy version.
method increment_weight_version
increment_weight_version() -> NoneAdvance the weight version. Called once per completed weight sync to the engines.
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:236-238
def increment_weight_version(self) -> None:
"""Advance the weight version. Called once per completed weight sync to the engines."""
self._weight_version += 1method abstractmethod get_endpoint_url
get_endpoint_url() -> strData-plane endpoint base URL (the router/proxy that load-balances requests).
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:249-251
def get_endpoint_url(self) -> str:
"""Data-plane endpoint base URL (the router/proxy that load-balances requests)."""
return self.proxy_urlmethod async generate
generate(input_batch: InferenceEngineInput, model: Optional[str] = None) -> InferenceEngineOutputGenerate completions via /v1/completions.
This is the interface for token-in-token-out workflows. Input will have token ids, and the output is token ids as well.
Each prompt is sent as a separate request to allow the router to route based on session_id. All requests are made in parallel.
With keep-mode pause, in-flight requests are frozen and resume transparently after /resume -- no client-side retry needed.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
input_batch | InferenceEngineInput | Contains prompt_token_ids, sampling_params, and optional session_ids. | required |
model | Optional[str] | Optional model identifier — the base model name or a loaded LoRA adapter name. When omitted, defaults to self.model_name if LoRA is not in use; raises ValueError if it is. | None |
Returns:
| Type | Description |
|---|---|
| InferenceEngineOutput | InferenceEngineOutput with responses, response_ids, and stop_reasons. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:365-458
async def generate(
self,
input_batch: InferenceEngineInput,
model: Optional[str] = None,
) -> InferenceEngineOutput:
"""
Generate completions via /v1/completions.
This is the interface for token-in-token-out workflows. Input will have
token ids, and the output is token ids as well.
Each prompt is sent as a separate request to allow the router to route
based on session_id. All requests are made in parallel.
With keep-mode pause, in-flight requests are frozen and resume
transparently after /resume -- no client-side retry needed.
Args:
input_batch: Contains prompt_token_ids, sampling_params, and optional session_ids.
model: Optional model identifier — the base model name or a loaded
LoRA adapter name. When omitted, defaults to ``self.model_name``
if LoRA is not in use; raises ``ValueError`` if it is.
Returns:
InferenceEngineOutput with responses, response_ids, and stop_reasons.
"""
model = self._resolve_model(model, "generate")
prompt_token_ids = input_batch.get("prompt_token_ids")
if prompt_token_ids is None:
raise ValueError("RemoteInferenceClient only accepts `prompt_token_ids`, not `prompts`.")
sampling_params = input_batch.get("sampling_params") or {}
if sampling_params.get("n", 1) > 1:
raise ValueError("n > 1 is not supported. Use `config.generator.n_samples_per_prompt` instead.")
session_ids = input_batch.get("session_ids")
mm_features = input_batch.get("mm_features")
cache_salt = input_batch.get("cache_salt")
get_logprobs = sampling_params.get("logprobs") is not None
# Two semaphores decouple the generate and detokenize stages:
# gen_sem: limits concurrent in-flight generate requests so we don't
# overwhelm the router/vLLM scheduler. Released as soon as
# generation finishes, so the GPU slot is freed immediately.
# detok_sem: limits concurrent detokenize calls independently. Uses the
# same concurrency limit so detokenize never starves generate.
# Semaphores are shared across all concurrent generate() calls on this client
# instance, so total in-flight requests are capped at
# SKYRL_GENERATE_CONCURRENCY_PER_ENGINE × num_engines regardless of how many
# callers invoke generate() simultaneously.
# TODO (sumanthrh) (RemoteInferenceClient data-plane-deprecation): We should move this outside of the client to a runner abstraction that will also parallelize client requests across processes.
gen_sem, detok_sem = self._get_semaphores()
batch_size = len(prompt_token_ids)
async def _throttled_generate(idx: int) -> Dict[str, Any]:
if gen_sem is None:
return await self._generate_single(
prompt_token_ids=prompt_token_ids[idx],
sampling_params=sampling_params,
session_id=session_ids[idx] if session_ids and idx < len(session_ids) else None,
mm_features=mm_features[idx] if mm_features and idx < len(mm_features) else None,
model=model,
cache_salt=cache_salt,
)
async with gen_sem:
return await self._generate_single(
prompt_token_ids=prompt_token_ids[idx],
sampling_params=sampling_params,
session_id=session_ids[idx] if session_ids and idx < len(session_ids) else None,
mm_features=mm_features[idx] if mm_features and idx < len(mm_features) else None,
model=model,
cache_salt=cache_salt,
)
async def _throttled_detokenize(token_ids: List[int]) -> str:
if detok_sem is None:
return (await self.detokenize([token_ids]))[0]
async with detok_sem:
return (await self.detokenize([token_ids]))[0]
raw_results = await asyncio.gather(*[_throttled_generate(idx) for idx in range(batch_size)])
responses = await asyncio.gather(*[_throttled_detokenize(r["response_ids"]) for r in raw_results])
rollout_expert_indices = [r.get("routed_experts") for r in raw_results]
has_routed_experts = any(x is not None for x in rollout_expert_indices)
return InferenceEngineOutput(
responses=responses,
stop_reasons=[r["stop_reason"] for r in raw_results],
response_ids=[r["response_ids"] for r in raw_results],
response_logprobs=[r["response_logprobs"] for r in raw_results] if get_logprobs else None,
rollout_expert_indices=rollout_expert_indices if has_routed_experts else None,
)method abstractmethod sample
sample(request_payload: SampleRequestPayload) -> SampleResponseSample completions via /inference/v1/generate (Tinker API).
Maps Tinker-style sample requests to the vLLM generate endpoint. Uses self._post() for automatic retry + backoff on transient errors.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
request_payload | SampleRequestPayload | SampleRequestPayload with {"json": <request-body>}. Expected keys in json: prompt, num_samples, sampling_params, session_id, include_prompt_logprobs (bool), topk_prompt_logprobs (int). model is optional and resolved via _resolve_model. | required |
Returns:
| Type | Description |
|---|---|
| SampleResponse | SampleResponse with type="sample", sequences list, prompt_logprobs, and topk_prompt_logprobs. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:609-733
async def sample(
self,
request_payload: SampleRequestPayload,
) -> SampleResponse:
"""
Sample completions via /inference/v1/generate (Tinker API).
Maps Tinker-style sample requests to the vLLM generate endpoint.
Uses self._post() for automatic retry + backoff on transient errors.
Args:
request_payload: SampleRequestPayload with {"json": <request-body>}.
Expected keys in json: prompt, num_samples, sampling_params,
session_id, include_prompt_logprobs (bool), topk_prompt_logprobs (int).
``model`` is optional and resolved via ``_resolve_model``.
Returns:
SampleResponse with type="sample", sequences list, prompt_logprobs, and topk_prompt_logprobs.
"""
session_id, body = _extract_session_id_and_body(request_payload)
model = self._resolve_model(body.get("model"), "sample")
body["model"] = model
prompt = body.get("prompt", {})
num_samples = body.get("num_samples", 1)
tinker_params = body.get("sampling_params", {})
# Note: Tinker SampleRequest uses "prompt_logprobs" (bool), while
# SamplingClient.sample() uses "include_prompt_logprobs".
include_prompt_logprobs = body.get("include_prompt_logprobs", body.get("prompt_logprobs", False))
topk_prompt_logprobs_k = body.get("topk_prompt_logprobs", 0)
# vLLM prompt logprob mapping
prompt_logprobs_sp = None
if include_prompt_logprobs:
prompt_logprobs_sp = topk_prompt_logprobs_k if topk_prompt_logprobs_k > 0 else 0
# Render prompt: flatten text tokens and, if images are present,
# call the render endpoint to get placeholder tokens + features.
token_ids, mm_features = await self._render_for_sample(prompt, session_id, model=model)
# Map Tinker SamplingParams → vLLM format
sampling_params: Dict[str, Any] = {
"n": num_samples,
"logprobs": 0,
"output_kind": 2,
"prompt_logprobs": prompt_logprobs_sp,
}
for tinker_key, vllm_key in _TINKER_SAMPLE_TO_VLLM_PARAM_MAP.items():
val = tinker_params.get(tinker_key)
if val is not None:
sampling_params[vllm_key] = val
payload: Dict[str, Any] = {
"sampling_params": sampling_params,
"model": model,
"token_ids": token_ids,
}
if mm_features is not None:
payload["features"] = mm_features
headers = {"Content-Type": "application/json"}
if session_id:
headers["X-Session-ID"] = str(session_id)
url = f"{self.proxy_url}/inference/v1/generate"
gen_sem, _ = self._get_semaphores()
if gen_sem is None:
response = await self._post(url, json=payload, headers=headers)
else:
async with gen_sem:
response = await self._post(url, json=payload, headers=headers)
# vLLM returns: list[dict[str(token_id) → {"logprob": float, ...}] | None]
result_prompt_logprobs: Optional[List[Optional[float]]] = None
result_topk_prompt_logprobs: Optional[List[Optional[List[Tuple[int, float]]]]] = None
raw_prompt_logprobs = response.get("prompt_logprobs")
if raw_prompt_logprobs is not None and include_prompt_logprobs:
result_prompt_logprobs = [
(pos_dict.get(str(tid)) or {}).get("logprob") if pos_dict is not None else None
for tid, pos_dict in zip(token_ids, raw_prompt_logprobs)
]
if topk_prompt_logprobs_k > 0:
# vLLM returns k or k+1 logprobs per position (the extra entry is the
# prompt token when it falls outside the top-k). Tinker always returns
# exactly top-k, so we sort and truncate below.
result_topk_prompt_logprobs = [
(
sorted(
[(int(tid), entry["logprob"]) for tid, entry in pos_dict.items()],
key=lambda x: x[1],
reverse=True,
)[:topk_prompt_logprobs_k]
if pos_dict is not None
else None
)
for _, pos_dict in zip(token_ids, raw_prompt_logprobs)
]
# Transform response choices → sequences
sequences = []
for choice in response.get("choices", []):
seq_logprobs: Optional[List[float]] = None
logprobs_data = choice.get("logprobs")
if logprobs_data is not None:
logprobs_content = logprobs_data.get("content", [])
if logprobs_content:
seq_logprobs = [lp["logprob"] for lp in logprobs_content]
sequences.append(
{
"tokens": choice["token_ids"],
"logprobs": seq_logprobs,
"stop_reason": choice.get("finish_reason"),
}
)
return {
"type": "sample",
"sequences": sequences,
"prompt_logprobs": result_prompt_logprobs,
"topk_prompt_logprobs": result_topk_prompt_logprobs,
}method abstractmethod async chat_completion
chat_completion(request_payload: Dict[str, Any]) -> Dict[str, Any]Chat completion via /v1/chat/completions.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
request_payload | Dict[str, Any] | Dict with {"json": <request-body>, "headers": <headers-dict>}. The request body must be an OpenAI-compatible chat completion request. model is optional and resolved via _resolve_model; if omitted the body is mutated to inject the resolved value before forwarding to vLLM. session_id can be included in the body for consistent routing. | required |
Returns:
| Type | Description |
|---|---|
| Dict[str, Any] | OpenAI-compatible chat completion response. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:735-766
async def chat_completion(
self,
request_payload: Dict[str, Any],
) -> Dict[str, Any]:
"""
Chat completion via /v1/chat/completions.
Args:
request_payload: Dict with {"json": <request-body>, "headers": <headers-dict>}.
The request body must be an OpenAI-compatible chat completion
request. ``model`` is optional and resolved via
``_resolve_model``; if omitted the body is mutated to inject the
resolved value before forwarding to vLLM. ``session_id`` can be
included in the body for consistent routing.
Returns:
OpenAI-compatible chat completion response.
"""
session_id, body = _extract_session_id_and_body(request_payload)
body["model"] = self._resolve_model(body.get("model"), "chat_completion")
headers = {"Content-Type": "application/json"}
if session_id:
headers["X-Session-ID"] = str(session_id)
url = f"{self.proxy_url}/v1/chat/completions"
gen_sem, _ = self._get_semaphores()
if gen_sem is None:
return await self._post(url, json=body, headers=headers)
else:
async with gen_sem:
return await self._post(url, json=body, headers=headers)method abstractmethod async render_chat_completion
render_chat_completion(request_payload: Dict[str, Any]) -> Dict[str, Any]Render a chat completion (apply chat template + tokenize) via /v1/chat/completions/render.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
request_payload | Dict[str, Any] | Dict with {"json": <request-body>}. The request body should be OpenAI-compatible chat completion request. model is optional and resolved via _resolve_model. session_id can be included in json for consistent routing. | required |
Returns:
| Type | Description |
|---|---|
| Dict[str, Any] | Rendered chat completion response (template-applied prompt and token IDs). |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:768-798
async def render_chat_completion(
self,
request_payload: Dict[str, Any],
) -> Dict[str, Any]:
"""
Render a chat completion (apply chat template + tokenize) via /v1/chat/completions/render.
Args:
request_payload: Dict with {"json": <request-body>}.
The request body should be OpenAI-compatible chat completion
request. ``model`` is optional and resolved via
``_resolve_model``. session_id can be included in json for
consistent routing.
Returns:
Rendered chat completion response (template-applied prompt and token IDs).
"""
session_id, body = _extract_session_id_and_body(request_payload)
body["model"] = self._resolve_model(body.get("model"), "render_chat_completion")
headers = {"Content-Type": "application/json"}
if session_id:
headers["X-Session-ID"] = str(session_id)
url = f"{self.proxy_url}/v1/chat/completions/render"
gen_sem, _ = self._get_semaphores()
if gen_sem is None:
return await self._post(url, json=body, headers=headers)
else:
async with gen_sem:
return await self._post(url, json=body, headers=headers)method abstractmethod async completion
completion(request_payload: Dict[str, Any]) -> Dict[str, Any]Completion via /v1/completions.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
request_payload | Dict[str, Any] | Dict with {"json": <request-body>, "headers": <headers-dict>}. The request body should be OpenAI-compatible completion request. model is optional and resolved via _resolve_model. session_id can be included in json for consistent routing. | required |
Returns:
| Type | Description |
|---|---|
| Dict[str, Any] | OpenAI-compatible completion response. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:800-830
async def completion(
self,
request_payload: Dict[str, Any],
) -> Dict[str, Any]:
"""
Completion via /v1/completions.
Args:
request_payload: Dict with {"json": <request-body>, "headers": <headers-dict>}.
The request body should be OpenAI-compatible completion
request. ``model`` is optional and resolved via
``_resolve_model``. session_id can be included in json for
consistent routing.
Returns:
OpenAI-compatible completion response.
"""
session_id, body = _extract_session_id_and_body(request_payload)
body["model"] = self._resolve_model(body.get("model"), "completion")
headers = {"Content-Type": "application/json"}
if session_id:
headers["X-Session-ID"] = str(session_id)
url = f"{self.proxy_url}/v1/completions"
gen_sem, _ = self._get_semaphores()
if gen_sem is None:
return await self._post(url, json=body, headers=headers)
else:
async with gen_sem:
return await self._post(url, json=body, headers=headers)method async tokenize
tokenize(texts: List[str], add_special_tokens: bool = True) -> List[List[int]]Tokenize texts.
Uses the local tokenizer if available, otherwise falls back to HTTP /tokenize.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
texts | List[str] | List of texts to tokenize. | required |
add_special_tokens | bool | Whether to add special tokens. | True |
Returns:
| Type | Description |
|---|---|
| List[List[int]] | List of token ID lists. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:832-865
async def tokenize(
self,
texts: List[str],
add_special_tokens: bool = True,
) -> List[List[int]]:
"""
Tokenize texts.
Uses the local tokenizer if available, otherwise falls back to HTTP /tokenize.
Args:
texts: List of texts to tokenize.
add_special_tokens: Whether to add special tokens.
Returns:
List of token ID lists.
"""
if self.tokenizer is not None:
return self.tokenizer(texts, add_special_tokens=add_special_tokens)["input_ids"]
url = f"{self.proxy_url}/tokenize"
# vLLM /tokenize expects individual requests, batch them
results = []
for text in texts:
payload = {
"model": self.model_name,
"prompt": text,
"add_special_tokens": add_special_tokens,
}
result = await self._post(url, json=payload)
results.append(result.get("tokens", []))
return resultsmethod async detokenize
detokenize(token_ids: List[List[int]]) -> List[str]Detokenize token IDs.
Uses the local tokenizer if available, otherwise falls back to HTTP /detokenize.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
token_ids | List[List[int]] | List of token ID lists. | required |
Returns:
| Type | Description |
|---|---|
| List[str] | List of decoded texts. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:867-897
async def detokenize(
self,
token_ids: List[List[int]],
) -> List[str]:
"""
Detokenize token IDs.
Uses the local tokenizer if available, otherwise falls back to HTTP /detokenize.
Args:
token_ids: List of token ID lists.
Returns:
List of decoded texts.
"""
if self.tokenizer is not None:
return self.tokenizer.batch_decode(token_ids)
url = f"{self.proxy_url}/detokenize"
# vLLM /detokenize expects individual requests, batch them
results = []
for ids in token_ids:
payload = {
"model": self.model_name,
"tokens": ids,
}
result = await self._post(url, json=payload)
results.append(result.get("prompt", ""))
return resultsmethod abstractmethod async finish_session
finish_session(session_id: str) -> NoneNotify the router that a session (trajectory) is complete.
Best-effort data-plane call to the router's /finish_session endpoint.
Session-aware routing policies (e.g. sticky_least_loaded)
use this to release the replica capacity held by the session so that new
trajectories are balanced onto less-busy engines.
Failures are logged but never raised: this runs in trajectory cleanup
paths (finally blocks, cancellation handlers) and must not mask the
original outcome. Routers/policies that don't track sessions treat this
as a no-op, and unknown session ids are ignored server-side.
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:899-932
async def finish_session(self, session_id: str) -> None:
"""Notify the router that a session (trajectory) is complete.
Best-effort data-plane call to the router's ``/finish_session`` endpoint.
Session-aware routing policies (e.g. ``sticky_least_loaded``)
use this to release the replica capacity held by the session so that new
trajectories are balanced onto less-busy engines.
Failures are logged but never raised: this runs in trajectory cleanup
paths (``finally`` blocks, cancellation handlers) and must not mask the
original outcome. Routers/policies that don't track sessions treat this
as a no-op, and unknown session ids are ignored server-side.
"""
if not session_id:
return
url = f"{self.proxy_url}/finish_session"
try:
session = await self._get_session()
# Bound this best-effort cleanup call: the shared session has no
# timeout (total=None), so an unresponsive router would otherwise
# hang the trajectory's finally block forever and wedge the loop.
async with session.post(
url,
params={"session_id": str(session_id)},
timeout=aiohttp.ClientTimeout(total=10.0),
) as resp:
# Drain the body so the keep-alive connection can be reused.
await resp.read()
if resp.status >= 400:
logger.warning(f"finish_session for session_id={session_id!r} returned HTTP {resp.status}")
except asyncio.TimeoutError:
logger.warning(f"finish_session for session_id={session_id!r} timed out after 10s (router unresponsive)")
except Exception as e:
logger.warning(f"finish_session for session_id={session_id!r} failed: {e}")method async pause
pause(mode: Union[PauseMode, str] = PauseMode.KEEP, clear_cache: bool = False) -> Dict[str, Any]Pause generation on all backends.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
mode | Union[PauseMode, str] | Pause mode determining how in-flight requests are handled. Can be a PauseMode enum or string ("abort", "keep", "wait"). - KEEP / "keep": Freeze in-flight requests in the scheduler. They resume where they left off on /resume. KV cache is preserved. No retry needed. (default) - ABORT / "abort": Abort in-flight requests immediately. Clients receive partial tokens and must retry with accumulated context. - WAIT / "wait": Wait for in-flight requests to complete before pausing. New requests are blocked. No retry needed. | KEEP |
clear_cache | bool | Whether to clear the KV cache on pause. Defaults to False. | False |
Returns:
| Type | Description |
|---|---|
| Dict[str, Any] | Dict mapping server_url to response. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:990-1014
async def pause(self, mode: Union[PauseMode, str] = PauseMode.KEEP, clear_cache: bool = False) -> Dict[str, Any]:
"""
Pause generation on all backends.
Args:
mode: Pause mode determining how in-flight requests are handled.
Can be a PauseMode enum or string ("abort", "keep", "wait").
- KEEP / "keep": Freeze in-flight requests in the scheduler.
They resume where they left off on /resume. KV cache is
preserved. No retry needed. (default)
- ABORT / "abort": Abort in-flight requests immediately. Clients
receive partial tokens and must retry with accumulated context.
- WAIT / "wait": Wait for in-flight requests to complete before
pausing. New requests are blocked. No retry needed.
clear_cache: Whether to clear the KV cache on pause. Defaults to False.
Returns:
Dict mapping server_url to response.
"""
if isinstance(mode, str):
mode = PauseMode(mode.lower())
params: Dict[str, Any] = {"mode": mode.value, "clear_cache": str(clear_cache).lower()}
return await self._call_all_servers("/pause", params=params)method async resume
resume() -> Dict[str, Any]Resume generation on all backends.
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1016-1018
async def resume(self) -> Dict[str, Any]:
"""Resume generation on all backends."""
return await self._call_all_servers("/resume")method abstractmethod async pause_generation
pause_generation(clear_cache: bool = False) -> Dict[str, Any]Pause using keep mode.
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1020-1022
async def pause_generation(self, clear_cache: bool = False) -> Dict[str, Any]:
"""Pause using keep mode."""
return await self.pause(mode=PauseMode.KEEP, clear_cache=clear_cache)method abstractmethod async resume_generation
resume_generation() -> Dict[str, Any]Resume after pause.
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1024-1026
async def resume_generation(self) -> Dict[str, Any]:
"""Resume after pause."""
return await self.resume()method abstractmethod async sleep
sleep(level: int = 2, tags: Optional[List[str]] = None) -> Dict[str, Any]Put all backends to sleep (offload weights to CPU).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
level | int | Sleep level (1 or 2). Level 2 offloads more aggressively. | 2 |
tags | Optional[List[str]] | Optional list of tags to sleep specific resources. Common tags: ["weights"], ["kv_cache"], or None for all. | None |
Returns:
| Type | Description |
|---|---|
| Dict[str, Any] | Dict mapping server_url to response. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1028-1052
async def sleep(self, level: int = 2, tags: Optional[List[str]] = None) -> Dict[str, Any]:
"""
Put all backends to sleep (offload weights to CPU).
Args:
level: Sleep level (1 or 2). Level 2 offloads more aggressively.
tags: Optional list of tags to sleep specific resources.
Common tags: ["weights"], ["kv_cache"], or None for all.
Returns:
Dict mapping server_url to response.
"""
# Mirror BaseVLLMInferenceEngine.sleep: when the trainer syncs LoRA adapters
# only, force level=1 so the base model survives via CPU backup. level=2
# discards weights with no source to restore from on wake_up(["weights"]).
if self.uses_lora_weight_sync and level != 1:
logger.info(
"Forcing sleep level=1 (uses_lora_weight_sync=True); requested level=%d would discard the base model.",
level,
)
level = 1
params: Dict[str, Any] = {"level": str(level)}
if tags:
params["tags"] = tags
return await self._call_all_servers("/sleep", params=params)method abstractmethod async wake_up
wake_up(tags: Optional[List[str]] = None) -> Dict[str, Any]Wake up all backends (load weights back to GPU).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
tags | Optional[List[str]] | Optional list of tags to wake up specific resources. Common tags: ["weights"], ["kv_cache"], or None for all. | None |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1054-1063
async def wake_up(self, tags: Optional[List[str]] = None) -> Dict[str, Any]:
"""
Wake up all backends (load weights back to GPU).
Args:
tags: Optional list of tags to wake up specific resources.
Common tags: ["weights"], ["kv_cache"], or None for all.
"""
params = {"tags": tags} if tags else {}
return await self._call_all_servers("/wake_up", params=params)method async sleep_for_weight_sync
sleep_for_weight_sync(offload_kv: bool = True) -> Dict[str, Any]Free GPU memory for weight sync while keeping in-flight requests frozen.
Sleeps the allocator via /collective_rpc without touching the scheduler
(unlike :meth:sleep, which preempts running requests and clears the prefix
cache). offload_kv offloads the KV cache to CPU so frozen requests resume
from it; set False to discard it (e.g. when the sync will reset it anyway).
Pair with :meth:wake_for_weight_sync; the caller must KEEP-pause before and
resume after.
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1065-1078
async def sleep_for_weight_sync(self, offload_kv: bool = True) -> Dict[str, Any]:
"""Free GPU memory for weight sync while keeping in-flight requests frozen.
Sleeps the allocator via ``/collective_rpc`` without touching the scheduler
(unlike :meth:`sleep`, which preempts running requests and clears the prefix
cache). ``offload_kv`` offloads the KV cache to CPU so frozen requests resume
from it; set False to discard it (e.g. when the sync will reset it anyway).
Pair with :meth:`wake_for_weight_sync`; the caller must KEEP-pause before and
``resume`` after.
"""
return await self._call_all_servers(
"/collective_rpc",
{"method": "skyrl_sleep_for_weight_sync", "kwargs": {"offload_kv": offload_kv}},
)method async wake_for_weight_sync
wake_for_weight_sync(tags: List[str]) -> Dict[str, Any]Restore allocator pools by tag (see :meth:sleep_for_weight_sync).
Wake ["weights"] before the broadcast and ["kv_cache"] after. Does
not resume generation -- call :meth:resume_generation once KV is back.
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1080-1089
async def wake_for_weight_sync(self, tags: List[str]) -> Dict[str, Any]:
"""Restore allocator pools by tag (see :meth:`sleep_for_weight_sync`).
Wake ``["weights"]`` before the broadcast and ``["kv_cache"]`` after. Does
not resume generation -- call :meth:`resume_generation` once KV is back.
"""
return await self._call_all_servers(
"/collective_rpc",
{"method": "skyrl_wake_for_weight_sync", "kwargs": {"tags": tags}},
)method abstractmethod async reset_prefix_cache
reset_prefix_cache(reset_running_requests: bool = False) -> Dict[str, Any]Reset KV cache on all backends.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
reset_running_requests | bool | Whether to reset running requests. | False |
Returns:
| Type | Description |
|---|---|
| Dict[str, Any] | Dict mapping server_url to response. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1091-1104
async def reset_prefix_cache(
self,
reset_running_requests: bool = False,
) -> Dict[str, Any]:
"""
Reset KV cache on all backends.
Args:
reset_running_requests: Whether to reset running requests.
Returns:
Dict mapping server_url to response.
"""
return await self._call_all_servers("/reset_prefix_cache", {"reset_running_requests": reset_running_requests})method abstractmethod async init_weight_update_communicator
init_weight_update_communicator(init_info: 'WeightSyncInitInfo') -> Dict[str, Any]Initialize weight sync via vLLM native /init_weight_transfer_engine.
Fetches per-server world sizes, expands init_info into per-server payloads (with correct NCCL rank offsets), and fans out to all servers.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
init_info | 'WeightSyncInitInfo' | A WeightSyncInitInfo (e.g. BroadcastInitInfo) that supports for_servers() and to_api_payload(). | required |
Returns:
| Type | Description |
|---|---|
| Dict[str, Any] | Dict mapping server_url to response. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1110-1137
async def init_weight_update_communicator(
self,
init_info: "WeightSyncInitInfo",
) -> Dict[str, Any]:
"""
Initialize weight sync via vLLM native /init_weight_transfer_engine.
Fetches per-server world sizes, expands init_info into per-server
payloads (with correct NCCL rank offsets), and fans out to all servers.
Args:
init_info: A WeightSyncInitInfo (e.g. BroadcastInitInfo) that supports
for_servers() and to_api_payload().
Returns:
Dict mapping server_url to response.
"""
_, world_size_per_server = await self.get_world_size()
num_servers = len(self.server_urls)
server_infos = init_info.for_servers(world_size_per_server, num_servers, dp_size=self.data_parallel_size)
payloads = [{"init_info": x.to_api_payload()} for x in server_infos]
results = await asyncio.gather(
*[
self._call_server(url, "/init_weight_transfer_engine", payload)
for url, payload in zip(self.server_urls, payloads)
]
)
return {url: resp for url, resp in results}method abstractmethod async update_named_weights
update_named_weights(update_info: Dict[str, Any]) -> Dict[str, Any]Update model weights via vLLM native /update_weights. Used for full parameter fine-tuning.
For LoRA weight sync, use load_lora_adapter() instead.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
update_info | Dict[str, Any] | Dict with keys expected by vLLM (names, dtype_names, shapes, packed, etc.) | required |
Returns:
| Type | Description |
|---|---|
| Dict[str, Any] | Dict mapping server_url to response. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1139-1157
async def update_named_weights(
self,
update_info: Dict[str, Any],
) -> Dict[str, Any]:
"""
Update model weights via vLLM native /update_weights. Used for full parameter fine-tuning.
For LoRA weight sync, use load_lora_adapter() instead.
Args:
update_info: Dict with keys expected by vLLM (names, dtype_names, shapes, packed, etc.)
Returns:
Dict mapping server_url to response.
"""
return await self._call_all_servers(
"/update_weights",
{"update_info": update_info},
)method async start_weight_update
start_weight_update(is_checkpoint_format: bool = True) -> Dict[str, Any]Start a new chunked weight update via /collective_rpc.
Calls the NewInferenceWorkerWrap.skyrl_start_weight_update method on all workers. For checkpoint-format weights this initializes layerwise reload. Must be called before any update_weights_ipc calls.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
is_checkpoint_format | bool | True if weights are in checkpoint format (need layerwise processing), False for kernel format. | True |
Returns:
| Type | Description |
|---|---|
| Dict[str, Any] | Dict mapping server_url to response. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1164-1188
async def start_weight_update(
self,
is_checkpoint_format: bool = True,
) -> Dict[str, Any]:
"""
Start a new chunked weight update via /collective_rpc.
Calls the NewInferenceWorkerWrap.skyrl_start_weight_update method on all
workers. For checkpoint-format weights this initializes layerwise
reload. Must be called before any update_weights_ipc calls.
Args:
is_checkpoint_format: True if weights are in checkpoint format
(need layerwise processing), False for kernel format.
Returns:
Dict mapping server_url to response.
"""
return await self._call_all_servers(
"/collective_rpc",
{
"method": "skyrl_start_weight_update",
"kwargs": {"is_checkpoint_format": is_checkpoint_format},
},
)method async update_weights_ipc
update_weights_ipc(update_info: Dict[str, Any]) -> Dict[str, Any]Send a single weight chunk via /collective_rpc.
Calls NewInferenceWorkerWrap.update_weights_ipc on all workers. Can be called multiple times between skyrl_start_weight_update and skyrl_finish_weight_update.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
update_info | Dict[str, Any] | Dict with backend-specific update info (names, dtype_names, shapes, ipc_handles_pickled or packed flag). | required |
Returns:
| Type | Description |
|---|---|
| Dict[str, Any] | Dict mapping server_url to response. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1190-1214
async def update_weights_ipc(
self,
update_info: Dict[str, Any],
) -> Dict[str, Any]:
"""
Send a single weight chunk via /collective_rpc.
Calls NewInferenceWorkerWrap.update_weights_ipc on all workers.
Can be called multiple times between skyrl_start_weight_update and
skyrl_finish_weight_update.
Args:
update_info: Dict with backend-specific update info (names,
dtype_names, shapes, ipc_handles_pickled or packed flag).
Returns:
Dict mapping server_url to response.
"""
return await self._call_all_servers(
"/collective_rpc",
{
"method": "update_weights_ipc",
"kwargs": {"update_info": update_info},
},
)method async update_weights_nccl
update_weights_nccl(update_info: Dict[str, Any]) -> Dict[str, Any]Send batched weight update via /collective_rpc to the broadcast receiver.
Calls NewInferenceWorkerWrap.update_weights_nccl on all workers, which routes weight_transfer_engine.receive_weights through the set_current_vllm_config wrap. Used by the broadcast (NCCL) sender as a temporary substitute for vLLM's native /update_weights endpoint until the upstream patch (vllm-project/vllm weight-sync-fix) lands.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
update_info | Dict[str, Any] | Dict with backend-specific update info (names, dtype_names, shapes, packed flag, etc.) — same shape vLLM's native /update_weights expects. | required |
Returns:
| Type | Description |
|---|---|
| Dict[str, Any] | Dict mapping server_url to response. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1216-1243
async def update_weights_nccl(
self,
update_info: Dict[str, Any],
) -> Dict[str, Any]:
"""
Send batched weight update via /collective_rpc to the broadcast receiver.
Calls NewInferenceWorkerWrap.update_weights_nccl on all workers,
which routes weight_transfer_engine.receive_weights through the
set_current_vllm_config wrap. Used by the broadcast (NCCL) sender as
a temporary substitute for vLLM's native /update_weights endpoint
until the upstream patch (vllm-project/vllm weight-sync-fix) lands.
Args:
update_info: Dict with backend-specific update info (names,
dtype_names, shapes, packed flag, etc.) — same shape vLLM's
native /update_weights expects.
Returns:
Dict mapping server_url to response.
"""
return await self._call_all_servers(
"/collective_rpc",
{
"method": "update_weights_nccl",
"kwargs": {"update_info": update_info},
},
)method async finish_weight_update
finish_weight_update() -> Dict[str, Any]Finish the current chunked weight update via /collective_rpc.
Calls NewInferenceWorkerWrap.skyrl_finish_weight_update on all workers. For checkpoint-format weights, runs layerwise postprocessing.
Returns:
| Type | Description |
|---|---|
| Dict[str, Any] | Dict mapping server_url to response. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1245-1258
async def finish_weight_update(self) -> Dict[str, Any]:
"""
Finish the current chunked weight update via /collective_rpc.
Calls NewInferenceWorkerWrap.skyrl_finish_weight_update on all workers.
For checkpoint-format weights, runs layerwise postprocessing.
Returns:
Dict mapping server_url to response.
"""
return await self._call_all_servers(
"/collective_rpc",
{"method": "skyrl_finish_weight_update"},
)method async load_lora_adapter
load_lora_adapter(lora_name: str, lora_path: str) -> Dict[str, Any]Load (or reload) a LoRA adapter on all backend servers via the SkyRL custom /skyrl/v1/load_lora_adapter endpoint.
After loading, generation/chat/completion requests can target this LoRA
by passing model=lora_name.
TODO(aaron): switch back to vLLM's /v1/load_lora_adapter once the upstream fix in https://github.com/vllm-project/vllm/pull/41482 lands in a vLLM release we depend on.
The custom endpoint (defined in vllm_server_actor.py) wraps add_lora with load_inplace=True (so the engine reloads the freshly-written safetensors) and then resets the cached LoRARequest's load_inplace=False (so subsequent generates don't reload from disk on every step). This avoids two vLLM 0.19.0 bugs that surface under colocate_all + tp=1 + num_engines>=2 — see vllm_server_actor.py:_skyrl_load_lora_adapter for the detailed explanation.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
lora_name | str | Name to register the adapter under on each server. | required |
lora_path | str | Path to the LoRA adapter on disk (must be accessible from servers). | required |
Returns:
| Type | Description |
|---|---|
| Dict[str, Any] | Dict mapping server_url to response. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1260-1306
async def load_lora_adapter(
self,
lora_name: str,
lora_path: str,
) -> Dict[str, Any]:
"""
Load (or reload) a LoRA adapter on all backend servers via the SkyRL
custom /skyrl/v1/load_lora_adapter endpoint.
After loading, generation/chat/completion requests can target this LoRA
by passing ``model=lora_name``.
TODO(aaron): switch back to vLLM's /v1/load_lora_adapter once the
upstream fix in https://github.com/vllm-project/vllm/pull/41482 lands
in a vLLM release we depend on.
The custom endpoint (defined in vllm_server_actor.py) wraps add_lora
with load_inplace=True (so the engine reloads the freshly-written
safetensors) and then resets the cached LoRARequest's load_inplace=False
(so subsequent generates don't reload from disk on every step). This
avoids two vLLM 0.19.0 bugs that surface under colocate_all + tp=1 +
num_engines>=2 — see vllm_server_actor.py:_skyrl_load_lora_adapter for
the detailed explanation.
Args:
lora_name: Name to register the adapter under on each server.
lora_path: Path to the LoRA adapter on disk (must be accessible from servers).
Returns:
Dict mapping server_url to response.
"""
session = await self._get_session()
async def _load_on_server(server_url: str):
url = f"{server_url}/skyrl/v1/load_lora_adapter"
payload = {"lora_name": lora_name, "lora_path": lora_path}
async with session.post(url, json=payload) as resp:
if resp.status >= 400:
body = await resp.json()
raise_for_status(resp, body)
return server_url, {"status": resp.status, "body": await resp.text()}
results = await asyncio.gather(*[_load_on_server(url) for url in self.server_urls])
logger.info(f"Loaded LoRA adapter '{lora_name}' from {lora_path}")
return {url: resp for url, resp in results}method async unload_lora_adapter
unload_lora_adapter(lora_name: str) -> Dict[str, Any]Unload a previously-loaded LoRA adapter on all backend servers via /v1/unload_lora_adapter.
After unloading, lora_name is no longer accepted as a model
target on any server. The underlying CPU/GPU LRU entries on vLLM age
out naturally as new adapters are loaded.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
lora_name | str | Name of the adapter to unload. | required |
Returns:
| Type | Description |
|---|---|
| Dict[str, Any] | Dict mapping server_url to response. |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1308-1340
async def unload_lora_adapter(self, lora_name: str) -> Dict[str, Any]:
"""
Unload a previously-loaded LoRA adapter on all backend servers via /v1/unload_lora_adapter.
After unloading, ``lora_name`` is no longer accepted as a ``model``
target on any server. The underlying CPU/GPU LRU entries on vLLM age
out naturally as new adapters are loaded.
Args:
lora_name: Name of the adapter to unload.
Returns:
Dict mapping server_url to response.
"""
payload = {"lora_name": lora_name}
# Mirror load_lora_adapter: vLLM returns plain text on success and JSON
# ErrorResponse (e.g. 404) on failure.
session = await self._get_session()
async def _unload_on_server(server_url: str):
url = f"{server_url}/v1/unload_lora_adapter"
async with session.post(url, json=payload) as resp:
if resp.status >= 400:
body = await resp.json()
raise_for_status(resp, body)
return server_url, {"status": resp.status, "body": await resp.text()}
results = await asyncio.gather(*[_unload_on_server(url) for url in self.server_urls])
logger.info(f"Unloaded LoRA adapter '{lora_name}'")
return {url: resp for url, resp in results}method abstractmethod async get_world_size
get_world_size() -> Tuple[int, int]Get total and per-server world size across all inference workers.
Fetches from vLLM's /get_world_size endpoint on each server. All servers are expected to have the same world size. Result is cached after first call.
When data_parallel_size > 1, server_urls contains num_engines * dp_size entries. vLLM reports the full DP * TP world size per server, which already covers all DP ranks in one deployment. To avoid double-counting, total_world_size = per_server_ws * num_deployments (not num_servers).
Returns:
| Type | Description |
|---|---|
| Tuple[int, int] | Tuple of (total_world_size, world_size_per_server). |
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1346-1388
async def get_world_size(self) -> Tuple[int, int]:
"""
Get total and per-server world size across all inference workers.
Fetches from vLLM's /get_world_size endpoint on each server.
All servers are expected to have the same world size.
Result is cached after first call.
When data_parallel_size > 1, server_urls contains num_engines * dp_size entries.
vLLM reports the full DP * TP world size per server, which already
covers all DP ranks in one deployment. To avoid double-counting,
total_world_size = per_server_ws * num_deployments (not num_servers).
Returns:
Tuple of (total_world_size, world_size_per_server).
"""
if self._world_size is not None:
return self._world_size
results = await self._call_all_servers("/get_world_size", {}, method="GET")
per_server = []
for server_url in self.server_urls:
resp = results.get(server_url)
if resp is None:
raise RuntimeError(f"No response for server {server_url}")
body = resp.get("body", {})
world_size = body.get("world_size")
if world_size is None:
raise RuntimeError(f"Missing world_size in response from {server_url}")
per_server.append(world_size)
assert all(
ws == per_server[0] for ws in per_server
), f"All servers must have the same world_size, got {per_server}"
# Each server is one DP rank. vLLM reports world_size = dp_size * tp_size * pp_size,
# which is the worker count across ALL DP ranks in one deployment.
# num_deployments = num_servers / dp_size (each deployment has dp_size servers).
# Total unique workers = per_server_ws * num_deployments.
num_deployments = len(self.server_urls) // self.data_parallel_size
self._world_size = (per_server[0] * num_deployments, per_server[0])
return self._world_sizemethod abstractmethod async teardown
teardown() -> NoneClose HTTP session.
Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1394-1398
async def teardown(self) -> None:
"""Close HTTP session."""
if self._session and not self._session.closed:
await self._session.close()
self._session = Nonemethod async aclose
aclose()Source code in skyrl/backends/skyrl_train/inference_servers/remote_inference_client.py:1429-1436
async def aclose(self):
if self._session is not None:
try:
await self._session.close()
except Exception as e:
logger.warning(f"Encountered exception {e} while closing client session")
pass
self._session = None