Skip to content

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
class DINOv2Model(nnx.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.
    """

    def __init__(
        self,
        img_size: int = 518,
        patch_size: int = 14,
        in_channels: int = 3,
        hidden_size: int = 384,
        num_layers: int = 12,
        num_heads: int = 6,
        mlp_dim: int = 1536,
        layer_scale_init: float = 1.0,
        use_gradient_checkpointing: bool = False,
        attention_fn: Callable[..., Any] | None = None,
        rngs: rnglib.Rngs | None = None,
        dtype: DTypeLike = jnp.float32,
        param_dtype: DTypeLike = jnp.float32,
        sharding: ShardingSpec = DINOv2Sharding,
    ) -> None:
        """Initialize DINOv2Model.

        Args:
            img_size (int, optional): Input image size (square). Defaults to 518.
            patch_size (int, optional): Patch size. Defaults to 14.
            in_channels (int, optional): Number of input channels. Defaults to 3.
            hidden_size (int, optional): Hidden dimension size. Defaults to 384 (small).
            num_layers (int, optional): Number of transformer layers. Defaults to 12.
            num_heads (int, optional): Number of attention heads. Defaults to 6.
            mlp_dim (int, optional): MLP intermediate dimension. Defaults to 1536.
            layer_scale_init (float, optional): Initial value for per-channel LayerScale parameters. Defaults to 1.0.
            use_gradient_checkpointing (bool, optional): Whether to use gradient checkpointing. Defaults to False.
            attention_fn (Callable[..., Any] | None, optional): 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).
            rngs (rnglib.Rngs | None, optional): The random number generator state. If None, initializes to nnx.Rngs(0).
            dtype (DTypeLike, optional): The data type for computations. Defaults to jnp.float32.
            param_dtype (DTypeLike, optional): The data type for parameters. Defaults to jnp.float32.
            sharding (ShardingSpec, optional): Sharding specification for parameters. Defaults to DINOv2Sharding.
        """
        if rngs is None:
            rngs = nnx.Rngs(0)
        self._original_config = None
        self._layerscale_value = layer_scale_init
        self.encoder = VisionTransformerBase(
            img_size=img_size,
            patch_size=patch_size,
            in_channels=in_channels,
            hidden_size=hidden_size,
            num_layers=num_layers,
            num_heads=num_heads,
            mlp_dim=mlp_dim,
            pooling_type="CLS",
            dropout_rate=0.0,
            use_pre_norm=False,
            use_patch_bias=True,
            layernorm_epsilon=1e-6,
            use_layer_scale=True,
            layer_scale_init=layer_scale_init,
            use_gradient_checkpointing=use_gradient_checkpointing,
            attention_fn=attention_fn,
            rngs=rngs,
            dtype=dtype,
            param_dtype=param_dtype,
            sharding=sharding,
        )

    def __call__(self, x: Float[Array, "batch height width channels"]) -> Float[Array, "batch hidden_size"]:
        """Run DINOv2 forward pass and return CLS token embedding.

        Args:
            x (Float[Array, "batch height width channels"]): Batch of input images in BHWC format.

        Returns:
            Float[Array, "batch hidden_size"]: CLS token embeddings after final LayerNorm.
        """
        return self.encoder(x)

    def save_pretrained(self, save_directory: str) -> None:
        """Save model weights and config in HuggingFace format.

        Args:
            save_directory (str): Directory path where the model will be saved.
        """
        from .params import save_pretrained as _save_pretrained

        _save_pretrained(self, save_directory)

    @classmethod
    def from_pretrained(
        cls,
        model_name_or_path: str,
        use_pytorch: bool = False,
        rngs: rnglib.Rngs | None = None,
        dtype: DTypeLike = jnp.float32,
        param_dtype: DTypeLike = jnp.float32,
        sharding: ShardingSpec = DINOv2Sharding,
        use_gradient_checkpointing: bool = False,
        attention_fn: Callable[..., Any] | None = None,
    ) -> "DINOv2Model":
        """Load a pretrained DINOv2 model from a local path or HuggingFace Hub.

        Args:
            model_name_or_path (str): Local directory or HuggingFace model ID (e.g. "facebook/dinov2-small").
            use_pytorch (bool): Whether to load from PyTorch weights. Defaults to False.
            rngs (rnglib.Rngs | None): The random number generator state. If None, initializes to nnx.Rngs(0).
            dtype (DTypeLike): The data type for computations. Defaults to jnp.float32.
            param_dtype (DTypeLike): The data type for parameters. Defaults to jnp.float32.
            sharding (ShardingSpec): Sharding specification for parameters. Defaults to DINOv2Sharding.
            use_gradient_checkpointing (bool): Whether to use gradient checkpointing. Defaults to False.
            attention_fn (Callable[..., Any] | None): Custom attention function (e.g. jimm.tokamax_attention or jimm.make_tokamax_attention("mosaic_tpu")). Defaults to None.

        Returns:
            DINOv2Model: Model with pretrained weights loaded.
        """
        from .params import load_from_pretrained

        return load_from_pretrained(
            cls,
            model_name_or_path,
            use_pytorch,
            rngs,
            dtype,
            param_dtype,
            sharding,
            use_gradient_checkpointing,
            attention_fn,
        )

    @classmethod
    def from_config(
        cls,
        config: dict[str, Any],
        *,
        rngs: rnglib.Rngs | None = None,
        dtype: DTypeLike = jnp.float32,
        param_dtype: DTypeLike = jnp.float32,
        sharding: ShardingSpec = DINOv2Sharding,
        use_gradient_checkpointing: bool = False,
        attention_fn: Callable[..., Any] | None = None,
    ) -> "DINOv2Model":
        """Create DINOv2Model from a HuggingFace-compatible config dict.

        Args:
            config (dict[str, Any]): Configuration dictionary in HuggingFace DINOv2 format.
            rngs (rnglib.Rngs | None): The random number generator state. If None, initializes to nnx.Rngs(0).
            dtype (DTypeLike): The data type for computations. Defaults to jnp.float32.
            param_dtype (DTypeLike): The data type for parameters. Defaults to jnp.float32.
            sharding (ShardingSpec): Sharding specification for parameters. Defaults to DINOv2Sharding.
            use_gradient_checkpointing (bool): Whether to use gradient checkpointing. Defaults to False.
            attention_fn (Callable[..., Any] | None): Custom attention function (e.g. jimm.tokamax_attention or jimm.make_tokamax_attention("mosaic_tpu")). Defaults to None.

        Returns:
            DINOv2Model: Model with randomly initialized weights matching the given config.
        """
        if rngs is None:
            rngs = nnx.Rngs(0)
        hidden_size = config["hidden_size"]
        mlp_dim = int(hidden_size * config.get("mlp_ratio", 4))
        return cls(
            img_size=config["image_size"],
            patch_size=config["patch_size"],
            in_channels=config.get("num_channels", 3),
            hidden_size=hidden_size,
            num_layers=config["num_hidden_layers"],
            num_heads=config["num_attention_heads"],
            mlp_dim=mlp_dim,
            layer_scale_init=config.get("layerscale_value", 1.0),
            use_gradient_checkpointing=use_gradient_checkpointing,
            attention_fn=attention_fn,
            rngs=rngs,
            dtype=dtype,
            param_dtype=param_dtype,
            sharding=sharding,
        )

__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
def __call__(self, x: Float[Array, "batch height width channels"]) -> Float[Array, "batch hidden_size"]:
    """Run DINOv2 forward pass and return CLS token embedding.

    Args:
        x (Float[Array, "batch height width channels"]): Batch of input images in BHWC format.

    Returns:
        Float[Array, "batch hidden_size"]: CLS token embeddings after final LayerNorm.
    """
    return self.encoder(x)

__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
def __init__(
    self,
    img_size: int = 518,
    patch_size: int = 14,
    in_channels: int = 3,
    hidden_size: int = 384,
    num_layers: int = 12,
    num_heads: int = 6,
    mlp_dim: int = 1536,
    layer_scale_init: float = 1.0,
    use_gradient_checkpointing: bool = False,
    attention_fn: Callable[..., Any] | None = None,
    rngs: rnglib.Rngs | None = None,
    dtype: DTypeLike = jnp.float32,
    param_dtype: DTypeLike = jnp.float32,
    sharding: ShardingSpec = DINOv2Sharding,
) -> None:
    """Initialize DINOv2Model.

    Args:
        img_size (int, optional): Input image size (square). Defaults to 518.
        patch_size (int, optional): Patch size. Defaults to 14.
        in_channels (int, optional): Number of input channels. Defaults to 3.
        hidden_size (int, optional): Hidden dimension size. Defaults to 384 (small).
        num_layers (int, optional): Number of transformer layers. Defaults to 12.
        num_heads (int, optional): Number of attention heads. Defaults to 6.
        mlp_dim (int, optional): MLP intermediate dimension. Defaults to 1536.
        layer_scale_init (float, optional): Initial value for per-channel LayerScale parameters. Defaults to 1.0.
        use_gradient_checkpointing (bool, optional): Whether to use gradient checkpointing. Defaults to False.
        attention_fn (Callable[..., Any] | None, optional): 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).
        rngs (rnglib.Rngs | None, optional): The random number generator state. If None, initializes to nnx.Rngs(0).
        dtype (DTypeLike, optional): The data type for computations. Defaults to jnp.float32.
        param_dtype (DTypeLike, optional): The data type for parameters. Defaults to jnp.float32.
        sharding (ShardingSpec, optional): Sharding specification for parameters. Defaults to DINOv2Sharding.
    """
    if rngs is None:
        rngs = nnx.Rngs(0)
    self._original_config = None
    self._layerscale_value = layer_scale_init
    self.encoder = VisionTransformerBase(
        img_size=img_size,
        patch_size=patch_size,
        in_channels=in_channels,
        hidden_size=hidden_size,
        num_layers=num_layers,
        num_heads=num_heads,
        mlp_dim=mlp_dim,
        pooling_type="CLS",
        dropout_rate=0.0,
        use_pre_norm=False,
        use_patch_bias=True,
        layernorm_epsilon=1e-6,
        use_layer_scale=True,
        layer_scale_init=layer_scale_init,
        use_gradient_checkpointing=use_gradient_checkpointing,
        attention_fn=attention_fn,
        rngs=rngs,
        dtype=dtype,
        param_dtype=param_dtype,
        sharding=sharding,
    )

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
@classmethod
def from_config(
    cls,
    config: dict[str, Any],
    *,
    rngs: rnglib.Rngs | None = None,
    dtype: DTypeLike = jnp.float32,
    param_dtype: DTypeLike = jnp.float32,
    sharding: ShardingSpec = DINOv2Sharding,
    use_gradient_checkpointing: bool = False,
    attention_fn: Callable[..., Any] | None = None,
) -> "DINOv2Model":
    """Create DINOv2Model from a HuggingFace-compatible config dict.

    Args:
        config (dict[str, Any]): Configuration dictionary in HuggingFace DINOv2 format.
        rngs (rnglib.Rngs | None): The random number generator state. If None, initializes to nnx.Rngs(0).
        dtype (DTypeLike): The data type for computations. Defaults to jnp.float32.
        param_dtype (DTypeLike): The data type for parameters. Defaults to jnp.float32.
        sharding (ShardingSpec): Sharding specification for parameters. Defaults to DINOv2Sharding.
        use_gradient_checkpointing (bool): Whether to use gradient checkpointing. Defaults to False.
        attention_fn (Callable[..., Any] | None): Custom attention function (e.g. jimm.tokamax_attention or jimm.make_tokamax_attention("mosaic_tpu")). Defaults to None.

    Returns:
        DINOv2Model: Model with randomly initialized weights matching the given config.
    """
    if rngs is None:
        rngs = nnx.Rngs(0)
    hidden_size = config["hidden_size"]
    mlp_dim = int(hidden_size * config.get("mlp_ratio", 4))
    return cls(
        img_size=config["image_size"],
        patch_size=config["patch_size"],
        in_channels=config.get("num_channels", 3),
        hidden_size=hidden_size,
        num_layers=config["num_hidden_layers"],
        num_heads=config["num_attention_heads"],
        mlp_dim=mlp_dim,
        layer_scale_init=config.get("layerscale_value", 1.0),
        use_gradient_checkpointing=use_gradient_checkpointing,
        attention_fn=attention_fn,
        rngs=rngs,
        dtype=dtype,
        param_dtype=param_dtype,
        sharding=sharding,
    )

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
@classmethod
def from_pretrained(
    cls,
    model_name_or_path: str,
    use_pytorch: bool = False,
    rngs: rnglib.Rngs | None = None,
    dtype: DTypeLike = jnp.float32,
    param_dtype: DTypeLike = jnp.float32,
    sharding: ShardingSpec = DINOv2Sharding,
    use_gradient_checkpointing: bool = False,
    attention_fn: Callable[..., Any] | None = None,
) -> "DINOv2Model":
    """Load a pretrained DINOv2 model from a local path or HuggingFace Hub.

    Args:
        model_name_or_path (str): Local directory or HuggingFace model ID (e.g. "facebook/dinov2-small").
        use_pytorch (bool): Whether to load from PyTorch weights. Defaults to False.
        rngs (rnglib.Rngs | None): The random number generator state. If None, initializes to nnx.Rngs(0).
        dtype (DTypeLike): The data type for computations. Defaults to jnp.float32.
        param_dtype (DTypeLike): The data type for parameters. Defaults to jnp.float32.
        sharding (ShardingSpec): Sharding specification for parameters. Defaults to DINOv2Sharding.
        use_gradient_checkpointing (bool): Whether to use gradient checkpointing. Defaults to False.
        attention_fn (Callable[..., Any] | None): Custom attention function (e.g. jimm.tokamax_attention or jimm.make_tokamax_attention("mosaic_tpu")). Defaults to None.

    Returns:
        DINOv2Model: Model with pretrained weights loaded.
    """
    from .params import load_from_pretrained

    return load_from_pretrained(
        cls,
        model_name_or_path,
        use_pytorch,
        rngs,
        dtype,
        param_dtype,
        sharding,
        use_gradient_checkpointing,
        attention_fn,
    )

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
def save_pretrained(self, save_directory: str) -> None:
    """Save model weights and config in HuggingFace format.

    Args:
        save_directory (str): Directory path where the model will be saved.
    """
    from .params import save_pretrained as _save_pretrained

    _save_pretrained(self, save_directory)