@@ -156,10 +156,11 @@ def forward(self, pixel_values):
156156
157157
158158class DINOv3ViTEmbeddings (nn .Module ):
159- def __init__ (self , hidden_size , num_register_tokens , num_channels , patch_size , dtype , device , operations ):
159+ def __init__ (self , hidden_size , num_register_tokens , num_channels , patch_size , dtype , device , operations , use_mask_token = True ):
160160 super ().__init__ ()
161161 self .cls_token = nn .Parameter (torch .empty (1 , 1 , hidden_size , device = device , dtype = dtype ))
162- self .mask_token = nn .Parameter (torch .empty (1 , 1 , hidden_size , device = device , dtype = dtype ))
162+ self .mask_token = nn .Parameter (torch .empty (1 , 1 , hidden_size , device = device , dtype = dtype )) if use_mask_token else None
163+
163164 self .register_tokens = nn .Parameter (torch .empty (1 , num_register_tokens , hidden_size , device = device , dtype = dtype ))
164165 self .patch_embeddings = operations .Conv2d (
165166 num_channels , hidden_size , kernel_size = patch_size , stride = patch_size , device = device , dtype = dtype
@@ -212,7 +213,7 @@ def forward(self, hidden_states, attention_mask=None, position_embeddings=None):
212213
213214
214215class DINOv3ViTModel (nn .Module ):
215- def __init__ (self , config , dtype , device , operations ):
216+ def __init__ (self , config , dtype , device , operations , use_mask_token = True ):
216217 super ().__init__ ()
217218 num_hidden_layers = config ["num_hidden_layers" ]
218219 hidden_size = config ["hidden_size" ]
@@ -228,7 +229,7 @@ def __init__(self, config, dtype, device, operations):
228229
229230 self .embeddings = DINOv3ViTEmbeddings (
230231 hidden_size , num_register_tokens , num_channels = num_channels , patch_size = patch_size ,
231- dtype = dtype , device = device , operations = operations
232+ dtype = dtype , device = device , operations = operations , use_mask_token = use_mask_token
232233 )
233234 self .rope_embeddings = DINOv3ViTRopePositionEmbedding (
234235 rope_theta , hidden_size , num_attention_heads , patch_size = patch_size , dtype = dtype , device = device
@@ -240,6 +241,10 @@ def __init__(self, config, dtype, device, operations):
240241 for _ in range (num_hidden_layers )])
241242 self .norm = operations .LayerNorm (hidden_size , eps = layer_norm_eps , dtype = dtype , device = device )
242243
244+ self .patch_size = patch_size
245+ self .embed_dim = self .embed_dims = hidden_size
246+ self .num_prefix_tokens = 1 + num_register_tokens # cls + register
247+
243248 def get_input_embeddings (self ):
244249 return self .embeddings .patch_embeddings
245250
@@ -257,3 +262,11 @@ def forward(self, pixel_values, bool_masked_pos=None, **kwargs):
257262 sequence_output = norm (hidden_states )
258263 pooled_output = sequence_output [:, 0 , :]
259264 return sequence_output , None , pooled_output , None
265+
266+ def forward_features (self , pixel_values , ** kwargs ):
267+ sequence_output = self .forward (pixel_values , ** kwargs )[0 ]
268+ b = pixel_values .shape [0 ]
269+ h = pixel_values .shape [- 2 ] // self .patch_size
270+ w = pixel_values .shape [- 1 ] // self .patch_size
271+ patches = sequence_output [:, self .num_prefix_tokens :, :]
272+ return patches .reshape (b , h , w , self .embed_dim ).permute (0 , 3 , 1 , 2 ).contiguous ()
0 commit comments