CanViT-B/16 Pretrained (JAX / Flax NNX)

JAX-native checkpoint for CanViT, converted from the PyTorch checkpoint.

Pretrained on ImageNet-21k via dense latent distillation from DINOv3 ViT-B.

Usage

uv add "canvit-nnx @ git+https://github.com/yberreby/CanViT-NNX.git"
import jax.numpy as jnp
from canvit_nnx import from_pretrained, Viewpoint, sample_at_viewpoint

model = from_pretrained("canvit/canvitb16-add-vpe-pretrain-g128px-s512px-in21k-dv3b16-2026-02-02-nnx")
state = model.init_state(batch_size=1, canvas_grid_size=32)

vp = Viewpoint.full_scene(batch_size=1)
glimpse = sample_at_viewpoint(spatial=image, viewpoint=vp, glimpse_size_px=128)
out = model(glimpse, state, vp)

# Canvas features should be layernormed before downstream use (PCA, probing, etc.)
canvas = model.get_spatial(out.state.canvas)
mean = canvas.mean(axis=-1, keepdims=True)
canvas = (canvas - mean) / jnp.sqrt(canvas.var(axis=-1, keepdims=True) + 1e-5)

Source: CanViT-NNX

Citation

@article{berreby2026canvit,
  title={CanViT: Toward Active-Vision Foundation Models},
  author={Berreby, Yoha{\"i}-Eliel and Du, Sabrina and Durand, Audrey and Krishna, B. Suresh},
  year={2026},
  eprint={2603.22570},
  archivePrefix={arXiv},
  primaryClass={cs.CV}
}
Downloads last month
36
Safetensors
Model size
95.2M params
Tensor type
F32
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Collection including canvit/canvitb16-add-vpe-pretrain-g128px-s512px-in21k-dv3b16-2026-02-02-nnx

Paper for canvit/canvitb16-add-vpe-pretrain-g128px-s512px-in21k-dv3b16-2026-02-02-nnx