Skip to content
Reference

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-client
import river_client as river

Core 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,
)
ParameterTypeDefaultDescription
api_keystrAPI key for authentication
endpointstr'api.river.ai'API endpoint hostname
portint443API port
timeoutfloat86400.0Default timeout for operations
use_sslboolTrueWhether to use SSL
enable_retriesboolTrueWhether 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_concurrencyint4Maximum 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() -> None

Drain 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,
) -> SessionContext

Create a session context manager.

Returns: SessionContext — Context manager that yields a Session

ParameterTypeDefaultDescription
timeoutfloat86400.0End-to-end session-creation timeout in seconds.
before_pollCallable[[], None] | NoneNone
on_session_creation_attemptedCallable[[], None] | NoneNone
on_session_createdCallable[[], None] | NoneNone
on_session_closedCallable[[], None] | NoneNone
tagsstr

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.

ParameterTypeDefaultDescription
promptsstr | list[str] | NoneNoneSingle prompt string or list of prompts. Mutually exclusive with prompt_token_ids.
base_modelstrBase model name (e.g. "Qwen/Qwen3.6-35B-A3B-FP8").
num_samplesint1Number of independent samples per prompt.
max_tokensint256Maximum tokens to generate per sample.
temperaturefloat1.0Sampling temperature.
top_pfloat1.0Nucleus sampling threshold.
top_kint-1Top-k sampling (-1 = disabled).
stoplist[str] | NoneNoneStop sequences.
seedint | NoneNoneRandom seed (varied per sample automatically).
return_prompt_logprobsboolFalseWhether to return prompt token logprobs.
logprobsint | NoneNoneIf set to K > 0, request the top-K alternative logprobs at each position. Off by default — enabling it roughly halves server throughput.
imageslist[Image] | list[list[Image]] | NoneNoneOptional raw image bytes for multimodal sampling. See sample for the per-prompt vs. broadcast semantics.
prompt_token_idslist[int] | list[list[int]] | NoneNonePre-tokenized prompt(s); mutually exclusive with prompts. See sample for details.
model_inputlist[dict] | list[list[dict]] | NoneNoneTraining-style chunk list(s); mutually exclusive with prompts / prompt_token_ids / images. See sample for details.
tokenizerAny | NoneNoneOptional tokenizer name or already-loaded tokenizer. Defaults to base_model after applying River model-alias resolution.
metrics_typestr''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.
timeoutfloat | NoneNoneTimeout in seconds.

Client.health_check

Client.health_check() -> bool

Check 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) -> ServerCapabilities

Read 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.

ParameterTypeDefault
model_namestr | NoneNone

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,
) -> Deployment

Create 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.

ParameterTypeDefault
checkpointstr | Checkpoint | NoneNone
base_modelstr | NoneNone
unified_replicasint | NoneNone
prefill_replicasint | NoneNone
decode_replicasint | NoneNone
idempotency_keystr | NoneNone
waitboolFalse
wait_timeoutfloat1800.0
timeoutfloat | NoneNone

Client.get_deployment

Client.get_deployment(deployment_id: str, *, timeout: float | None = None) -> Deployment

Fetch a deployment by id, including its tombstone after deletion.

Dedicated deployments are gated and disabled by default. See create_deployment for access requirements.

ParameterTypeDefault
deployment_idstr
timeoutfloat | NoneNone

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.

ParameterTypeDefault
include_deletedboolFalse
timeoutfloat | NoneNone

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,
) -> Deployment

Set 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.

ParameterTypeDefault
deployment_idstr
unified_replicasint | NoneNone
prefill_replicasint | NoneNone
decode_replicasint | NoneNone
idempotency_keystr | NoneNone
timeoutfloat | NoneNone

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,
) -> Deployment

Delete 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.

ParameterTypeDefault
deployment_idstr
idempotency_keystr | NoneNone
waitboolFalse
wait_timeoutfloat600.0
timeoutfloat | NoneNone

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.

ParameterTypeDefault
deployment_idstr
timeoutfloat | NoneNone

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'),
) -> Deployment

Wait 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.

ParameterTypeDefault
deployment_idstr
timeoutfloat1800.0
poll_intervalfloat5.0
untiltuple[str, ...]('ready', 'degraded')

Client.chat_complete

Client.chat_complete(
    messages: list[dict],
    *,
    base_model: str,
    timeout: float | None = None,
    **kwargs,
) -> ChatCompleteResult

Chat 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.

ParameterTypeDefaultDescription
messageslist[dict]OpenAI-format messages list.
base_modelstrBase model name for routing.
timeoutfloat | NoneNoneTimeout 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,
) -> ChatCompleteResult

Chat completion from a saved checkpoint.

Returns: ChatCompleteResult — ChatCompleteResult with response_json and status_code.

ParameterTypeDefaultDescription
messageslist[dict]OpenAI-format messages list.
checkpoint_pathstrriver:// checkpoint path.
base_modelstr''Base model name (optional; resolved from DB if empty).
timeoutfloat | NoneNoneTimeout in seconds.
kwargs

Client.chat_complete_from_training

Client.chat_complete_from_training(
    messages: list[dict],
    *,
    model_id: str,
    timeout: float | None = None,
    **kwargs,
) -> ChatCompleteResult

Chat completion from in-memory training weights.

Returns: ChatCompleteResult — ChatCompleteResult with response_json and status_code.

ParameterTypeDefaultDescription
messageslist[dict]OpenAI-format messages list.
model_idstrTraining model ID (e.g. session_id:model:seq).
timeoutfloat | NoneNoneTimeout in seconds.
kwargs

Client.close

Client.close() -> None

Close 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,
) -> ImageHandle

Upload into this session; bytes remain recoverable in the local image cache.

ParameterTypeDefault
databytes
idempotency_keystr | NoneNone
timeoutfloat | NoneNone

Session.upload_image_async

async Session.upload_image_async(
    data: bytes,
    *,
    idempotency_key: str | None = None,
    timeout: float | None = None,
) -> ImageHandle

Bounded non-blocking upload, automatically released when this session ends.

ParameterTypeDefault
databytes
idempotency_keystr | NoneNone
timeoutfloat | NoneNone

Session.release_image

Session.release_image(
    image: ImageHandle | None = None,
    *,
    idempotency_key: str | None = None,
) -> None
ParameterTypeDefault
imageImageHandle | NoneNone
idempotency_keystr | NoneNone

Session.release_image_async

async Session.release_image_async(
    image: ImageHandle | None = None,
    *,
    idempotency_key: str | None = None,
) -> None
ParameterTypeDefault
imageImageHandle | NoneNone
idempotency_keystr | NoneNone

Session.restore_images

async Session.restore_images(value: Any, *, image_store: ImageStore) -> Any

Re-upload checkpoint references once per content hash and remap every occurrence.

ParameterType
valueAny
image_storeImageStore

Session.get_server_capabilities

Session.get_server_capabilities(*, model_name: str | None = None) -> ServerCapabilities

Read protocol support, optionally for one authorized model.

ParameterTypeDefault
model_namestr | NoneNone

Session.attest_training_data

Session.attest_training_data(
    artifacts: list[TrainingDataArtifact],
    timeout: float = 86400.0,
) -> TrainingDataAttestation

Ask 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.

ParameterTypeDefault
artifactslist[TrainingDataArtifact]
timeoutfloat86400.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,
) -> Model

Create a new model for training.

Returns: Model — Model object for training

ParameterTypeDefaultDescription
base_modelstrBase model name (e.g., "Qwen/Qwen3.6-35B-A3B-FP8")
loraLoraConfigLoRA 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.
tokenizerstr | Any | NoneNoneTokenizer name (defaults to base_model) or an already-loaded tokenizer object
checkpointstr | Checkpoint | NoneNoneOptional 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.
timeoutfloat86400.0Timeout in seconds
training_data_attestationTrainingDataAttestation | str | NoneNoneOptional 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.

ParameterTypeDefaultDescription
promptsstr | list[str] | NoneNoneSingle prompt string or list of prompts. Mutually exclusive with prompt_token_ids.
base_modelstrBase model name (e.g. "Qwen/Qwen3.6-35B-A3B-FP8").
checkpointstr | Checkpoint | NoneNoneOptional river:// path or Checkpoint object. If provided, samples from that checkpoint's LoRA weights.
num_samplesint1Number of independent samples per prompt.
max_tokensint256Maximum tokens to generate per sample.
temperaturefloat1.0Sampling temperature.
top_pfloat1.0Nucleus sampling threshold.
top_kint-1Top-k sampling (-1 = disabled).
stoplist[str] | NoneNoneStop sequences.
seedint | NoneNoneRandom seed (varied per sample automatically).
seedslist[int] | NoneNoneExact per-prompt/sample seeds, mutually exclusive with seed.
return_prompt_logprobsboolFalseWhether to return prompt token logprobs.
logprobsint | NoneNoneIf set to K > 0, request the top-K alternative logprobs at each position. Off by default — enabling it roughly halves server throughput.
return_expert_routingboolFalse
imageslist[Image] | list[list[Image]] | NoneNoneOptional raw image bytes for multimodal sampling. See sample for the per-prompt vs. broadcast semantics.
prompt_token_idslist[int] | list[list[int]] | NoneNonePre-tokenized prompt(s); mutually exclusive with prompts. See sample for details.
model_inputlist[dict] | list[list[dict]] | NoneNoneTraining-style chunk list(s); mutually exclusive with prompts / prompt_token_ids / images. See sample for details.
tokenizerAny | NoneNoneOptional already-loaded tokenizer. Passing this avoids repeated Hugging Face cache/network resolution in tight loops.
metrics_typestr''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.
timeoutfloat86400.0Timeout 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,
) -> PendingSamplingBatch

Submit 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.

ParameterTypeDefault
promptsstr | list[str] | NoneNone
base_modelstr
checkpointstr | Checkpoint | NoneNone
num_samplesint1
max_tokensint256
temperaturefloat1.0
top_pfloat1.0
top_kint-1
stoplist[str] | NoneNone
seedint | NoneNone
seedslist[int] | NoneNone
return_prompt_logprobsboolFalse
return_prompt_token_idsboolFalse
logprobsint | NoneNone
return_expert_routingboolFalse
imageslist[Image] | list[list[Image]] | NoneNone
prompt_token_idslist[int] | list[list[int]] | NoneNone
model_inputlist[dict] | list[list[dict]] | NoneNone
tokenizerAny | NoneNone
idempotency_keystr | NoneNone
metrics_typestr''
timeoutfloat86400.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,
) -> PendingSample

Submit sampling from a base model or checkpoint without waiting.

See sample for the full kwarg reference.

ParameterTypeDefault
promptsstr | list[str] | NoneNone
base_modelstr
checkpointstr | Checkpoint | NoneNone
num_samplesint1
max_tokensint256
temperaturefloat1.0
top_pfloat1.0
top_kint-1
stoplist[str] | NoneNone
seedint | NoneNone
seedslist[int] | NoneNone
return_prompt_logprobsboolFalse
logprobsint | NoneNone
return_expert_routingboolFalse
imageslist[Image] | list[list[Image]] | NoneNone
prompt_token_idslist[int] | list[list[int]] | NoneNone
model_inputlist[dict] | list[list[dict]] | NoneNone
tokenizerAny | NoneNone
metrics_typestr''
timeoutfloat86400.0

Session.session_id

Session.session_id: str

SessionContext

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() -> ServerCapabilities

Read the endpoint's public protocol support (used by RL preflight).

Model.get_policy_version

Model.get_policy_version() -> PolicyVersion | None

Read 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,
) -> ForwardResult

Forward pass only (compute loss, no gradients).

Returns: ForwardResult — ForwardResult with metrics

ParameterTypeDefaultDescription
datalist[dict]List of training samples, each with "input_ids" and "labels"
loss_fnstr'cross_entropy'Loss function name
timeoutfloat86400.0Timeout in seconds
expected_policy_idstr | NoneNone
loss_configfloat

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,
) -> ForwardResult

Forward + 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.

ParameterTypeDefaultDescription
datalist[dict]List of training samples, each with "input_ids" and "labels"
loss_fnstr'cross_entropy'Loss function name
timeoutfloat86400.0Timeout in seconds
return_logprobsboolFalseDeprecated 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_outboolTrueWhen 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_metricboolFalseWhen 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_replayboolFalseWhen 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_idstr | NoneNone
loss_configfloat

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,
) -> OptimStepResult

Apply gradients with Adam optimizer.

Returns: OptimStepResult — OptimStepResult with metrics (step, lr, grad_norm, grad_norm_finite)

ParameterTypeDefaultDescription
lrfloatLearning rate
beta1float0.9Adam beta1
beta2float0.999Adam beta2
epsfloat1e-08Adam epsilon
weight_decayfloat0.0Weight decay
grad_clip_normfloat | NoneNoneGradient clipping norm (None to disable)
timeoutfloat86400.0Timeout in seconds
expected_policy_idstr | NoneNone
idempotency_keystr | NoneNone
gradient_scalefloat | NoneNoneMultiply 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).

ParameterTypeDefaultDescription
datalist[dict]List of training samples, each with "input_ids" and "labels"
lrfloatLearning rate
loss_fnstr'cross_entropy'
beta1float0.9
beta2float0.999
epsfloat1e-08
weight_decayfloat0.0
grad_clip_normfloat | NoneNone
compute_expert_flip_metricboolFalse
force_routing_replayboolFalse
timeoutfloat86400.0
loss_configfloat

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,
) -> PendingOp

Submit 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.

ParameterTypeDefault
datalist[dict]
loss_fnstr'cross_entropy'
timeoutfloat86400.0
return_logprobsboolFalse
zero_outboolTrue
compute_expert_flip_metricboolFalse
force_routing_replayboolFalse
expected_policy_idstr | NoneNone
loss_configfloat

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,
) -> PendingOp

Submit 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.

ParameterTypeDefault
lrfloat
beta1float0.9
beta2float0.999
epsfloat1e-08
weight_decayfloat0.0
grad_clip_normfloat | NoneNone
timeoutfloat86400.0
expected_policy_idstr | NoneNone
idempotency_keystr | NoneNone
gradient_scalefloat | NoneNone

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.

ParameterTypeDefault
datalist[dict]
lrfloat
loss_fnstr'cross_entropy'
beta1float0.9
beta2float0.999
epsfloat1e-08
weight_decayfloat0.0
grad_clip_normfloat | NoneNone
compute_expert_flip_metricboolFalse
force_routing_replayboolFalse
timeoutfloat86400.0
loss_configfloat

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.

ParameterTypeDefaultDescription
promptsstr | list[str] | NoneNoneSingle prompt string or list of prompts. Mutually exclusive with prompt_token_ids.
num_samplesint1Number of independent samples per prompt.
max_tokensint256Maximum tokens to generate per sample.
temperaturefloat1.0Sampling temperature.
top_pfloat1.0Nucleus sampling threshold.
top_kint-1Top-k sampling (-1 = disabled).
stoplist[str] | NoneNoneStop sequences.
seedint | NoneNoneRandom seed (varied per sample automatically).
seedslist[int] | NoneNoneExact per-prompt/sample seeds. Mutually exclusive with seed and ordered by prompt then sample.
return_prompt_logprobsboolFalseWhether to return prompt token logprobs.
logprobsint | NoneNoneIf 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.
imageslist[Image] | list[list[Image]] | NoneNoneOptional 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_routingboolFalseCapture 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_idslist[int] | list[list[int]] | NoneNonePre-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_inputlist[dict] | list[list[dict]] | NoneNoneTraining-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_typestr''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.
timeoutfloat86400.0Timeout in seconds for the entire operation (includes server-side wait for LoRA slot availability).
poll_intervalfloat1.0Seconds 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,
) -> PendingSamplingBatch

Submit 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.

ParameterTypeDefault
promptsstr | list[str] | NoneNone
num_samplesint1
max_tokensint256
temperaturefloat1.0
top_pfloat1.0
top_kint-1
stoplist[str] | NoneNone
seedint | NoneNone
seedslist[int] | NoneNone
return_prompt_logprobsboolFalse
return_prompt_token_idsboolFalse
logprobsint | NoneNone
imageslist[Image] | list[list[Image]] | NoneNone
return_expert_routingboolFalse
prompt_token_idslist[int] | list[list[int]] | NoneNone
model_inputlist[dict] | list[list[dict]] | NoneNone
idempotency_keystr | NoneNone
pinned_policy_idstr | NoneNone
policy_selectionstr'ordered'
retained_kv_groupslist[str] | NoneNone
kv_cache_policydict | NoneNone
metrics_typestr''
timeoutfloat86400.0
poll_intervalfloat0.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,
) -> PendingSample

Submit sampling from current training weights without waiting.

See sample for the full kwarg reference.

ParameterTypeDefault
promptsstr | list[str] | NoneNone
num_samplesint1
max_tokensint256
temperaturefloat1.0
top_pfloat1.0
top_kint-1
stoplist[str] | NoneNone
seedint | NoneNone
seedslist[int] | NoneNone
return_prompt_logprobsboolFalse
logprobsint | NoneNone
imageslist[Image] | list[list[Image]] | NoneNone
return_expert_routingboolFalse
prompt_token_idslist[int] | list[list[int]] | NoneNone
model_inputlist[dict] | list[list[dict]] | NoneNone
metrics_typestr''
timeoutfloat86400.0
poll_intervalfloat1.0

Model.chat_complete

Model.chat_complete(
    messages: list[dict],
    *,
    timeout: float | None = None,
    **kwargs,
) -> ChatCompleteResult

OpenAI-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.

ParameterTypeDefaultDescription
messageslist[dict]OpenAI-format messages list (e.g. [{"role": "user", "content": "Hello"}]).
timeoutfloat | NoneNoneTimeout 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,
) -> Checkpoint

Save a checkpoint of the current model weights.

Returns: Checkpoint — Checkpoint object with river:// path.

ParameterTypeDefaultDescription
namestrCheckpoint name (e.g. "final" or "step_000100").
modestr'training'"training" saves optimizer state (for training continuation), "inference" saves PEFT format only (for sampling/inference).
timeoutfloat86400.0Timeout in seconds.
ttldatetime.timedelta | NoneNoneLifetime 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.
immutableboolFalse
expected_policy_idstr | NoneNone

Model.load_weights

Model.load_weights(
    checkpoint: str | Checkpoint,
    load_optimizer: bool = True,
    timeout: float = 86400.0,
) -> None

Load weights from a checkpoint.

ParameterTypeDefaultDescription
checkpointstr | CheckpointA river:// path string or a Checkpoint object. If a Checkpoint is passed, its step is restored on the model.
load_optimizerboolTrueWhether to load optimizer state.
timeoutfloat86400.0Timeout in seconds.

Model.model_id

Model.model_id: str

Model.base_model

Model.base_model: str

Model.training_run_id

Model.training_run_id: str

Model.step

Model.step: int

Current 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.

FieldTypeDefault
rankint16
train_attnboolTrue
train_mlpboolTrue
train_unembedboolFalse
seedint | NoneNone

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.

FieldTypeDefault
tokenslist[int]
textstr
logprobslist[float]
stop_reasonstr
model_stepint
prompt_logprobslist[float] | NoneNone
request_idstr''
prompt_token_idslist[int] | NoneNone
top_logprobslist[list[TopLogprob]] | NoneNone
prompt_top_logprobslist[list[TopLogprob]] | NoneNone
expert_routingExpertRouting | NoneNone
metricsdict[str, float]{}
token_data_is_exactboolTrue
policy_versionPolicyVersion | NoneNone
kv_cache_policy_versionPolicyVersion | NoneNone
retained_kvboolFalse
cached_prompt_tokensint0
prompt_tokensint0
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)
ParameterTypeDefault
requiredboolFalse

ChatCompleteResult

Result of a chat completion request.

FieldType
response_jsonstr
status_codeint

ForwardResult

Result of forward or forward_backward pass.

FieldTypeDefault
metricsdict[str, float]
logprobslist | NoneNone

OptimStepResult

Result of an optimizer step.

FieldTypeDefault
metricsdict[str, float]
policy_versionPolicyVersion | NoneNone

Checkpoint

A saved model checkpoint.

FieldTypeDefault
pathstr
stepint
checkpoint_typestr
policy_versionPolicyVersion | NoneNone

TopLogprob

One top-K candidate token at a single position.

FieldTypeDefault
logprobfloat
token_idint
tokenstr''

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.

FieldType
request_idstr
PendingOp.result
PendingOp.result(
    *,
    before_poll: Callable[[], None] | None = None,
) -> ForwardResult | OptimStepResult

Block until the operation completes and return the result.

ParameterTypeDefault
before_pollCallable[[], None] | NoneNone

PendingSample

A submitted but not-yet-resolved sampling operation.

FieldType
request_idstr
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.

ParameterTypeDefault
retry_connection_errorsboolTrue
before_pollCallable[[], None] | NoneNone
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.

ParameterTypeDefault
executorExecutor | NoneNone
retry_connection_errorsboolTrue
before_pollCallable[[], None] | NoneNone

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.

FieldTypeDefault
topk_idsbytesb''
num_tokensint0
num_decoder_layersint0
top_kint0
layer_indiceslist[int][]
handlestr''

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.

FieldTypeDefault
idstr
modelstr
base_urlstr
checkpointstr
base_modelstr
topologystr
replicasdict[str, DeploymentReplicas]
phasestr
created_atstr
updated_atstr
phase_reasonstr | NoneNone
operation_idstr | NoneNone
deleted_atstr | NoneNone
is_servingbool
is_terminalbool

DeploymentReplicas

Desired, allocated, ready and pending counts for one role of a deployment.

pending are allocated replicas the cell cannot place yet (no capacity).

FieldTypeDefault
desiredint
allocatedint
readyint
pendingint0

ServerCapabilities

Supported public protocols and model names, without server internals.

FieldTypeDefault
supported_modelstuple[str, ...]
featuresfrozenset[str]
model_featuresdict[str, frozenset[str]]{}
ServerCapabilities.require_model
ServerCapabilities.require_model(model: str, *features: str) -> None
ParameterType
modelstr
featuresstr
ServerCapabilities.require
ServerCapabilities.require(*features: str) -> None
ParameterType
featuresstr

PolicyVersion

Server identity of a committed logical weight state (not a byte hash).

FieldTypeDefault
idstr
lineage_idstr
stepint
parent_idstr | NoneNone
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.

FieldType
idstr
sha256str
byte_countint
widthint
heightint
session_idstr

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.

FieldType
request_idstuple[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.

ParameterTypeDefault
executorExecutor | NoneNone
collectorSamplingResultCollector | NoneNone

SamplingCompletion

One sample or failure, identified in the original prompt/sample grid.

FieldTypeDefault
prompt_indexint
sample_indexint
request_idstr
sampleSample | NoneNone
errorException | NoneNone

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]
ParameterType
batchPendingSamplingBatch
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) -> None

Queue one metric point and return without waiting for the network.

ParameterType
metricstr
stepint
valuefloat
RunMetricsLogger.flush
RunMetricsLogger.flush() -> None

Request prompt delivery of the points currently in the buffer.

RunMetricsLogger.close
RunMetricsLogger.close() -> None

Drain 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.

FieldType
namestr
expected_sha256str
contentbytes

AttestedTrainingDataArtifact

Server-observed source-artifact digest and size.

FieldType
namestr
sha256str
size_bytesint

TrainingDataAttestation

A server-verified source-artifact manifest bound to one session.

FieldType
training_data_attestation_idstr
artifactslist[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.

FieldTypeDescription
messageHuman-readable error message
status_codegRPC status code name (e.g., "UNAVAILABLE", "DEADLINE_EXCEEDED")
detailsAdditional error details from the server
original_errorThe original gRPC RpcError for debugging
error_codestr | NoneStructured server reason, such as IMAGE_EXPIRED
image_idstr | NoneImage requiring re-upload, when supplied by the server
RiverConnectionError.from_grpc_error
RiverConnectionError.from_grpc_error(
    error: Exception,
    context: str = 'API call',
) -> RiverConnectionError

Create a RiverConnectionError from a gRPC RpcError.

ParameterTypeDefault
errorException
contextstr'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.

ParameterTypeDefault
tokenizerstr | Any | NoneNone
base_modelstr | NoneNone
revisionstr | NoneNone
local_files_onlyboolFalse
resolve_aliasesboolTrue

resolve_tokenizer_name

resolve_tokenizer_name(model_name: str) -> str

Return 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.

ParameterType
model_namestr

MODEL_TOKENIZER_ALIASES

MODEL_TOKENIZER_ALIASES: dict[str, str]
KeyValue
Qwen/Qwen3.6-35B-A3B-FP8Qwen/Qwen3.6-35B-A3B
Qwen/Qwen3.5-397B-A17B-FP8Qwen/Qwen3.5-397B-A17B-FP8
nvidia/Kimi-K2.6-NVFP4nvidia/Kimi-K2.6-NVFP4
moonshotai/Kimi-K3moonshotai/Kimi-K3
nvidia/GLM-5.1-NVFP4nvidia/GLM-5.1-NVFP4
nvidia/GLM-5.2-NVFP4nvidia/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 rl

rl.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,
)
ParameterTypeDefaultDescription
engineRollout engine sharing the training model.
optimizerAdamOptimizer settings for each update.
completionGroupCompletionWhen a rollout group can be consumed.
normalizestrLoss scaling: "token", "sequence", "batch", or "sum".
groups_per_stepintReward groups consumed per training batch.
group_sizeint16Samples requested for each dataset row.
advantageNoneAdvantage estimator; defaults to GroupCentered().
lossstr'cispo'"cispo", "ppo", "importance_sampling", or "decoupled_ppo".
max_stalenessint2Maximum training-policy age in optimizer steps; zero gives synchronous admission.
stale_policystr'wait'"wait" holds complete groups within the bound; "mask_span" masks old spans; "keep" ignores the bound.
allow_trajectory_lossboolFalseAcknowledge options that may drop or fully mask trajectories.
allow_unbounded_stalenessboolFalseAcknowledge stale_policy="keep".
truncationTruncationTruncation()Reward or drop policy for truncated trajectories.
checkpointCheckpointing | NoneNoneDurable trainer-state checkpoint settings.
init_checkpointNoneOptional starting weights checkpoint.
evaluatorNoneOptional checkpoint-based evaluation runner.
run_configNoneSerializable recipe values checked on resume.
max_importance_weightfloat5.0Importance-ratio limit for decoupled PPO.
routing_replayboolFalseReplay captured expert routing during training.
repeat_datasetboolTrueCycle dataset rows when more groups are needed.
batch_orderstr'completion'Consume groups by "completion" or "admission".
validate_batchNoneOptional callback before a full batch trains.
forward_backward_groupsint | NoneNoneSubmit forward/backward after this many ready groups; requires synchronous group-centered training.
forward_backward_batchForwardBackwardBatch | Literal['auto'] | None'auto'"auto" enables threshold-based submission when eligible; set a ForwardBackwardBatch or None explicitly to override.
max_pending_forward_backwardint16Limit concurrent submitted operations.
loss_configLoss-specific numeric options such as eps_max.

Methods: progress, run

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.

ParameterTypeDefaultDescription
datasetRows used to create rollout groups.
stepsintNumber of batches to complete. A batch with no usable training data still counts, but does not update the model.
after_recoveryNoneOptional 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,
)
ParameterTypeDefaultDescription
modelTraining model or checkpoint-pinned sampler.
envEnv 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.
rendererModel-specific chat renderer that supports continuation.
budgetBudgetBudget()Per-trajectory token, turn, and image limits.
scheduleScheduleSchedule()Client-side admission and sampling submission limits.
sampling_policystr'trajectory'"trajectory" pins one policy for every turn; "segment" allows a newer policy between turns.
kv_cacheKVCacheKVCache()Maximum retained KV age and behavior at that limit.
temperaturefloat1.0Sampling temperature; training requires a positive value.
top_pfloat1.0Nucleus sampling threshold.
top_kint-1Top-k sampling limit; -1 disables it.
seedint0Base sampling seed.
sample_timeoutfloat1800.0Timeout for one sampling request in seconds.
environment_timeoutfloat | NoneNoneOptional timeout for an environment call in seconds.
return_expert_routingboolFalseRequest 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.

ParameterTypeDefault
row
group_sizeint
completionGroupCompletion
idstr | NoneNone
priorityint0

rl.RolloutEngine.next_group

async rl.RolloutEngine.next_group() -> CompletedGroup

Wait 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.

ParameterType
rows
group_sizeint
completionGroupCompletion

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.

AttributeTypeDescription
toolstuple[Tool, ...] | list[Tool]Tools available to the model during each turn.
recoveryLiteral['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) -> float

Return the final reward for a completed trajectory.

Parameter
traj
row

rl.Env.on_turn

async rl.Env.on_turn(traj) -> list[Message] | None

Return 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) -> float

Override for task-specific partial credit. Default is explicit zero.

ParameterType
traj
row
causestr

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
FieldTypeDefaultDescription
max_turnsint8Maximum assistant turns.
max_generated_tokensint32768Total generated-token budget across all turns.
max_context_tokensint131072Maximum context tokens in the current conditioning run.
segment_tokensint4096Maximum tokens requested in one sampling segment.
final_answer_reserveint0Tokens reserved for a final answer after tool use.
tool_output_tokensint4096Per-tool-output cap before truncation.
max_imagesint | NoneNoneOptional cap on images in a trajectory.
max_turn_tokensint | NoneNoneOptional 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
FieldTypeDefaultDescription
windowfloat0.01Seconds to gather ready prompts into a sampling batch.
max_batchint128Maximum prompts in one sampling submission.
max_batch_bytesint256 * 1024 * 1024Target upper bound on an expanded submission envelope.
concurrencyint256Maximum admitted trajectories.
max_sample_requestsint | NoneNoneMaximum outstanding sampling prompts; defaults to concurrency.
max_rpc_workersint16Threads used for submit, poll, and decode calls.
admit_aheadint2Batches of rollout groups admitted ahead of training.
priorityLiteral['fifo', 'predicted_length']'fifo'Ready-queue ordering policy.
oversample_factorfloat1.0Extra 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
FieldTypeDefaultDescription
modeLiteral['wait', 'deadline']"wait" completes every member; "deadline" bounds stragglers by generated tokens.
min_membersint2Minimum completed members needed for a usable group.
max_straggler_tokensint | NoneNoneToken allowance for deadline mode.
on_stragglersLiteral['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
FieldTypeDefaultDescription
max_stalenessint0Maximum policy steps between the retained KV's anchor weights and the weights used for a new turn.
on_limitLiteral['hold', 'refill']'hold'"hold" keeps sampling at a compatible policy; "refill" moves to current weights and rebuilds the prefix.
allow_reprefillboolFalseExplicit 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
FieldTypeDefaultDescription
trainLiteral['zero_reward', 'reward', 'drop', 'drop_and_exclude_from_baseline']'zero_reward'Default handling of a truncated trajectory.
by_causedict[str, str]{}Overrides keyed by the trajectory's truncation cause.
rl.Truncation.policy
rl.Truncation.policy(cause: str | None) -> str | None
ParameterType
causestr | 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
FieldTypeDefaultDescription
standardizebool | strFalseFalse keeps centered rewards; True divides by each group's standard deviation; "batch" divides all centered rewards by their shared batch standard deviation.
epsfloat1e-08Lower 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
FieldTypeDefaultDescription
standardizeboolFalseDivide by the batch standard deviation when true.
epsfloat1e-08Lower 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
FieldTypeDefaultDescription
min_sequencesint | None64Submit after this many useful sequences, if set.
min_tokensint | None131072Submit after this many full input tokens, if set.
rl.ForwardBackwardBatch.reached
rl.ForwardBackwardBatch.reached(sequences: int, tokens: int) -> bool
ParameterType
sequencesint
tokensint

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
FieldTypeDefaultDescription
lrfloatBase learning rate before the optional schedule multiplier.
beta1float0.9First-moment decay.
beta2float0.95Second-moment decay.
epsfloat1e-08Numerical stability constant.
weight_decayfloat0.0Decoupled weight decay.
grad_clip_normfloat | NoneNoneOptional global gradient-norm cap.
scheduleCallable[[int, int], float] | NoneNoneOptional 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
FieldTypeDefaultDescription
warmupint0Number of initial updates used for linear warmup.
floorfloat0.0Minimum multiplier after cosine decay.
rl.Cosine.__call__
rl.Cosine.__call__(step: int, steps: int) -> float
ParameterType
stepint
stepsint

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
FieldTypeDefaultDescription
run_dirstr | PathLocal or shared filesystem directory for this run's state.
weights_everyint20Optimizer updates between full weight checkpoints.
rollout_everyfloat60.0Seconds between rollout-state snapshots.
on_signaltuple[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.

ParameterDefaultDescription
warmup0Number of initial updates used for linear warmup.
floor0.0Minimum 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,
)
ParameterTypeDefault
groups
estimator
normalize
truncationTruncation
min_members
current_step
max_staleness
defer_normalizationFalse
advantagesNone

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

ParameterTypeDescription
dataTraining records containing behavior logprobs and advantages.
proximal_logprobsLogprobs from the policy used for this update.
max_importance_weightfloatUpper 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)
ParameterTypeDefault
chunkslist[ModelInputChunk]
messagesNone
idNone
FieldType
id
messages
final_text
last_stop_reason
rewardfloat | None
truncatedstr | None
done
turns
metricsdict[str, float]
pending_tokenslist[int]
phase
elapsed
wrap_up
sample_index
spanstuple[Span, ...]
runint
generated_tokensint
context_tokensint
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)
ParameterTypeDefault
sampleSample
policy_stepint | NoneNone
rl.Trajectory.append_framing
rl.Trajectory.append_framing(chunks: list[ModelInputChunk])
ParameterType
chunkslist[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.

ParameterType
messages
chunkslist[ModelInputChunk]
rl.Trajectory.to_data
rl.Trajectory.to_data(
    advantage: float,
    *,
    current_step: int | None = None,
    max_staleness: int | None = None,
) -> list[dict]
ParameterTypeDefault
advantagefloat
current_stepint | NoneNone
max_stalenessint | NoneNone
rl.Trajectory.state_dict
rl.Trajectory.state_dict()
rl.Trajectory.from_state_dict
rl.Trajectory.from_state_dict(state, *, recovering=False)
ParameterDefault
state
recoveringFalse

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,
)
FieldType
kindLiteral['prompt', 'generated', 'framing']
policy_stepint | None
policy_versionPolicyVersion | None
kv_cache_policy_versionPolicyVersion | None
retained_kvbool
routing_handlestr | None
routing_num_tokensint | None
runint
logprobsnp.ndarray | None
chunkslist[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
FieldTypeDefaultDescription
idstrGroup identifier.
rowobjectSource dataset row.
trajectorieslist[Trajectory]Completed or truncated member trajectories.
finalboolTrueWhether this group has no further members to emit.
closed_by_deadlineboolFalseWhether the group's token deadline ended collection.
wait_secondsfloat0.0Time 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
FieldTypeDescription
functionobjectAsync Python function called when the tool is selected.
specdictTool description and JSON-compatible argument schema.
signatureinspect.SignaturePython signature used to bind tool arguments.
rl.Tool.__call__
async rl.Tool.__call__(**kwargs)
Parameter
kwargs

rl.tool function

rl.tool(function) -> Tool

Derive 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,
)
ParameterDefaultDescription
holdoutRows reserved for evaluation and excluded from training.
engine_factoryCallable (checkpoint, variant) returning an evaluation rollout engine.
every20Training batches between periodic evaluations; zero disables periodic evaluation.
group_size4Samples per holdout row in periodic evaluations.
final_group_size8Samples per row in the final evaluation.
variants('default',)Named evaluation configurations to run per checkpoint.
sinkNoneOptional callback receiving each Evaluation result.
FieldTypeDescription
holdoutDataset rows reserved for evaluation.
variantsNames of evaluation configurations.
resultslist[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)
ParameterDefault
model
step
finalFalse
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
FieldTypeDescription
stepintTraining step represented by the checkpoint.
variantstrName of the evaluation configuration.
checkpointCheckpointSaved weights sampled by this evaluation.
metricsdict[str, float]Aggregate evaluation measurements.
trajectorieslistIndividual evaluation rollouts.

rl.CheckpointSampler

Sample from a saved checkpoint on a separate evaluation session.

rl.CheckpointSampler(session, *, base_model, checkpoint: Checkpoint, tokenizer)
ParameterTypeDescription
sessionSession allocated for evaluation.
base_modelModel used to create the checkpoint.
checkpointCheckpointSaved weights and training step to evaluate.
tokenizerTokenizer used to decode evaluation samples.
FieldDescription
stepTraining step recorded by the checkpoint.
submit_sampling_batchBatch 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)
ParameterDefaultDescription
runActive W&B run.
table_samples8Maximum 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
FieldTypeDescription
nintNumber of batches completed by this trainer.
model_stepintNumber of optimizer updates applied to the model. A batch with no usable training data advances n but not model_step.
metricsdict[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.
trajectorieslistTrajectories 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.

ParameterTypeDefaultDescription
trainerAsyncTrainerConfigured trainer to run.
datasetRows used to create rollout groups.
stepsintNumber of training batches to complete.
on_stepNoneOptional 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.