Frozen-features evaluation harness for JEPA-family encoders (I-JEPA,
V-JEPA 2, V-JEPA 2.1). Mirrors the structure of fde_train
but is built around encoders + probes instead of chat-completion +
scoring, because JEPA models are vision/video feature extractors with no
generative head.
jepa/
├── configs/
│ ├── encoder/ # one YAML per registered encoder
│ ├── task/ # one YAML per registered probe task
│ ├── all_image.yaml # multi-task suite for an image encoder
│ ├── all_video.yaml # multi-task suite for a video encoder
│ └── smoke.yaml # CPU-only synthetic smoke
├── src/
│ ├── encoders/ # JEPAEncoder + I-JEPA, V-JEPA 2, synth
│ ├── probes/ # LinearProbe, AttentiveProbe
│ ├── tasks/ # ProbeTask + cifar100, imagenet1k, ssv2, synth
│ ├── features.py # extract + cache (encoder, task, split) features
│ └── runner.py # ProbeRunner: encoder × task → metrics
├── scripts/
│ ├── extract_features.py # CLI: pre-compute features and cache to disk
│ ├── run_probe.py # CLI: one (encoder, task, probe) end-to-end
│ └── run_evals.py # CLI: YAML-driven multi-task suite
├── data/ # cached features (.pt) live here by default
├── results/ # one folder per (encoder, task) run
└── requirements.txt
The fde_train benchmark adapters all assume an autoregressive LLM:
build_messages → vllm.chat → score(response). JEPA models don't have a
generate() and don't speak text. Their canonical evaluation protocol is
instead:
- Run the frozen encoder once over the dataset → cache
(N, T, D)token features and(N, D)pooled features. - Train a small probe on top: linear (logistic regression) or attentive (a single learned query attending over tokens). Probes are cheap, so multiple probe variants per encoder are common.
- Score the probe on the held-out split.
This repo gives you that loop, plus registries that mirror fde_train's
benchmark registry so the two harnesses feel similar to use.
You can re-use the conda env from fde_train (it already has torch
2.9.1+cu130 on cu13 H200s). Just add the harness deps:
conda activate fde
cd /data/wenli/jepa
pip install -r requirements.txttransformers >= 4.53 is required for VJEPA2Model / AutoVideoProcessor
to be importable. If you're on the older transformers pinned by the verl
install, pin-bump it: pip install -U 'transformers>=4.53'. (vLLM and verl
both tolerate that range.)
End-to-end pipeline check using a synthetic encoder + synthetic task. Should finish in < 5 seconds on CPU and report ~100% accuracy:
cd /data/wenli/jepa
python scripts/run_probe.py --encoder synth_image --task synth_image --probe linearOr via the YAML driver:
python scripts/run_evals.py --config configs/smoke.yamlCIFAR-100 is small (~150 MB), unrestricted, and a long-standing image SSL sanity check. This run should fit on one GPU:
python scripts/run_probe.py \
--encoder ijepa_vith14_1k \
--task cifar100_linear \
--probe linear \
--cache-dir data/features \
--output-dir results/ijepa_vith14_1kThe first call extracts ~50k train + 10k test features and caches them
under data/features/ijepa_vith14_1k/cifar100_linear/. Subsequent calls
(e.g. with different probe hyper-parameters) reuse the cache and skip the
encoder forward pass.
SSv2 is the canonical video benchmark for V-JEPA 2. The dataset is not
shipped here; point ssv2_root at a local mount that contains
{train,validation}.json annotation files and a videos/ subdir.
python scripts/run_probe.py \
--encoder vjepa2_vitl_256 \
--task ssv2_attentive \
--probe attentive \
--task-kwargs ssv2_root=/data/ssv2 \
--cache-dir data/features \
--output-dir results/vjepa2_vitl_256The bundled ssv2.py is intentionally a sketch; expect to replace its
clip sampler with whatever works on your storage layout (decord, pyav,
pre-extracted RGB tensors, ...).
- Create
src/encoders/<name>.pywith a subclass ofJEPAEncoder. - Implement
load,preprocess,forward. Both pooled and token outputs should be returned (seeEncoderOutput). - Decorate the class with
@register_encoder("my_jepa"). - Import the module in
src/encoders/__init__.pyso registration runs.
Template:
from .base import EncoderOutput, JEPAEncoder
from .registry import register_encoder
@register_encoder("my_jepa")
class MyJEPA(JEPAEncoder):
modality = "image" # or "video"
default_dtype = "bfloat16"
def load(self):
from transformers import AutoModel, AutoProcessor
self._model = AutoModel.from_pretrained("...").to(self.device).eval()
self._processor = AutoProcessor.from_pretrained("...")
def preprocess(self, batch):
out = self._processor(images=list(batch), return_tensors="pt")
return {k: v.to(self.device) for k, v in out.items()}
def forward(self, model_inputs):
outputs = self._model(**model_inputs)
tokens = outputs.last_hidden_state
return EncoderOutput(tokens=tokens.float(), pooled=tokens.mean(1).float())- Create
src/tasks/<name>.pywith a subclass ofProbeTask. - Set
modality,num_classes,metric,default_probe. - Implement
load_train/load_test, returningTaskSplit(items=..., labels=...). - Decorate with
@register_task("my_task")and import it insrc/tasks/__init__.py.
build_probe is inherited from the base class and supports linear /
attentive out of the box; override it only if you need a regression head
or a custom probe.
- DINOv2 / MAE: the harness is built around JEPA, but the encoder
interface is generic — adding a
dinov2.pyadapter would take ~30 LOC. - World-model rollouts (V-JEPA 2-AC robot control, action anticipation, dense depth): these need different probe heads / evaluation loops than classification. The runner is the place to extend for those — keep encoder + task abstractions, swap out the probe.
- Distributed feature extraction: features are extracted on a single
device. For ImageNet-1K full train (1.28M images) that's hours on one
H200; for production you'd shard with
accelerate launchor write a multi-process script that round-robins indices.
run_evals.py writes one folder per task under output_dir:
results/ijepa_vith14_1k/
├── summary.json # merged metrics across tasks
├── cifar100_linear/
│ └── summary.json
└── imagenet1k/
└── summary.json