API reference
Reference for river_client and its rl library, generated from river-client 0.11.0. The guide covers the API as a walkthrough; this page lists public classes, methods, and types.
pip install river-clientimport river_client as riverCore client
Use these primitives for sampling, sessions, models, and training operations.
Client
River API client.
Connects to the River API server over gRPC with automatic retry on transient failures and connection keepalive.
Client(
api_key: str,
endpoint: str = 'api.river.ai',
port: int = 443,
timeout: float = 86400.0,
use_ssl: bool = True,
enable_retries: bool = True,
image_upload_concurrency: int = 4,
)| Parameter | Type | Default | Description |
|---|---|---|---|
api_key | str | API key for authentication | |
endpoint | str | 'api.river.ai' | API endpoint hostname |
port | int | 443 | API port |
timeout | float | 86400.0 | Default timeout for operations |
use_ssl | bool | True | Whether to use SSL |
enable_retries | bool | True | Whether submit RPCs and gRPC transport may retry transient failures. Disable this for fail-closed evaluation protocols that require one server submission per model turn. |
image_upload_concurrency | int | 4 | Maximum simultaneous async image RPCs (default 4). |
Methods: aclose, session, sample, health_check, get_capabilities, get_server_capabilities, create_deployment, get_deployment, list_deployments, scale_on_target, delete_deployment, get_deployment_usage, wait_for_deployment, chat_complete, chat_complete_from_checkpoint, chat_complete_from_training, close
Client.aclose
async Client.aclose() -> NoneDrain started image uploads and close connections off the event loop.
Client.session
Client.session(
*,
timeout: float = 86400.0,
before_poll: Callable[[], None] | None = None,
on_session_creation_attempted: Callable[[], None] | None = None,
on_session_created: Callable[[], None] | None = None,
on_session_closed: Callable[[], None] | None = None,
**tags: str,
) -> SessionContextCreate a session context manager.
Returns: SessionContext — Context manager that yields a Session
| Parameter | Type | Default | Description |
|---|---|---|---|
timeout | float | 86400.0 | End-to-end session-creation timeout in seconds. |
before_poll | Callable[[], None] | None | None | |
on_session_creation_attempted | Callable[[], None] | None | None | |
on_session_created | Callable[[], None] | None | None | |
on_session_closed | Callable[[], None] | None | None | |
tags | str |
Client.sample
Client.sample(
prompts: str | list[str] | None = None,
*,
base_model: str,
num_samples: int = 1,
max_tokens: int = 256,
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = -1,
stop: list[str] | None = None,
seed: int | None = None,
return_prompt_logprobs: bool = False,
logprobs: int | None = None,
images: list[Image] | list[list[Image]] | None = None,
prompt_token_ids: list[int] | list[list[int]] | None = None,
model_input: list[dict] | list[list[dict]] | None = None,
tokenizer: Any | None = None,
metrics_type: str = '',
timeout: float | None = None,
) -> list[Sample]Sample from a base model (no session required).
Returns: list[Sample] — Flat list[Sample] — all samples across all prompts. For a single prompt with num_samples=1 (the default), this is a list with one element: result[0].text.
| Parameter | Type | Default | Description |
|---|---|---|---|
prompts | str | list[str] | None | None | Single prompt string or list of prompts. Mutually exclusive with prompt_token_ids. |
base_model | str | Base model name (e.g. "Qwen/Qwen3.6-35B-A3B-FP8"). | |
num_samples | int | 1 | Number of independent samples per prompt. |
max_tokens | int | 256 | Maximum tokens to generate per sample. |
temperature | float | 1.0 | Sampling temperature. |
top_p | float | 1.0 | Nucleus sampling threshold. |
top_k | int | -1 | Top-k sampling (-1 = disabled). |
stop | list[str] | None | None | Stop sequences. |
seed | int | None | None | Random seed (varied per sample automatically). |
return_prompt_logprobs | bool | False | Whether to return prompt token logprobs. |
logprobs | int | None | None | If set to K > 0, request the top-K alternative logprobs at each position. Off by default — enabling it roughly halves server throughput. |
images | list[Image] | list[list[Image]] | None | None | Optional raw image bytes for multimodal sampling. See sample for the per-prompt vs. broadcast semantics. |
prompt_token_ids | list[int] | list[list[int]] | None | None | Pre-tokenized prompt(s); mutually exclusive with prompts. See sample for details. |
model_input | list[dict] | list[list[dict]] | None | None | Training-style chunk list(s); mutually exclusive with prompts / prompt_token_ids / images. See sample for details. |
tokenizer | Any | None | None | Optional tokenizer name or already-loaded tokenizer. Defaults to base_model after applying River model-alias resolution. |
metrics_type | str | '' | Opaque server-interpreted token enabling extra scalar metrics on the response. Unrecognized values are silently ignored; when recognized, per-result metrics appear on Sample.metrics. |
timeout | float | None | None | Timeout in seconds. |
Client.health_check
Client.health_check() -> boolCheck API health.
Returns: bool — True if healthy
Client.get_capabilities
Client.get_capabilities() -> list[str]Get supported model names; preserves the original list-returning API.
Client.get_server_capabilities
Client.get_server_capabilities(*, model_name: str | None = None) -> ServerCapabilitiesRead public feature contracts, optionally for one authorized model.
This checks protocol support, not current model capacity or queue state. Results are fetched fresh so a new run detects an endpoint change.
| Parameter | Type | Default |
|---|---|---|
model_name | str | None | None |
Client.create_deployment
Client.create_deployment(
checkpoint: str | Checkpoint | None = None,
*,
base_model: str | None = None,
unified_replicas: int | None = None,
prefill_replicas: int | None = None,
decode_replicas: int | None = None,
idempotency_key: str | None = None,
wait: bool = False,
wait_timeout: float = 1800.0,
timeout: float | None = None,
) -> DeploymentCreate a dedicated deployment of a checkpoint or base model.
Supply exactly one of checkpoint or base_model. Base-model
creation requires a server advertising base_model_creation.
This is a gated feature, disabled by default. Contact River to enable dedicated deployments for your team and the checkpoint's base model before calling this method. Use a team API key with that access; personal API keys cannot create deployments. Installing or upgrading the client does not enable access. The server rejects creation when the feature is disabled or the required team/model access is absent.
The checkpoint path may be republished with new weights. New or restarted workers load the available weights; serving workers keep their loaded adapter. Republishing does not automatically reload the fleet or provide an atomic update, including across prefill/decode workers. For a consistent update at the same URL, scale all roles to zero, wait for workers to drain and stop, finish saving, then restore the targets. Use a fresh checkpoint name and a new deployment when the old deployment must keep serving during the update.
Supply exactly one capacity group: unified_replicas for unified
serving, or both prefill_replicas and decode_replicas for
prefill/decode disaggregated serving. The topology is inferred from
the group and is immutable afterwards. The backend selects the
hardware specification, placement and ports from the checkpoint's
base model.
Provisioning is asynchronous: the returned deployment is usually
phase="accepted". Use the standard OpenAI client with
base_url=deployment.base_url once it is ready; the URL binds
the checkpoint, so the existing model argument can stay unchanged.
Pass wait=True to
block until it serves (or fails). An idempotency_key makes a
retried call return the same accepted deployment; each explicit call
without one creates a new deployment.
| Parameter | Type | Default |
|---|---|---|
checkpoint | str | Checkpoint | None | None |
base_model | str | None | None |
unified_replicas | int | None | None |
prefill_replicas | int | None | None |
decode_replicas | int | None | None |
idempotency_key | str | None | None |
wait | bool | False |
wait_timeout | float | 1800.0 |
timeout | float | None | None |
Client.get_deployment
Client.get_deployment(deployment_id: str, *, timeout: float | None = None) -> DeploymentFetch a deployment by id, including its tombstone after deletion.
Dedicated deployments are gated and disabled by default. See
create_deployment for access requirements.
| Parameter | Type | Default |
|---|---|---|
deployment_id | str | |
timeout | float | None | None |
Client.list_deployments
Client.list_deployments(
*,
include_deleted: bool = False,
timeout: float | None = None,
) -> list[Deployment]List the caller's deployments with their per-role replica counts.
Dedicated deployments are gated and disabled by default. See
create_deployment for access requirements.
| Parameter | Type | Default |
|---|---|---|
include_deleted | bool | False |
timeout | float | None | None |
Client.scale_on_target
Client.scale_on_target(
deployment_id: str,
*,
unified_replicas: int | None = None,
prefill_replicas: int | None = None,
decode_replicas: int | None = None,
idempotency_key: str | None = None,
timeout: float | None = None,
) -> DeploymentSet the absolute desired count of each supplied role.
Dedicated deployments are gated and disabled by default. Increasing capacity, including resuming from zero, requires active access for your team and this deployment's base model; contact River to request it. Reducing capacity remains available after that access is revoked.
Omitted roles keep their current target. A unified deployment accepts
only unified_replicas; a prefill/decode deployment accepts
prefill_replicas and/or decode_replicas, and a target with one
role at zero and the other positive is rejected — scale both to zero
together to stop, and both up to resume.
| Parameter | Type | Default |
|---|---|---|
deployment_id | str | |
unified_replicas | int | None | None |
prefill_replicas | int | None | None |
decode_replicas | int | None | None |
idempotency_key | str | None | None |
timeout | float | None | None |
Client.delete_deployment
Client.delete_deployment(
deployment_id: str,
*,
idempotency_key: str | None = None,
wait: bool = False,
wait_timeout: float = 600.0,
timeout: float | None = None,
) -> DeploymentDelete a deployment in any live state; repeating is a no-op success.
Dedicated deployments are gated and disabled by default. Deleting an existing deployment remains available after your team's access to create or increase capacity for its model is revoked.
| Parameter | Type | Default |
|---|---|---|
deployment_id | str | |
idempotency_key | str | None | None |
wait | bool | False |
wait_timeout | float | 600.0 |
timeout | float | None | None |
Client.get_deployment_usage
Client.get_deployment_usage(
deployment_id: str,
*,
timeout: float | None = None,
) -> dict[str, Any]Requested replica/GPU-hours per role and accepted-operation history.
Dedicated deployments are gated and disabled by default. See
create_deployment for access requirements.
| Parameter | Type | Default |
|---|---|---|
deployment_id | str | |
timeout | float | None | None |
Client.wait_for_deployment
Client.wait_for_deployment(
deployment_id: str,
*,
timeout: float = 1800.0,
poll_interval: float = 5.0,
until: tuple[str, ...] = ('ready', 'degraded'),
) -> DeploymentWait for a requested phase and its serving capacity (or failure).
Dedicated deployments are gated and disabled by default. See
create_deployment for access requirements.
The default accepts partially available capacity. Pass
until=("ready",) to wait for every requested replica to be ready.
| Parameter | Type | Default |
|---|---|---|
deployment_id | str | |
timeout | float | 1800.0 |
poll_interval | float | 5.0 |
until | tuple[str, ...] | ('ready', 'degraded') |
Client.chat_complete
Client.chat_complete(
messages: list[dict],
*,
base_model: str,
timeout: float | None = None,
**kwargs,
) -> ChatCompleteResultChat completion from a base model (no LoRA).
Builds an OpenAI-format request body and sends it through the
gRPC ChatCompleteFromBase RPC.
Returns: ChatCompleteResult — ChatCompleteResult with response_json and status_code.
| Parameter | Type | Default | Description |
|---|---|---|---|
messages | list[dict] | OpenAI-format messages list. | |
base_model | str | Base model name for routing. | |
timeout | float | None | None | Timeout in seconds. |
kwargs |
Client.chat_complete_from_checkpoint
Client.chat_complete_from_checkpoint(
messages: list[dict],
*,
checkpoint_path: str,
base_model: str = '',
timeout: float | None = None,
**kwargs,
) -> ChatCompleteResultChat completion from a saved checkpoint.
Returns: ChatCompleteResult — ChatCompleteResult with response_json and status_code.
| Parameter | Type | Default | Description |
|---|---|---|---|
messages | list[dict] | OpenAI-format messages list. | |
checkpoint_path | str | river:// checkpoint path. | |
base_model | str | '' | Base model name (optional; resolved from DB if empty). |
timeout | float | None | None | Timeout in seconds. |
kwargs |
Client.chat_complete_from_training
Client.chat_complete_from_training(
messages: list[dict],
*,
model_id: str,
timeout: float | None = None,
**kwargs,
) -> ChatCompleteResultChat completion from in-memory training weights.
Returns: ChatCompleteResult — ChatCompleteResult with response_json and status_code.
| Parameter | Type | Default | Description |
|---|---|---|---|
messages | list[dict] | OpenAI-format messages list. | |
model_id | str | Training model ID (e.g. session_id:model:seq). | |
timeout | float | None | None | Timeout in seconds. |
kwargs |
Client.close
Client.close() -> NoneClose connections after any started image uploads finish.
Session
A training session with GPU allocation.
Entered through Client.session; it owns the models you train.
Methods: upload_image, upload_image_async, release_image, release_image_async, restore_images, get_server_capabilities, attest_training_data, create_model, sample, submit_sampling_batch, submit_sample
Session.upload_image
Session.upload_image(
data: bytes,
*,
idempotency_key: str | None = None,
timeout: float | None = None,
) -> ImageHandleUpload into this session; bytes remain recoverable in the local image cache.
| Parameter | Type | Default |
|---|---|---|
data | bytes | |
idempotency_key | str | None | None |
timeout | float | None | None |
Session.upload_image_async
async Session.upload_image_async(
data: bytes,
*,
idempotency_key: str | None = None,
timeout: float | None = None,
) -> ImageHandleBounded non-blocking upload, automatically released when this session ends.
| Parameter | Type | Default |
|---|---|---|
data | bytes | |
idempotency_key | str | None | None |
timeout | float | None | None |
Session.release_image
Session.release_image(
image: ImageHandle | None = None,
*,
idempotency_key: str | None = None,
) -> None| Parameter | Type | Default |
|---|---|---|
image | ImageHandle | None | None |
idempotency_key | str | None | None |
Session.release_image_async
async Session.release_image_async(
image: ImageHandle | None = None,
*,
idempotency_key: str | None = None,
) -> None| Parameter | Type | Default |
|---|---|---|
image | ImageHandle | None | None |
idempotency_key | str | None | None |
Session.restore_images
async Session.restore_images(value: Any, *, image_store: ImageStore) -> AnyRe-upload checkpoint references once per content hash and remap every occurrence.
| Parameter | Type |
|---|---|
value | Any |
image_store | ImageStore |
Session.get_server_capabilities
Session.get_server_capabilities(*, model_name: str | None = None) -> ServerCapabilitiesRead protocol support, optionally for one authorized model.
| Parameter | Type | Default |
|---|---|---|
model_name | str | None | None |
Session.attest_training_data
Session.attest_training_data(
artifacts: list[TrainingDataArtifact],
timeout: float = 86400.0,
) -> TrainingDataAttestationAsk the API to hash source artifacts and retain their manifest.
The source bytes are discarded by the API after hashing. Pass the
result to create_model to make the server fail closed before
every forward/backward request if that manifest disappears or no longer
belongs to the model's session.
| Parameter | Type | Default |
|---|---|---|
artifacts | list[TrainingDataArtifact] | |
timeout | float | 86400.0 |
Session.create_model
Session.create_model(
base_model: str,
lora: LoraConfig,
tokenizer: str | Any | None = None,
checkpoint: str | Checkpoint | None = None,
timeout: float = 86400.0,
training_data_attestation: TrainingDataAttestation | str | None = None,
) -> ModelCreate a new model for training.
Returns: Model — Model object for training
| Parameter | Type | Default | Description |
|---|---|---|---|
base_model | str | Base model name (e.g., "Qwen/Qwen3.6-35B-A3B-FP8") | |
lora | LoraConfig | LoRA configuration for the run (required). River training is LoRA-only — no base model supports full fine-tuning — so the server rejects create_model without one. | |
tokenizer | str | Any | None | None | Tokenizer name (defaults to base_model) or an already-loaded tokenizer object |
checkpoint | str | Checkpoint | None | None | Optional checkpoint to load after creation. Can be a river:// path string or a Checkpoint object. If a Checkpoint is passed, its step is restored and load_optimizer is set automatically based on checkpoint type. |
timeout | float | 86400.0 | Timeout in seconds |
training_data_attestation | TrainingDataAttestation | str | None | None | Optional server-verified source-artifact manifest. When supplied, the API rejects forward/backward requests if its manifest is missing or no longer belongs to this session. |
Session.sample
Session.sample(
prompts: str | list[str] | None = None,
*,
base_model: str,
checkpoint: str | Checkpoint | None = None,
num_samples: int = 1,
max_tokens: int = 256,
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = -1,
stop: list[str] | None = None,
seed: int | None = None,
seeds: list[int] | None = None,
return_prompt_logprobs: bool = False,
logprobs: int | None = None,
return_expert_routing: bool = False,
images: list[Image] | list[list[Image]] | None = None,
prompt_token_ids: list[int] | list[list[int]] | None = None,
model_input: list[dict] | list[list[dict]] | None = None,
tokenizer: Any | None = None,
metrics_type: str = '',
timeout: float = 86400.0,
) -> list[list[Sample]]Sample from a base model or checkpoint.
When checkpoint is provided, the server loads the LoRA from the
saved checkpoint, generates text, then unloads.
Returns: list[list[Sample]] — list[list[Sample]] — outer list is per-prompt, inner list is per-sample.
| Parameter | Type | Default | Description |
|---|---|---|---|
prompts | str | list[str] | None | None | Single prompt string or list of prompts. Mutually exclusive with prompt_token_ids. |
base_model | str | Base model name (e.g. "Qwen/Qwen3.6-35B-A3B-FP8"). | |
checkpoint | str | Checkpoint | None | None | Optional river:// path or Checkpoint object. If provided, samples from that checkpoint's LoRA weights. |
num_samples | int | 1 | Number of independent samples per prompt. |
max_tokens | int | 256 | Maximum tokens to generate per sample. |
temperature | float | 1.0 | Sampling temperature. |
top_p | float | 1.0 | Nucleus sampling threshold. |
top_k | int | -1 | Top-k sampling (-1 = disabled). |
stop | list[str] | None | None | Stop sequences. |
seed | int | None | None | Random seed (varied per sample automatically). |
seeds | list[int] | None | None | Exact per-prompt/sample seeds, mutually exclusive with seed. |
return_prompt_logprobs | bool | False | Whether to return prompt token logprobs. |
logprobs | int | None | None | If set to K > 0, request the top-K alternative logprobs at each position. Off by default — enabling it roughly halves server throughput. |
return_expert_routing | bool | False | |
images | list[Image] | list[list[Image]] | None | None | Optional raw image bytes for multimodal sampling. See sample for the per-prompt vs. broadcast semantics. |
prompt_token_ids | list[int] | list[list[int]] | None | None | Pre-tokenized prompt(s); mutually exclusive with prompts. See sample for details. |
model_input | list[dict] | list[list[dict]] | None | None | Training-style chunk list(s); mutually exclusive with prompts / prompt_token_ids / images. See sample for details. |
tokenizer | Any | None | None | Optional already-loaded tokenizer. Passing this avoids repeated Hugging Face cache/network resolution in tight loops. |
metrics_type | str | '' | Opaque server-interpreted token enabling extra scalar metrics on the response. Unrecognized values are silently ignored; when recognized, per-result metrics appear on Sample.metrics. |
timeout | float | 86400.0 | Timeout in seconds. |
Session.submit_sampling_batch
Session.submit_sampling_batch(
prompts: str | list[str] | None = None,
*,
base_model: str,
checkpoint: str | Checkpoint | None = None,
num_samples: int = 1,
max_tokens: int = 256,
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = -1,
stop: list[str] | None = None,
seed: int | None = None,
seeds: list[int] | None = None,
return_prompt_logprobs: bool = False,
return_prompt_token_ids: bool = False,
logprobs: int | None = None,
return_expert_routing: bool = False,
images: list[Image] | list[list[Image]] | None = None,
prompt_token_ids: list[int] | list[list[int]] | None = None,
model_input: list[dict] | list[list[dict]] | None = None,
tokenizer: Any | None = None,
idempotency_key: str | None = None,
metrics_type: str = '',
timeout: float = 86400.0,
) -> PendingSamplingBatchSubmit up to 128 independent samples in one RPC.
Uses the new sampling transport; requires server support. Consume
pending.as_completed() to receive individual successes/failures.
Other arguments match sample. Existing submit_sample retains
its whole-batch future behavior.
| Parameter | Type | Default |
|---|---|---|
prompts | str | list[str] | None | None |
base_model | str | |
checkpoint | str | Checkpoint | None | None |
num_samples | int | 1 |
max_tokens | int | 256 |
temperature | float | 1.0 |
top_p | float | 1.0 |
top_k | int | -1 |
stop | list[str] | None | None |
seed | int | None | None |
seeds | list[int] | None | None |
return_prompt_logprobs | bool | False |
return_prompt_token_ids | bool | False |
logprobs | int | None | None |
return_expert_routing | bool | False |
images | list[Image] | list[list[Image]] | None | None |
prompt_token_ids | list[int] | list[list[int]] | None | None |
model_input | list[dict] | list[list[dict]] | None | None |
tokenizer | Any | None | None |
idempotency_key | str | None | None |
metrics_type | str | '' |
timeout | float | 86400.0 |
Session.submit_sample
Session.submit_sample(
prompts: str | list[str] | None = None,
*,
base_model: str,
checkpoint: str | Checkpoint | None = None,
num_samples: int = 1,
max_tokens: int = 256,
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = -1,
stop: list[str] | None = None,
seed: int | None = None,
seeds: list[int] | None = None,
return_prompt_logprobs: bool = False,
logprobs: int | None = None,
return_expert_routing: bool = False,
images: list[Image] | list[list[Image]] | None = None,
prompt_token_ids: list[int] | list[list[int]] | None = None,
model_input: list[dict] | list[list[dict]] | None = None,
tokenizer: Any | None = None,
metrics_type: str = '',
timeout: float = 86400.0,
) -> PendingSampleSubmit sampling from a base model or checkpoint without waiting.
See sample for the full kwarg reference.
| Parameter | Type | Default |
|---|---|---|
prompts | str | list[str] | None | None |
base_model | str | |
checkpoint | str | Checkpoint | None | None |
num_samples | int | 1 |
max_tokens | int | 256 |
temperature | float | 1.0 |
top_p | float | 1.0 |
top_k | int | -1 |
stop | list[str] | None | None |
seed | int | None | None |
seeds | list[int] | None | None |
return_prompt_logprobs | bool | False |
logprobs | int | None | None |
return_expert_routing | bool | False |
images | list[Image] | list[list[Image]] | None | None |
prompt_token_ids | list[int] | list[list[int]] | None | None |
model_input | list[dict] | list[list[dict]] | None | None |
tokenizer | Any | None | None |
metrics_type | str | '' |
timeout | float | 86400.0 |
Session.session_id
Session.session_id: strSessionContext
Context manager for Session with auto-heartbeat.
The context manager returned by Client.session.
Model
A training model with mutable in-memory weights.
Created by Session.create_model.
Methods: get_server_capabilities, get_policy_version, forward, forward_backward, optim_step, train_step, submit_forward_backward, submit_optim_step, submit_train_step, sample, submit_sampling_batch, submit_sample, chat_complete, save_weights, load_weights
Model.get_server_capabilities
Model.get_server_capabilities() -> ServerCapabilitiesRead the endpoint's public protocol support (used by RL preflight).
Model.get_policy_version
Model.get_policy_version() -> PolicyVersion | NoneRead the last server-committed policy; legacy servers return None.
Model.forward
Model.forward(
data: list[dict],
loss_fn: str = 'cross_entropy',
timeout: float = 86400.0,
expected_policy_id: str | None = None,
**loss_config: float,
) -> ForwardResultForward pass only (compute loss, no gradients).
Returns: ForwardResult — ForwardResult with metrics
| Parameter | Type | Default | Description |
|---|---|---|---|
data | list[dict] | List of training samples, each with "input_ids" and "labels" | |
loss_fn | str | 'cross_entropy' | Loss function name |
timeout | float | 86400.0 | Timeout in seconds |
expected_policy_id | str | None | None | |
loss_config | float |
Model.forward_backward
Model.forward_backward(
data: list[dict],
loss_fn: str = 'cross_entropy',
timeout: float = 86400.0,
return_logprobs: bool = False,
zero_out: bool = True,
compute_expert_flip_metric: bool = False,
force_routing_replay: bool = False,
expected_policy_id: str | None = None,
**loss_config: float,
) -> ForwardResultForward + backward pass (compute gradients).
Returns: ForwardResult — ForwardResult with metrics and logprobs when returned by the worker.
Requests larger than the 1 GiB upload limit are split at datum
boundaries into sequential gradient-accumulation sub-batches. The
first sub-batch honors zero_out; later sub-batches retain its
gradients. A failed sub-batch stops the sequence and this method does
not submit an optimizer step. timeout applies to each submitted
sub-batch, so the total operation can take up to the number of chunks
times timeout. Split requests return only metrics the client can
combine exactly across sub-batches; per-sub-batch-only metrics are
omitted.
| Parameter | Type | Default | Description |
|---|---|---|---|
data | list[dict] | List of training samples, each with "input_ids" and "labels" | |
loss_fn | str | 'cross_entropy' | Loss function name |
timeout | float | 86400.0 | Timeout in seconds |
return_logprobs | bool | False | Deprecated no-op. Training losses return per-token logprobs when the worker includes them in the result; this argument is accepted for older callers but is not sent to the server as a loss configuration key. |
zero_out | bool | True | When gradient_accumulation is enabled, clear existing gradients before this call. Use True for the first micro-batch and False for subsequent micro-batches. |
compute_expert_flip_metric | bool | False | When True, compares sampled expert routing against the training-time routing per token and MoE layer, then emits one scalar into ForwardResult.metrics: * expert_flip/per_token_expert_rate ∈ [0, 1] — the fraction of individual top-k expert slots that differ, (top_k − |intersection|) / top_k over (token, layer). Independent of force_routing_replay. Requires per-datum routing keys from Sample.routing_datum_keys(required=True). |
force_routing_replay | bool | False | When true, replay the sampled expert selection while recomputing routing weights at those experts with the trainer's live gate. Every datum in data must include the keys returned by Sample.routing_datum_keys(required=True). |
expected_policy_id | str | None | None | |
loss_config | float |
Model.optim_step
Model.optim_step(
lr: float,
beta1: float = 0.9,
beta2: float = 0.999,
eps: float = 1e-08,
weight_decay: float = 0.0,
grad_clip_norm: float | None = None,
timeout: float = 86400.0,
expected_policy_id: str | None = None,
idempotency_key: str | None = None,
gradient_scale: float | None = None,
) -> OptimStepResultApply gradients with Adam optimizer.
Returns: OptimStepResult — OptimStepResult with metrics (step, lr, grad_norm, grad_norm_finite)
| Parameter | Type | Default | Description |
|---|---|---|---|
lr | float | Learning rate | |
beta1 | float | 0.9 | Adam beta1 |
beta2 | float | 0.999 | Adam beta2 |
eps | float | 1e-08 | Adam epsilon |
weight_decay | float | 0.0 | Weight decay |
grad_clip_norm | float | None | None | Gradient clipping norm (None to disable) |
timeout | float | 86400.0 | Timeout in seconds |
expected_policy_id | str | None | None | |
idempotency_key | str | None | None | |
gradient_scale | float | None | None | Multiply accumulated gradients before clipping/Adam. |
Model.train_step
Model.train_step(
data: list[dict],
lr: float,
*,
loss_fn: str = 'cross_entropy',
beta1: float = 0.9,
beta2: float = 0.999,
eps: float = 1e-08,
weight_decay: float = 0.0,
grad_clip_norm: float | None = None,
compute_expert_flip_metric: bool = False,
force_routing_replay: bool = False,
timeout: float = 86400.0,
**loss_config: float,
) -> tuple[ForwardResult, OptimStepResult]Complete training step: forward+backward plus optimizer update.
Submits forward+backward and the optimizer step back-to-back —
the server runs them in order as one pipelined unit, without a
client round trip in between — then waits for both results. See
forward_backward and optim_step for the full
parameter reference.
The error path differs from calling forward_backward() then
optim_step(): the optimizer step is already submitted when
forward-backward resolves, so if forward-backward fails, this
call raises its error while the optimizer step still runs
server-side and Model.step has already advanced. Callers
that need to inspect both outcomes should use
submit_train_step.
A train step is a complete step: gradients are always cleared
first. For micro-batch gradient accumulation, use
submit_forward_backward(zero_out=...) and
submit_optim_step directly.
For a request larger than the 1 GiB upload limit, the client waits for each accumulated forward/backward sub-batch before submitting the optimizer step. Smaller requests retain the normal pipelined path. The forward/backward result from a split request includes only metrics the client can combine exactly across its sub-batches.
Returns: tuple[ForwardResult, OptimStepResult] — Tuple of (ForwardResult, OptimStepResult).
| Parameter | Type | Default | Description |
|---|---|---|---|
data | list[dict] | List of training samples, each with "input_ids" and "labels" | |
lr | float | Learning rate | |
loss_fn | str | 'cross_entropy' | |
beta1 | float | 0.9 | |
beta2 | float | 0.999 | |
eps | float | 1e-08 | |
weight_decay | float | 0.0 | |
grad_clip_norm | float | None | None | |
compute_expert_flip_metric | bool | False | |
force_routing_replay | bool | False | |
timeout | float | 86400.0 | |
loss_config | float |
Model.submit_forward_backward
Model.submit_forward_backward(
data: list[dict],
loss_fn: str = 'cross_entropy',
timeout: float = 86400.0,
return_logprobs: bool = False,
zero_out: bool = True,
compute_expert_flip_metric: bool = False,
force_routing_replay: bool = False,
expected_policy_id: str | None = None,
**loss_config: float,
) -> PendingOpSubmit forward+backward without blocking. Returns a PendingOp.
Call pending.result() later to get the ForwardResult.
This enables pipelining: submit step N+1 while step N is still running.
See forward_backward for the full kwarg reference.
When submitting multiple micro-batches before submit_optim_step, use
zero_out=True for the first one and False for later submissions.
timeout bounds the submit RPC and, separately, the wait inside
pending.result().
Requests larger than the 1 GiB upload limit must use the synchronous
forward_backward method so the client can keep later
sub-batches and any optimizer step behind successful earlier ones.
| Parameter | Type | Default |
|---|---|---|
data | list[dict] | |
loss_fn | str | 'cross_entropy' |
timeout | float | 86400.0 |
return_logprobs | bool | False |
zero_out | bool | True |
compute_expert_flip_metric | bool | False |
force_routing_replay | bool | False |
expected_policy_id | str | None | None |
loss_config | float |
Model.submit_optim_step
Model.submit_optim_step(
lr: float,
beta1: float = 0.9,
beta2: float = 0.999,
eps: float = 1e-08,
weight_decay: float = 0.0,
grad_clip_norm: float | None = None,
timeout: float = 86400.0,
expected_policy_id: str | None = None,
idempotency_key: str | None = None,
gradient_scale: float | None = None,
) -> PendingOpSubmit optimizer step without blocking. Returns a PendingOp.
Call pending.result() later to get the OptimStepResult.
Model.step advances at submit time, even if the operation
later fails.
| Parameter | Type | Default |
|---|---|---|
lr | float | |
beta1 | float | 0.9 |
beta2 | float | 0.999 |
eps | float | 1e-08 |
weight_decay | float | 0.0 |
grad_clip_norm | float | None | None |
timeout | float | 86400.0 |
expected_policy_id | str | None | None |
idempotency_key | str | None | None |
gradient_scale | float | None | None |
Model.submit_train_step
Model.submit_train_step(
data: list[dict],
lr: float,
*,
loss_fn: str = 'cross_entropy',
beta1: float = 0.9,
beta2: float = 0.999,
eps: float = 1e-08,
weight_decay: float = 0.0,
grad_clip_norm: float | None = None,
compute_expert_flip_metric: bool = False,
force_routing_replay: bool = False,
timeout: float = 86400.0,
**loss_config: float,
) -> tuple[PendingOp, PendingOp]Submit a complete training step without blocking.
Fires forward+backward and the optimizer step back-to-back; the
server runs them in submission order per model. Returns the two
PendingOps as (forward_backward, optim_step).
Because both are submitted up front, a failed forward-backward
does not cancel the already-submitted optimizer step.
Model.step advances at submit time. See
train_step for the full kwarg reference.
Requests larger than the 1 GiB upload limit must use
train_step, which waits for all accumulated sub-batches before
submitting the optimizer step.
| Parameter | Type | Default |
|---|---|---|
data | list[dict] | |
lr | float | |
loss_fn | str | 'cross_entropy' |
beta1 | float | 0.9 |
beta2 | float | 0.999 |
eps | float | 1e-08 |
weight_decay | float | 0.0 |
grad_clip_norm | float | None | None |
compute_expert_flip_metric | bool | False |
force_routing_replay | bool | False |
timeout | float | 86400.0 |
loss_config | float |
Model.sample
Model.sample(
prompts: str | list[str] | None = None,
*,
num_samples: int = 1,
max_tokens: int = 256,
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = -1,
stop: list[str] | None = None,
seed: int | None = None,
seeds: list[int] | None = None,
return_prompt_logprobs: bool = False,
logprobs: int | None = None,
images: list[Image] | list[list[Image]] | None = None,
return_expert_routing: bool = False,
prompt_token_ids: list[int] | list[list[int]] | None = None,
model_input: list[dict] | list[list[dict]] | None = None,
metrics_type: str = '',
timeout: float = 86400.0,
poll_interval: float = 1.0,
) -> list[list[Sample]]Sample from the current in-memory training weights.
Generates text using the model's current weights, without needing to save a checkpoint first. Per-token logprobs are always returned.
Returns: list[list[Sample]] — list[list[Sample]] — outer list is per-prompt, inner list is per-sample. Each Sample has .tokens, .text, .logprobs, and .stop_reason.
| Parameter | Type | Default | Description |
|---|---|---|---|
prompts | str | list[str] | None | None | Single prompt string or list of prompts. Mutually exclusive with prompt_token_ids. |
num_samples | int | 1 | Number of independent samples per prompt. |
max_tokens | int | 256 | Maximum tokens to generate per sample. |
temperature | float | 1.0 | Sampling temperature. |
top_p | float | 1.0 | Nucleus sampling threshold. |
top_k | int | -1 | Top-k sampling (-1 = disabled). |
stop | list[str] | None | None | Stop sequences. |
seed | int | None | None | Random seed (varied per sample automatically). |
seeds | list[int] | None | None | Exact per-prompt/sample seeds. Mutually exclusive with seed and ordered by prompt then sample. |
return_prompt_logprobs | bool | False | Whether to return prompt token logprobs. |
logprobs | int | None | None | If set to K > 0, request the top-K alternative logprobs at each output position (and, when return_prompt_logprobs=True, at each prompt position). Off by default — enabling it roughly halves server throughput for small serialization gain, so it's opt-in. |
images | list[Image] | list[list[Image]] | None | None | Optional raw image bytes (PNG / JPEG) for multimodal sampling. Accepts list[Image] (broadcast the same image set to every prompt) or list[list[Image]] (per-prompt explicit). Bytes are sent to the inference backend as image data. Most ergonomic source: **Qwen35VLRenderer.build_sample_prompt(messages).to_kwargs(), which emits {"prompt", "images"}. The image format is inferred from the bytes' magic header, so no separate format hint is sent over the wire. |
return_expert_routing | bool | False | Capture per-token MoE expert routing during this sampling call. When available, each Sample exposes an .expert_routing object that can be round-tripped into training data with sample.routing_datum_keys(required=True) before calling forward_backward(force_routing_replay=...) or forward_backward(compute_expert_flip_metric=True). |
prompt_token_ids | list[int] | list[list[int]] | None | None | Pre-tokenized prompt(s) — a flat list[int] (one prompt) or list[list[int]] (one entry per prompt). Mutually exclusive with prompts. Ids are forwarded verbatim for sampling, bypassing server-side tokenization, so the sampled continuation is conditioned on exactly these ids (no training/sampling tokenization skew). Ids must be valid for this model's vocabulary. May be combined with images using the same single-placeholder convention as text prompts: one un-expanded <|image_pad|>-style token id per image, in images order — the placeholder count must match the image count exactly (a surplus of images is silently dropped by the backend otherwise). |
model_input | list[dict] | list[list[dict]] | None | None | Training-style chunk list(s) — the same [{"type": "text", "tokens": [...]}, {"type": "image", "data": bytes, ...}, ...] shape forward_backward accepts, for one prompt (list[dict]) or a batch (list[list[dict]]). Lowered client-side to prompt_token_ids + images (each image chunk becomes one un-expanded placeholder token). Mutually exclusive with prompts / prompt_token_ids / images. expected_tokens and format on image chunks are accepted and ignored; to validate the backend's image expansion against expected_tokens, pass return_prompt_logprobs=True and count placeholder ids in the echoed Sample.prompt_token_ids. |
metrics_type | str | '' | Opaque server-interpreted token enabling extra scalar metrics on the response. Unrecognized values are silently ignored; when recognized, per-result metrics appear on Sample.metrics. |
timeout | float | 86400.0 | Timeout in seconds for the entire operation (includes server-side wait for LoRA slot availability). |
poll_interval | float | 1.0 | Seconds between completion polls once the request is in flight. The default (1s) suits ad-hoc sampling; tight RL loops that immediately consume the results can lower it to shave the post-completion notice lag off every step. |
Model.submit_sampling_batch
Model.submit_sampling_batch(
prompts: str | list[str] | None = None,
*,
num_samples: int = 1,
max_tokens: int = 256,
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = -1,
stop: list[str] | None = None,
seed: int | None = None,
seeds: list[int] | None = None,
return_prompt_logprobs: bool = False,
return_prompt_token_ids: bool = False,
logprobs: int | None = None,
images: list[Image] | list[list[Image]] | None = None,
return_expert_routing: bool = False,
prompt_token_ids: list[int] | list[list[int]] | None = None,
model_input: list[dict] | list[list[dict]] | None = None,
idempotency_key: str | None = None,
pinned_policy_id: str | None = None,
policy_selection: str = 'ordered',
retained_kv_groups: list[str] | None = None,
kv_cache_policy: dict | None = None,
metrics_type: str = '',
timeout: float = 86400.0,
poll_interval: float = 0.1,
) -> PendingSamplingBatchSubmit up to 128 independent samples in one RPC.
Uses the new sampling transport; requires server support. Consume
pending.as_completed() to receive individual successes/failures.
policy_selection="ordered" preserves training-queue ordering.
policy_selection="latest_snapshot" samples the newest available
immutable committed snapshot without waiting for queued training work.
Publish one with save_weights(mode="inference", immutable=True) first;
missing snapshots fail explicitly. This mode may lag the committed head.
pinned_policy_id continues a previously returned policy snapshot and
cannot be combined with latest_snapshot.
Alternatively, retained_kv_groups supplies one trajectory UUID per
independent sample (including the num_samples expansion) and permits
earlier-policy KV within kv_cache_policy: max_staleness (updates),
anchor_policy_id, previous_policy_id, and refill_on_limit (bool).
Continuations echo the prior result's cache origin and sampling policy.
Check Sample.retained_kv for worker acknowledgement.
return_prompt_token_ids echoes canonical IDs without requesting
prompt logprobs, preserving prefix-cache eligibility.
Other arguments match sample. Existing submit_sample retains
its whole-batch future behavior.
| Parameter | Type | Default |
|---|---|---|
prompts | str | list[str] | None | None |
num_samples | int | 1 |
max_tokens | int | 256 |
temperature | float | 1.0 |
top_p | float | 1.0 |
top_k | int | -1 |
stop | list[str] | None | None |
seed | int | None | None |
seeds | list[int] | None | None |
return_prompt_logprobs | bool | False |
return_prompt_token_ids | bool | False |
logprobs | int | None | None |
images | list[Image] | list[list[Image]] | None | None |
return_expert_routing | bool | False |
prompt_token_ids | list[int] | list[list[int]] | None | None |
model_input | list[dict] | list[list[dict]] | None | None |
idempotency_key | str | None | None |
pinned_policy_id | str | None | None |
policy_selection | str | 'ordered' |
retained_kv_groups | list[str] | None | None |
kv_cache_policy | dict | None | None |
metrics_type | str | '' |
timeout | float | 86400.0 |
poll_interval | float | 0.1 |
Model.submit_sample
Model.submit_sample(
prompts: str | list[str] | None = None,
*,
num_samples: int = 1,
max_tokens: int = 256,
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = -1,
stop: list[str] | None = None,
seed: int | None = None,
seeds: list[int] | None = None,
return_prompt_logprobs: bool = False,
logprobs: int | None = None,
images: list[Image] | list[list[Image]] | None = None,
return_expert_routing: bool = False,
prompt_token_ids: list[int] | list[list[int]] | None = None,
model_input: list[dict] | list[list[dict]] | None = None,
metrics_type: str = '',
timeout: float = 86400.0,
poll_interval: float = 1.0,
) -> PendingSampleSubmit sampling from current training weights without waiting.
See sample for the full kwarg reference.
| Parameter | Type | Default |
|---|---|---|
prompts | str | list[str] | None | None |
num_samples | int | 1 |
max_tokens | int | 256 |
temperature | float | 1.0 |
top_p | float | 1.0 |
top_k | int | -1 |
stop | list[str] | None | None |
seed | int | None | None |
seeds | list[int] | None | None |
return_prompt_logprobs | bool | False |
logprobs | int | None | None |
images | list[Image] | list[list[Image]] | None | None |
return_expert_routing | bool | False |
prompt_token_ids | list[int] | list[list[int]] | None | None |
model_input | list[dict] | list[list[dict]] | None | None |
metrics_type | str | '' |
timeout | float | 86400.0 |
poll_interval | float | 1.0 |
Model.chat_complete
Model.chat_complete(
messages: list[dict],
*,
timeout: float | None = None,
**kwargs,
) -> ChatCompleteResultOpenAI-compatible chat completion from current training weights.
Like model.sample() but using the OpenAI chat-completions format
instead of raw prompts.
Returns: ChatCompleteResult — ChatCompleteResult with response_json (full OpenAI-format JSON string) and status_code.
| Parameter | Type | Default | Description |
|---|---|---|---|
messages | list[dict] | OpenAI-format messages list (e.g. [{"role": "user", "content": "Hello"}]). | |
timeout | float | None | None | Timeout in seconds. |
kwargs |
Model.save_weights
Model.save_weights(
name: str,
mode: str = 'training',
timeout: float = 86400.0,
ttl: datetime.timedelta | None = None,
immutable: bool = False,
expected_policy_id: str | None = None,
) -> CheckpointSave a checkpoint of the current model weights.
Returns: Checkpoint — Checkpoint object with river:// path.
| Parameter | Type | Default | Description |
|---|---|---|---|
name | str | Checkpoint name (e.g. "final" or "step_000100"). | |
mode | str | 'training' | "training" saves optimizer state (for training continuation), "inference" saves PEFT format only (for sampling/inference). |
timeout | float | 86400.0 | Timeout in seconds. |
ttl | datetime.timedelta | None | None | Lifetime before the checkpoint is reaped. Applies to explicit user-saved checkpoints in both modes; when omitted, the server default is 1 year. 1 year is also the maximum — the server rejects a longer ttl. |
immutable | bool | False | |
expected_policy_id | str | None | None |
Model.load_weights
Model.load_weights(
checkpoint: str | Checkpoint,
load_optimizer: bool = True,
timeout: float = 86400.0,
) -> NoneLoad weights from a checkpoint.
| Parameter | Type | Default | Description |
|---|---|---|---|
checkpoint | str | Checkpoint | A river:// path string or a Checkpoint object. If a Checkpoint is passed, its step is restored on the model. | |
load_optimizer | bool | True | Whether to load optimizer state. |
timeout | float | 86400.0 | Timeout in seconds. |
Model.model_id
Model.model_id: strModel.base_model
Model.base_model: strModel.training_run_id
Model.training_run_id: strModel.step
Model.step: intCurrent training step.
Advances when an optimizer step is submitted (so pipelined submissions observe a consistent value), not when it resolves — a failed or timed-out optimizer step leaves it advanced.
Data types
Values returned by the client and the configuration objects you pass to it.
LoraConfig
LoRA adapter configuration.
| Field | Type | Default |
|---|---|---|
rank | int | 16 |
train_attn | bool | True |
train_mlp | bool | True |
train_unembed | bool | False |
seed | int | None | None |
Sample
A generated sample.
tokens contains generated token IDs when the response includes them;
older responses may fall back to retokenizing text client-side.
prompt_token_ids / top_logprobs / prompt_top_logprobs are
None when the server did not provide them or the feature was not
requested.
expert_routing is populated when return_expert_routing=True
was passed on the request.
metrics carries per-result scalar metrics; empty unless the
originating request passed a recognized metrics_type token.
| Field | Type | Default |
|---|---|---|
tokens | list[int] | |
text | str | |
logprobs | list[float] | |
stop_reason | str | |
model_step | int | |
prompt_logprobs | list[float] | None | None |
request_id | str | '' |
prompt_token_ids | list[int] | None | None |
top_logprobs | list[list[TopLogprob]] | None | None |
prompt_top_logprobs | list[list[TopLogprob]] | None | None |
expert_routing | ExpertRouting | None | None |
metrics | dict[str, float] | {} |
token_data_is_exact | bool | True |
policy_version | PolicyVersion | None | None |
kv_cache_policy_version | PolicyVersion | None | None |
retained_kv | bool | False |
cached_prompt_tokens | int | 0 |
prompt_tokens | int | 0 |
Sample.routing_datum_keys
Sample.routing_datum_keys(*, required: bool = False) -> dict[str, bytes | str]Return the per-datum keys to splat into a forward_backward
datum when enabling force_routing_replay or
compute_expert_flip_metric.
Returns an empty dict when this sample carries no routing capture and
required is false. Raises ValueError when required is true
and no routing handle is available.
Example:
data = []
for sample in samples:
datum = {
"input_ids": ...,
"advantages": ...,
**sample.routing_datum_keys(required=True),
}
data.append(datum)| Parameter | Type | Default |
|---|---|---|
required | bool | False |
ChatCompleteResult
Result of a chat completion request.
| Field | Type |
|---|---|
response_json | str |
status_code | int |
ForwardResult
Result of forward or forward_backward pass.
| Field | Type | Default |
|---|---|---|
metrics | dict[str, float] | |
logprobs | list | None | None |
OptimStepResult
Result of an optimizer step.
| Field | Type | Default |
|---|---|---|
metrics | dict[str, float] | |
policy_version | PolicyVersion | None | None |
Checkpoint
A saved model checkpoint.
| Field | Type | Default |
|---|---|---|
path | str | |
step | int | |
checkpoint_type | str | |
policy_version | PolicyVersion | None | None |
TopLogprob
One top-K candidate token at a single position.
| Field | Type | Default |
|---|---|---|
logprob | float | |
token_id | int | |
token | str | '' |
PendingOp
A submitted but not-yet-resolved async operation.
Returned by Model.submit_forward_backward() and
Model.submit_optim_step(). Call .result() to block until complete.
| Field | Type |
|---|---|
request_id | str |
PendingOp.result
PendingOp.result(
*,
before_poll: Callable[[], None] | None = None,
) -> ForwardResult | OptimStepResultBlock until the operation completes and return the result.
| Parameter | Type | Default |
|---|---|---|
before_poll | Callable[[], None] | None | None |
PendingSample
A submitted but not-yet-resolved sampling operation.
| Field | Type |
|---|---|
request_id | str |
PendingSample.result
PendingSample.result(
*,
retry_connection_errors: bool = True,
before_poll: Callable[[], None] | None = None,
) -> list[list[Sample]]Block until sampling completes and return grouped samples.
retry_connection_errors=False exposes the first RetrieveFuture
transport failure to callers that own a sealed retry policy.
| Parameter | Type | Default |
|---|---|---|
retry_connection_errors | bool | True |
before_poll | Callable[[], None] | None | None |
PendingSample.result_async
async PendingSample.result_async(
*,
executor: Executor | None = None,
retry_connection_errors: bool = True,
before_poll: Callable[[], None] | None = None,
) -> list[list[Sample]]Wait with bounded RPC workers, releasing them between polls.
Pass a shared executor to bound simultaneous RPCs and result decoding independently from the number of outstanding sampling operations.
| Parameter | Type | Default |
|---|---|---|
executor | Executor | None | None |
retry_connection_errors | bool | True |
before_poll | Callable[[], None] | None | None |
ExpertRouting
Captured MoE expert routing for one sample.
Populated when the caller passed return_expert_routing=True on
the sample request and routing capture was available.
To enable router replay / flip-rate metrics, splat the canonical
per-datum keys via Sample.routing_datum_keys() rather than
building them by hand:
datum = {..., **sample.routing_datum_keys(required=True)}Current servers return routing captures as an opaque handle; the
splat yields expert_routing_handle. topk_ids is retained for
inspection and proto round-trip compatibility only; replay always
recomputes the routing weights in the trainer.
The shape header fields (num_decoder_layers / top_k /
layer_indices) are informational only and let the user inspect
capture metadata.
Byte layout (when ids are present): [num_tokens, num_decoder_layers, top_k] row-major int16 little-endian. num_tokens = seqlen - 1
because routing capture omits the trailing position.
| Field | Type | Default |
|---|---|---|
topk_ids | bytes | b'' |
num_tokens | int | 0 |
num_decoder_layers | int | 0 |
top_k | int | 0 |
layer_indices | list[int] | [] |
handle | str | '' |
Deployment
A dedicated streaming deployment: one checkpoint on its own engines.
Dedicated deployments are gated and disabled by default. Contact River to enable access for your team and base model before creating capacity.
base_url binds the checkpoint when used with the standard OpenAI
client, so its existing model argument can stay unchanged. model is
the deployment identity (equal to id), also accepted by global /v1.
replicas is keyed by role: unified for unified serving, prefill
and decode for disaggregated serving. phase is one of
accepted, provisioning, ready, degraded, unavailable,
scaled_to_zero, failed, deleting, deleted.
| Field | Type | Default |
|---|---|---|
id | str | |
model | str | |
base_url | str | |
checkpoint | str | |
base_model | str | |
topology | str | |
replicas | dict[str, DeploymentReplicas] | |
phase | str | |
created_at | str | |
updated_at | str | |
phase_reason | str | None | None |
operation_id | str | None | None |
deleted_at | str | None | None |
is_serving | bool | |
is_terminal | bool |
DeploymentReplicas
Desired, allocated, ready and pending counts for one role of a deployment.
pending are allocated replicas the cell cannot place yet (no capacity).
| Field | Type | Default |
|---|---|---|
desired | int | |
allocated | int | |
ready | int | |
pending | int | 0 |
ServerCapabilities
Supported public protocols and model names, without server internals.
| Field | Type | Default |
|---|---|---|
supported_models | tuple[str, ...] | |
features | frozenset[str] | |
model_features | dict[str, frozenset[str]] | {} |
ServerCapabilities.require_model
ServerCapabilities.require_model(model: str, *features: str) -> None| Parameter | Type |
|---|---|
model | str |
features | str |
ServerCapabilities.require
ServerCapabilities.require(*features: str) -> None| Parameter | Type |
|---|---|
features | str |
PolicyVersion
Server identity of a committed logical weight state (not a byte hash).
| Field | Type | Default |
|---|---|---|
id | str | |
lineage_id | str | |
step | int | |
parent_id | str | None | None |
PolicyVersion.from_proto
PolicyVersion.from_proto(value)| Parameter |
|---|
value |
ImageHandle
Immutable metadata for an image owned by one live API session.
Checkpoints retain bytes in a client-side ImageStore and re-upload on resume.
| Field | Type |
|---|---|
id | str |
sha256 | str |
byte_count | int |
width | int |
height | int |
session_id | str |
ImageStore
Content-addressed image bytes on a local/shared durable filesystem.
| Field |
|---|
directory |
ImageStore.put
ImageStore.put(data)| Parameter |
|---|
data |
ImageStore.read
ImageStore.read(image)| Parameter |
|---|
image |
ImageStore.copy_images
ImageStore.copy_images(value, source)| Parameter |
|---|
value |
source |
ImageStore.prune
ImageStore.prune(retained)| Parameter |
|---|
retained |
PendingSamplingBatch
Independent operations accepted together; results arrive in completion order.
as_completed reports failures per sample. Cancelling local waiting does
not cancel accepted server operations. Retain request_ids to retrieve them.
| Field | Type |
|---|---|
request_ids | tuple[str, ...] |
PendingSamplingBatch.as_completed
async PendingSamplingBatch.as_completed(
*,
executor: Executor | None = None,
collector: SamplingResultCollector | None = None,
) -> AsyncIterator[SamplingCompletion]Consume independently; share a collector across concurrent batches.
| Parameter | Type | Default |
|---|---|---|
executor | Executor | None | None |
collector | SamplingResultCollector | None | None |
SamplingCompletion
One sample or failure, identified in the original prompt/sample grid.
| Field | Type | Default |
|---|---|---|
prompt_index | int | |
sample_index | int | |
request_id | str | |
sample | Sample | None | None |
error | Exception | None | None |
SamplingResultCollector
Coalesce result snapshots across batches on one event loop.
At most max_poll_rpcs short streams run concurrently, each querying up to 128 handles. Backpressure permits one decoded result buffered per stream. Executors are caller-owned; no thread is occupied between snapshots.
| Field |
|---|
executor |
max_poll_rpcs |
snapshot_rpcs |
polled_handles |
SamplingResultCollector.register
SamplingResultCollector.register(batch: PendingSamplingBatch) -> list[asyncio.Future]| Parameter | Type |
|---|---|
batch | PendingSamplingBatch |
SamplingResultCollector.wake
SamplingResultCollector.wake()SamplingResultCollector.aclose
async SamplingResultCollector.aclose()RunMetricsLogger
Buffers and sends run metrics without interrupting training.
RunMetricsLogger.log
RunMetricsLogger.log(metric: str, step: int, value: float) -> NoneQueue one metric point and return without waiting for the network.
| Parameter | Type |
|---|---|
metric | str |
step | int |
value | float |
RunMetricsLogger.flush
RunMetricsLogger.flush() -> NoneRequest prompt delivery of the points currently in the buffer.
RunMetricsLogger.close
RunMetricsLogger.close() -> NoneDrain buffered points when possible, then stop the worker.
TrainingDataArtifact
Source bytes and expected digest for server-side integrity verification.
The API hashes content itself and retains only the resulting manifest.
Callers should verify their source files locally before constructing these
objects, then bind the returned TrainingDataAttestation to the
model they create.
| Field | Type |
|---|---|
name | str |
expected_sha256 | str |
content | bytes |
AttestedTrainingDataArtifact
Server-observed source-artifact digest and size.
| Field | Type |
|---|---|
name | str |
sha256 | str |
size_bytes | int |
TrainingDataAttestation
A server-verified source-artifact manifest bound to one session.
| Field | Type |
|---|---|
training_data_attestation_id | str |
artifacts | list[AttestedTrainingDataArtifact] |
Errors
Every server and transport failure derives from RiverError, so a single except river.RiverError catches all of them. Invalid arguments still raise the standard ValueError and TypeError.
RiverError
Inherits Exception.
Base exception for River client errors.
AuthenticationError
Inherits RiverError.
Authentication failed.
CapacityError
Inherits RiverError.
No capacity available for the operation.
ModelNotFoundError
Inherits RiverError.
Model not found.
RiverConnectionError
Inherits RiverError.
Connection or communication error with the River API server.
This wraps gRPC errors with a more user-friendly message while preserving the original error details for debugging.
| Field | Type | Description |
|---|---|---|
message | Human-readable error message | |
status_code | gRPC status code name (e.g., "UNAVAILABLE", "DEADLINE_EXCEEDED") | |
details | Additional error details from the server | |
original_error | The original gRPC RpcError for debugging | |
error_code | str | None | Structured server reason, such as IMAGE_EXPIRED |
image_id | str | None | Image requiring re-upload, when supplied by the server |
RiverConnectionError.from_grpc_error
RiverConnectionError.from_grpc_error(
error: Exception,
context: str = 'API call',
) -> RiverConnectionErrorCreate a RiverConnectionError from a gRPC RpcError.
| Parameter | Type | Default |
|---|---|---|
error | Exception | |
context | str | 'API call' |
RiverTimeoutError
Inherits RiverError.
Operation timed out while retaining its recoverable future ID.
| Field |
|---|
request_id |
SessionHeartbeatError
Inherits RiverConnectionError.
Session heartbeat was lost or rejected while a session was active.
Tokenizers
Helpers for tokenizing prompts yourself, for example when you build prompt_token_ids for training data.
load_tokenizer
load_tokenizer(
tokenizer: str | Any | None = None,
*,
base_model: str | None = None,
revision: str | None = None,
local_files_only: bool = False,
resolve_aliases: bool = True,
)Load or return a tokenizer for River client result parsing.
base_model remains the public River model name used for routing. When
it is a hyphen-suffixed deployment alias of a known canonical model, this
helper resolves it to the underlying Hugging Face tokenizer id before
loading. local_files_only keeps sealed jobs from reaching Hugging Face
after their tokenizer revision has been frozen. resolve_aliases=False
retains a deployment's own tokenizer source for unpinned compatibility
paths.
| Parameter | Type | Default |
|---|---|---|
tokenizer | str | Any | None | None |
base_model | str | None | None |
revision | str | None | None |
local_files_only | bool | False |
resolve_aliases | bool | True |
resolve_tokenizer_name
resolve_tokenizer_name(model_name: str) -> strReturn the tokenizer for a canonical River model or one of its aliases.
Deployment aliases must append a hyphen-delimited suffix to a canonical model name. Longest-root matching keeps the result deterministic if a future canonical model name extends another root.
| Parameter | Type |
|---|---|
model_name | str |
MODEL_TOKENIZER_ALIASES
MODEL_TOKENIZER_ALIASES: dict[str, str]| Key | Value |
|---|---|
Qwen/Qwen3.6-35B-A3B-FP8 | Qwen/Qwen3.6-35B-A3B |
Qwen/Qwen3.5-397B-A17B-FP8 | Qwen/Qwen3.5-397B-A17B-FP8 |
nvidia/Kimi-K2.6-NVFP4 | nvidia/Kimi-K2.6-NVFP4 |
moonshotai/Kimi-K3 | moonshotai/Kimi-K3 |
nvidia/GLM-5.1-NVFP4 | nvidia/GLM-5.1-NVFP4 |
nvidia/GLM-5.2-NVFP4 | nvidia/GLM-5.2-NVFP4 |
Reinforcement learning
The river_client.rl library manages multi-turn rollouts, rewards, policy staleness, and optimizer updates on top of Client, Session, and Model.
from river_client import rlrl.AsyncTrainer
Train a model from rollout groups while sampling continues.
Computes advantages from completed groups, submits forward/backward work,
and applies optimizer updates. Set max_staleness=0 to train
synchronously or allow bounded policy lag for asynchronous training.
Use one trainer to update a given model.
rl.AsyncTrainer(
*,
engine,
optimizer: Adam,
completion: GroupCompletion,
normalize: str,
groups_per_step: int,
group_size: int = 16,
advantage=None,
loss: str = 'cispo',
max_staleness: int = 2,
stale_policy: str = 'wait',
allow_trajectory_loss: bool = False,
allow_unbounded_staleness: bool = False,
truncation: Truncation = Truncation(),
checkpoint: Checkpointing | None = None,
init_checkpoint=None,
evaluator=None,
run_config=None,
max_importance_weight: float = 5.0,
routing_replay: bool = False,
repeat_dataset: bool = True,
batch_order: str = 'completion',
validate_batch=None,
forward_backward_groups: int | None = None,
forward_backward_batch: ForwardBackwardBatch | Literal['auto'] | None = 'auto',
max_pending_forward_backward: int = 16,
**loss_config,
)| Parameter | Type | Default | Description |
|---|---|---|---|
engine | Rollout engine sharing the training model. | ||
optimizer | Adam | Optimizer settings for each update. | |
completion | GroupCompletion | When a rollout group can be consumed. | |
normalize | str | Loss scaling: "token", "sequence", "batch", or "sum". | |
groups_per_step | int | Reward groups consumed per training batch. | |
group_size | int | 16 | Samples requested for each dataset row. |
advantage | None | Advantage estimator; defaults to GroupCentered(). | |
loss | str | 'cispo' | "cispo", "ppo", "importance_sampling", or "decoupled_ppo". |
max_staleness | int | 2 | Maximum training-policy age in optimizer steps; zero gives synchronous admission. |
stale_policy | str | 'wait' | "wait" holds complete groups within the bound; "mask_span" masks old spans; "keep" ignores the bound. |
allow_trajectory_loss | bool | False | Acknowledge options that may drop or fully mask trajectories. |
allow_unbounded_staleness | bool | False | Acknowledge stale_policy="keep". |
truncation | Truncation | Truncation() | Reward or drop policy for truncated trajectories. |
checkpoint | Checkpointing | None | None | Durable trainer-state checkpoint settings. |
init_checkpoint | None | Optional starting weights checkpoint. | |
evaluator | None | Optional checkpoint-based evaluation runner. | |
run_config | None | Serializable recipe values checked on resume. | |
max_importance_weight | float | 5.0 | Importance-ratio limit for decoupled PPO. |
routing_replay | bool | False | Replay captured expert routing during training. |
repeat_dataset | bool | True | Cycle dataset rows when more groups are needed. |
batch_order | str | 'completion' | Consume groups by "completion" or "admission". |
validate_batch | None | Optional callback before a full batch trains. | |
forward_backward_groups | int | None | None | Submit forward/backward after this many ready groups; requires synchronous group-centered training. |
forward_backward_batch | ForwardBackwardBatch | Literal['auto'] | None | 'auto' | "auto" enables threshold-based submission when eligible; set a ForwardBackwardBatch or None explicitly to override. |
max_pending_forward_backward | int | 16 | Limit concurrent submitted operations. |
loss_config | Loss-specific numeric options such as eps_max. |
rl.AsyncTrainer.progress
rl.AsyncTrainer.progress()Return current batch, model-update, and activity counters.
model_step reports the last completed optimizer update, including
while another update is in flight.
rl.AsyncTrainer.run
async rl.AsyncTrainer.run(dataset, *, steps: int, after_recovery=None)Yield a Step after each completed training batch.
| Parameter | Type | Default | Description |
|---|---|---|---|
dataset | Rows used to create rollout groups. | ||
steps | int | Number of batches to complete. A batch with no usable training data still counts, but does not update the model. | |
after_recovery | None | Optional async callback receiving the number of restored batches before new training starts. It also runs if all requested batches were already complete. Make it idempotent, because recovery may retry it. |
rl.RolloutEngine
Generate multi-turn trajectories from an environment and sampling model.
Coordinates rollout groups, sampling concurrency, and policy/KV-cache
continuity. Use it directly for custom algorithms or pass it to
AsyncTrainer for managed training.
rl.RolloutEngine(
model,
*,
env,
renderer,
budget: Budget = Budget(),
schedule: Schedule = Schedule(),
sampling_policy: str = 'trajectory',
kv_cache: KVCache = KVCache(),
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = -1,
seed: int = 0,
sample_timeout: float = 1800.0,
environment_timeout: float | None = None,
return_expert_routing: bool = False,
)| Parameter | Type | Default | Description |
|---|---|---|---|
model | Training model or checkpoint-pinned sampler. | ||
env | Env instance for stateless tasks, or a factory that makes one instance per trajectory for stateful tasks. A shared instance must manage its own state by trajectory id. | ||
renderer | Model-specific chat renderer that supports continuation. | ||
budget | Budget | Budget() | Per-trajectory token, turn, and image limits. |
schedule | Schedule | Schedule() | Client-side admission and sampling submission limits. |
sampling_policy | str | 'trajectory' | "trajectory" pins one policy for every turn; "segment" allows a newer policy between turns. |
kv_cache | KVCache | KVCache() | Maximum retained KV age and behavior at that limit. |
temperature | float | 1.0 | Sampling temperature; training requires a positive value. |
top_p | float | 1.0 | Nucleus sampling threshold. |
top_k | int | -1 | Top-k sampling limit; -1 disables it. |
seed | int | 0 | Base sampling seed. |
sample_timeout | float | 1800.0 | Timeout for one sampling request in seconds. |
environment_timeout | float | None | None | Optional timeout for an environment call in seconds. |
return_expert_routing | bool | False | Request expert routing for supported models. |
Methods: submit, next_group, accept, acknowledge, rollout, progress, sampling_metrics, state_dict, load_state_dict
rl.RolloutEngine.submit
rl.RolloutEngine.submit(
row,
*,
group_size: int,
completion: GroupCompletion,
id: str | None = None,
priority: int = 0,
)Start a rollout group for one dataset row and return its group id.
Call inside async with engine. Each member gets an independent
sampling seed.
| Parameter | Type | Default |
|---|---|---|
row | ||
group_size | int | |
completion | GroupCompletion | |
id | str | None | None |
priority | int | 0 |
rl.RolloutEngine.next_group
async rl.RolloutEngine.next_group() -> CompletedGroupWait for the next completed group or deadline partial result.
rl.RolloutEngine.accept
rl.RolloutEngine.accept(result)Mark a delivered result as consumed by the caller.
| Parameter |
|---|
result |
rl.RolloutEngine.acknowledge
rl.RolloutEngine.acknowledge(group_id)Release a group only after all of its emitted portions are consumed.
| Parameter |
|---|
group_id |
rl.RolloutEngine.rollout
async rl.RolloutEngine.rollout(rows, *, group_size: int, completion: GroupCompletion)Yield completed trajectory groups; custom algorithms can own updates.
| Parameter | Type |
|---|---|
rows | |
group_size | int |
completion | GroupCompletion |
rl.RolloutEngine.progress
rl.RolloutEngine.progress()Return current rollout and environment activity counters.
Operation seconds include in-flight work and sum across concurrent calls. They must not be added together as elapsed wall-clock time.
rl.RolloutEngine.sampling_metrics
rl.RolloutEngine.sampling_metrics()Cumulative sampling counters and current admission pressure.
rl.RolloutEngine.state_dict
async rl.RolloutEngine.state_dict()Capture rollout state for a checkpoint.
An in-flight sample may replay one segment after restore. Stateful environments are snapshotted at stable boundaries; work inside a tool or reset call cannot be resumed from its partial side effects.
rl.RolloutEngine.load_state_dict
rl.RolloutEngine.load_state_dict(state)Restore saved rollout state into a fresh engine before starting it.
| Parameter |
|---|
state |
rl.Env
Define the prompts, tools, rewards, and turn transitions for a rollout.
Pass a shared instance to RolloutEngine for stateless tasks, or a
factory that creates a private environment for each trajectory. Override
reset and reward; the default on_turn executes declared tools.
| Attribute | Type | Description |
|---|---|---|
tools | tuple[Tool, ...] | list[Tool] | Tools available to the model during each turn. |
recovery | Literal['drop', 'stateless', 'snapshot'] | How unfinished trajectories resume after a checkpoint: "drop" (default), "stateless", or "snapshot". |
Methods: reset, reward, on_turn, on_truncated, snapshot, restore, close
rl.Env.reset
async rl.Env.reset(row) -> list[Message]Return the initial conversation messages for a dataset row.
| Parameter |
|---|
row |
rl.Env.reward
async rl.Env.reward(traj, row) -> floatReturn the final reward for a completed trajectory.
| Parameter |
|---|
traj |
row |
rl.Env.on_turn
async rl.Env.on_turn(traj) -> list[Message] | NoneReturn new environment messages after the latest model turn.
The default executes tool calls and returns their results, or None
when the model has no more tool calls. Overrides should return only
new messages, without re-rendering prior model output. Use
traj.rewrite explicitly if you need to compact the conversation.
| Parameter |
|---|
traj |
rl.Env.on_truncated
async rl.Env.on_truncated(traj, row, cause: str) -> floatOverride for task-specific partial credit. Default is explicit zero.
| Parameter | Type |
|---|---|
traj | |
row | |
cause | str |
rl.Env.snapshot
async rl.Env.snapshot(traj)Capture private environment state for recovery="snapshot".
| Parameter |
|---|
traj |
rl.Env.restore
async rl.Env.restore(traj, state)Restore private state captured by snapshot after a restart.
| Parameter |
|---|
traj |
state |
rl.Env.close
async rl.Env.close()Release resources owned by a per-trajectory environment factory.
Rollout configuration
Limits and policies passed to rl.RolloutEngine and rl.AsyncTrainer.
rl.Budget
Limits for one trajectory.
rl.Budget(
max_turns: int = 8,
max_generated_tokens: int = 32768,
max_context_tokens: int = 131072,
segment_tokens: int = 4096,
final_answer_reserve: int = 0,
tool_output_tokens: int = 4096,
max_images: int | None = None,
max_turn_tokens: int | None = None,
) -> None| Field | Type | Default | Description |
|---|---|---|---|
max_turns | int | 8 | Maximum assistant turns. |
max_generated_tokens | int | 32768 | Total generated-token budget across all turns. |
max_context_tokens | int | 131072 | Maximum context tokens in the current conditioning run. |
segment_tokens | int | 4096 | Maximum tokens requested in one sampling segment. |
final_answer_reserve | int | 0 | Tokens reserved for a final answer after tool use. |
tool_output_tokens | int | 4096 | Per-tool-output cap before truncation. |
max_images | int | None | None | Optional cap on images in a trajectory. |
max_turn_tokens | int | None | None | Optional cap on generated tokens in one turn. |
rl.Budget.next_segment
rl.Budget.next_segment(traj) -> int| Parameter |
|---|
traj |
rl.Schedule
Client-side rollout admission and sampling submission limits.
rl.Schedule(
window: float = 0.01,
max_batch: int = 128,
max_batch_bytes: int = 256 * 1024 * 1024,
concurrency: int = 256,
max_sample_requests: int | None = None,
max_rpc_workers: int = 16,
admit_ahead: int = 2,
priority: Literal['fifo', 'predicted_length'] = 'fifo',
oversample_factor: float = 1.0,
) -> None| Field | Type | Default | Description |
|---|---|---|---|
window | float | 0.01 | Seconds to gather ready prompts into a sampling batch. |
max_batch | int | 128 | Maximum prompts in one sampling submission. |
max_batch_bytes | int | 256 * 1024 * 1024 | Target upper bound on an expanded submission envelope. |
concurrency | int | 256 | Maximum admitted trajectories. |
max_sample_requests | int | None | None | Maximum outstanding sampling prompts; defaults to concurrency. |
max_rpc_workers | int | 16 | Threads used for submit, poll, and decode calls. |
admit_ahead | int | 2 | Batches of rollout groups admitted ahead of training. |
priority | Literal['fifo', 'predicted_length'] | 'fifo' | Ready-queue ordering policy. |
oversample_factor | float | 1.0 | Extra groups admitted to replace uninformative groups. |
sample_request_limit |
rl.GroupCompletion
Choose when a reward group can be consumed.
A deadline is measured in generated tokens and takes effect between sampling segments; it cannot interrupt a sample already in progress.
rl.GroupCompletion(
mode: Literal['wait', 'deadline'],
min_members: int = 2,
max_straggler_tokens: int | None = None,
on_stragglers: Literal['truncate', 'carry_over', 'discard'] = 'truncate',
) -> None| Field | Type | Default | Description |
|---|---|---|---|
mode | Literal['wait', 'deadline'] | "wait" completes every member; "deadline" bounds stragglers by generated tokens. | |
min_members | int | 2 | Minimum completed members needed for a usable group. |
max_straggler_tokens | int | None | None | Token allowance for deadline mode. |
on_stragglers | Literal['truncate', 'carry_over', 'discard'] | 'truncate' | What to do with unfinished members at that deadline. |
rl.KVCache
Set how long retained KV may be reused as sampling weights change.
This policy applies to sampling_policy="segment". Trajectory-pinned
sampling uses one policy throughout and keeps same-policy KV.
rl.KVCache(
max_staleness: int = 0,
on_limit: Literal['hold', 'refill'] = 'hold',
allow_reprefill: bool = False,
) -> None| Field | Type | Default | Description |
|---|---|---|---|
max_staleness | int | 0 | Maximum policy steps between the retained KV's anchor weights and the weights used for a new turn. |
on_limit | Literal['hold', 'refill'] | 'hold' | "hold" keeps sampling at a compatible policy; "refill" moves to current weights and rebuilds the prefix. |
allow_reprefill | bool | False | Explicit acknowledgement of the prefix-cache miss caused by on_limit="refill". |
rl.KVCache.validate_sampling_policy
rl.KVCache.validate_sampling_policy(sampling_policy)| Parameter |
|---|
sampling_policy |
rl.Truncation
How truncated trajectories contribute to training.
rl.Truncation(
train: Literal['zero_reward', 'reward', 'drop', 'drop_and_exclude_from_baseline'] = 'zero_reward',
by_cause: dict[str, str] = dict(),
) -> None| Field | Type | Default | Description |
|---|---|---|---|
train | Literal['zero_reward', 'reward', 'drop', 'drop_and_exclude_from_baseline'] | 'zero_reward' | Default handling of a truncated trajectory. |
by_cause | dict[str, str] | {} | Overrides keyed by the trajectory's truncation cause. |
rl.Truncation.policy
rl.Truncation.policy(cause: str | None) -> str | None| Parameter | Type |
|---|---|
cause | str | None |
Training configuration
Advantage, optimizer, and forward/backward options.
rl.GroupCentered
Subtract each reward group's mean before training.
rl.GroupCentered(standardize: bool | str = False, eps: float = 1e-08) -> None| Field | Type | Default | Description |
|---|---|---|---|
standardize | bool | str | False | False keeps centered rewards; True divides by each group's standard deviation; "batch" divides all centered rewards by their shared batch standard deviation. |
eps | float | 1e-08 | Lower bound on a standard deviation used for division. |
rl.GroupCentered.__call__
rl.GroupCentered.__call__(rewards, *, min_members)| Parameter |
|---|
rewards |
min_members |
rl.Batchwise
Subtract one mean across the complete training batch.
rl.Batchwise(standardize: bool = False, eps: float = 1e-08) -> None| Field | Type | Default | Description |
|---|---|---|---|
standardize | bool | False | Divide by the batch standard deviation when true. |
eps | float | 1e-08 | Lower bound on that standard deviation. |
rl.Batchwise.__call__
rl.Batchwise.__call__(rewards, *, min_members)| Parameter |
|---|
rewards |
min_members |
rl.ForwardBackwardBatch
Submit forward/backward work while more synchronous rollouts finish.
Work is submitted when either configured threshold is reached; the final smaller chunk is submitted at the step boundary. Complete reward groups stay together, so a chunk can exceed a threshold. Token counts include the full model input, including prompt context.
rl.ForwardBackwardBatch(
min_sequences: int | None = 64,
min_tokens: int | None = 131072,
) -> None| Field | Type | Default | Description |
|---|---|---|---|
min_sequences | int | None | 64 | Submit after this many useful sequences, if set. |
min_tokens | int | None | 131072 | Submit after this many full input tokens, if set. |
rl.ForwardBackwardBatch.reached
rl.ForwardBackwardBatch.reached(sequences: int, tokens: int) -> bool| Parameter | Type |
|---|---|
sequences | int |
tokens | int |
rl.Adam
Optimizer parameters for one trainer update.
rl.Adam(
lr: float,
beta1: float = 0.9,
beta2: float = 0.95,
eps: float = 1e-08,
weight_decay: float = 0.0,
grad_clip_norm: float | None = None,
schedule: Callable[[int, int], float] | None = None,
) -> None| Field | Type | Default | Description |
|---|---|---|---|
lr | float | Base learning rate before the optional schedule multiplier. | |
beta1 | float | 0.9 | First-moment decay. |
beta2 | float | 0.95 | Second-moment decay. |
eps | float | 1e-08 | Numerical stability constant. |
weight_decay | float | 0.0 | Decoupled weight decay. |
grad_clip_norm | float | None | None | Optional global gradient-norm cap. |
schedule | Callable[[int, int], float] | None | None | Optional multiplier called with the update index and total planned updates; use rl.cosine() for checkpointable runs. |
rl.Adam.kwargs
rl.Adam.kwargs(step, steps)| Parameter |
|---|
step |
steps |
rl.Cosine
Cosine learning-rate multiplier with optional linear warmup.
rl.Cosine(warmup: int = 0, floor: float = 0.0) -> None| Field | Type | Default | Description |
|---|---|---|---|
warmup | int | 0 | Number of initial updates used for linear warmup. |
floor | float | 0.0 | Minimum multiplier after cosine decay. |
rl.Cosine.__call__
rl.Cosine.__call__(step: int, steps: int) -> float| Parameter | Type |
|---|---|
step | int |
steps | int |
rl.Checkpointing
Store recoverable trainer and rollout state in a durable directory.
rl.Checkpointing(
run_dir: str | Path,
weights_every: int = 20,
rollout_every: float = 60.0,
on_signal: tuple[str, ...] = (),
) -> None| Field | Type | Default | Description |
|---|---|---|---|
run_dir | str | Path | Local or shared filesystem directory for this run's state. | |
weights_every | int | 20 | Optimizer updates between full weight checkpoints. |
rollout_every | float | 60.0 | Seconds between rollout-state snapshots. |
on_signal | tuple[str, ...] | () | Signals that request a final checkpoint before shutdown. |
rl.cosine function
rl.cosine(*, warmup=0, floor=0.0)Create a checkpointable cosine learning-rate schedule.
| Parameter | Default | Description |
|---|---|---|
warmup | 0 | Number of initial updates used for linear warmup. |
floor | 0.0 | Minimum multiplier after cosine decay. |
rl.build_batch function
rl.build_batch(
groups,
*,
estimator,
normalize,
truncation: Truncation,
min_members,
current_step,
max_staleness,
defer_normalization=False,
advantages=None,
)| Parameter | Type | Default |
|---|---|---|
groups | ||
estimator | ||
normalize | ||
truncation | Truncation | |
min_members | ||
current_step | ||
max_staleness | ||
defer_normalization | False | |
advantages | None |
rl.decouple_ppo function
rl.decouple_ppo(data, proximal_logprobs, *, max_importance_weight: float)Correct PPO advantages for the policy used in an async update.
Scales each advantage by the proximal-to-behavior policy ratio, capped
by max_importance_weight. See AReaL's decoupled objective:
https://areal-ai.io/docs/en/algorithms/async.html
| Parameter | Type | Description |
|---|---|---|
data | Training records containing behavior logprobs and advantages. | |
proximal_logprobs | Logprobs from the policy used for this update. | |
max_importance_weight | float | Upper bound on the correction ratio. |
Trajectories and tools
The environment exchanges these values with the rollout engine.
rl.Trajectory
A rollout's conversation, exact model inputs, and generated spans.
The messages and final_text views support tools, rewards, and
inspection. Training uses the recorded token spans, preserving the exact
inputs and outputs sampled by the model. rewrite starts a new
conditioning run when the conversation must be compacted.
rl.Trajectory(chunks: list[ModelInputChunk], *, messages=None, id=None)| Parameter | Type | Default |
|---|---|---|
chunks | list[ModelInputChunk] | |
messages | None | |
id | None |
| Field | Type |
|---|---|
id | |
messages | |
final_text | |
last_stop_reason | |
reward | float | None |
truncated | str | None |
done | |
turns | |
metrics | dict[str, float] |
pending_tokens | list[int] |
phase | |
elapsed | |
wrap_up | |
sample_index | |
spans | tuple[Span, ...] |
run | int |
generated_tokens | int |
context_tokens | int |
rl.Trajectory.model_input
rl.Trajectory.model_input() -> list[ModelInputChunk]rl.Trajectory.prompt_ids
rl.Trajectory.prompt_ids() -> list[int]rl.Trajectory.append_span
rl.Trajectory.append_span(sample: Sample, *, policy_step: int | None = None)| Parameter | Type | Default |
|---|---|---|
sample | Sample | |
policy_step | int | None | None |
rl.Trajectory.append_framing
rl.Trajectory.append_framing(chunks: list[ModelInputChunk])| Parameter | Type |
|---|---|
chunks | list[ModelInputChunk] |
rl.Trajectory.rewrite
rl.Trajectory.rewrite(messages, *, chunks: list[ModelInputChunk])Replace conversation context while retaining earlier training spans.
Supply freshly rendered chunks for the replacement messages. The next model turn must prefill this new context; earlier samples are unchanged.
| Parameter | Type |
|---|---|
messages | |
chunks | list[ModelInputChunk] |
rl.Trajectory.to_data
rl.Trajectory.to_data(
advantage: float,
*,
current_step: int | None = None,
max_staleness: int | None = None,
) -> list[dict]| Parameter | Type | Default |
|---|---|---|
advantage | float | |
current_step | int | None | None |
max_staleness | int | None | None |
rl.Trajectory.state_dict
rl.Trajectory.state_dict()rl.Trajectory.from_state_dict
rl.Trajectory.from_state_dict(state, *, recovering=False)| Parameter | Default |
|---|---|
state | |
recovering | False |
rl.Span
rl.Span(
kind,
chunks,
*,
logprobs=None,
policy_step=None,
run=0,
policy_version=None,
retained_kv=False,
kv_cache_policy_version=None,
routing_handle=None,
routing_num_tokens=None,
)| Field | Type |
|---|---|
kind | Literal['prompt', 'generated', 'framing'] |
policy_step | int | None |
policy_version | PolicyVersion | None |
kv_cache_policy_version | PolicyVersion | None |
retained_kv | bool |
routing_handle | str | None |
routing_num_tokens | int | None |
run | int |
logprobs | np.ndarray | None |
chunks | list[ModelInputChunk] |
rl.Span.state_dict
rl.Span.state_dict()rl.CompletedGroup
Rollout trajectories collected for one dataset row.
rl.CompletedGroup(
id: str,
row: object,
trajectories: list[Trajectory],
final: bool = True,
closed_by_deadline: bool = False,
wait_seconds: float = 0.0,
) -> None| Field | Type | Default | Description |
|---|---|---|---|
id | str | Group identifier. | |
row | object | Source dataset row. | |
trajectories | list[Trajectory] | Completed or truncated member trajectories. | |
final | bool | True | Whether this group has no further members to emit. |
closed_by_deadline | bool | False | Whether the group's token deadline ended collection. |
wait_seconds | float | 0.0 | Time spent waiting for the group to complete. |
rl.Tool
Callable tool with a model-facing argument schema.
rl.Tool(function: object, spec: dict, signature: inspect.Signature) -> None| Field | Type | Description |
|---|---|---|
function | object | Async Python function called when the tool is selected. |
spec | dict | Tool description and JSON-compatible argument schema. |
signature | inspect.Signature | Python signature used to bind tool arguments. |
rl.Tool.__call__
async rl.Tool.__call__(**kwargs)| Parameter |
|---|
kwargs |
rl.tool function
rl.tool(function) -> ToolDerive a River ToolSpec from an async function's signature and docstring.
| Parameter |
|---|
function |
Evaluation and logging
Evaluate checkpoints and record training results.
rl.Evaluator
Run a holdout evaluation against saved model checkpoints.
Each variant uses a separate rollout engine built from a
CheckpointSampler. Evaluation can continue while the training model
advances; allocate its session on evaluation capacity.
rl.Evaluator(
holdout,
*,
engine_factory,
every=20,
group_size=4,
final_group_size=8,
variants=('default',),
sink=None,
)| Parameter | Default | Description |
|---|---|---|
holdout | Rows reserved for evaluation and excluded from training. | |
engine_factory | Callable (checkpoint, variant) returning an evaluation rollout engine. | |
every | 20 | Training batches between periodic evaluations; zero disables periodic evaluation. |
group_size | 4 | Samples per holdout row in periodic evaluations. |
final_group_size | 8 | Samples per row in the final evaluation. |
variants | ('default',) | Named evaluation configurations to run per checkpoint. |
sink | None | Optional callback receiving each Evaluation result. |
| Field | Type | Description |
|---|---|---|
holdout | Dataset rows reserved for evaluation. | |
variants | Names of evaluation configurations. | |
results | list[Evaluation] | Completed evaluation results. |
rl.Evaluator.config_dict
rl.Evaluator.config_dict()rl.Evaluator.exclude_holdout
rl.Evaluator.exclude_holdout(rows)| Parameter |
|---|
rows |
rl.Evaluator.launch
async rl.Evaluator.launch(model, step, *, final=False)| Parameter | Default |
|---|---|
model | |
step | |
final | False |
rl.Evaluator.wait
async rl.Evaluator.wait()rl.Evaluator.close
async rl.Evaluator.close()rl.Evaluator.state_dict
rl.Evaluator.state_dict()rl.Evaluator.load_state_dict
rl.Evaluator.load_state_dict(state)| Parameter |
|---|
state |
rl.Evaluation
Metrics and trajectories from evaluating one checkpoint variant.
rl.Evaluation(
step: int,
variant: str,
checkpoint: Checkpoint,
metrics: dict[str, float],
trajectories: list,
) -> None| Field | Type | Description |
|---|---|---|
step | int | Training step represented by the checkpoint. |
variant | str | Name of the evaluation configuration. |
checkpoint | Checkpoint | Saved weights sampled by this evaluation. |
metrics | dict[str, float] | Aggregate evaluation measurements. |
trajectories | list | Individual evaluation rollouts. |
rl.CheckpointSampler
Sample from a saved checkpoint on a separate evaluation session.
rl.CheckpointSampler(session, *, base_model, checkpoint: Checkpoint, tokenizer)| Parameter | Type | Description |
|---|---|---|
session | Session allocated for evaluation. | |
base_model | Model used to create the checkpoint. | |
checkpoint | Checkpoint | Saved weights and training step to evaluate. |
tokenizer | Tokenizer used to decode evaluation samples. |
| Field | Description |
|---|---|
step | Training step recorded by the checkpoint. |
submit_sampling_batch | Batch submission bound to this checkpoint. |
rl.CheckpointSampler.get_server_capabilities
rl.CheckpointSampler.get_server_capabilities()rl.CheckpointSampler.sample
rl.CheckpointSampler.sample(**kwargs)| Parameter |
|---|
kwargs |
rl.CheckpointSampler.submit_sample
rl.CheckpointSampler.submit_sample(**kwargs)| Parameter |
|---|
kwargs |
rl.WandbSink
Log evaluation metrics and sampled trajectories to a W&B run.
Metrics use the evaluated checkpoint's step, even if an evaluation finishes after newer training batches.
rl.WandbSink(run, *, table_samples=8)| Parameter | Default | Description |
|---|---|---|
run | Active W&B run. | |
table_samples | 8 | Maximum trajectories to include in the sample table. |
rl.WandbSink.__call__
rl.WandbSink.__call__(result)| Parameter |
|---|
result |
Results and helpers
Step results and the synchronous entry point.
rl.Step
Result of one completed training batch.
rl.Step(n: int, model_step: int, metrics: dict[str, float], trajectories: list) -> None| Field | Type | Description |
|---|---|---|
n | int | Number of batches completed by this trainer. |
model_step | int | Number of optimizer updates applied to the model. A batch with no usable training data advances n but not model_step. |
metrics | dict[str, float] | Training metrics for this batch, plus sampling activity since the previous result. Sampling metrics include work on other groups admitted ahead of this batch. |
trajectories | list | Trajectories selected for this batch. |
rl.run function
rl.run(trainer: AsyncTrainer, dataset, *, steps: int, on_step=None)Run training from synchronous Python code and return the final step.
In an async program, iterate over trainer.run(...) instead.
| Parameter | Type | Default | Description |
|---|---|---|---|
trainer | AsyncTrainer | Configured trainer to run. | |
dataset | Rows used to create rollout groups. | ||
steps | int | Number of training batches to complete. | |
on_step | None | Optional callback called with each completed Step. |
RL errors and warnings
Operational failures propagate instead of becoming task rewards.
rl.InfrastructureError
Inherits RuntimeError.
A tool or environment failure that should stop the rollout.
Raise this for unavailable external services instead of returning a tool
error to the model. TimeoutError and ConnectionError also stop the
rollout rather than becoming task rewards.
rl.RLConfigurationWarning
Inherits UserWarning.
An explicitly acknowledged configuration can waste rollout compute.