VISReg: Variance-Invariance-Sketching Regularization for JEPA training

arXiv Project Page GitHub

Key results:

  • 💪 Strong collapse prevention: High gradient when embedding collapse
  • Friendly to scale training: Linear complexity to scaling factors
  • 🧩 Easy to train: Similar to LeJEPA, it is a heuristic-free method
  • 🏆 Best OOD performance: Achieve the best accuracy on 6 OOD datasets
  • 📉 Data efficiency: Achieving a similar average accuracy to DINOv2 with 90% less data
  • 🧬 Robust to low-quality datasets: It is robust to long-tailed and sparse datasets

Available Checkpoints

File Architecture Patch Size Embed Dim Backbone Params Pre-training Data
visreg-vit-b-inet1k.pth ViT-Base 16 768 86M ImageNet-1K
visreg-vit-l-inet1k.pth ViT-Large 14 1024 304M ImageNet-1K
visreg-vit-l-inet22k.pth ViT-Large 14 1024 304M ImageNet-22K

What is in each file

Every checkpoint is a single flat state_dict containing the ViT backbone and the projection head used during pretraining. No optimizer, scheduler, online-probe or training-config state is included, so the files load with weights_only=True.

cls_token, pos_embed, patch_embed.*, blocks.*, norm.*   # timm ViT backbone (unprefixed keys)
proj.0 ... proj.8                                       # projection head, MLP(embed_dim -> 2048 -> 2048 -> proj_dim)
File proj_dim Projection activation Head params
visreg-vit-b-inet1k.pth 256 ReLU 6.3M
visreg-vit-l-inet1k.pth 384 GELU 7.1M
visreg-vit-l-inet22k.pth 384 GELU 7.1M

The head is published so the models can be fine-tuned, or SSL pretraining continued, with the projector that was actually trained. For frozen-feature use (linear probing, segmentation, retrieval) the backbone alone is enough.

Usage

Load the backbone with timm

The proj.* entries have no counterpart in a bare timm ViT, so drop them before loading:

import timm
import torch

# ViT-Base/16
state_dict = torch.load("visreg-vit-b-inet1k.pth", map_location="cpu", weights_only=True)
model = timm.create_model("vit_base_patch16_224", pretrained=False, num_classes=0, dynamic_img_size=True)
model.load_state_dict({k: v for k, v in state_dict.items() if not k.startswith("proj.")})

# ViT-Large/14 (ImageNet-22K)
state_dict = torch.load("visreg-vit-l-inet22k.pth", map_location="cpu", weights_only=True)
model = timm.create_model("vit_large_patch14_224", pretrained=False, num_classes=0, dynamic_img_size=True)
model.load_state_dict({k: v for k, v in state_dict.items() if not k.startswith("proj.")})

Load the backbone and projection head

Using the GitHub repo, which rebuilds the pretraining encoder with the correct proj_dim and activation and loads it with strict=True:

from downstream.model_zoo import load_visreg_encoder

encoder = load_visreg_encoder("visreg_vit_l_inet22k")  # downloads from this repo on first use
emb, proj = encoder(images)                             # same interface as pretraining

Download with huggingface_hub

from huggingface_hub import hf_hub_download

path = hf_hub_download(repo_id="BooBooWu/visreg", filename="visreg-vit-b-inet1k.pth")
path = hf_hub_download(repo_id="BooBooWu/visreg", filename="visreg-vit-l-inet1k.pth")
path = hf_hub_download(repo_id="BooBooWu/visreg", filename="visreg-vit-l-inet22k.pth")

Feature extraction

from PIL import Image
from torchvision import transforms

transform = transforms.Compose([
    transforms.Resize(256),
    transforms.CenterCrop(224),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])

img = transform(Image.open("image.jpg")).unsqueeze(0)

with torch.no_grad():
    features = model(img)  # [1, embed_dim]

Evaluation

Full evaluation suite (linear probe, segmentation, fine-tuning) is available in the GitHub repo. The scripts accept a release name and download the weights automatically:

# Linear probe on 10+ datasets
python downstream/linear_prob/run_evaluation.py \
    --checkpoint visreg_vit_l_inet22k \
    --model vit_l \
    --datasets all

A local .pth path works too:

python downstream/linear_prob/run_evaluation.py \
    --checkpoint visreg-vit-b-inet1k.pth \
    --model vit_b \
    --datasets all

Citation

@inproceedings{wu2026visreg,
  title     = {VISReg: Variance-Invariance-Sketching Regularization for JEPA training},
  author    = {Wu, Haiyu and Balestriero, Randall and Levine, Morgan},
  booktitle = {arXiv},
  year      = {2026}
}

License

This project (code and pretrained weights) is released under CC BY-NC 4.0 for non-commercial use only.

Downloads last month
-
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Datasets used to train BooBooWu/visreg

Paper for BooBooWu/visreg