A Joint-Embedding Predictive Architecture (JEPA) world model that learns the dynamics of a two-room navigation environment entirely in representation space — no pixel reconstruction — in 88,997 parameters.
Given an initial observation and a sequence of actions, the model unrolls predicted latent states
many steps into the future. Representation quality is measured by probing: a 2-layer MLP is trained
to recover the agent's true (x, y) coordinates from the predicted latents.
| Probe | MSE ↓ | What it tests |
|---|---|---|
probe_normal |
4.40 | In-distribution trajectories |
probe_wall |
7.58 | Learning to stop at walls / pass through doors |
probe_wall_other |
7.62 | Generalization to wall positions never seen in training |
probe_expert |
84.84 | Long-horizon rollout |
Trainable parameters: 88,997. Training data: 2.5M frames of exploratory trajectories.
Final project for CSCI-GA 2572 (Deep Learning, NYU Courant — Yann LeCun / Alfredo Canziani).
JEPA training minimizes the distance between predicted and target representations. The trivial solution is for the encoder to output a constant — energy goes to zero, and the representation carries no information. Reconstruction-based fixes were disallowed, so preventing collapse had to come from regularization and architecture.
We attack it from two directions.
1. VICReg regularization (loss.py)
Standard variance–invariance–covariance regularization: a hinge on the per-dimension standard deviation of the batch keeps every latent dimension active, and an off-diagonal covariance penalty decorrelates them.
2. Action regularization — an inverse-dynamics auxiliary objective (ActionRegularizer2D)
This is the piece we found mattered most. A small conv head consumes the difference between the predicted next latent and the current latent, and must predict the action that caused it:
â = f(s̃ₙ − sₙ₋₁) loss += λ · ‖â − aₙ₋₁‖²
A collapsed representation makes this impossible — if all latents are identical, the difference is zero and carries no action information. So the auxiliary loss is not just an anti-collapse penalty; it actively forces the latent space to encode the environment's dynamics rather than any arbitrary information-preserving code.
observation (2, 65, 65) channel 0: agent channel 1: walls + border
│
▼
Encoder2D stacked stride-2 convs, channels 16 → 32 → ... (capped at 256),
│ 1×1 conv to out_c. Number of blocks derived from the target
▼ latent side length, so resolution is a config knob.
latent (C', H, W) latents stay SPATIAL — no flattening
│
├──────────────► Predictor2D ──► s̃ₙ conditioned on action aₙ₋₁
│ │
│ ├──► VICReg vs. target encoder output
│ └──► ActionRegularizer2D ──► â (inverse dynamics)
▼
(recurrent: feed s̃ₙ back in for multi-step rollout)
Two decisions worth calling out:
- Spatial latents. Most JEPA implementations flatten to a vector. Keeping the latent as
(C', H, W)preserves the spatial correspondence between the latent and the room layout, which matters because the model must perceive where the wall and door are in each episode — layouts differ across trajectories. - Teacher forcing as a curriculum. Training starts teacher-forced (predict one step from ground-truth encodings) and moves to fully recurrent rollout, which stabilizes early training before compounding error becomes the dominant signal.
models.py contains ~15 variants explored along the way, all runnable via config:
| Variant | Idea | Outcome |
|---|---|---|
JEPA |
Baseline recurrent JEPA, flat latents | Collapsed without strong regularization |
AdversarialJEPA / ...WithRegularization |
Discriminator pushes latents off a degenerate distribution | Trained, but unstable and parameter-hungry |
InfoMaxJEPA |
Information-maximization objective | Worked; harder to tune than VICReg |
ActionRegularizationJEPA |
Inverse-dynamics auxiliary loss, flat latents | Strong collapse resistance |
ActionRegularizationJEPA2D |
The above with spatial latents | Clear improvement on wall probes |
ActionRegularizationJEPA2DFlexibleEncoder |
Swappable timm backbone (ResNet-18, SE-ResNeXt-26) |
Final submission |
...2Dv0 … ...2Dv3 |
Encoder depth/width, separate wall encoder, embedding dim sweeps | Incremental |
Configs for every variant live in config/.
pip install -r requirements.txt
# train (config selected in configs.py / config/*.yaml)
python -m main
# evaluate: probes the trained model and prints MSE per validation set
python main.pyDataset paths default to the NYU HPC layout (/scratch/DL24FA/...); override data_path in the
config to point at your own copy. States have shape
(num_trajectories, trajectory_length, 2, 64, 64); actions are (Δx, Δy) vectors.
Submitted weights: model_weights.pth. Reported metrics: metrics.txt.
models.py all JEPA variants + encoders, predictors, regularizers, probers
loss.py VICReg (invariance / variance / covariance)
engine.py train + eval loops
train.py, main.py entry points; main.py also runs the probing evaluation
evaluator.py probing harness (provided by the course, unmodified)
config/ ~20 experiment configs
optimized_dataset.py memory-mapped trajectory loading
Kevin Mathew T · Siddharth Agarwal · Allen George Ajith
- LeCun, A Path Towards Autonomous Machine Intelligence (2022) — JEPA
- Bardes, Ponce, LeCun, VICReg (2022)
- Zbontar et al., Barlow Twins (2021)
- Grill et al., BYOL (2020)
Original course assignment specification
The text below is the original CSCI-GA 2572 final project handout, preserved for context.
JEPA
Joint embedding prediction architecture (JEPA) is an energy based architecture for self supervised learning first proposed by LeCun (2022). Essentially, it works by asking the model to predict its own representations of future observations.
More formally, in the context of this problem, given an agent trajectory
$$ \begin{align} \text{Encoder}: &\tilde{s}_0 = s_0 = \text{Enc}_\theta(o_0) \ \text{Predictor}: &\tilde{s}_n = \text{Pred}_\phi(\tilde{s}_{n-1}, u_{n-1}) \end{align} $$
Where
The architecture may also be teacher-forcing (non-recurrent):
$$ \begin{align} \text{Encoder}: &s_n = \text{Enc}_\theta(o_n) \ \text{Predictor}: &\tilde{s}_n = \text{Pred}_\phi(s_{n-1}, u_{n-1}) \end{align} $$
The JEPA training objective would be to minimize the energy for the observation-action sequence
$$ \begin{align} \text{Target Encoder}: &s'_n = \text{Enc}_{\psi}(o_n) \ \text{System energy}: &F(\tau) = \sum_{n=1}^{N}D(\tilde{s}_n, s'_n) \end{align} $$
Where the Target Encoder
Environment and data set
The dataset consists of random trajectories collected from a toy environment consisting of an agent (dot) in two rooms separated by a wall. There's a door in a wall. The agent cannot travel through the border wall or middle wall (except through the door). Different trajectories may have different wall and door positions. Thus our JEPA model needs to be able to perceive and distinguish environment layouts.
Task
Implement and train a JEPA architecture on a dataset of 2.5M frames of exploratory trajectories.
The model is evaluated on how well the predicted representations capture the true
Constraints:
- It has to be a JEPA architecture — trained by minimizing the distance between predictions and targets in the representation space, while preventing collapse.
- Any method of preventing collapse is allowed except image reconstruction (e.g. MAE).
- Only the provided data may be used; image augmentation is allowed.
Evaluation
Quality is measured by probing. The JEPA world model is unrolled recurrently
$$ \begin{align} F(x,y) &= \sum_{n=1}^{N} C[y_n, \text{Prober}(\tilde{s}_n)]\ C(y, \tilde{y}) &= \lVert y - \tilde{y} \rVert _2^2 \end{align} $$
The prober trains on 170k frames from probe_normal/train. Two known validation sets are
probe_normal/val (similar to training) and probe_wall/val (agent running straight towards the
wall and sometimes the door, testing whether the model learned the dynamics of stopping at walls).
Two further held-out validation sets test long-horizon prediction and generalization to novel
layouts — during training the wall is excluded from a range of x-axes, and the held-out set places
it there.
Competition criteria
Teams were evaluated on 5 criteria: MSE on probe_normal (weight 1), MSE on probe_wall
(weight 1), MSE on long horizon probing (weight 1), MSE on out-of-domain wall probing (weight 1),
and model parameter count — fewer is better (weight 0.25).
