I-JEPA: a context encoder, an averaged target encoder, a predictor.
Parameters
configIJEPAConfigAttributes
Notes
Reference: Assran, Mahmoud, et al., "Self-Supervised Learning from Images with a Joint-Embedding Predictive Architecture", CVPR 2023 (arXiv:2301.08243), Section 3 and Appendix A.1.
The target encoder is not updated by the optimiser. Call
update_target after each step, with
momentum for the value the schedule asks for.
Two details come from the released code rather than the paper: the
target representations are layer-normalised (without affine terms)
before the loss reads them, and that loss is smooth L1 where the paper
writes . config.objective selects between them.
Examples
>>> import lucid
>>> from lucid.models.vision.ijepa import IJEPAConfig, IJEPAModel
>>> config = IJEPAConfig(
... image_size=32, patch_size=8, dim=16, depth=1, num_heads=2,
... predictor_dim=8, predictor_depth=1, min_keep=1,
... target_scale=(0.1, 0.2), context_scale=(0.8, 1.0))
>>> model = IJEPAModel(config)
>>> out = model(lucid.rand(2, 3, 32, 32))
>>> out.prediction.shape == out.target.shape
True
>>> float(out.loss.item()) >= 0.0
True
The representation a downstream task uses is the target encoder's,
average-pooled:
>>> model.encode(lucid.rand(2, 3, 32, 32)).shape
(2, 16)Used by 2
Constructors
1Properties
1Instance methods
6The representation downstream tasks use.
Appendix A.1 evaluates the target encoder, average-pooled over patches — not the context encoder, which is a silent accuracy loss.
Parameters
(B, C, H, W).Returns
Tensor(B, dim).
Notes
Not wrapped in no_grad: the target encoder's parameters are
frozen already, so nothing accumulates into them, and a method a
probe reads through should not sever the gradient a probe needs.
The stop-gradient that matters is in forward, where the
targets are built.
forward(x: Tensor)Run one pretraining step's worth of computation.
Parameters
(B, C, H, W) at config.image_size.Returns
IJEPAOutputThe loss and what it was computed from.
The average-pooled target encoder, which is encode.
A backbone's features here are the representation the paper evaluates — the target encoder's, not the context encoder's.
The target encoder's momentum at a training step.
Parameters
stepinttotal_stepsintReturns
floatLinear between config.ema[0] and config.ema[1] — 0.996
to 1 at the paper's values, so the target stops moving exactly
as training ends.
Everything an optimiser should be given.
Returns
list of ParameterThe context encoder and the predictor. The target encoder is excluded: it follows by moving average, and handing it to an optimiser would train the thing that defines the target.
Move the target encoder toward the context encoder.
Call it after the optimiser step: this writes parameters in place, and doing that while a graph that read them is alive would sever the two.
Parameters
momentumfloatmomentum.