--- license: mit tags: - world-model - jepa - planning - robotics - lewm --- # EpiWM: LeWorldModel trained with epiplexity instead of SIGReg [Blog post](https://the-puzzler.github.io/blog/epiwm/) · [Code and results](https://github.com/the-puzzler/epijepa/tree/main/worldmodel) EpiWM: these are the best checkpoints from replacing [LeWorldModel](https://github.com/lucas-maes/le-wm)'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](https://github.com/the-puzzler/epijepa/tree/main/worldmodel). 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](https://huggingface.co/basilboy/epiwm/blob/main/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 ```bash # 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 ``` ```python 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_.csv` | 3000 random frames per environment: episode, step, true state, PCA and t-SNE coordinates for EpiWM and released LeWM | | `embeddings/trajectories_.csv` + `embeddings/videos/_ep.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/_ep_pca.mp4` | the episode video side by side with the moving point in both PCA spaces | | `embeddings/pca_basis_.npz`, `pca_variance_.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 |