epiwm / README.md
basilboy's picture
Link EpiWM blog and correct score data reference
326fac0 verified
|
Raw History Blame Contribute Delete
4.12 kB
---
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_<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 |