RL training¶
Reinforcement learning from a reward: SimEnv over any SimEngine, the PPO, FastSAC and FastTD3 trainers, every RLTrainSpec field, the checkpoint the rl provider reads.
By the end of this page you have trained a PPO actor on a MuJoCo SimEnv, read its checkpoint back, evaluated it, and know every field the three trainers accept.
import tempfile
from strands_robots.simulation import create_simulation
from strands_robots.simulation.predicates import make_predicate
from strands_robots.training import create_trainer
from strands_robots.training.rl import RLTrainSpec, SimEnv, load_deployable_actor, read_checkpoint_meta
def make_env() -> SimEnv:
sim = create_simulation("mujoco")
sim.create_world()
sim.add_robot("so101")
return SimEnv(
sim,
actor_obs_keys=["1", "2", "3", "4", "5", "6"],
reward_terms=[make_predicate("joint_progress", joint="1", target=0.5)],
success_fn=make_predicate("joint_above", joint="1", value=0.3),
max_episode_steps=50,
action_scale=0.15,
)
spec = RLTrainSpec(env_factory=make_env, output_dir=tempfile.mkdtemp(), total_timesteps=96, rollout_steps=24, learning_rate=3e-4)
trainer = create_trainer("ppo")
result = trainer.train(spec)
print(result.status, sorted(result.metrics))
meta = read_checkpoint_meta(result.checkpoint_dir)
print({k: meta[k] for k in ("provider", "num_actor_obs", "num_actions", "actor_obs_keys", "action_keys", "hidden_dims")})
print(sorted(trainer.evaluate(num_episodes=2)))
actor = load_deployable_actor(result.checkpoint_dir)
print(type(actor).__name__, actor.actor_obs_keys)
You should see (a few [sim] action value ... outside the range warnings on the gripper are the untrained actor):
success ['entropy', 'iteration', 'iterations_recorded', 'latest_loss', 'latest_step', 'mean_episode_return', 'mean_reward', 'metrics_path', 'surrogate_loss', 'value_loss']
{'provider': 'ppo', 'num_actor_obs': 6, 'num_actions': 6, 'actor_obs_keys': ['1', '2', '3', '4', '5', '6'], 'action_keys': ['1', '2', '3', '4', '5', '6'], 'hidden_dims': [128, 128]}
['episodes_successful_at_reset', 'max_return', 'mean_length', 'mean_return', 'min_return', 'num_episodes', 'returns', 'std_return', 'success_measured', 'success_rate']
DeployableActor ['1', '2', '3', '4', '5', '6']
Real runs need total_timesteps in the hundreds of thousands. create_policy("rl", checkpoint_dir=result.checkpoint_dir) drives a robot with it (rl).
SimEnv¶
SimEnv(engine, actor_obs_keys, reward_terms, *, action_dim=None, robot_name=None, critic_obs_keys=None, max_episode_steps=200, action_scale=1.0, n_substeps=5, success_fn=None, reset_fn=None, device="cpu", skip_images=True) wraps a live SimEngine as one environment of (1, D) tensors.
actor_obs_keys: ordered scalar keys fromget_observation(joint names,.velcompanions, floating-base keys); the order is part of the weights.reward_terms:(sim) -> floatcallables summed per step. Build them withmake_predicatefrom the predicates.critic_obs_keys: privileged sim-only keys for an asymmetric critic.action_dimdefaults tolen(engine.robot_action_keys(robot)), the actuator count, not always the joint count.action_scalebounds what the actor commands;0disconnects it and is refused.n_substeps=5: a position servo needs several physics steps per target.success_fnends an episode as a real terminal;max_episode_stepsis a truncation, value-bootstrapped by the trainers. Withoutsuccess_fn,evaluatereportssuccess_measured=Falseand asuccess_rateof zero measuring nothing.
VecSimEnv(env_factory, num_envs) steps N independent SimEnv through one thread pool, stacks to (N, D), keeping the terminal observation in infos[i]["terminal_obs"] across autoreset. GymSimEnv(sim_env) is the gymnasium.Env wrapper.
Trainers¶
| provider | class | family | own fields |
|---|---|---|---|
ppo |
PpoTrainer |
on-policy, GAE, clipped surrogate | gamma, lam, clip_param, num_learning_epochs, num_mini_batches, entropy_coef, value_loss_coef, max_grad_norm, init_noise_std, normalize_advantage |
fast_sac |
FastSacTrainer |
off-policy, replay buffer, entropy temperature | buffer_size, batch_size, learning_starts, gradient_steps, tau, autotune_alpha, init_alpha, alpha_lr, target_entropy |
fast_td3 |
FastTd3Trainer |
off-policy, twin critics, delayed actor | the SAC buffer fields plus policy_delay, exploration_noise_std, target_noise_std, target_noise_clip |
All three share setup, save_checkpoint, load_checkpoint, latest_checkpoint and evaluate(spec=None, checkpoint_dir=None, num_episodes=10); train fails closed via validate. evaluate updates nothing: mean action, no gradients, normalizers frozen.
RLTrainSpec¶
Extends TrainSpec (so output_dir, learning_rate, seed and the rest are there) with:
| field | default | read by |
|---|---|---|
env_factory |
required | all; a zero-arg callable returning a fresh SimEnv |
total_timesteps |
100_000 |
all |
rollout_steps |
24 |
all |
num_envs |
1 |
ppo (>1 wraps VecSimEnv); fast_sac and fast_td3 refuse anything but 1 |
actor_obs_keys, critic_obs_keys |
[] (from the env) |
all |
gamma |
0.99 |
all |
lam |
0.95 |
ppo |
clip_param |
0.2 |
ppo |
num_learning_epochs, num_mini_batches |
5, 4 |
ppo |
entropy_coef, value_loss_coef, max_grad_norm |
0.0, 1.0, 1.0 |
ppo |
hidden_dims |
(128, 128) |
all |
init_noise_std |
1.0 |
ppo |
normalize_obs, normalize_advantage |
True, True |
all; ppo |
device |
None (auto) |
all |
log_interval |
10 |
all; checkpoint cadence in iterations |
buffer_size, batch_size, learning_starts, gradient_steps, tau |
100_000, 256, 1_000, 1, 0.005 |
fast_sac, fast_td3 |
autotune_alpha, init_alpha, alpha_lr, target_entropy |
True, 1.0, 3e-4, None |
fast_sac |
policy_delay, exploration_noise_std, target_noise_std, target_noise_clip |
2, 0.1, 0.2, 0.5 |
fast_td3 |
learning_rate |
1e-4 |
all |
Booleans are checked, not read by truthiness; counts are positive integers; the two loss coefficients accept any finite real.
The checkpoint¶
save_checkpoint writes policy.pt (the state_dict, the frozen EmpiricalNormalization, provider) and policy_meta.json with provider, num_actor_obs, num_critic_obs, num_actions, actor_obs_keys, action_keys, hidden_dims, iteration. read_checkpoint_meta refuses a file missing any of the first six, by name. load_deployable_actor(checkpoint_dir, device) rebuilds the network the provider names (PPO raw means, FastTD3 a tanh, FastSAC a squashed mean/log-std pair), restores weights and normalizer, and returns the DeployableActor whose act(obs) create_policy("rl") calls.
Limits¶
- CPU MuJoCo is the only in-process batched path (
VecSimEnvthreads N engines); GPU-parallel RL: isaaclab. - No image observations:
actor_obs_keysare scalars,skip_images=Trueby default. - Three algorithms, one MLP shape each, no recurrent actor; curriculum is whatever
reset_fnand the terraindifficultyknob give.