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 exactly 0.000 on 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.

Downloads last month

-

Downloads are not tracked for this model. How to track
Video Preview
loading