This post extends the I-JEPA tutorial to video. We'll implement V-JEPA — the video version of I-JEPA — and train it on Moving MNIST in about 194 lines of PyTorch.
- Source:
vjepa.py - Paper: Bardes et al., V-JEPA: Latent Video Prediction for Visual Representation Learning (arXiv 2404.08471)
If you already followed the I-JEPA post, V-JEPA is a one-sentence change: swap 2D image patches for 3D video tubelets, swap rectangle masks for spatial-tube masks. Everything else — the EMA target encoder, the predictor with mask-token queries, the latent-space loss — survives unchanged.
This post focuses on the two things that do change: the tubelet patcher and the tube masking.
We train on Moving MNIST — 10,000 video clips, 20 frames each, 64×64 pixels. Two MNIST digits drift across each frame and bounce off the walls. We subsample to 10 frames per clip to keep things fast.
The video is cut into tubelets: 3D patches of shape (2, 16, 16) — 2 frames × 16×16 pixels. With 10 frames and 64×64 spatial resolution we get a 5×4×4 = 80-token grid per video.
spatial 4x4 per tubelet
┌────┬────┬────┬────┐
│ p0 │ p1 │ p2 │ p3 │ ← one of 5 temporal slices
├────┼────┼────┼────┤ each slice covers 2 frames
│ p4 │ p5 │ p6 │ p7 │
├────┼────┼────┼────┤
│ p8 │ p9 │p10 │p11 │
├────┼────┼────┼────┤
│p12 │p13 │p14 │p15 │
└────┴────┴────┴────┘
I-JEPA masks rectangles in 2D space. V-JEPA masks tubes — pick a spatial rectangle, then extend it through every temporal slice. A masked tube hides the same patch positions across all frames.
The official recipe uses two mask groups per batch:
- short-range: 8 small tubes at spatial scale 0.15 each
- long-range: 2 large tubes at spatial scale 0.7 each
Both groups span the full temporal axis. Each group produces its own (context, targets) pair, and the predictor is called once per group.
The mask-grid figures below are real outputs from one training step:
Top row: the original video frames (we show 5 evenly-spaced timesteps). Middle row: the context — what the encoder sees. Bottom row: the targets — what the predictor must reconstruct. Notice that the masks are identical across timesteps: that's the "tube" property. Short-range tubes nibble the corners; long-range tubes hide most of the frame.
Animated, the tube property is clearer — the mask stays put in space while the digits move:
Original frames on the left, context (visible to the encoder) in the middle, targets (what the predictor reconstructs) on the right. The black rectangles are the tube footprints — same (row, col) cells at every timestep.
A 3D ViT. A single Conv3d does both patchification and linear embedding in one op:
class VideoEncoder(nn.Module): # f_theta (context encoder)
def __init__(self, num_frames=10, t_patch=2, img_size=64, patch_size=16,
in_chans=1, dim=128, depth=6, heads=4):
super().__init__()
self.t_grid = num_frames // t_patch # 10 / 2 = 5 temporal slices
self.s_grid = img_size // patch_size # 64 / 16 = 4x4 spatial grid
self.n_patches = self.t_grid * self.s_grid * self.s_grid # 5*4*4 = 80 tokens
self.t_patch = t_patch; self.patch_size = patch_size; self.dim = dim
self.tubelet_proj = nn.Conv3d( # (t_patch, patch, patch) kernel
in_chans, dim,
kernel_size=(t_patch, patch_size, patch_size),
stride=(t_patch, patch_size, patch_size))
self.register_buffer("pos", sincos_3d( # 3D pos = 1D-t concat 2D-spatial
self.t_grid, self.s_grid, self.s_grid, dim))
self.blocks = nn.ModuleList([Block(dim, heads) for _ in range(depth)])
self.norm = nn.LayerNorm(dim, eps=1e-6)
def forward(self, videos, idx=None):
tokens = self.tubelet_proj(videos).flatten(2).transpose(1, 2) # (B, 80, dim)
if idx is None: # full pass: encode all 80
idx = torch.arange(tokens.size(1), device=videos.device).expand(tokens.size(0), -1)
x = tokens + self.pos[idx]
else: # subset: encode only context tokens
x = tokens.gather(1, idx.unsqueeze(-1).expand(-1, -1, tokens.size(-1))) + self.pos[idx]
for blk in self.blocks: x = blk(x)
return self.norm(x)Two differences vs. I-JEPA's Encoder:
Conv3dinstead ofConv2d: the kernel spans 2 frames × 16 spatial pixels.- 3D positional embedding:
sincos_3dconcatenates a 1D temporal sin-cos for time with a 2D sin-cos for spatial. The paper just says "3D sin-cos" without prescribing how the dimensions are split; the 25/75 temporal/spatial split here follows the official repo.
Everything downstream of this — the EMA target encoder, the predictor with mask-token queries — is structurally identical to I-JEPA.
The mask sampler runs once per batch and produces two groups:
MASK_GROUPS = [("short", 8, 0.15), ("long", 2, 0.7)] # (label, n_blocks, spatial_scale)
def sample_vjepa_masks(B, t_grid, s_grid, rng=None, min_ctx=8, ar_range=(0.75, 1.5)):
rng = rng or random
min_visible_cells = max(1, math.ceil(min_ctx / t_grid))
groups = []
for label, n_blocks, scale in MASK_GROUPS:
h, w = _bsize(s_grid, scale, rng.uniform(*ar_range))
ctx_spatial, pred_spatial = [], []
for _ in range(B):
masked, visible = _sample_spatial_tubes(n_blocks, h, w, s_grid, rng, min_visible_cells)
ctx_spatial.append(sorted(visible))
pred_spatial.append(sorted(masked))
# Keep batch tensors rectangular without breaking the tube property:
# trim whole spatial cells, then expand every selected cell across all time steps.
Lc, Lp = min(len(c) for c in ctx_spatial), min(len(p) for p in pred_spatial)
groups.append({
"label": label, "n_blocks": n_blocks, "block_hw": (h, w),
"ctx": [_expand_tubes(sorted(rng.sample(c, Lc)), t_grid, s_grid) for c in ctx_spatial],
"pred": [_expand_tubes(sorted(rng.sample(p, Lp)), t_grid, s_grid) for p in pred_spatial],
})
return groupsThe "tube" structure comes from sampling spatial cells first and expanding each selected cell across every temporal slice. The sampler never moves target tokens back into context; if it must trim for rectangular batch tensors, it trims whole spatial tubes rather than individual tokens. On the tiny Moving-MNIST grid this is still a toy approximation of the paper's large-grid masking, but the key invariant — same spatial footprint over time — is preserved.
The objective is the same as I-JEPA's, just applied per mask group:
with
Code map in train():
with torch.no_grad():
full = F.layer_norm(tgt_enc(videos), (D,)) # LN(s_y); no_grad = stop-gradient
per = {} # per-group losses
for g in groups:
ci = torch.tensor(g["ctx"], device=device)
pi = torch.tensor(g["pred"], device=device)
tgt = full.gather(1, pi.unsqueeze(-1).expand(-1, -1, D)) # [LN(s_y)]_{B_g}
per[g["label"]] = (pred(ctx_enc(videos, ci), ci, pi) - tgt).abs().mean()
# L1: mean |hat_s_y - s_y|
loss = sum(per.values()) / len(per) # (1/|G|) sum over groupsTwo paper-vs-code notes carry over from the I-JEPA tutorial:
- LayerNorm on targets is a code-only detail.
- L1 (
abs().mean()) is what the official config uses (loss_exp: 1.0). I-JEPA's code uses smooth-L1; V-JEPA's uses straight L1.
def train(epochs=5, batch_size=32, lr=3e-4, wd=0.05,
ema_start=0.998, ema_end=1.0, device=None):
ds = MovingMNISTVideos(num_frames=10)
loader = DataLoader(ds, batch_size=batch_size, shuffle=True, drop_last=True)
ctx_enc = VideoEncoder().to(device) # f_theta
tgt_enc = copy.deepcopy(ctx_enc).to(device) # f_theta_bar
for p in tgt_enc.parameters(): p.requires_grad_(False)
pred = Predictor(t_grid=ctx_enc.t_grid, s_grid=ctx_enc.s_grid).to(device)
opt = torch.optim.AdamW(param_groups([ctx_enc, pred], wd), lr=lr)
total = epochs * len(loader); rng = random.Random(0); step = 0
for epoch in range(epochs):
for videos in loader:
videos = videos.to(device); B = videos.size(0)
groups = sample_vjepa_masks(B, ctx_enc.t_grid, ctx_enc.s_grid, rng=rng)
with torch.no_grad():
full = F.layer_norm(tgt_enc(videos), (ctx_enc.dim,)) # LN(s_y)
per = {}
for g in groups: # short + long groups
ci = torch.tensor(g["ctx"], device=device)
pi = torch.tensor(g["pred"], device=device)
tgt = full.gather(1, pi.unsqueeze(-1).expand(-1, -1, ctx_enc.dim))
per[g["label"]] = (pred(ctx_enc(videos, ci), ci, pi) - tgt).abs().mean()
loss = sum(per.values()) / len(per) # average over groups
opt.zero_grad(); loss.backward(); opt.step()
m = ema_start + (ema_end - ema_start) * (step / max(1, total - 1)) # 0.998 -> 1.0
ema_update(tgt_enc, ctx_enc, m); step += 1 # update f_theta_bar- Learning rate —
3e-4, constant. (No warmup/cosine — V-JEPA's training schedule is short enough at our scale.) - Weight decay —
0.05, 2D+ params only. - Batch size —
32(videos are bigger than CIFAR images). - Epochs —
5. - EMA momentum —
0.998 → 1.0linear. Matches thevitl16.yamlreference config. - Mask groups —
[("short", 8, 0.15), ("long", 2, 0.7)]. Long scale stays0.7; the sampler trims whole spatial tubes only when needed for rectangular batches. - Encoder — ViT-tiny: dim 128, depth 6, heads 4. Tubelet
(2, 16, 16). - Predictor — dim 64, depth 4, heads 4.
python vjepa.py # train only
python vjepa_extras.py # train + write mask grids + per-group loss curves5 epochs of training take about a minute on an M-series Mac.
The loss is reported per mask group. Long-range tubes are harder (less context) so the long curve sits slightly above the short curve:
Both curves drop steeply for the first ~100 steps, then plateau in the 0.05–0.06 L1 range. The plateau is the noise floor: with random Conv3d features as targets and tight context, there's a limit to how well a tiny predictor can hit them.
Moving MNIST has no class labels, so we don't run a linear probe here. (The I-JEPA tutorial does that on CIFAR-10.) The proof that the algorithm trains is the loss curve plus the masks rendering correctly.
Five things V-JEPA gets right, in roughly the order they matter:
-
Predict embeddings, not pixels. Inherited from I-JEPA. The encoder is free to discard pixel-level texture and lighting and spend capacity on whatever's predictable across spacetime. There is no pixel decoder anywhere in the model.
-
Tubes, not random patches. Tube masking — a spatial block held constant across every temporal slice — kills the cheapest temporal shortcut (copy from the previous frame). The predictor must use spatially distant context, which pushes the encoder toward higher-level features.
-
Two mask scales, not one. The 8×0.15 short-range group covers small targets with rich context (easy locally, lots of detail). The 2×0.7 long-range group hides most of the frame and forces global reasoning. Training on both at once gives a curriculum-by-batch effect — the encoder has to do both jobs.
-
EMA target encoder is load-bearing. The target encoder is a slow-moving copy of the context encoder. Late in training it tracks a recent-history-mean of the context encoder rather than its current state. Without this, the loss collapses — the network learns to output whatever the target encoder happens to output, which is whatever the network just output.
-
L1, not L2. The paper switched from squared error to absolute error and reports it's more stable. Latent-space targets carry occasional outliers; L1 cares less about them.
The first three are about what you predict; the last two are about how you train. Both matter equally — strip any one of them and the method stops working.
vjepa2.pyadds a second post-training phase: freeze the V-JEPA encoder, train an action-conditioned predictor that does latent-state rollouts. That's V-JEPA 2-AC — the world-model variant Meta uses for robotic planning.cjepa.pykeeps the latent-prediction core but masks at the object-trajectory level instead of patch-rectangle level, using an identity anchor and no EMA. That's C-JEPA.
The masking strategy gets stranger across the family; the predict-embeddings-not-pixels core stays.




