Skip to content

Latest commit

 

History

361 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

Action-Conditioned JEPA World Model

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).


The core problem: collapse

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.

Architecture

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.

What we tried

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/.

Reproducing

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.py

Dataset 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.

Repository layout

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

Contributors

Kevin Mathew T · Siddharth Agarwal · Allen George Ajith

References


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 $\tau$, i.e. an observation-action sequence $\tau = (o_0, u_0, o_1, u_1, \ldots, o_{N-1}, u_{N-1}, o_N)$, we specify a recurrent JEPA architecture as:

$$ \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 $\tilde{s}_n$ is the predicted state at time index $n$, and $s_n$ is the encoder output at time index $n$.

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 $\tau$, given to us by the sum of the distance between predicted states $\tilde{s}_n$ and the target states $s'_n$, where:

$$ \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 $\text{Enc}_\psi$ may be identical to Encoder $\text{Enc}_\theta$ (VicReg, Barlow Twins), or not (BYOL).

$D(\tilde{s}_n, s'_n)$ is some "distance" function. However, minimizing the energy naively is problematic because it can lead to representation collapse. There are techniques (such as ones mentioned above) to prevent this collapse by adding regularisers, contrastive samples, or specific architectural choices.

Recurrent JEPA over 4 timesteps

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.

Two room environment

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 $(x, y)$ coordinate of the agent.

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 $N$ times into the future, conditioned on initial observation $o_0$ and action sequence $u_0, u_1, \ldots, u_{N-1}$, generating predicted representations $\tilde{s}_1, \ldots, \tilde{s}_N$. A 2-layer FC prober is then trained to extract ground truth agent coordinates $y = (y_1,y_2)$:

$$ \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).

About

Action-conditioned JEPA world model for two-room navigation in 89K parameters. VICReg plus an inverse-dynamics auxiliary loss for collapse resistance; 4.40 MSE on normal probes.

Topics

Resources

Stars

2 stars

Watchers

2 watching

Forks

Releases

Packages

Used by

Contributors

Languages