V-JEPA: a context encoder, an averaged target encoder, a predictor.
Parameters
configVJEPAConfigAttributes
Notes
Reference: Bardes, Adrien, et al., "Revisiting Feature Prediction for Learning Visual Representations from Video", arXiv:2404.08471, 2024, Sections 3.1–3.3.
Clips are (B, T, C, H, W) with T = config.num_frames. Call
update_target after each optimiser step, with the value
momentum gives.
The target features are layer-normalised without affine terms before any block is taken from them — that is in the released code and not in the paper, and it is the usual way a reimplementation fails.
Examples
>>> import lucid
>>> from lucid.models.vision.vjepa import VJEPAConfig, VJEPAModel
>>> config = VJEPAConfig(
... image_size=32, patch_size=8, tubelet_size=2, num_frames=4,
... dim=24, depth=1, num_heads=2, predictor_dim=12,
... predictor_depth=1, short_range_blocks=2, long_range_blocks=1)
>>> model = VJEPAModel(config)
>>> out = model(lucid.rand(2, 4, 3, 32, 32))
>>> out.short_prediction.shape == out.short_target.shape
True
>>> float(out.loss.item()) >= 0.0
True
The representation a downstream task reads is the target encoder's:
>>> model.encode(lucid.rand(2, 4, 3, 32, 32)).shape
(2, 24)Used by 2
Constructors
1Properties
1Instance methods
7forward(x: Tensor)Run one pretraining step's worth of computation.
Parameters
(B, T, C, H, W) at config.image_size.Returns
VJEPAOutputThe averaged loss, and what each collection was asked and answered.
The average-pooled target encoder, which is encode.
The target encoder's momentum at a training step.
Parameters
stepinttotal_stepsintReturns
floatLinear from config.ema[0] toward config.ema[1] over
total_steps * config.ema_schedule_scale steps. The
released runs stretch the schedule by a quarter and stop
early, so the momentum is still short of 1 when training
ends — reproducing it with a schedule the length of the run
is a different experiment.
The target encoder's whole feature map, (B, N, dim).
What the attentive probe reads, and what average-pooling in
encode throws away.
Everything an optimiser should be given.
Returns
list of ParameterThe context encoder and the predictor. The target encoder follows by moving average and must not be optimised.
Move the target encoder toward the context encoder.
Call it after the optimiser step: this writes parameters in place.
Parameters
momentumfloatmomentum.