mmdiff / models /dino_fusion.py
multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
2267636 verified
Raw
History Blame Contribute Delete
3.78 kB
"""torch.hub DINOv3/DINOv2 feature extractor (fallback for the HF extractor).
Used by `core.build_dino_extractor` only when a backbone name is not in the
HuggingFace repo map. Extracts intermediate transformer layers as spatial maps.
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
class DINOv3FeatureExtractor(nn.Module):
"""Extract intermediate ViT layers from a DINOv3/DINOv2 model loaded via torch.hub.
``take_indices`` selects the layers to return (a list of block indices, matching
the HuggingFace extractor); the backbone is frozen.
"""
def __init__(self, model_name="dinov3_vitb16", take_indices=(2, 5, 8, 11)):
super().__init__()
self.model_name = model_name
self.take_indices = list(take_indices)
if model_name.startswith('facebook/'):
hf_to_hub = {
'facebook/dinov2-base': 'dinov2_vitb14',
'facebook/dinov2-small': 'dinov2_vits14',
'facebook/dinov2-large': 'dinov2_vitl14',
'facebook/dinov2-giant': 'dinov2_vitg14',
'facebook/dinov3-vitb16-pretrain-lvd1689m': 'dinov3_vitb16',
'facebook/dinov3-vits16-pretrain-lvd1689m': 'dinov3_vits16',
'facebook/dinov3-vitl16-pretrain-lvd1689m': 'dinov3_vitl16',
}
model_name = hf_to_hub.get(model_name, 'dinov2_vitb14')
self.model_name = model_name
if 'dinov2' in model_name:
self.dino = torch.hub.load('facebookresearch/dinov2', model_name)
self.patch_size = 14
elif 'dinov3' in model_name:
try:
self.dino = torch.hub.load('facebookresearch/dinov3', model_name)
except Exception as e:
# torch.hub gated/unavailable -> fall back to HuggingFace weights.
print(f"[DINO] torch.hub failed ({e}); loading from HuggingFace")
from transformers import AutoModel
hf_map = {
'dinov3_vits16': 'facebook/dinov3-vits16-pretrain-lvd1689m',
'dinov3_vitb16': 'facebook/dinov3-vitb16-pretrain-lvd1689m',
'dinov3_vitl16': 'facebook/dinov3-vitl16-pretrain-lvd1689m',
}
self.dino = AutoModel.from_pretrained(
hf_map.get(model_name, 'facebook/dinov3-vitb16-pretrain-lvd1689m'),
trust_remote_code=True)
self.patch_size = 16
else:
raise ValueError(f"Unsupported model name: {model_name}. Use dinov2_* or dinov3_*")
self.dino.eval()
for p in self.dino.parameters():
p.requires_grad = False
@torch.no_grad()
def forward(self, images):
"""images: [B, 3, H, W] in [0, 1] -> list of feature maps [B, C, H//p, W//p]."""
h, w = images.shape[-2:]
if h % self.patch_size != 0 or w % self.patch_size != 0:
new_h = ((h + self.patch_size - 1) // self.patch_size) * self.patch_size
new_w = ((w + self.patch_size - 1) // self.patch_size) * self.patch_size
images = F.interpolate(images, size=(new_h, new_w), mode='bilinear', align_corners=False)
h, w = new_h, new_w
features = self.dino.get_intermediate_layers(images, self.take_indices, return_class_token=False)
patch_h, patch_w = h // self.patch_size, w // self.patch_size
outs = []
for feat in features:
B, N, C = feat.shape
outs.append(feat.permute(0, 2, 1).reshape(B, C, patch_h, patch_w).contiguous())
return outs
def create_dino_extractor(model_name="dinov3_vitb16", take_indices=(2, 5, 8, 11)):
return DINOv3FeatureExtractor(model_name=model_name, take_indices=take_indices)