lejepa-control
All trained weights for the lejepa_control research repo: amortized latent planners on top of frozen LeWM world models (arXiv:2603.19312 β JEPA-style: ViT-tiny encoder β 192-dim latent, 6-layer predictor, frameskip 5 so one latent step = 5 simulator steps).
The world models are never trained by the planners. Each planner/controller below was trained to plan
inside the frozen latent space of one of the world_models/ checkpoints (or the original
quentinll PushT/Reacher/TwoRooms LeWM models), and is evaluated in
the simulator, which it never sees during planning β the measured quantity is how well latent-space
planning transfers to real rollouts.
Repository layout
βββ world_models/ locally-trained LeWM world models (config.json + final weights)
β βββ pointmaze/ 5 variants (see table below)
β βββ humanoid/ full 40-epoch model + 2-epoch smoke
βββ controllers/ phase-1 IterativeController runs, one folder per environment
β βββ pusht/ headline runs + arrival/hold and objective ablations + exp7 variants
β βββ pointmaze/ humanoid/ reacher/ tworooms/
β βββ <env>/scratch/ smoke / dry-run checkpoints (sanity checks only)
βββ planners/ phase-2 planners (PushT), one folder per architecture
β βββ recursive/ RecursivePlanner β causal f/g recursion, staged bring-up AβE
β βββ cross_attention/ cross-attention planner β joint refinement, causal influence mask
β βββ scratch/ smoke runs
βββ density_models/ behavior-density (support-constraint) models, one per environment
βββ decoder/ latentβpixel visualization decoder (side channel, never in the loop)
βββ manifold_transfer/ E3 encoder-transfer study checkpoints
βββ e3_encoders/ trained encoder best/final checkpoints
βββ adapters/ linear-map fits between latent spaces (linear/orthogonal/mlp Γ bias)
βββ e3lite/ E3-lite encoder checkpoints
World models
All trained for 40 epochs with the LeWM recipe. config.json next to each weights_*.pt matches the
LeWM config schema (lejepa_control/world_model.py::load_lewm).
Path (world_models/β¦) |
Environment | Notes |
|---|---|---|
pointmaze/lewm-pointmaze-r3 |
PointMaze, 3-room | base 3-room dataset |
pointmaze/lewm-pointmaze-r3g |
PointMaze, 3-room gated | gate (door) variant |
pointmaze/lewm-pointmaze-r3g-p7 |
PointMaze, 3-room gated | p7 variant β pairs with the p7 CEM-planner gate/video evals in the repo |
pointmaze/lewm-pointmaze-v1-astar |
PointMaze | dataset variant with A*-generated goal pairs |
pointmaze/lewm-pointmaze-v2-contact |
PointMaze | contact-rich dataset variant |
humanoid/lewm-humanoid |
Humanoid | full run, weights_epoch_40.pt |
humanoid/lewm-humanoid-smoke |
Humanoid | 2-epoch smoke, sanity checks only |
Not hosted here: the PushT / Reacher / TwoRooms LeWM world models β those are
quentinll's, fetch via the swm CLI (see the GitHub repo's
docs/SIMULATOR_GUIDE.md).
Controllers (phase-1, IterativeController)
Non-causal controller: refines all H action blocks jointly via full self-attention, trained on the
arrival+hold objective. One folder per run, file is always controller.pt.
Run (controllers/<env>/β¦) |
What it is |
|---|---|
pusht/controller |
original terminal-only baseline |
pusht/ah_hold{0.0,0.5,1.0} |
arrival+hold weight ablation β ah_hold0.5 is the headline PushT model (94% closed-loop / 88% open-loop at rh=1) |
pusht/abl_no_support, pusht/abl_terminal_only |
objective ablations (support term off / terminal-only) |
pusht/exp7/{base_r2,fused192,fused192_s2,fused256,w192np_split} |
one-operator (fused) controller variants β outcome inconclusive, see docs/exp7_analysis.md |
pointmaze/controller_pointmaze |
PointMaze port |
pointmaze/controller_pointmaze_v1-astar |
PointMaze, v1-astar world model + dataset |
humanoid/controller_humanoid |
Humanoid port |
reacher/controller_reacher |
Reacher port |
tworooms/controller_tworoom |
TwoRooms port |
<env>/scratch/* |
dry-run / smoke checkpoints, not results |
Paired with each controller is a behavior-density model under density_models/<env>/ (density.pt) β
it supplies the calibrated support threshold used by the support_loss term.
Planners (phase-2, PushT)
Run (planners/<arch>/β¦) |
What it is |
|---|---|
recursive/planner_{A,B,C,D} |
staged bring-up of the RecursivePlanner (stages AβD) |
recursive/planner_D10, planner_D10_extended, planner_E |
later training stages / extensions |
recursive/planner_curriculum (+ stage_A..D/) |
chained curriculum AβBβCβD (docs/COMBINED_TRAINING.md) |
recursive/planner_combined{,_c5,_c8} (+ stage_A..D/) |
combined-objective curriculum runs at two chunk sizes |
cross_attention/planner_xa_{A,B,C,D} |
cross-attention planner staged bring-up |
cross_attention/planner_xa_{E1_all,E2_full} |
full bring-up runs of the cross-attention planner |
scratch/* |
smoke runs |
Architecture context: RecursivePlanner commits block k and never revisits it; the cross-attention
planner (planner_xa.py) keeps the causal influence mask but refines the whole plan jointly, like the
phase-1 controller. The controller-vs-planner crossover sits between rh=2 and rh=3 β always check the
rh a number was measured at before comparing (see experiments/README.md in the GitHub repo).
Manifold transfer
Checkpoints for the encoder-transfer study in manifold_transfer/ (see results/e3_full_findings.md):
e3_encoders/ holds best/final trained encoders (e.g. pusht-resnet18, tworooms-self-resnet18,
tworooms-self-resnet18-ext), adapters/ holds linear-map fits between latent spaces
({tworooms,reacher}_{linear,orthogonal,mlp}_{noB,withB}), and e3lite/ the E3-lite controls.
Usage
from huggingface_hub import snapshot_download
import torch
p = snapshot_download("SaltedLemon/lejepa-control")
wm_state = torch.load(f"{p}/world_models/pointmaze/lewm-pointmaze-r3g-p7/weights_epoch_40.pt",
map_location="cpu", weights_only=False)
Evaluate with the repo's harnesses (config/env flags in the GitHub docs):
# phase-1 controller, PushT headline
py scripts/eval_controller.py --controller <path>/controllers/pusht/ah_hold0.5/controller.pt \
--receding-horizon 1 --num-eval 50 --seed 42
# phase-2 planner
py lejepa_control_2/scripts/eval_planner.py --planner planner \
--checkpoint <path>/planners/recursive/planner_D10/planner.pt --horizon 5 --num-eval 50
Notes and caveats
.ptfiles are plaintorch.savearchives (dicts withstate_dicts), not safetensors.- All real-env numbers in this project use 50 held-out episodes, seed 42, identical start/goal pairs;
one episode = 2 percentage points β differences of 6β14 points are usually noise, so use
scripts/paired_stats.pyrather than raw deltas. - Excluded on purpose: intermediate training snapshots (world-model
_old_epochs/, E3snap*step checkpoints), regenerable latent caches, and third-partymodels--quentinll--*downloads. - The decoder is a visualization side channel β it is never part of the control loop.