lewm-resolution-collapse — auxiliary checkpoints
Small auxiliary networks trained for the generality experiments behind Resolution Collapse in Latent-Space CEM, a study of why CEM planners ranking candidates through a pretrained LeWM world model lose all ranking signal exactly at the terminal frame they optimize — and how to fix it.
Planning itself goes through the officially released LeWM checkpoints
(quentinll/lewm-tworooms, -pusht, -reacher, -cube) — nothing about
the base world model is touched. The two sets of networks here are the only
things actually trained as part of this work, both operating on LeWM's
frozen latent space.
What's in this repo
RND novelty nets ({env}_rnd.pt + {env}_rnd.json) — one Random
Network Distillation predictor per environment, trained on LeWM latents
from the real expert dataset. Used as the external, "someone else would
plausibly have proposed this" novelty baseline in the replacement-control
experiment. The accompanying JSON records in-distribution vs. out-of-
distribution reconstruction error and the resulting novelty ratio, e.g. for
TwoRoom the predictor's error on dimension-shuffled latents is ~114x its
error on real ones.
Latent-space state probes ({env}_{quantity}_probe.pt +
matching .json) — a linear probe trained to decode ground-truth physical
state (proprioception / object pose) directly from LeWM's latent
embedding, held-out R² reported in the JSON. On TwoRoom, agent position
decodes from the latent space at R² ≈ 0.99; on PushT, state decodes at
R² ≈ 0.73 on average across dimensions.
| file | env | role | held-out metric |
|---|---|---|---|
tworoom_rnd.pt |
TwoRoom | RND novelty | 113.8x shuffled / 9.5x scaled novelty ratio |
pusht_rnd.pt |
PushT | RND novelty | see pusht_rnd.json |
reacher_rnd.pt |
Reacher | RND novelty | see reacher_rnd.json |
cube_rnd.pt |
Cube | RND novelty | see cube_rnd.json |
tworoom_proprio_probe.pt |
TwoRoom | state probe | R² = 0.993 |
pusht_state_probe.pt |
PushT | state probe | R² = 0.732 |
Loading
import torch
rnd = torch.load("tworoom_rnd.pt", map_location="cpu", weights_only=False)
probe = torch.load("tworoom_proprio_probe.pt", map_location="cpu", weights_only=False)
See exp_generality.py (RND) and eval_objective_variants.py (probes) in
the code repository
for how each is built and consumed.
License
MIT, consistent with the base LeWM repository these checkpoints plan on top of.