Views
No views yet
uv add "canvit-nnx @ git+https://github.com/yberreby/CanViT-NNX.git"1import jax.numpy as jnp
2from canvit_nnx import from_pretrained, Viewpoint, sample_at_viewpoint
3
4model = from_pretrained("canvit/canvitb16-add-vpe-pretrain-g128px-s512px-in1k-dv3b16-2026-06-22-nnx")
5state = model.init_state(batch_size=1, canvas_grid_size=32)
6
7vp = Viewpoint.full_scene(batch_size=1)
8glimpse = sample_at_viewpoint(spatial=image, viewpoint=vp, glimpse_size_px=128)
9out = model(glimpse, state, vp)
10
11# Canvas features should be layernormed before downstream use (PCA, probing, etc.)
12canvas = model.get_spatial(out.state.canvas)
13mean = canvas.mean(axis=-1, keepdims=True)
14canvas = (canvas - mean) / jnp.sqrt(canvas.var(axis=-1, keepdims=True) + 1e-5)1@article{berreby2026canvit,
2 title={CanViT: Toward Active-Vision Foundation Models},
3 author={Berreby, Yoha{\"i}-Eliel and Du, Sabrina and Durand, Audrey and Krishna, B. Suresh},
4 year={2026},
5 eprint={2603.22570},
6 archivePrefix={arXiv},
7 primaryClass={cs.CV}
8}