TL;DR. This repository studies an RSSM-style latent world model for LunarLander-v3: it trains the world model, runs CEM-based MPC with it, trains a model-based A2C policy in latent imagination, trains a model-free A2C baseline, compares all three, and investigates how to pick the best world-model checkpoint for closed-loop use from offline diagnostics alone using Jacobian-based metrics we propose. Full white paper describing it is here:
Predicting Closed-Loop Performance of Latent World Models: Offline Checkpoint Selection for MPC and Model-Based RL Under Non-Markovian Rewards in LunarLander.
The repository contains a complete pipeline. A latent world model (RSSM-style, after DreamerV2-V3) is trained on an offline dataset of human-piloted LunarLander-v3 episodes. The same model is then reused in two ways: as the dynamics model for a CEM-based MPC controller, and as the imagination engine for a model-based actor--critic (A2C) policy trained entirely in latent space. A model-free A2C agent is also trained on the live environment as a baseline. World-model-based approaches are much more data efficient: both MPC and WM-based A2C reach a strong policy using only the offline dataset (~181K real env transitions across 872 episodes), whereas the model-free RL policy needs ~11.8M real env transitions to reach its peak — roughly 65× more environment interaction. Peak mean returns over 100 deterministic episodes: eval_rl.py) for systematically evaluating actor-critic checkpoints, and a validation-time metrics suite (CROF) for picking best world-model checkpoints offline without touching the real environment; both are described in detail in the white paper on Arxiv is here.
World-model-based path (model-based RL):
- Collect Dataset → 2. Train World Model → 3. Test World Model → 4. Train Actor-Critic (WM-based) → 5. Test Policy
Model-free path (baseline comparison):
- Train Model-Free Actor-Critic → 2. Test Policy
MPC path (no actor needed):
- Collect Dataset → 2. Train World Model → 3. WM MPC Policy (CEM)
Collect human demonstrations or random episodes using keyboard control. Controls: LEFT/RIGHT/UP/DOWN arrows; ESC to end and save.
# Basic usage (saves to lunarlander_dataset.npz)
python collect_dataset.py
# Specify output path
python collect_dataset.py --dataset lunarlander_train_dataset.npz
# With custom seed for reproducibility
python collect_dataset.py --dataset lunarlander_train_dataset.npz --seed 12345| Option | Default | Description |
|---|---|---|
--dataset |
lunarlander_dataset.npz |
Output dataset .npz path |
--seed |
0 | Base RNG seed for episode surfaces |
Replay recorded episodes in the LunarLander environment to verify the dataset.
# Replay 20 episodes (default)
python replay_dataset.py
# Replay with custom dataset and settings
python replay_dataset.py --dataset lunarlander_train_dataset.npz --episodes 50 --seed 12345
# Slower replay (0.1 seconds per step)
python replay_dataset.py --dataset lunarlander_train_dataset.npz --episodes 10 --sleep 0.1| Option | Default | Description |
|---|---|---|
--dataset |
lunarlander_dataset.npz |
Path to dataset .npz |
--episodes |
20 | Number of episodes to replay |
--seed |
0 | Random seed for selecting episodes |
--sleep |
0.02 | Sleep between steps (seconds) |
Train the RSSM world model on sequences from real episodes. Uses SequenceDataset with teacher forcing.
# Basic usage (reads config.yaml, default datasets)
python train_models.py --phase world_model
# With explicit dataset paths
python train_models.py --phase world_model --config config.yaml \
--train_dataset lunarlander_train_dataset.npz \
--val_dataset lunarlander_val_dataset.npz| Option | Default | Description |
|---|---|---|
--phase |
(required) | world_model or actor_critic |
--config |
config.yaml |
Path to config file |
--train_dataset |
lunarlander_train_dataset.npz |
Training dataset path |
--val_dataset |
lunarlander_val_dataset.npz |
Validation dataset path |
--seed |
12345 | Random seed for reproducibility |
Config is in config.yaml: sequence_length, batch_size, lr, epochs, etc.
Evaluate the trained world model. Modes: teacher (posterior reconstruction), open (prior rollout), sim (constant action rollout).
# Teacher mode (default): posterior reconstruction with GT actions
python test_worldmodel.py --mode teacher --max_episodes 5
# Open mode: prior rollout from first obs, compare to GT
python test_worldmodel.py --mode open --dataset lunarlander_val_dataset.npz --max_episodes 10
# Save plots and animations
python test_worldmodel.py --mode open --max_episodes 5 --plot_dir plots --animate
# Sim mode: constant action rollout
python test_worldmodel.py --mode sim --constant_action 0 --max_episodes 3
# Use specific checkpoint
python test_worldmodel.py --mode teacher --checkpoint world_model.pt --max_episodes 5| Option | Default | Description |
|---|---|---|
--dataset |
lunarlander_val_dataset.npz |
Validation dataset path |
--config |
config.yaml |
Config file |
--checkpoint |
world_model.pt |
World model checkpoint path |
--mode |
teacher |
teacher / open / sim |
--max_episodes |
5 | Episodes to evaluate |
--seed |
0 | Seed for episode selection |
--plot_dir |
None | Save plots to directory |
--animate |
False | Save animations (.mp4) |
--constant_action |
0 | Action id for sim mode |
Train the policy (actor) and value function (critic) via imagined rollouts in the world model's latent space. Requires a trained world model. This is the model-based RL approach.
# Basic usage
python train_models.py --phase actor_critic
# With explicit dataset path
python train_models.py --phase actor_critic --config config.yaml \
--train_dataset lunarlander_train_dataset.npz| Option | Default | Description |
|---|---|---|
--phase |
(required) | actor_critic |
--config |
config.yaml |
Path to config file |
--train_dataset |
lunarlander_train_dataset.npz |
Training dataset path |
Config: horizon, past_horizon, future_horizon (must satisfy P+F=H), batch_size, lr, epochs, etc.
Run a trained policy in the LunarLander environment. Supports both WM-based (latent) and model-free (obs) actors. Runs headless by default for fast evaluation; pass --render to open a Pygame window with a stats overlay (episode, return, action probs, entropy, FPS). At the end of the run, a scorecard summary is printed (mean/worst return, perfect/negative/catastrophic counts, avg entropy, avg steps).
# Test WM-based actor headless (requires world model + latent actor)
python test_policy.py --actor_type latent --world_model world_model.pt --actor actor.pt --episodes 20
# Test model-free actor headless (no world model needed)
python test_policy.py --actor_type obs --actor actor_mf.pt --episodes 20
# Watch with rendering (window with overlay; press R to toggle, Q/ESC to quit)
python test_policy.py --actor_type latent --world_model world_model.pt --actor actor.pt --render --render_fps 30
# Stochastic action sampling instead of deterministic argmax
python test_policy.py --actor_type obs --actor actor_mf.pt --stochastic| Option | Default | Description |
|---|---|---|
--actor_type |
latent |
latent (WM-based, uses RSSM) or obs (model-free, raw observations) |
--config |
config.yaml |
Config for model dimensions |
--world_model |
world_model.pt |
World model checkpoint (only for latent actor) |
--actor |
actor.pt / actor_mf.pt |
Actor checkpoint (default depends on --actor_type) |
--episodes |
20 | Number of episodes to run |
--max_steps |
600 | Max steps per episode |
--render |
off | Open a Pygame window with stats overlay during rollout |
--render_fps |
30.0 | Visual FPS cap when rendering (sleep-based throttle) |
--deterministic |
(default) | Use argmax action selection |
--stochastic |
Sample actions from the policy distribution | |
--seed |
12345 | Random seed for reproducibility |
Render-window keyboard controls: R toggles rendering on/off, Q/ESC quits.
Run model-predictive control directly in Gymnasium using the world model only (no actor network). At each step it plans over discrete action sequences with CEM, executes the best first action, then replans.
# Basic run (with render)
python wm_mpc_policy.py --config config.yaml --world_model world_model.pt --render
# Faster/no-render run
python wm_mpc_policy.py --config config.yaml --world_model world_model.pt --episodes 10Keyboard controls (render window):
R: toggle render on/offQorESC: quit
Useful knobs:
--horizon,--population,--elites,--cem_itersfor planner strength/speed--w_angle,--w_ang_vel,--w_vx,--w_vy_downfor landing objective weights
Train a policy directly from on-policy environment interactions without a world model. Uses the same A2C algorithm (GAE, entropy regularization) as the WM-based trainer, but operates on raw observations instead of latent states. Serves as a baseline for comparing sample efficiency and compute cost against the world-model-based approach.
# Train from scratch
python train_modelfree_actorcritic.py --seed 12345
# Resume from latest checkpoint
python train_modelfree_actorcritic.py --resume --seed 12345
# Train with rendering (slower, shows one episode per epoch)
python train_modelfree_actorcritic.py --render --seed 12345| Option | Default | Description |
|---|---|---|
--epochs |
1000 | Number of training epochs |
--episodes_per_epoch |
50 | On-policy episodes collected per epoch |
--max_steps |
600 | Max steps per episode |
--lr |
3e-4 | Learning rate (AdamW) |
--gamma |
0.99 | Discount factor |
--lambda_gae |
0.95 | GAE lambda |
--entropy_coeff |
0.2 | Initial entropy coefficient |
--entropy_coeff_end |
0.01 | Final entropy coefficient (linearly decayed) |
--hidden_dim |
256 | Hidden dim for ActorObs / CriticObs |
--checkpoint_freq |
10 | Save checkpoints every N epochs |
--resume |
Resume from latest actor_mf / critic_mf checkpoint pair |
|
--render |
Render one episode per epoch | |
--seed |
12345 | Random seed |
Checkpoints are saved as actor_mf_<date>_<time>_epoch_<N>.pt and critic_mf_<date>_<time>_epoch_<N>.pt. Training logs are written to train_modelfree_actorcritic_logs.txt.
Run test_policy.py as a subprocess for every actor checkpoint in a folder and aggregate the results into a single log. Supports both WM-based AC (--actor_type latent, requires --world_model) and model-free AC (--actor_type obs).
# WM-based AC sweep (CROF-selected world model held fixed)
python eval_rl.py \
--checkpoints_dir checkpoints \
--actor_type latent \
--world_model world_model.pt \
--output rl_eval_logs.txt
# Model-free AC sweep (no world model)
python eval_rl.py \
--checkpoints_dir checkpoints \
--actor_type obs \
--epoch_stride 5 \
--output rl_eval_logs_mfAC.txt| Option | Default | Description |
|---|---|---|
--checkpoints_dir |
(required) | Folder containing actor checkpoints |
--actor_type |
(required) | latent (WM-based) or obs (model-free) |
--world_model |
None | World model checkpoint (required when --actor_type latent) |
--actor_pattern |
actor.*epoch_(\d+)\.pt$ |
Regex matching actor filenames; group 1 must be the integer epoch |
--episodes |
20 | Episodes per checkpoint |
--seed |
12345 | Random seed (passed through to test_policy.py) |
--max_steps |
600 | Max steps per episode (matches MPC sweep) |
--epoch_min |
None | Skip checkpoints below this epoch |
--epoch_max |
None | Skip checkpoints above this epoch |
--epoch_stride |
1 | Take every N-th checkpoint (use 5 for model-free actors saved every 10 epochs) |
--output |
rl_eval_logs.txt |
Output log file |
--config |
config.yaml |
Path to config |
--append |
off | Append to output instead of overwriting |
--top_k |
10 | Print top-K checkpoints by mean return at the end |
Runs are deterministic at the given seed. The sweep produces:
- a per-checkpoint scorecard block (full
test_policy.pystdout) for each actor - a sweep summary table (sorted by epoch) listing mean/worst return, perfect/negative/catastrophic counts, avg steps, avg entropy
- a top-K table sorted by mean return
- a
Best by mean_returnline identifying the winning checkpoint
This is the script used in the paper to evaluate WM-based AC (trained with the CROF-selected world model, WM 280) and the model-free AC baseline; the full paper sweep covers 9 world-model checkpoints with 100 deterministic episodes per actor save (--episodes 100).
The repository includes sample datasets and trained checkpoints:
| File | Description |
|---|---|
lunarlander_train_dataset.npz |
Training dataset (750 episodes, 158,685 transitions) |
lunarlander_val_dataset.npz |
Validation dataset (122 episodes, 22,231 transitions) |
world_model.pt |
Pretrained world model checkpoint — epoch 280 (raw min CROF, both CROF-A and CROF-B variants) |
actor.pt |
WM-based actor checkpoint — epoch 800 (best-by-mean A2C save trained on WM 280; mean +217.5 over 100 deterministic episodes, 79/100 perfect, 2/100 catastrophic) |
critic.pt |
WM-based critic checkpoint — epoch 800 (paired with the actor above) |
actor_mf.pt |
Model-free actor checkpoint — epoch 760 (peak of the model-free sweep; mean +193.0 over 100 deterministic episodes, 54/100 perfect, 0/100 catastrophic) |
critic_mf.pt |
Model-free critic checkpoint — epoch 760 (paired with the actor above) |
Use --checkpoint world_model.pt when testing the world model.
# World-model-based pipeline
python collect_dataset.py --dataset lunarlander_train_dataset.npz --seed 12345
python replay_dataset.py --dataset lunarlander_train_dataset.npz --episodes 5
python train_models.py --phase world_model --train_dataset lunarlander_train_dataset.npz --val_dataset lunarlander_val_dataset.npz
python test_worldmodel.py --mode open --max_episodes 2 --plot_dir plots
python train_models.py --phase actor_critic --train_dataset lunarlander_train_dataset.npz
python test_policy.py --actor_type latent --world_model world_model.pt --actor actor.pt --episodes 20
# MPC (no actor needed, uses world model directly)
python wm_mpc_policy.py --config config.yaml --world_model world_model.pt --render --episodes 20
# Model-free baseline
python train_modelfree_actorcritic.py --seed 12345
python test_policy.py --actor_type obs --actor actor_mf.pt --episodes 20
# Sweep all actor-critic checkpoints in a folder (WM-based or model-free)
python eval_rl.py --checkpoints_dir checkpoints --actor_type latent --world_model world_model.pt --output rl_eval_logs.txt
python eval_rl.py --checkpoints_dir checkpoints --actor_type obs --epoch_stride 5 --output rl_eval_logs_mfAC.txt- Decoder design to multi-head outputs:
- Shared decoder backbone from
[h, z] - Physics head (6 continuous dims:
x, y, vx, vy, angle, ang_vel) - Contact head (2 binary dims: leg contacts)
- Done head (1 binary dim)
- Shared decoder backbone from
- Loss design to match output types:
- Physics: MSE
- Contact: BCE-with-logits
- Done: BCE-with-logits
- (plus existing reward and KL terms)
- Model design:
Linear -> LayerNorm -> SiLUare good setup for physics modelling
- Posterior design:
- Posterior uses
obs_tonly, whiledoneis a training target (predicted).
- Posterior uses
- The main gain comes from separating continuous and binary prediction heads/losses, which reduced blurry terminal/contact dynamics and improved landing quality under MPC optimizer.
- Capacity, normalization (LinearNorm) and activation (SiLU) changes further improve optimization stability.
- Standard training-time metrics (validation loss, reward RMSE, multi-step open-loop RMSE) all keep improving monotonically past the point of best MPC performance, and pick deeply overfit checkpoints (epoch ≥460) where MPC closed-loop return has collapsed. MAE-based variants are even less informative because they suppress the rare large prediction errors that matter for control.
- A separate set of validation-time structural metrics — Jacobian-based Reward Observability Fraction averaged over curated "good" and "bad" states, plus controllability/observability rank fractions and multi-step observation RMSE, combined into a composite score (CROF) — turn out to be the strongest predictors of MPC return (best Spearman ρ ≈ −0.71 in our sweep) and select checkpoints inside the high-MPC plateau without ever touching the real environment. Epoch 280 was selected via this procedure (raw minimum of both CROF-A and CROF-B).
- The full analysis (definitions, sweeps, and selection results) will be published with the accompanying paper. The
world_model.ptcheckpoint shipped here is the one selected by that procedure.
-
World model (
world_model.pt): epoch 280, the CROF-selected world-model checkpoint (raw minimum of both CROF-A and CROF-B variants). Identifiable purely from offline validation-time diagnostics — no real-environment evaluation needed. -
WM-based AC (
actor.pt+critic.pt): epoch 800, the best-by-mean A2C save when trained on world model epoch 280. Real-env mean return$+217.5$ over 100 deterministic episodes (79/100 perfect landings, 2/100 catastrophic). -
Model-free AC (
actor_mf.pt+critic_mf.pt): epoch 760, the peak of the model-free sweep. Real-env mean return$+193.0$ over 100 deterministic episodes (54/100 perfect landings, 0/100 catastrophic), trained on $\approx 11.8$M real environment transitions to reach its peak. - Older actor/critic checkpoints are kept for reproducibility of earlier experiments. Full checkpoint sweeps live under
archive/.
