--- library_name: flax pipeline_tag: image-feature-extraction tags: - vision - jax - flax - canvit - active-vision license: mit --- # CanViT-B/16 Pretrained (JAX / Flax NNX) JAX-native checkpoint for [CanViT](https://arxiv.org/abs/2603.22570), converted from the [PyTorch checkpoint](https://huggingface.co/canvit/canvitb16-add-vpe-pretrain-g128px-s512px-in21k-dv3b16-2026-02-02). Pretrained on ImageNet-21k via dense latent distillation from DINOv3 ViT-B. ## Usage ```bash uv add "canvit-nnx @ git+https://github.com/yberreby/CanViT-NNX.git" ``` ```python 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](https://github.com/yberreby/CanViT-NNX) ## Citation ```bibtex @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} } ```