EpiWM: LeWorldModel trained with epiplexity instead of SIGReg
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 |