EpiWM: LeWorldModel trained with epiplexity instead of SIGReg

Blog post · Code and results

EpiWM: these are the best checkpoints from replacing LeWorldModel's SIGReg anti-collapse term with the epiplexity score from EpiJEPA. The code, the full analysis and our replication of LeWM's SIGReg recipe are in the-puzzler/epijepa.

There is one folder per environment, laid out like LeWM's own releases:

Folder Environment λ Steps Seed paper-50 n=200 n=500 Released LeWM (paper-50 / n=200 / n=500)
tworoom/ TwoRoom 0.03 30k 1 100 100 100.0 86 / 85.0 / 82.8
pusht/ Push-T 0.1 60k 3 90 90.0 89.0 96 / 83.5 / 84.6
cube/ Cube (OGBench, single) 0.03 60k 3 74 75.5 72.4 68 / 63.0 / 66.0
reacher/ Reacher (DMC) 0.3 200k 1 76 72.0 76.6 52 / 62.0 / 60.8

These are planning success rates (%) with LeWM's CEM planner and evaluation (eval.py). The 3-seed means of the same recipe are TwoRoom 99.9, Push-T 88.5, Cube 71.5 and Reacher 73.5 on n=500. LeWM's SIGReg recipe retrained by us at the same budget (2 seeds) gives 87.9, 88.6, 65.5 and 62.2. The released LeWM column is our measurement with LeWM's released evaluation; Reacher evaluations use random-policy data, while the paper describes SAC-collected data. Each folder holds the seed with the best n=500 score; every seed's number is in analysis/scores/all_scores.csv.

Model

The architecture is LeWM's released one: a ViT-tiny encoder (patch 14, 224 px), an AdaLN transformer predictor with history 3, and MLP projectors, with 18M parameters. The only change is a non-affine BatchNorm on the projector output (module.ProjectorBN, included here). It is trained with

loss = ||pred(z_t, a_t) - z_t+1||^2 - lambda * S(z) / S0

where S is the epiplexity of the embeddings with respect to a frozen random CNN reservoir of the same frames.

Usage

# with the repo's worldmodel/ folder on PYTHONPATH (config.json refers to module.ProjectorBN)
hf download basilboy/epiwm --local-dir $STABLEWM_HOME/checkpoints/epiwm
cd epijepa/worldmodel
python eval.py --config-name=pusht.yaml policy=epiwm/pusht        # LeWM's eval: paper-50
import stable_worldmodel as swm
model = swm.wm.utils.load_pretrained("epiwm/pusht")   # resolved under $STABLEWM_HOME/checkpoints

The weights are plain state dicts (weights.pt, fp32, about 72 MB), saved with stable-worldmodel's save_pretrained (transformers-5 ViT key layout).

Analysis data (analysis/)

These are the data behind the figures in the GitHub repo's worldmodel/analysis/. Column descriptions are in that folder's README, and the scripts there regenerate everything.

Path What it is
embeddings/embeddings_<env>.csv 3000 random frames per environment: episode, step, true state, PCA and t-SNE coordinates for EpiWM and released LeWM
embeddings/trajectories_<env>.csv + embeddings/videos/<env>_ep<N>.mp4 one full episode per environment. Row step == i is frame i of the video, with its true state and coordinates in both models' PCA spaces
embeddings/animations/<env>_ep<N>_pca.mp4 the episode video side by side with the moving point in both PCA spaces
embeddings/pca_basis_<env>.npz, pca_variance_<env>.csv the PCA bases (mean, 50×192 components) and explained variance
probes/probes.json representation probes, both models, all environments; variant_search_*.json are the probes from method selection
training_logs/*.csv training and validation curves of the reported config, 3 seeds per environment
scores/all_scores.csv every planning result: all runs, checkpoints and evaluation sets
Downloads last month

-

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