Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions docs/source/en/model_doc/dinov3.md
Original file line number Diff line number Diff line change
Expand Up @@ -169,6 +169,9 @@ print("Pooled output shape:", pooled_output.shape)
[[autodoc]] DINOv3ViTModel
- forward

## DINOv3ViTBackbone
[[autodoc]] DINOv3ViTBackbone

## DINOv3ConvNextModel

[[autodoc]] DINOv3ConvNextModel
Expand Down
1 change: 1 addition & 0 deletions src/transformers/models/auto/modeling_auto.py
Original file line number Diff line number Diff line change
Expand Up @@ -1700,6 +1700,7 @@ class _BaseModelWithGenerate(PreTrainedModel, GenerationMixin):
("dinov2", "Dinov2Backbone"),
("dinov2_with_registers", "Dinov2WithRegistersBackbone"),
("dinov3_convnext", "DINOv3ConvNextBackbone"),
("dinov3_vit", "DINOv3ViTBackbone"),
("focalnet", "FocalNetBackbone"),
("hgnet_v2", "HGNetV2Backbone"),
("hiera", "HieraBackbone"),
Expand Down
29 changes: 28 additions & 1 deletion src/transformers/models/dinov3_vit/configuration_dinov3_vit.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,12 +18,13 @@

from ...configuration_utils import PreTrainedConfig
from ...utils import logging
from ...utils.backbone_utils import BackboneConfigMixin, get_aligned_output_features_output_indices


logger = logging.get_logger(__name__)


class DINOv3ViTConfig(PreTrainedConfig):
class DINOv3ViTConfig(BackboneConfigMixin, PreTrainedConfig):
r"""
This is the configuration class to store the configuration of a [`DINOv3Model`]. It is used to instantiate an
DINOv3 model according to the specified arguments, defining the model architecture. Instantiating a configuration
Expand Down Expand Up @@ -86,6 +87,16 @@ class DINOv3ViTConfig(PreTrainedConfig):
pos_embed_rescale (`float`, *optional*, defaults to 2.0):
Amount to randomly rescale position embedding coordinates in log-uniform value in [1/rescale, rescale],
applied only in training mode if not `None`.
out_features (`list[str]`, *optional*):
If used as backbone, list of features to output. Can be any of `"stem"`, `"stage1"`, `"stage2"`, etc.
(depending on how many stages the model has). Will default to the last stage if unset.
out_indices (`list[int]`, *optional*):
If used as backbone, list of indices of features to output. Can be any of 0, 1, 2, etc.
(depending on how many stages the model has). Will default to the last stage if unset.
apply_layernorm (`bool`, *optional*, defaults to `True`):
Whether to apply layer normalization to the feature maps when used as backbone.
reshape_hidden_states (`bool`, *optional*, defaults to `True`):
Whether to reshape the hidden states to spatial dimensions when used as backbone.

Example:

Expand Down Expand Up @@ -131,6 +142,10 @@ def __init__(
pos_embed_shift: Optional[float] = None,
pos_embed_jitter: Optional[float] = None,
pos_embed_rescale: Optional[float] = 2.0,
out_features: Optional[list[str]] = None,
out_indices: Optional[list[int]] = None,
apply_layernorm: bool = True,
reshape_hidden_states: bool = True,
**kwargs,
):
super().__init__(**kwargs)
Expand Down Expand Up @@ -161,6 +176,18 @@ def __init__(
self.pos_embed_shift = pos_embed_shift
self.pos_embed_jitter = pos_embed_jitter
self.pos_embed_rescale = pos_embed_rescale
# Initialize backbone-specific configuration
self.apply_layernorm = apply_layernorm
self.reshape_hidden_states = reshape_hidden_states

# Initialize backbone stage names
stage_names = ["stem"] + [f"stage{i}" for i in range(1, num_hidden_layers + 1)]
self.stage_names = stage_names

# Initialize backbone features/indices
self._out_features, self._out_indices = get_aligned_output_features_output_indices(
out_features=out_features, out_indices=out_indices, stage_names=stage_names
)


__all__ = ["DINOv3ViTConfig"]
84 changes: 77 additions & 7 deletions src/transformers/models/dinov3_vit/modeling_dinov3_vit.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,11 +29,12 @@

from ...activations import ACT2FN
from ...modeling_layers import GradientCheckpointingLayer
from ...modeling_outputs import BaseModelOutputWithPooling
from ...modeling_outputs import BackboneOutput, BaseModelOutputWithPooling
from ...modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel
from ...processing_utils import Unpack
from ...pytorch_utils import compile_compatible_method_lru_cache
from ...utils import TransformersKwargs, auto_docstring
from ...utils import TransformersKwargs, auto_docstring, can_return_tuple
from ...utils.backbone_utils import BackboneMixin
from ...utils.generic import check_model_inputs
from .configuration_dinov3_vit import DINOv3ViTConfig

Expand Down Expand Up @@ -522,10 +523,79 @@ def forward(
sequence_output = self.norm(hidden_states)
pooled_output = sequence_output[:, 0, :]

return BaseModelOutputWithPooling(
last_hidden_state=sequence_output,
pooler_output=pooled_output,
)
return BaseModelOutputWithPooling(last_hidden_state=sequence_output, pooler_output=pooled_output)


@auto_docstring
class DINOv3ViTBackbone(DINOv3ViTPreTrainedModel, BackboneMixin):
def __init__(self, config):
super().__init__(config)
super()._init_backbone(config)

self.embeddings = DINOv3ViTEmbeddings(config)
self.rope_embeddings = DINOv3ViTRopePositionEmbedding(config)
self.layer = nn.ModuleList([DINOv3ViTLayer(config) for _ in range(config.num_hidden_layers)])
self.norm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
self.gradient_checkpointing = False

self.num_features = [config.hidden_size for _ in range(config.num_hidden_layers + 1)]
self.post_init()

def get_input_embeddings(self):
return self.embeddings.patch_embeddings

@check_model_inputs()
@can_return_tuple
def forward(
self,
pixel_values: torch.Tensor,
output_hidden_states: Optional[bool] = None,
**kwargs: Unpack[TransformersKwargs],
) -> BackboneOutput:
pixel_values = pixel_values.to(self.embeddings.patch_embeddings.weight.dtype)
hidden_states = self.embeddings(pixel_values)
position_embeddings = self.rope_embeddings(pixel_values)

stage_hidden_states: list[torch.Tensor] = [hidden_states]

for layer_module in self.layer:
hidden_states = layer_module(hidden_states, position_embeddings=position_embeddings)
stage_hidden_states.append(hidden_states)

batch_size, _, image_height, image_width = pixel_values.shape
patch_size = self.config.patch_size
num_patches_height = image_height // patch_size
num_patches_width = image_width // patch_size

num_prefix = 1 + getattr(self.config, "num_register_tokens", 0)

feature_maps = []
sequence_output = None
last_stage_idx = len(self.stage_names) - 1
for idx, (stage_name, hidden_state) in enumerate(zip(self.stage_names, stage_hidden_states)):
if idx == last_stage_idx:
hidden_state = self.norm(hidden_state)
sequence_output = hidden_state
elif self.config.apply_layernorm:
hidden_state = self.norm(hidden_state)

if stage_name in self.out_features:
patch_tokens = hidden_state[:, num_prefix:, :]
if self.config.reshape_hidden_states:
fmap = (
patch_tokens.reshape(batch_size, num_patches_height, num_patches_width, patch_tokens.shape[-1])
.permute(0, 3, 1, 2)
.contiguous()
)
else:
fmap = patch_tokens

feature_maps.append(fmap)

output = BackboneOutput(feature_maps=tuple(feature_maps))
output.last_hidden_state = sequence_output

return output


__all__ = ["DINOv3ViTModel", "DINOv3ViTPreTrainedModel"]
__all__ = ["DINOv3ViTModel", "DINOv3ViTPreTrainedModel", "DINOv3ViTBackbone"]
84 changes: 77 additions & 7 deletions src/transformers/models/dinov3_vit/modular_dinov3_vit.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,11 +33,12 @@
from transformers.models.pixtral.modeling_pixtral import PixtralAttention, rotate_half

from ...modeling_layers import GradientCheckpointingLayer
from ...modeling_outputs import BaseModelOutputWithPooling
from ...modeling_outputs import BackboneOutput, BaseModelOutputWithPooling
from ...modeling_utils import ALL_ATTENTION_FUNCTIONS
from ...processing_utils import Unpack
from ...pytorch_utils import compile_compatible_method_lru_cache
from ...utils import TransformersKwargs, auto_docstring, logging
from ...utils import TransformersKwargs, auto_docstring, can_return_tuple, logging
from ...utils.backbone_utils import BackboneMixin
from ...utils.generic import check_model_inputs
from .configuration_dinov3_vit import DINOv3ViTConfig

Expand Down Expand Up @@ -417,10 +418,79 @@ def forward(
sequence_output = self.norm(hidden_states)
pooled_output = sequence_output[:, 0, :]

return BaseModelOutputWithPooling(
last_hidden_state=sequence_output,
pooler_output=pooled_output,
)
return BaseModelOutputWithPooling(last_hidden_state=sequence_output, pooler_output=pooled_output)


@auto_docstring
class DINOv3ViTBackbone(DINOv3ViTPreTrainedModel, BackboneMixin):
def __init__(self, config):
super().__init__(config)
super()._init_backbone(config)

self.embeddings = DINOv3ViTEmbeddings(config)
self.rope_embeddings = DINOv3ViTRopePositionEmbedding(config)
self.layer = nn.ModuleList([DINOv3ViTLayer(config) for _ in range(config.num_hidden_layers)])
self.norm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
self.gradient_checkpointing = False

self.num_features = [config.hidden_size for _ in range(config.num_hidden_layers + 1)]
self.post_init()

def get_input_embeddings(self):
return self.embeddings.patch_embeddings

@check_model_inputs()
@can_return_tuple
def forward(
self,
pixel_values: torch.Tensor,
output_hidden_states: Optional[bool] = None,
**kwargs: Unpack[TransformersKwargs],
) -> BackboneOutput:
pixel_values = pixel_values.to(self.embeddings.patch_embeddings.weight.dtype)
hidden_states = self.embeddings(pixel_values)
position_embeddings = self.rope_embeddings(pixel_values)

stage_hidden_states: list[torch.Tensor] = [hidden_states]

for layer_module in self.layer:
hidden_states = layer_module(hidden_states, position_embeddings=position_embeddings)
stage_hidden_states.append(hidden_states)

batch_size, _, image_height, image_width = pixel_values.shape
patch_size = self.config.patch_size
num_patches_height = image_height // patch_size
num_patches_width = image_width // patch_size

num_prefix = 1 + getattr(self.config, "num_register_tokens", 0)

feature_maps = []
sequence_output = None
last_stage_idx = len(self.stage_names) - 1
for idx, (stage_name, hidden_state) in enumerate(zip(self.stage_names, stage_hidden_states)):
if idx == last_stage_idx:
hidden_state = self.norm(hidden_state)
sequence_output = hidden_state
elif self.config.apply_layernorm:
hidden_state = self.norm(hidden_state)

if stage_name in self.out_features:
patch_tokens = hidden_state[:, num_prefix:, :]
if self.config.reshape_hidden_states:
fmap = (
patch_tokens.reshape(batch_size, num_patches_height, num_patches_width, patch_tokens.shape[-1])
.permute(0, 3, 1, 2)
.contiguous()
)
else:
fmap = patch_tokens

feature_maps.append(fmap)

output = BackboneOutput(feature_maps=tuple(feature_maps))
output.last_hidden_state = sequence_output

return output


__all__ = ["DINOv3ViTModel", "DINOv3ViTPreTrainedModel"]
__all__ = ["DINOv3ViTModel", "DINOv3ViTPreTrainedModel", "DINOv3ViTBackbone"]
54 changes: 52 additions & 2 deletions tests/models/dinov3_vit/test_modeling_dinov3_vit.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@
import torch
from torch import nn

from transformers import DINOv3ViTModel
from transformers import DINOv3ViTBackbone, DINOv3ViTModel


if is_vision_available():
Expand Down Expand Up @@ -112,8 +112,53 @@ def get_config(self):
is_decoder=False,
initializer_range=self.initializer_range,
num_register_tokens=self.num_register_tokens,
stage_names=["embeddings"] + [f"stage{i}" for i in range(1, self.num_hidden_layers + 1)],
out_indices=[0, 1],
reshape_hidden_states=True,
)

def create_and_check_backbone(self, config, pixel_values, labels):
config.out_features = ["stage1", "stage2"]
config.reshape_hidden_states = True

model = DINOv3ViTBackbone(config)
model.to(torch_device)
model.eval()

with torch.no_grad():
outputs = model(pixel_values)

self.parent.assertEqual(len(outputs.feature_maps), 2)
for fm in outputs.feature_maps:
b, c, h, w = fm.shape
self.parent.assertEqual(b, self.batch_size)
self.parent.assertEqual(c, self.hidden_size)
self.parent.assertGreater(h, 0)
self.parent.assertGreater(w, 0)

def test_output_hidden_states(self):
config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()

for model_class in self.all_model_classes:
model = model_class(config)
model.to(torch_device)
model.eval()

with torch.no_grad():
outputs = model(**inputs_dict, output_hidden_states=True)

self.assertIsNotNone(outputs.hidden_states)
expected_num_hidden_states = config.num_hidden_layers + 1
self.assertEqual(len(outputs.hidden_states), expected_num_hidden_states)

for hidden_state in outputs.hidden_states:
expected_shape = (
self.model_tester.batch_size,
self.model_tester.seq_length,
self.model_tester.hidden_size,
)
self.assertEqual(hidden_state.shape, expected_shape)

def create_and_check_model(self, config, pixel_values, labels):
model = DINOv3ViTModel(config=config)
model.to(torch_device)
Expand Down Expand Up @@ -142,7 +187,7 @@ class Dinov3ModelTest(ModelTesterMixin, PipelineTesterMixin, unittest.TestCase):
attention_mask and seq_length.
"""

all_model_classes = (DINOv3ViTModel,) if is_torch_available() else ()
all_model_classes = (DINOv3ViTModel, DINOv3ViTBackbone) if is_torch_available() else ()
pipeline_model_mapping = (
{
"image-feature-extraction": DINOv3ViTModel,
Expand All @@ -153,11 +198,16 @@ class Dinov3ModelTest(ModelTesterMixin, PipelineTesterMixin, unittest.TestCase):

test_resize_embeddings = False
test_torch_exportable = True
test_attention_outputs = False

def setUp(self):
self.model_tester = DINOv3ViTModelTester(self)
self.config_tester = ConfigTester(self, config_class=DINOv3ViTConfig, has_text_modality=False, hidden_size=37)

def test_backbone(self):
config, pixel_values, labels = self.model_tester.prepare_config_and_inputs()
self.model_tester.create_and_check_backbone(config, pixel_values, labels)

def test_config(self):
self.config_tester.run_common_tests()

Expand Down