title: “Models”
description: “Start with a backend Model subclass. Put weight loading in load() and implement the prediction method your policy needs, usually predict(). The same class works in a local evaluation, behind rlmesh.serve, and in a managed container. A ModelSpec declares its inputs and actions for adapters and managed probing.
A model in RLMesh is a policy you serve or drive against an environment. Wrap a predict callable for a small evaluation, or subclass a backend Model when you need to load weights, keep state, or serve the policy. Declare a ModelSpec when RLMesh needs to adapt an environment’s observation and action layout to the model’s own format.
Authoring a model is independent of authoring an environment. The two sides meet only through roles, resolved by Adapters. See Author an Environment for the environment side.
Use the Model Reference for prediction method signatures, batching, device placement, and lifecycle details.
Construction styles
Use a subclass for a policy you will reuse or package. Wrapping a function is convenient for a short local experiment.
| Style | What you write | Use it when |
|---|---|---|
Subclass a backend Model |
set the spec class attribute, implement load() plus one predict corner |
the form you serve and ship |
| Wrap a predict callable | rlmesh.numpy.Model(lambda obs: ...) |
a baseline or a one-file script |
Pick the backend by the array type your policy speaks. The backend only changes how observation leaves are decoded before predict and how returned actions are encoded; the four predict corners and the lifecycle are identical across all of them.
| Class | Import | Observation it receives | Action it returns |
|---|---|---|---|
rlmesh.Model |
native (no extras) | RLMesh-native values | RLMesh-native values |
rlmesh.numpy.Model |
pip install "rlmesh[numpy]" |
NumPy arrays + primitives | NumPy arrays |
rlmesh.torch.Model |
pip install "rlmesh[torch]" |
Torch tensors (device-aware) | Torch tensors |
rlmesh.jax.Model |
pip install "rlmesh[jax]" |
JAX arrays | JAX arrays |
See Framework Backends for the backend helpers, and Framework Backends for choosing one.
The minimal model
A subclass loads its weights in load() and maps an observation to an action in predict(). A minimal policy looks like this:
import rlmesh
class MyPolicy(rlmesh.numpy.Model):
def load(self):
self.bias = 0 # load weights INTO self
def predict(self, observation):
return self.bias
model = MyPolicy()
Set the spec class attribute to a ModelSpec and the adapter resolves from the environment’s published tags, so predict works in your model’s own conventions regardless of which environment it runs against. Set it to rlmesh.NO_ADAPTER to skip resolution and have predict receive the raw observation.
import rlmesh
import rlmesh.adapters as adapt
class MyPolicy(rlmesh.torch.Model):
spec = adapt.ModelSpec(
input={
"image": adapt.Image(adapt.IMAGE_PRIMARY, size=224),
"state": adapt.Concat(adapt.EEF_POS, adapt.GRIPPER_POS),
},
output=adapt.Action(adapt.Actuator(adapt.ACTION_DELTA_POS, dim=3)),
)
def load(self):
self.device = "cuda"
self.net = load_weights().to(self.device)
def predict(self, observation):
return self.net(observation["image"], observation["state"])
size=224 sets both height and width. The spec is the model side of the contract in Adapters; every field on every spec leaf is in Adapter Reference.
Lifecycle
A model has four lifecycle hooks. They fire identically on the local run / session loop and the served wire path, so a model behaves the same whether you drive it in-process or dial it over a socket.
| Seam | When it fires | What to do in it |
|---|---|---|
load(**config) |
once, at construction | Load weights into self; keep heavy imports here, not at module top. |
on_episode_end([episode_id]) |
when an episode ends | Clear per-episode state (RNN hidden state, chunk replay). |
close() |
when you call it | Release resources held by a model instance you passed to a run. |
on_episode_end / on_close |
constructor callbacks | The same edges for a wrapped callable that cannot override methods. |
Construct a subclass with MyPolicy() to use load() defaults, or call MyPolicy.from_config(**config) to pass validated keywords to load(). A served model uses the same configuration path. A caller-owned model instance or environment stays open after run() or session(); call close() when you are done. RLMesh closes models and environments it constructs from a class or factory for that call.
There is no episode-begin hook. Per-episode state is lazy-seeded on the first predict, so a stateful model clears its state at episode end via on_episode_end().
One served model fans several environments in, so episodes interleave and on_episode_end fires once per episode that ends, naming it: declare def on_episode_end(self, episode_id="") when you keep state keyed by id. def on_episode_end(self) stays valid when you do not.
For torch and jax models, set self.device inside load() when you move your weights onto it. That is the one source of truth: RLMesh moves every observation tensor leaf onto self.device before predict, so you never call .to(device) yourself.
def load(self):
self.device = "cuda"
self.policy = Policy.from_pretrained("org/checkpoint").to(self.device)
self.stats = load_norm_stats() # normalization stats load here too
The device and framework mechanics, and what happens when you set device on a numpy or native model, are in Model Reference.
Choose a prediction method
Start with predict(observation) when your policy returns one action per call. If it returns a chunk of future actions, implement predict_chunk instead. Policies that run one forward pass across several environment lanes can implement predict_batch or predict_chunk_batch.
| Your policy | Implement |
|---|---|
| One action per observation | predict |
| A chunk of actions per observation | predict_chunk |
| One forward pass over several lanes | predict_batch |
| Batched action chunks | predict_chunk_batch |
The runtime derives simpler forms from the method you implement where possible. The Model Reference has the exact signatures, derivation rules, execution horizon, per-episode context, and batch shapes. For a complete VLA model, see the VLA examples.
Run it
Use run to evaluate complete episodes, or session to call reset, predict, and step yourself:
model = MyPolicy()
result = model.run(env, seeds=range(10))
print(result.mean_reward, result.success_rate)
run returns a RunResult with .episodes, .mean_reward, and .success_rate (None when the env does not report every episode’s outcome). env may be a local env, an EnvFactory, a RemoteEnv, or an address string. The module-level rlmesh.run(model, env, ...) and rlmesh.session(model, env, ...) accept a Model subclass or instance or a served handle. Wrap a prediction function in the backend Model it expects, such as rlmesh.numpy.Model(fn).
The full run / session / read story (seeds, instruction injection, the execution horizon end to end, and reading canonical roles off an observation) is in Running Evaluations.
Build into a container
Save your model class in an importable module and use the standard serving command:
python -m rlmesh.serve my_pkg:MyPolicy
rlmesh.serve loads the model and serves it at RLMESH_ADDRESS (default 0.0.0.0:50051). It also publishes the description used by managed image probes. Construction parameters are read from RLMESH_MAKE_KWARGS and passed to load(**binding).
For managed evaluations, declare a ModelSpec so the probe can construct test inputs. The Bring Your Own Container example includes a complete model, a matching EnvFactory, Dockerfiles, and a local test before upload.
Connect to a server you started with rlmesh.RemoteModel(address). To let RLMesh start a prebuilt local container, use rlmesh.SandboxModel("image://my-model:latest"):
with rlmesh.session(rlmesh.SandboxModel("image://my-model:latest"), env) as sess:
obs, _ = sess.reset(seed=0)
while not sess.done:
obs, *_ = sess.step(sess.predict(obs))
Closing the session stops the model container it started. Sandbox helpers are experimental; see Sandbox Environments. A local sandbox does not submit a managed evaluation.