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.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support