TI-JEPA checkpoints
Trained weights for every arm in the TI-JEPA project: the matched-memory baseline, TI-JEPA itself, a recurrent (GRU) ablation, and the official ViT-Tiny + AdaLN-transformer scale runs, across three physics environments (InertiaBall, Pendulum, CartPole).
Code, environments, evaluation protocols, and demo videos live in the GitHub repo. The short version of why this exists: a single rendered frame can't carry velocity information no matter how large the encoder is — only an explicit finite-difference motion code, predicted alongside pose, fixes that. These checkpoints are the direct evidence for that claim across six environment/scale combinations.
Same starting frame, opposite initial velocity — CartPole, real dynamics forward.
What's in each folder
| folder | environment | scale | what you get |
|---|---|---|---|
inertia_ball/ |
InertiaBall | small CNN | baseline_k1, baseline_k3, baseline_rnn, tijepa |
pendulum/ |
Pendulum | small CNN | baseline_k1, baseline_k3, baseline_rnn, tijepa |
cartpole/ |
CartPole | small CNN | baseline_k1, baseline_k3, baseline_rnn, tijepa, plus the big-backbone/big-q ablations used in the scaling analysis |
official_scale/ |
CartPole, Pendulum | ViT-Tiny/14 + AdaLN | {env}_baseline_k1, {env}_baseline_k3, {env}_tijepa |
official_scale_seed2/ |
CartPole, Pendulum | ViT-Tiny/14 + AdaLN | same arms, independent training seed, used to confirm the official-scale results reproduce |
baseline_k1 is the single-frame-target baseline with no history (k=1,
the setting Corollary 1 is actually about). baseline_k3 is the same
recipe given a 3-frame window — the fair, memory-matched opponent for
TI-JEPA everywhere else in the paper. baseline_rnn swaps the aggregator
for a GRU at the same parameter footprint, for the RSSM-style comparison.
Each .pt ships with a _history.json of its training curve.
Loading a checkpoint
Needs the ti_jepa package from the GitHub repo:
import torch
from ti_jepa.models import TIJEPAEncoder, TIJEPAPredictor
ckpt = torch.load("cartpole/tijepa.pt", map_location="cpu", weights_only=False)
encoder = TIJEPAEncoder(**ckpt["args"])
encoder.load_state_dict(ckpt["encoder"])
predictor = TIJEPAPredictor(**ckpt["args"])
predictor.load_state_dict(ckpt["predictor"])
For the official-scale checkpoints (official_scale/, official_scale_seed2/),
swap in ti_jepa.vit_backbone and ti_jepa.official_predictor instead — see
scripts/run_official_scale_cartpole.sh in the GitHub repo for the exact
loading path used to produce the paper's numbers.
Numbers these checkpoints reproduce
cartpole/baseline_k1,cartpole/baseline_k3,cartpole/tijepa: the closed-loop planning gap (64% lower final distance for TI-JEPA, $p=5.1\times10^{-15}$) and the probe-ladder/kill-experiment table.official_scale/*: the memoryless baseline's branch-separation ratio at exactly0.000on both CartPole and Pendulum, and the CartPole velocity probe moving from $r=0.30$ (small CNN) to $r=0.91$ (ViT-Tiny) once you swap the backbone and keep everything else fixed.official_scale_seed2/*: the same two numbers above, from an independent seed, confirming they aren't a training-noise artifact.
License
MIT, same as the code.