Skip to content

Latest commit

 

History

1 Commit

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

jepa

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

Why a separate harness from fde_train?

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:

  1. Run the frozen encoder once over the dataset → cache (N, T, D) token features and (N, D) pooled features.
  2. 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.
  3. 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.

Install

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

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

Smoke test (no GPU, no model download)

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 linear

Or via the YAML driver:

python scripts/run_evals.py --config configs/smoke.yaml

Real probe: I-JEPA on CIFAR-100 (linear)

CIFAR-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_1k

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

Real probe: V-JEPA 2 on Something-Something V2 (attentive)

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_256

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

Adding a new encoder

  1. Create src/encoders/<name>.py with a subclass of JEPAEncoder.
  2. Implement load, preprocess, forward. Both pooled and token outputs should be returned (see EncoderOutput).
  3. Decorate the class with @register_encoder("my_jepa").
  4. Import the module in src/encoders/__init__.py so 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())

Adding a new task

  1. Create src/tasks/<name>.py with a subclass of ProbeTask.
  2. Set modality, num_classes, metric, default_probe.
  3. Implement load_train / load_test, returning TaskSplit(items=..., labels=...).
  4. Decorate with @register_task("my_task") and import it in src/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.

What's not here

  • DINOv2 / MAE: the harness is built around JEPA, but the encoder interface is generic — adding a dinov2.py adapter 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 launch or write a multi-process script that round-robins indices.

Output layout

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

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages