DINOv2
DINOv2 is a self-supervised vision transformer trained on a large curated image dataset without manual annotation. It uses a self-distillation objective and produces general-purpose image representations that transfer well across downstream tasks.
Paper: "DINOv2: Learning Robust Visual Features without Supervision" (Oquab et al., 2023) Code: github.com/facebookresearch/dinov2
The jimm implementation returns the CLS token embedding after the final LayerNorm. The key architectural difference from standard ViT is LayerScale: each transformer block multiplies the attention and MLP residuals by a learnable per-channel scalar vector before the residual add, stabilizing training at scale.
Supported models
| HuggingFace ID | hidden_size |
Layers | Heads | Patch | Image size |
|---|---|---|---|---|---|
facebook/dinov2-small |
384 | 12 | 6 | 14 | 518 |
facebook/dinov2-base |
768 | 12 | 12 | 14 | 518 |
facebook/dinov2-large |
1024 | 24 | 16 | 14 | 518 |
facebook/dinov2-giant |
1536 | 40 | 24 | 14 | 518 |
Basic usage
import jimm
import numpy as np
model = jimm.DINOv2Model.from_pretrained("facebook/dinov2-small")
images = np.random.rand(4, 518, 518, 3).astype(np.float32)
embeddings = model(images) # shape: (4, 384)
Flash / Splash Attention
DINOv2 supports hardware-accelerated attention via Tokamax:
| Backend | Hardware | Notes |
|---|---|---|
"mosaic" |
NVIDIA H100 (SM90) / B100 (SM100) | Pallas Mosaic GPU kernel |
"triton" |
Any NVIDIA GPU | Pallas Triton kernel |
"cudnn" |
NVIDIA GPU | Via JAX-NN / cuDNN |
"mosaic_tpu" |
TPU v5+ (all generations) | Splash attention (block-sparse) |
"xla_chunked" |
GPU / TPU | Flash-style chunked XLA |
"xla" |
Any | Standard XLA fallback |
import jimm
# GPU: try H100 Mosaic kernel, fall back to Triton, then XLA
model = jimm.DINOv2Model.from_pretrained("facebook/dinov2-base",
attention_fn=jimm.make_tokamax_attention(["mosaic", "triton", "xla"]))
# TPU: try Splash attention, fall back to chunked XLA
model = jimm.DINOv2Model.from_pretrained("facebook/dinov2-base",
attention_fn=jimm.make_tokamax_attention(["mosaic_tpu", "xla_chunked"]))
Note: DINOv2 processes 1370 tokens at 518x518 with patch size 14, which is substantially longer than typical ViT sequences. Flash attention provides a meaningful memory reduction at this sequence length.
FSDP / Explicit Sharding
DINOv2 supports JAX explicit sharding (FSDP-style) via DINOv2Sharding, which extends ViTSharding with additional sharding for the LayerScale vectors.
from jax.experimental import mesh_utils
from jax.sharding import AxisType, Mesh
import jax
n_devices = jax.device_count()
mesh = Mesh(
mesh_utils.create_device_mesh((1, n_devices)),
("data", "fsdp"),
axis_types=(AxisType.Explicit, AxisType.Explicit),
)
jax.set_mesh(mesh)
model = jimm.DINOv2Model.from_pretrained("facebook/dinov2-base")
# model params are automatically sharded across fsdp axis
DINOv2Sharding specs represent per-layer shapes. The Transformer stack prepends None for the scan axis to Variable metadata after nnx.vmap, so the optimizer receives the correct stacked spec natively.
To disable sharding, pass sharding=NoSharding() (import: from jimm.common.sharding import NoSharding).
jimm.models.dinov2.DINOv2Model
Bases: Module
DINOv2 vision transformer feature extractor.
This implements the DINOv2 architecture from "DINOv2: Learning Robust Visual Features without Supervision" (Oquab et al., 2023). Returns CLS token embeddings. Unlike standard ViT, each transformer block applies per-channel LayerScale to the attention and MLP residuals.
Source code in src/jimm/models/dinov2/dinov2_model.py
15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 | |
__call__(x)
Run DINOv2 forward pass and return CLS token embedding.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
Float[Array, 'batch height width channels']
|
Batch of input images in BHWC format. |
required |
Returns:
| Type | Description |
|---|---|
Float[Array, 'batch hidden_size']
|
Float[Array, "batch hidden_size"]: CLS token embeddings after final LayerNorm. |
Source code in src/jimm/models/dinov2/dinov2_model.py
88 89 90 91 92 93 94 95 96 97 | |
__init__(img_size=518, patch_size=14, in_channels=3, hidden_size=384, num_layers=12, num_heads=6, mlp_dim=1536, layer_scale_init=1.0, use_gradient_checkpointing=False, attention_fn=None, rngs=None, dtype=jnp.float32, param_dtype=jnp.float32, sharding=DINOv2Sharding)
Initialize DINOv2Model.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
img_size
|
int
|
Input image size (square). Defaults to 518. |
518
|
patch_size
|
int
|
Patch size. Defaults to 14. |
14
|
in_channels
|
int
|
Number of input channels. Defaults to 3. |
3
|
hidden_size
|
int
|
Hidden dimension size. Defaults to 384 (small). |
384
|
num_layers
|
int
|
Number of transformer layers. Defaults to 12. |
12
|
num_heads
|
int
|
Number of attention heads. Defaults to 6. |
6
|
mlp_dim
|
int
|
MLP intermediate dimension. Defaults to 1536. |
1536
|
layer_scale_init
|
float
|
Initial value for per-channel LayerScale parameters. Defaults to 1.0. |
1.0
|
use_gradient_checkpointing
|
bool
|
Whether to use gradient checkpointing. Defaults to False. |
False
|
attention_fn
|
Callable[..., Any] | None
|
Custom attention function compatible with nnx.MultiHeadAttention's attention_fn interface (e.g. jimm.tokamax_attention or jimm.make_tokamax_attention("mosaic_tpu")). Defaults to None (uses nnx.dot_product_attention). |
None
|
rngs
|
Rngs | None
|
The random number generator state. If None, initializes to nnx.Rngs(0). |
None
|
dtype
|
DTypeLike
|
The data type for computations. Defaults to jnp.float32. |
float32
|
param_dtype
|
DTypeLike
|
The data type for parameters. Defaults to jnp.float32. |
float32
|
sharding
|
ShardingSpec
|
Sharding specification for parameters. Defaults to DINOv2Sharding. |
DINOv2Sharding
|
Source code in src/jimm/models/dinov2/dinov2_model.py
24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 | |
from_config(config, *, rngs=None, dtype=jnp.float32, param_dtype=jnp.float32, sharding=DINOv2Sharding, use_gradient_checkpointing=False, attention_fn=None)
classmethod
Create DINOv2Model from a HuggingFace-compatible config dict.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
config
|
dict[str, Any]
|
Configuration dictionary in HuggingFace DINOv2 format. |
required |
rngs
|
Rngs | None
|
The random number generator state. If None, initializes to nnx.Rngs(0). |
None
|
dtype
|
DTypeLike
|
The data type for computations. Defaults to jnp.float32. |
float32
|
param_dtype
|
DTypeLike
|
The data type for parameters. Defaults to jnp.float32. |
float32
|
sharding
|
ShardingSpec
|
Sharding specification for parameters. Defaults to DINOv2Sharding. |
DINOv2Sharding
|
use_gradient_checkpointing
|
bool
|
Whether to use gradient checkpointing. Defaults to False. |
False
|
attention_fn
|
Callable[..., Any] | None
|
Custom attention function (e.g. jimm.tokamax_attention or jimm.make_tokamax_attention("mosaic_tpu")). Defaults to None. |
None
|
Returns:
| Name | Type | Description |
|---|---|---|
DINOv2Model |
DINOv2Model
|
Model with randomly initialized weights matching the given config. |
Source code in src/jimm/models/dinov2/dinov2_model.py
150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 | |
from_pretrained(model_name_or_path, use_pytorch=False, rngs=None, dtype=jnp.float32, param_dtype=jnp.float32, sharding=DINOv2Sharding, use_gradient_checkpointing=False, attention_fn=None)
classmethod
Load a pretrained DINOv2 model from a local path or HuggingFace Hub.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
model_name_or_path
|
str
|
Local directory or HuggingFace model ID (e.g. "facebook/dinov2-small"). |
required |
use_pytorch
|
bool
|
Whether to load from PyTorch weights. Defaults to False. |
False
|
rngs
|
Rngs | None
|
The random number generator state. If None, initializes to nnx.Rngs(0). |
None
|
dtype
|
DTypeLike
|
The data type for computations. Defaults to jnp.float32. |
float32
|
param_dtype
|
DTypeLike
|
The data type for parameters. Defaults to jnp.float32. |
float32
|
sharding
|
ShardingSpec
|
Sharding specification for parameters. Defaults to DINOv2Sharding. |
DINOv2Sharding
|
use_gradient_checkpointing
|
bool
|
Whether to use gradient checkpointing. Defaults to False. |
False
|
attention_fn
|
Callable[..., Any] | None
|
Custom attention function (e.g. jimm.tokamax_attention or jimm.make_tokamax_attention("mosaic_tpu")). Defaults to None. |
None
|
Returns:
| Name | Type | Description |
|---|---|---|
DINOv2Model |
DINOv2Model
|
Model with pretrained weights loaded. |
Source code in src/jimm/models/dinov2/dinov2_model.py
109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 | |
save_pretrained(save_directory)
Save model weights and config in HuggingFace format.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
save_directory
|
str
|
Directory path where the model will be saved. |
required |
Source code in src/jimm/models/dinov2/dinov2_model.py
99 100 101 102 103 104 105 106 107 | |