temporal-flow-world-model checkpoints
Checkpoints for temporal-flow-world-model: an encoder trained from video alone, so that the change dz between two frames paints the optical flow between them through a pointwise read, F_p = M(dz)ᵀ g_p.
Each file is {"model": state_dict, "args": FlowWM constructor args, "meta": {world, source run, step, training}}.
Load one with the repo:
from tfwm import models
model, encode = models.load({"type": "flow", "checkpoint": "hf://MasonJK99/temporal-flow-world-model/pusht/tfwm_rf91_lags1248.pt"})
| file | world | encoder | receptive field | lags | training |
|---|---|---|---|---|---|
pusht/tfwm_rf91_lags1248.pt |
Push-T | ViT-tiny, 224 px | 91 px (φ kernels 7, 7, 19) | {1, 2, 4, 8} | 76 210 steps (SMWM's budget), batch 256, AdamW 1e-4 |
pusht/ablations/tfwm_rf91_lags12.pt |
Push-T | ViT-tiny | 91 px | {1, 2} | the same; the lag-set ablation |
reacher/tfwm_rf91_lags1248.pt |
Reacher | ViT-tiny | 91 px | {1, 2, 4, 8} | 76 210 steps, batch 256, AdamW 1e-4 |
tworoom/tfwm_rf7_lags12.pt |
TwoRoom | ViT-tiny, 224 px | 7 px (φ kernels 3, 3, 1) | {1, 2} | 25 600 of a 29 200-step schedule (SMWM's budget), batch 256, AdamW 1e-4; an earlier recipe, see below |
toy/independent_rf7.pt |
toy dot world: two independent dots | CNN, 64 px | 7 px (3, 3) | {1, 2} | 50 epochs, Adam 5e-4 |
toy/coupled_rf7.pt |
toy dot world: one coupled pair | CNN, 64 px | 7 px (3, 3) | {1, 2} | the same |
toy/combined_rf7.pt |
toy "combined" dot world | CNN, 64 px | 7 px (3, 3) | {1, 2} | 50 epochs, Adam 5e-4 |
toy/combined_rf19.pt |
toy "combined" | CNN | 19 px (3, 9) | {1, 2} | the same |
toy/combined_rf27.pt |
toy "combined" | CNN | 27 px (3, 13) | {1, 2} | the same |
toy/sprite_rf7.pt |
toy sprite world (one rotating arrow) | CNN, 64 px | 7 px (3, 3) | {1, 2} | 50 epochs, Adam 5e-4 |
toy/sprite_rf19.pt |
toy sprite | CNN | 19 px (3, 9) | {1, 2} | the same |
toy/sprite_rf27.pt |
toy sprite | CNN | 27 px (3, 13) | {1, 2} | the same |
index.json lists every file's args and training settings. The Push-T model was trained in two segments (to ¼×, then
resumed to 1× on the same cosine schedule). The TwoRoom file is the research chart, trained with an earlier recipe than Push-T and Reacher: receptive field 7 px instead of 91, lags {1, 2}, stopped at step 25 600 of 29 200, and a flow constant of 40.6865 px from a sampled measurement (the lag-1 maximum over its training windows is 40.8993). A TwoRoom model with the Push-T / Reacher recipe (RF 91) will be added after retraining.
SMWM baselines
SMWM world models we trained ourselves (inverse-dynamics objective, seed 0), used as the comparison baseline. Each is a
PyTorch Lightning checkpoint saved as is (state_dict plus optimizer state), and the .yaml file with the same name
holds its full training config.
| file | world | λ (inverse weight) | training |
|---|---|---|---|
smwm/smwm_reacher_inverse_lambda5_seed0.ckpt |
Reacher | 5 | 10 epochs (67 500 steps), batch 256, AdamW 1e-4, ViT-tiny 224 px |
smwm/smwm_cube_inverse_lambda1_seed0.ckpt |
OGBench Cube | 1 | 10 epochs (67 500 steps), the same |
smwm/smwm_tworoom_inverse_lambda0p1_seed0.ckpt |
TwoRoom | 0.1 | 10 epochs (29 200 steps), the same |
smwm/smwm_tworoom_inverse_lambda0p1_seed0_v2.ckpt |
TwoRoom | 0.1 | 10 epochs (29 200 steps), the same; use this one for TwoRoom (see below) |
smwm/smwm_pusht_inverse_lambda30_seed0.ckpt |
Push-T | 30 | 10 epochs (76 210 steps), the same |
The research run names were <world>_inverse_lambda_<λ>_seed0, except for Reacher, whose run was named
reacher_inverse_lambda_1_seed0 but trained with λ = 5.
TwoRoom: use _v2. smwm_tworoom_inverse_lambda0p1_seed0.ckpt (sha256 effaea99…) encodes the agent's y
linearly but folds its x into a zigzag over two principal components (turning at the wall), so x is not linear in the
embedding: on SMWM's eval split its projector embedding has PCA explained-variance ratios 0.744 / 0.111 / 0.104 and a
linear probe recovers x with R² ≈ 0.36. smwm_tworoom_inverse_lambda0p1_seed0_v2.ckpt (sha256 570282…) is another
run of the same code (SMWM planning/train.py, commit 9d22bcc) and config, which gives the two-component spectrum
SMWM reports (0.512 / 0.455 / 0.009) and x, y R² 1.00; a third local run agrees (0.490 / 0.480). The old file is kept
for reproducibility of earlier numbers.