PushT world models
Action-conditioned world models from S. Bateman et al., Training Controllable World Models via Action Influence Maximization, CoRL 2026.
PushT provides a controlled wind tunnel for studying how data scale and composition affect world-model prediction and generalization. These models compare large passive collections with smaller, task-specific active exploration collections using a shared architecture. The paired Future Sight PushT dataset provides the training and evaluation data.
Ground truth (left) and model predictions (right), using identical recorded actions. Thirteen hand-selected examples show Recovery1 + transferred AIM500, PPO1 + AIM500 (fixed reward estimate), and Recovery100 on BC, BC-failure, and PushAway evaluation cases. Model names appear in the video; each clip shows 33 predicted transitions at 256 × 256 resolution. These examples illustrate model behavior; they are not an aggregate benchmark or a ranking of models.
Models at a glance
16 checkpoints · 256 × 256 RGB · approximately 580 MB each · 9.28 GB total
- Inputs: RGB observations and continuous 6D end-effector delta actions.
- Outputs: predicted RGB frames at 256 × 256 resolution.
- Format: native Orbax inference bundles, including the frozen VAE.
- Weights: exponential moving average (EMA); no optimizer or training-resume state.
Choose a model
PPO1 and Recovery1 use the first 1% of the corresponding training metadata;
PPO100 and Recovery100 use the full corpus. The 500 suffix denotes the
additional collection budget. AIM500 uses 100 trajectories from each of five
collection iterations, with a 2:1 PPO-to-AIM sampling mixture.
| Model | Training step | Download |
|---|---|---|
| PPO1 + AIM500 | 36,000 | 580.1 MB |
| PPO1 | 26,000 | 580.0 MB |
| PPO1 + AIM500 (fixed reward estimate) | 42,000 | 580.0 MB |
| PPO1 + AIM500 (MI only) | 39,000 | 580.0 MB |
| PPO1 + AIM500 (novelty only) | 42,000 | 580.1 MB |
| PPO1 + DADS500 | 29,000 | 580.0 MB |
| PPO1 + DIAYN500 | 34,000 | 580.1 MB |
| PPO1 + ICM500 | 29,000 | 579.9 MB |
| PPO1 + METRA500 | 26,000 | 580.0 MB |
| PPO1 + random walk500 | 28,000 | 580.0 MB |
| PPO1 + Recovery500 | 27,000 | 580.1 MB |
| PPO1 + uniform random500 | 33,000 | 580.0 MB |
| PPO100 | 44,000 | 580.1 MB |
| Recovery1 | 37,000 | 580.0 MB |
| Recovery1 + transferred AIM500 | 42,000 | 580.1 MB |
| Recovery100 | 40,000 | 580.1 MB |
The fixed reward estimate ablation uses AIM data without reward re-estimation. Novelty-only and MI-only models ablate the reward components. Transferred AIM adds exploration collected from the PPO-based setting to a Recovery1 training mixture. See the experiment-specific recipes in Future Sight for exact subset and mixture definitions; the budget suffix alone does not specify sampling probabilities.
Download and load one model
Run from a Future Sight checkout with its Pixi environment installed and one
visible GPU, for example CUDA_VISIBLE_DEVICES=0 pixi run python example.py
after saving the following snippet as example.py. The example
uses existing huggingface_hub and Future Sight APIs. No Transformers model
conversion or trust_remote_code is required.
import os
from pathlib import Path
import jax
from huggingface_hub import snapshot_download
from future_sight.models.components.sharding import create_device_mesh
from future_sight.models.loading import load_world_model
cache = Path(os.environ.get("FUTURE_SIGHT_CACHE_DIR", "cache")).expanduser()
snapshot = Path(snapshot_download(
repo_id="smbml/future-sight-pusht-models",
revision="main", # Use a full HF commit hash to pin an experiment.
allow_patterns=["models.json", "models/ppo1-aim500/*"],
cache_dir=cache / "models",
))
jax.devices() # Initialize JAX before any TensorFlow-backed data loader.
mesh = create_device_mesh(data=1, model=1)
model, manifest = load_world_model(str(snapshot / "models/ppo1-aim500"), mesh)
print(manifest["weights"]["source_training_step"])
The example explicitly honors FUTURE_SIGHT_CACHE_DIR; the HF library itself
does not interpret that Future Sight variable. Change the folder in
allow_patterns and the load path to select another model. HF uses your existing
login when available; these public files are also downloadable anonymously.
Completed files are reused from the HF cache, and Orbax reads them in place.
The restored model can be used with Future Sight's Python inference session, standalone PushT evaluator, or local interactive application. Model-specific HF CLI integration and a hosted interactive Space are not provided by this repo.
Architecture and provenance
All models use the latent diffusion forcing B/2 architecture: 12 transformer layers, width 768, 12 attention heads, and 2 × 2 latent patches. The VAE maps 256 × 256 images to 32 × 32 × 4 latents. The published manifests specify the exact action bounds, x-prediction / v-space training objective, BF16 runtime, 128 inference steps, four-frame sequence window and three cached predecessors.
Each bundle includes a versioned manifest with its training step, source commit,
and hashes of every restored tensor. models.json additionally
records downloadable file sizes and SHA-256 digests. Public metadata removes
private storage paths and uses the future_sight import namespace; checkpoint
payloads and tensor identities are unchanged. Historical source repository names
remain in provenance. Private training-contract files are not distributed.
Our base VAE is
stabilityai/sd-vae-ft-mse,
used through pcuenq's Flax conversion
and included as a frozen StableVAE in every inference bundle. Both upstream
model cards specify MIT.
The original upstream HF revision was not recorded; the bundle records the
converted VAE's content identity. The trained Future Sight weights are also
released under the MIT license. See upstream attribution
for the bundled VAE.
Validation and limitations
All 16 source checkpoints passed strict restoration and finite-output checks on two selected evaluation cases each. They also ran in the local interactive model switcher. Published model files are listed above. The PPO1 + AIM500 bundle additionally passed an authenticated Hugging Face upload/download round trip, file-hash verification, offline cache reuse, and exact raw-frame agreement with its local source on BC and PushAway when independently restored states use the same compiled inference graph. Separately compiled GPU runs were not pixel-identical; this check does not establish bitwise reproducibility across runs. Other models have local inference checks and remote file verification, not separate HF round-trip inference runs.
The local checks used JAX/JAXLIB 0.10.2, Flax 0.12.7 and Orbax 0.12.1 on an NVIDIA RTX PRO 6000 Blackwell GPU. They do not establish compatibility with every hardware or dependency combination. Small local checks and videos do not reproduce the complete paper evaluation. Samples may drift or produce implausible dynamics, particularly over longer horizons or unfamiliar actions.
Citation
If you use these models or the paired dataset, please cite:
@inproceedings{bateman2026aim,
title = {Training Controllable World Models via Action Influence Maximization},
author = {Bateman, Samuel M. and Yin, Tenny and Zheng, Chongyi and Huang, Lei
and Wang, Brian and Eysenbach, Ben and Fisac, Jaime Fern\'andez and Shah, Dhruv},
booktitle = {Conference on Robot Learning},
year = {2026}
}
Paper link forthcoming.