multimodalart HF Staff commited on
Commit
ec46c56
·
verified ·
1 Parent(s): ebd8815

Upload models/dinov3_hf_extractor.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. models/dinov3_hf_extractor.py +52 -58
models/dinov3_hf_extractor.py CHANGED
@@ -3,39 +3,49 @@
3
  import os
4
  import torch
5
  import torch.nn as nn
6
- from transformers import AutoImageProcessor, AutoModel
7
 
8
 
9
  class DINOv3HFExtractor(nn.Module):
10
  """
11
  Extracts intermediate features from DINOv3 via HuggingFace transformers.
12
-
13
- Properly handles DINOv3 register tokens:
14
- - Output shape: [B, 1 + P + R, C] where:
15
- - 1 = CLS token
16
- - P = spatial patches (1024 for 512x512 with 16x16 patches)
17
- - R = register tokens (typically 4 for DINOv3)
18
-
19
- Returns 4 feature maps of shape [B, C_dino, 32, 32] from last 4 layers.
20
- Input images must be [B, 3, 512, 512] in [0, 1] range.
21
  """
22
-
23
- def __init__(self, repo_id="facebook/dinov3-vitb16-pretrain-lvd1689m", take_last=None, take_indices=None, trainable=False):
 
 
 
 
24
  super().__init__()
25
 
26
- token = os.environ.get("HF_TOKEN")
27
- self.proc = AutoImageProcessor.from_pretrained(repo_id, token=token)
28
-
29
- # Disable resizing/cropping so native 512x512 maps to 32x32 patches
30
  for k in ("do_resize", "do_center_crop"):
31
  if hasattr(self.proc, k):
32
  setattr(self.proc, k, False)
33
-
34
- self.model = AutoModel.from_pretrained(repo_id, token=token)
 
 
 
 
 
 
 
 
 
 
 
35
  self.model.config.output_hidden_states = True
36
-
37
- # trainable=True is used by the `dino_only` ablation (fine-tune the backbone);
38
- # otherwise the backbone is frozen and kept in eval mode.
39
  self._frozen = not trainable
40
  if self._frozen:
41
  self.model.eval()
@@ -45,93 +55,77 @@ class DINOv3HFExtractor(nn.Module):
45
  self.model.train()
46
  for p in self.model.parameters():
47
  p.requires_grad = True
48
-
49
- # ImageNet normalization stats: mean [0.485, 0.456, 0.406], std [0.229, 0.224, 0.225]
50
  mean = torch.tensor(self.proc.image_mean).view(1, 3, 1, 1)
51
  std = torch.tensor(self.proc.image_std).view(1, 3, 1, 1)
52
  self.register_buffer("mean", mean, persistent=False)
53
  self.register_buffer("std", std, persistent=False)
54
-
55
  if take_indices is not None:
56
  self.take_indices = take_indices
57
  self.take_last = None
58
  else:
59
  self.take_last = take_last if take_last is not None else 4
60
  self.take_indices = None
61
-
62
  self.patch_size = getattr(self.model.config, "patch_size", 16)
63
  self.num_register_tokens = getattr(self.model.config, "num_register_tokens", 0)
64
-
65
  hidden_size = getattr(self.model.config, "hidden_size", 768)
66
-
67
  layers = self.take_indices if self.take_indices is not None else f"last {self.take_last}"
68
  trainable_str = "trainable" if not self._frozen else "frozen"
69
- print(f"[DINOv3] {repo_id.split('/')[-1]}: dim={hidden_size}, patch={self.patch_size}, "
70
  f"layers={layers} ({trainable_str})")
71
-
72
  def train(self, mode: bool = True):
73
- """Keep a frozen backbone in eval mode; otherwise follow `mode`."""
74
  self.training = mode
75
  if self._frozen:
76
  self.model.eval()
77
  else:
78
  self.model.train(mode)
79
  return self
80
-
81
  def forward(self, images_512: torch.Tensor):
82
- """Extract DINOv3 features from [B, 3, H, W] images in [0, 1]; returns maps [B, C, H//16, W//16]."""
83
- # Disable grad only when the backbone is frozen; otherwise allow fine-tuning
84
  with torch.set_grad_enabled(not self._frozen):
85
  return self._forward(images_512)
86
 
87
  def _forward(self, images_512: torch.Tensor):
88
  x = (images_512 - self.mean) / self.std
89
-
90
  out = self.model(pixel_values=x, output_hidden_states=True)
91
- hidden_states = out.hidden_states # Tuple: [emb, blk1, ..., blkL]
92
-
93
  B, _, H, W = images_512.shape
94
  H_patches = H // self.patch_size
95
  W_patches = W // self.patch_size
96
- P = H_patches * W_patches # total spatial patches
97
  R = self.num_register_tokens
98
-
99
- # Each hidden state is [B, 1 + P + R, C], laid out as [CLS][P spatial patches][R register tokens]
100
  maps = []
101
  if self.take_indices is not None:
102
  for idx in self.take_indices:
103
  hidden = hidden_states[idx]
104
- spatial = hidden[:, 1:1+P, :] # [B, P, C], drop CLS and register tokens
105
  C = spatial.shape[-1]
106
  spatial_map = spatial.transpose(1, 2).reshape(B, C, H_patches, W_patches).contiguous()
107
  maps.append(spatial_map)
108
  else:
109
  for hidden in hidden_states[-self.take_last:]:
110
- spatial = hidden[:, 1:1+P, :] # [B, P, C], drop CLS and register tokens
111
  C = spatial.shape[-1]
112
  spatial_map = spatial.transpose(1, 2).reshape(B, C, H_patches, W_patches).contiguous()
113
  maps.append(spatial_map)
114
-
115
  return maps
116
 
117
 
118
  def create_dinov3_hf_extractor(repo_id="facebook/dinov3-vitb16-pretrain-lvd1689m", take_last=None, take_indices=None, trainable=False):
119
  """
120
  Factory for DINOv3HFExtractor (frozen in eval mode unless trainable=True).
121
- repo_id options: vits16 (384-dim), vitb16 (768-dim, recommended), vitl16 (1024-dim).
122
  """
123
- return DINOv3HFExtractor(repo_id=repo_id, take_last=take_last, take_indices=take_indices, trainable=trainable)
124
-
125
-
126
- if __name__ == "__main__":
127
- extractor = create_dinov3_hf_extractor().cuda()
128
-
129
- dummy_input = torch.randn(2, 3, 512, 512).cuda()
130
- print(f"\nInput: {dummy_input.shape}")
131
-
132
- features = extractor(dummy_input)
133
-
134
- print(f"\nExtracted {len(features)} feature maps:")
135
- for i, feat in enumerate(features):
136
- print(f" Layer {i}: {feat.shape}")
137
-
 
3
  import os
4
  import torch
5
  import torch.nn as nn
6
+ from transformers import DINOv3ViTConfig, DINOv3ViTModel, DINOv3ViTImageProcessorFast
7
 
8
 
9
  class DINOv3HFExtractor(nn.Module):
10
  """
11
  Extracts intermediate features from DINOv3 via HuggingFace transformers.
12
+
13
+ Builds the model from config (no gated download required); weights are loaded
14
+ from the MMDiff checkpoint which bundles the DINOv3 backbone.
15
+
16
+ Returns 4 feature maps of shape [B, C_dino, H//16, W//16] from selected layers.
17
+ Input images must be [B, 3, H, W] in [0, 1] range.
 
 
 
18
  """
19
+
20
+ def __init__(self, repo_id="facebook/dinov3-vitb16-pretrain-lvd1689m",
21
+ take_last=None, take_indices=None, trainable=False,
22
+ hidden_size=768, num_hidden_layers=12, num_attention_heads=12,
23
+ intermediate_size=3072, patch_size=16, image_size=512,
24
+ num_register_tokens=4):
25
  super().__init__()
26
 
27
+ # Build image processor from default config (no gated download needed)
28
+ self.proc = DINOv3ViTImageProcessorFast()
29
+
30
+ # Disable resizing/cropping so native resolution maps to patches
31
  for k in ("do_resize", "do_center_crop"):
32
  if hasattr(self.proc, k):
33
  setattr(self.proc, k, False)
34
+
35
+ # Build model from config (random weights; real weights loaded from checkpoint)
36
+ config = DINOv3ViTConfig(
37
+ hidden_size=hidden_size,
38
+ num_hidden_layers=num_hidden_layers,
39
+ num_attention_heads=num_attention_heads,
40
+ intermediate_size=intermediate_size,
41
+ patch_size=patch_size,
42
+ image_size=image_size,
43
+ num_register_tokens=num_register_tokens,
44
+ hidden_act="gelu",
45
+ )
46
+ self.model = DINOv3ViTModel(config)
47
  self.model.config.output_hidden_states = True
48
+
 
 
49
  self._frozen = not trainable
50
  if self._frozen:
51
  self.model.eval()
 
55
  self.model.train()
56
  for p in self.model.parameters():
57
  p.requires_grad = True
58
+
59
+ # ImageNet normalization stats
60
  mean = torch.tensor(self.proc.image_mean).view(1, 3, 1, 1)
61
  std = torch.tensor(self.proc.image_std).view(1, 3, 1, 1)
62
  self.register_buffer("mean", mean, persistent=False)
63
  self.register_buffer("std", std, persistent=False)
64
+
65
  if take_indices is not None:
66
  self.take_indices = take_indices
67
  self.take_last = None
68
  else:
69
  self.take_last = take_last if take_last is not None else 4
70
  self.take_indices = None
71
+
72
  self.patch_size = getattr(self.model.config, "patch_size", 16)
73
  self.num_register_tokens = getattr(self.model.config, "num_register_tokens", 0)
74
+
75
  hidden_size = getattr(self.model.config, "hidden_size", 768)
76
+
77
  layers = self.take_indices if self.take_indices is not None else f"last {self.take_last}"
78
  trainable_str = "trainable" if not self._frozen else "frozen"
79
+ print(f"[DINOv3] Built from config: dim={hidden_size}, patch={self.patch_size}, "
80
  f"layers={layers} ({trainable_str})")
81
+
82
  def train(self, mode: bool = True):
 
83
  self.training = mode
84
  if self._frozen:
85
  self.model.eval()
86
  else:
87
  self.model.train(mode)
88
  return self
89
+
90
  def forward(self, images_512: torch.Tensor):
 
 
91
  with torch.set_grad_enabled(not self._frozen):
92
  return self._forward(images_512)
93
 
94
  def _forward(self, images_512: torch.Tensor):
95
  x = (images_512 - self.mean) / self.std
96
+
97
  out = self.model(pixel_values=x, output_hidden_states=True)
98
+ hidden_states = out.hidden_states
99
+
100
  B, _, H, W = images_512.shape
101
  H_patches = H // self.patch_size
102
  W_patches = W // self.patch_size
103
+ P = H_patches * W_patches
104
  R = self.num_register_tokens
105
+
 
106
  maps = []
107
  if self.take_indices is not None:
108
  for idx in self.take_indices:
109
  hidden = hidden_states[idx]
110
+ spatial = hidden[:, 1:1+P, :]
111
  C = spatial.shape[-1]
112
  spatial_map = spatial.transpose(1, 2).reshape(B, C, H_patches, W_patches).contiguous()
113
  maps.append(spatial_map)
114
  else:
115
  for hidden in hidden_states[-self.take_last:]:
116
+ spatial = hidden[:, 1:1+P, :]
117
  C = spatial.shape[-1]
118
  spatial_map = spatial.transpose(1, 2).reshape(B, C, H_patches, W_patches).contiguous()
119
  maps.append(spatial_map)
120
+
121
  return maps
122
 
123
 
124
  def create_dinov3_hf_extractor(repo_id="facebook/dinov3-vitb16-pretrain-lvd1689m", take_last=None, take_indices=None, trainable=False):
125
  """
126
  Factory for DINOv3HFExtractor (frozen in eval mode unless trainable=True).
127
+ Builds from config no gated download required.
128
  """
129
+ return DINOv3HFExtractor(
130
+ repo_id=repo_id, take_last=take_last, take_indices=take_indices, trainable=trainable,
131
+ )