rlmesh.jax.Model
Experimental JAX-backed model: predict works in JAX values.
Model(
source: Callable[..., object] | object | None = None,
*,
spec: object | None = None,
on_episode_end: LifecycleCallback | None = None,
on_close: LifecycleCallback | None = None,
trust_entrypoints: bool = False,
) -> NoneThe JAX-typed ModelBase: wrap a predict
callable (Model(fn, spec=...)) or subclass and override predict;
run(env, seeds=[...]) returns a typed RunResult. Observations
arrive as immutable JAX arrays. See
ModelBase.
Generic over the observation/action types, defaulting to JaxValue:
wrapping an annotated predict callable infers them (Model(predict) with
def predict(obs: X) -> Y is a Model[X, Y], and its session a
Session[X, Y]); subclasses and unannotated sources bind JaxValue.
Examples
>>> from rlmesh.jax import Model
>>> Model(lambda observation: 0).run(
... "127.0.0.1:5555", seeds=[0]
... ).mean_reward
0.0Attributes
Methods
- describe()Return this model's full metadata envelope (see
rlmesh.describe()). - from_config()Build this model with
load(**config): the configured constructor. - load()Load weights into
self(from_pretrainedetc.); heavy imports here. - predict()Map one observation to an action (or an action chunk per
spec). - predict_chunk()Optional: map one observation to a CHUNK of actions (leading axis = chunk).
- predict_batch()Optional: one batched forward over all N vectorized lanes.
- predict_chunk_batch()Optional: one batched forward returning an action CHUNK per lane.
- on_episode_end()Optional: called when an episode ends (no-op by default).
- close()Release this model's resources: the explicit end of its life.
- run()Drive this model against an env on the native runtime loop.
- session()Bind this model to an env and return a
Sessionto drive by hand. - serve()Host this model as an endpoint (blocking).
spec
attribute[source]spec: object | Noneparams
attribute[source]params: ParamSpec | Noneworkflow_edition
attribute[source]workflow_edition: str | Nonedevice
attribute[source]device: object | Noneallow_fusion
attribute[source]allow_fusion: boolnative_chunk
attribute[source]native_chunk: int | Nonedescribe
function[source]describe() -> dict[str, Any]Return this model’s full metadata envelope (see rlmesh.describe()).
from_config
function[source]from_config(**config: Any) -> _ModelTBuild this model with load(**config): the configured constructor.
MyModel() loads with load’s own defaults; from_config is the
same construction with keywords – every keyword configures load()
and is validated first against the signature and the declared
params (unknown names, missing required ones, and out-of-domain
values raise ParamError). The instance is built
without the automatic load, then load(**resolved) runs exactly once;
a served model’s binding takes the identical path, so a configuration
that works locally serves the same. Only an authored subclass qualifies:
a framework base class with no predict corner of its own is refused
(wrap an existing policy with Model(predict, spec=...) instead).
load
function[source]load(**kwargs: Any) -> NoneLoad weights into self (from_pretrained etc.); heavy imports here.
Optional subclass hook; a no-op by default. Runs exactly once per
instance, before any predict: with its own defaults from MyModel(),
or with the keywords from_config() / a served binding resolved
against the signature and params.
predict
function[source]predict(observation: ObsT) -> ActTMap one observation to an action (or an action chunk per spec).
Override when subclassing Model; the default raises. A model built by
wrapping a predict callable uses that callable instead of this method.
predict_chunk
function[source]predict_chunk(observation: ObsT) -> ActTOptional: map one observation to a CHUNK of actions (leading axis = chunk).
Override this alongside predict() when the policy emits an action
chunk in one forward pass (ACT, diffusion, flow, VLA action heads). Return
your model’s native chunk; the runtime owns the replay – it executes the
first execution_horizon actions one per step without re-calling the model,
then re-plans. Defining this method is the chunk capability: a single-action
model leaves it unimplemented and the runtime re-plans every step.
Most policies ignore the execution horizon – their chunk length is fixed by
the trained weights, and the runtime simply uses a prefix of the native chunk.
An autoregressive decoder that can stop early may add an optional
execution_horizon: int = 1 parameter – as a second positional parameter
or keyword-only (keep the default, so it stays a compatible override of this
one-arg base); the runtime fills it with how many actions it will execute,
so the head can decode exactly that many instead of its full natural length.
The default raises.
predict_batch
function[source]predict_batch(observations: ObsT) -> ActTOptional: one batched forward over all N vectorized lanes.
Override (alongside predict()) to run a single forward pass for a
vectorized route instead of one call per lane. The runtime fuses the N
per-lane observations into one batched observation – every leaf gains a
leading batch axis, so a Dict observation arrives as {key: array[N, ...]}
(a NumPy/Torch/JAX array per leaf, stacked for this model’s framework), not a
list of N dicts. Return the batched action the same way: one value whose
leaves carry the leading batch axis (e.g. array[N, action_dim]); the
runtime splits it back per lane. The engine prefers this corner for a
vectorized route; the default raises and the route is driven per-lane through
predict().
(The dependency-free rlmesh.Model over raw Value trees can’t fuse
opaque tensors, so it instead receives the per-lane list and returns one.)
predict_chunk_batch
function[source]predict_chunk_batch(observations: ObsT) -> ActTOptional: one batched forward returning an action CHUNK per lane.
The batched counterpart of predict_chunk(): receives the fused batched
observation (leaves [N, ...]; see predict_batch()) and returns the
batched native chunk – leaves [N, chunk, ...] (batch axis first, then the
per-lane chunk axis). The runtime splits the batch axis back per lane, executes
a prefix (execution_horizon) of each, then re-plans. The engine prefers it
for a vectorized chunked route. Like predict_chunk(), an autoregressive
decoder may add an optional execution_horizon: int = 1 second parameter to
decode exactly that many. The default raises.
on_episode_end
function[source]on_episode_end(episode_id: str = '') -> NoneOptional: called when an episode ends (no-op by default).
The only per-episode boundary both the local loop and the served wire path
signal, so a stateful model clears its state here. episode_id is the
episode that just ended – declare it when the model keeps state keyed by
id (several episodes interleave through one served model);
def on_episode_end(self) remains a valid override when it does not.
close
function[source]close() -> NoneRelease this model’s resources: the explicit end of its life.
Never called by run() or a session close – a model instance is
borrowed there and stays usable. Override it to unload weights; on a
wrapped policy the base fires the on_close callback (or the
policy’s own close) once. Serving calls it when the server stops.
run
function[source]run(
env_or_address: LocalEnvTarget,
*,
seeds: Sequence[int] | None = None,
episodes: int | None = None,
max_episode_steps: int | None = None,
max_episode_seconds: float | None = None,
hooks: RunHooks | None = None,
instruction: str | None = None,
close_env: bool = False,
trust_entrypoints: bool | None = None,
execution_horizon: int = 1,
prefetch_lead: int = 0,
trial_index_base: int = 0,
workflow_edition: str | None = None,
) -> RunResultDrive this model against an env on the native runtime loop.
One loop for every env shape: a single env and a vectorized
(num_envs > 1) one run through the same native runtime the served
path uses – the route resolves at connect (adapter, per-episode frame
buffers), the engine dispatches the most specific predict corner
(predict_batch() / predict_chunk_batch() for a vectorized
route), and execution_horizon (> 1) executes that many actions of
each predicted chunk before re-planning (needs predict_chunk()).
prefetch_lead (> 0) turns on async inference over that replay:
with that many (or fewer) frames of the current chunk left, the
runtime predicts the next chunk while they execute, so the chunk is
conditioned on an observation up to prefetch_lead steps stale and
the result is not comparable to a synchronous run; a chunk prefetched
across an episode boundary is discarded.
env_or_address is a bare address string the loop dials, an object
with an address (EnvServer, RemoteEnv /
RemoteVectorEnv), or a local env object – served on a loopback
port for the duration of the run (tag it via
rlmesh.adapters.tag() for a spec’d model). Both the model and a
caller’s env are borrowed: the run releases what it created (the
loopback server, a factory-built env, the runtime session) and leaves
them usable, on success, failure, or interrupt alike; close_env
opts a caller’s env into shutdown, and close() is the model’s own
release.
episodes is the exact number of episodes the result holds: one by
default, the length of seeds when only seeds are given (one reset
seed per episode), and the two must agree when both are; 0 returns
an empty result. On a vectorized env the budget bounds episode
starts, so the scored set is fixed by the budget alone – a lane the
env rolls past it (NEXT_STEP autoreset steps the whole vector in
lockstep, so that lane cannot be paused) runs unscored: never counted,
reported, or seen by hooks. max_episode_steps /
max_episode_seconds truncate an episode at a step / wall-clock cap,
and a non-terminating env is bounded regardless: without an explicit
cap the runtime truncates any episode at 100,000 steps (the same
built-in bound as Session.run). Explicit
seeds and caps need the runtime to own resets, so they are refused
before any episode starts on an autoresetting vector env (drive it
with episodes; seeds on a driver-reset vector env must be a
multiple of num_envs). Every episode walks a trial
ordinal, trial_index_base + i for episode i: the runtime
delivers it as reset(options={"trial_index": ...}) to an env that
declared the key in :attr:EnvFactory.reset_options \<rlmesh.EnvFactory.reset_options> (an env that did not never sees
it) and reports it on :attr:EpisodeResult.trial \<rlmesh.EpisodeResult.trial>, so a local eval walks the same states
as a platform shard given the same base (a non-zero base needs the
runtime to own resets, like seeds). A per-call
trust_entrypoints override applies to this run only. On the
result, each episode’s predict_ms / step_ms are that episode’s
own per-step means (the runtime times every predict and env step it
issues; see StepEvent for how chunk replay and
prefetch are charged), and RunResult.telemetry carries the run-wide
aggregate – every measured op/metric series with count and avg/p50/p95/p99
(print(result.format_telemetry()) for a table) – so a slow run can
be attributed to the model forward, the env step, serialization, or
queueing without a profiler.
workflow_edition pins the semantics this run is evaluated under –
the runtime’s own declaration, above every other surface. It defaults to
RLMESH_WORKFLOW_EDITION, then workflow_edition on this model,
then [tool.rlmesh] workflow_edition; declaring none floats the run to
this build’s newest edition (reported once per process). An edition
neither side can run is refused before any episode starts, naming what
each tier wants and can do.
hooks observes the loop: a RunHooks receives the
runtime’s own events – episode starts, every step as a
StepEvent with the terminal flags the runtime knows at
that step, and episode ends – in the same per-episode order as
Session.run, with
RunHooks.on_run_start receiving a
RunContext for role reads and frame discovery. Hooks
never change the result: it is the runtime’s report either way, and
the EpisodeResult handed to on_episode_end is the
same record the result holds. On a vectorized env episodes
interleave. instruction overrides the text input of a spec’d model
in its declared shape, exactly as on the session loop. The live viewer
(view=) is a session() option.
session
function[source]session(
env_or_address: LocalEnvTarget,
*,
instruction: str | None = None,
close_env: bool = False,
trust_entrypoints: bool | None = None,
execution_horizon: int = 1,
view: ViewArg = None,
workflow_edition: str | None = None,
) -> Session[ObsT, ActT]Bind this model to an env and return a Session to drive by hand.
The manual counterpart of run(): drive reset / predict / step
yourself, or call Session.run() to pump whole episodes – as many times
as you like; the caller-held session (its connection, viewer, and adapter
state) stays open until you close it (close() or the with block).
Closing it never closes this model: it is yours, usable for the next
session or run until you call close().
env_or_address is an env object, an EnvFactory, a
remote-env handle, or an address string (see run()).
execution_horizon (> 1) executes that many actions per predicted chunk, one
per env step, when this model defines predict_chunk() (see run()).
workflow_edition declares the semantics this session runs under, with
the same precedence as run(); it reaches the wire only for a dialed
address (a local env object negotiates nothing).
serve
function[source]serve(address: str, *, options: ServeOptions | None = None) -> NoneHost this model as an endpoint (blocking).
A spec’d model resolves its adapter per env from the env contract the
resolve_adapter handshake delivers, then applies it around predict; a
spec-less / NO_ADAPTER model serves its own predict directly.
The endpoint declares a workflow edition on every handshake: options’
own workflow_edition if it sets one, else the resolved declaration
(RLMESH_WORKFLOW_EDITION, workflow_edition, [tool.rlmesh]).