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

  • .pt files are plain torch.save archives (dicts with state_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.py rather than raw deltas.
  • Excluded on purpose: intermediate training snapshots (world-model _old_epochs/, E3 snap* step checkpoints), regenerable latent caches, and third-party models--quentinll--* downloads.
  • The decoder is a visualization side channel β€” it is never part of the control loop.
Downloads last month

-

Downloads are not tracked for this model. How to track
Video Preview
loading

Paper for SaltedLemon/lejepa-control