Spaces:
Sleeping
Sleeping
Upload models/dinov3_hf_extractor.py with huggingface_hub
Browse files- 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
|
| 7 |
|
| 8 |
|
| 9 |
class DINOv3HFExtractor(nn.Module):
|
| 10 |
"""
|
| 11 |
Extracts intermediate features from DINOv3 via HuggingFace transformers.
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
|
| 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",
|
|
|
|
|
|
|
|
|
|
|
|
|
| 24 |
super().__init__()
|
| 25 |
|
| 26 |
-
|
| 27 |
-
self.proc =
|
| 28 |
-
|
| 29 |
-
# Disable resizing/cropping so native
|
| 30 |
for k in ("do_resize", "do_center_crop"):
|
| 31 |
if hasattr(self.proc, k):
|
| 32 |
setattr(self.proc, k, False)
|
| 33 |
-
|
| 34 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 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]
|
| 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
|
| 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
|
| 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, :]
|
| 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, :]
|
| 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 |
-
|
| 122 |
"""
|
| 123 |
-
return DINOv3HFExtractor(
|
| 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 |
+
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|