Spaces:
Running
Running
File size: 8,696 Bytes
e16aadc 1c927c2 e16aadc 1c927c2 e16aadc 1c927c2 e16aadc 1c927c2 31d7432 1c927c2 31d7432 1c927c2 31d7432 1c927c2 8ab04fc 1c927c2 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 | """DINOv2 特征提取器 — vit_small_patch14_reg4_dinov2 (带寄存器令牌)"""
import torch
import torch.nn.functional as F
from torchvision import transforms
from PIL import Image
import numpy as np
from module.config import DEVICE
class DINOv2Extractor:
"""DINOv2 全局语义特征提取器
使用 vit_small_patch14_reg4_dinov2 模型(4个寄存器令牌),
相比无寄存器版本,CLS token 特征质量更高,对背景噪声更鲁棒。
输出 384 维归一化特征向量。
支持提取 patch 特征用于可视化匹配。
"""
TIMM_MODEL_NAME = 'vit_small_patch14_reg4_dinov2'
PATCH_SIZE = 14
def __init__(self, device=None):
self.device = device or DEVICE
self.model = None
self._use_timm = False
# DINOv2 标准预处理:518x518,ImageNet 归一化
self.transform = transforms.Compose([
transforms.Resize(518, interpolation=transforms.InterpolationMode.BICUBIC),
transforms.CenterCrop(518),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]),
])
try:
import timm
print(f"通过 timm 加载 DINOv2 (reg4): {self.TIMM_MODEL_NAME}")
self.model = timm.create_model(self.TIMM_MODEL_NAME, pretrained=True)
self._use_timm = True
self.model = self.model.to(self.device).eval()
print("DINOv2 reg4 模型加载完成")
except Exception as e:
print(f"timm 加载 DINOv2 失败: {e}")
self.model = None
@property
def is_available(self):
return self.model is not None
@torch.no_grad()
def _get_cls_feature(self, img_tensor):
"""从 forward_features 提取 CLS token,适配 timm 的 dict 返回格式"""
features = self.model.forward_features(img_tensor)
# timm 0.9+ 返回 dict
if isinstance(features, dict):
if 'x_norm_clstoken' in features:
return features['x_norm_clstoken']
# 兜底:取第一个值的 CLS 位置
for v in features.values():
if isinstance(v, torch.Tensor) and v.dim() >= 2:
return v[:, 0]
# 旧版 timm 返回 tensor (B, N, C)
if isinstance(features, torch.Tensor):
return features[:, 0]
raise ValueError(f"无法从 forward_features 提取 CLS token: {type(features)}")
@torch.no_grad()
def extract_feature(self, img_path):
"""提取单张图片的 DINOv2 CLS 全局特征向量(384维,已归一化)"""
if self.model is None:
return None
img = Image.open(img_path).convert('RGB')
tensor = self.transform(img).unsqueeze(0).to(self.device)
cls_token = self._get_cls_feature(tensor)
return F.normalize(cls_token, dim=-1).squeeze(0)
@torch.no_grad()
def compute_similarity(self, img_path1, img_path2):
"""计算两张图片的 DINOv2 CLS 特征余弦相似度(返回 0~1)"""
feat1 = self.extract_feature(img_path1)
feat2 = self.extract_feature(img_path2)
if feat1 is None or feat2 is None:
return None
sim = F.cosine_similarity(feat1.unsqueeze(0), feat2.unsqueeze(0))
return float(sim.item())
@torch.no_grad()
def extract_patch_features(self, img_path):
"""提取图片的 patch 级别特征(用于可视化匹配)
Returns:
dict with keys:
- 'cls': CLS token (384维)
- 'patches': patch token 特征 (n_patches, 384)
- 'patch_grid': (n_h, n_w) patch 网格尺寸
- 'image_size': 原始图片尺寸 (w, h)
"""
if self.model is None:
return None
img_pil = Image.open(img_path).convert('RGB')
orig_w, orig_h = img_pil.size
tensor = self.transform(img_pil).unsqueeze(0).to(self.device)
features = self.model.forward_features(tensor)
if isinstance(features, dict):
cls_token = features.get('x_norm_clstoken', None)
patch_tokens = features.get('x_norm_patchtokens', None)
# 如果没有 x_norm_patchtokens,从 x_norm_tokens 中提取
if patch_tokens is None and 'x_norm_tokens' in features:
all_tokens = features['x_norm_tokens']
# 第0个是 CLS,后面是 patches + registers
# reg4 模型有4个 register tokens 在末尾
patch_tokens = all_tokens[:, 1:-4] if all_tokens.shape[1] > 1 else None
if cls_token is None and 'x_norm_tokens' in features:
cls_token = features['x_norm_tokens'][:, 0]
elif isinstance(features, torch.Tensor):
# 旧版 timm: (B, N, C),N = 1(cls) + n_patches + n_registers
cls_token = features[:, 0]
# reg4 有4个 register tokens 在末尾
patch_tokens = features[:, 1:-4] if features.shape[1] > 5 else features[:, 1:]
else:
return None
if cls_token is None or patch_tokens is None:
return None
cls_feat = F.normalize(cls_token.squeeze(0), dim=-1)
patch_feats = F.normalize(patch_tokens.squeeze(0), dim=-1)
# DINOv2 预处理:Resize(518) → CenterCrop(518)
# 需要记录 CenterCrop 的偏移量,用于可视化时将 patch 坐标映射回原图
input_size = 518 # CenterCrop 尺寸
n_patches = patch_feats.shape[0]
n_side = int(n_patches ** 0.5)
if n_side * n_side != n_patches:
n_h = input_size // self.PATCH_SIZE
n_w = n_h
else:
n_h = n_w = n_side
# 计算 CenterCrop 在 Resize 后图像上的偏移
# Resize(518) 将短边缩放到518,长边按比例缩放
resize_w, resize_h = orig_w, orig_h
if orig_w <= orig_h:
resize_w = 518
resize_h = int(orig_h * 518 / orig_w)
else:
resize_h = 518
resize_w = int(orig_w * 518 / orig_h)
crop_x = (resize_w - input_size) // 2
crop_y = (resize_h - input_size) // 2
return {
'cls': cls_feat.cpu().numpy(),
'patches': patch_feats.cpu().numpy(),
'patch_grid': (n_h, n_w),
'image_size': (orig_w, orig_h),
'crop_offset': (crop_x, crop_y), # CenterCrop 偏移
'resize_size': (resize_w, resize_h), # Resize 后尺寸
}
def compute_patch_matches(self, img_path1, img_path2, top_k=50):
"""计算两张图片之间的 patch 级别匹配(用于可视化)
Returns:
dict with keys:
- 'cls_similarity': CLS token 余弦相似度
- 'matches': list of (query_idx, hist_idx, similarity) 最优匹配
- 'query_grid': (n_h, n_w)
- 'hist_grid': (n_h, n_w)
- 'query_image_size': (w, h)
- 'hist_image_size': (w, h)
"""
feat1 = self.extract_patch_features(img_path1)
feat2 = self.extract_patch_features(img_path2)
if feat1 is None or feat2 is None:
return None
# CLS 相似度
cls_sim = float(np.dot(feat1['cls'], feat2['cls']))
# Patch 间余弦相似度矩阵
sim_matrix = feat1['patches'] @ feat2['patches'].T # (n1, n2)
# 双向最近邻匹配(类似 MASt3R 的互为最近邻)
matches = []
q_best = sim_matrix.argmax(axis=1) # 每个 query patch 的最佳匹配
h_best = sim_matrix.argmax(axis=0) # 每个 hist patch 的最佳匹配
for qi, hi in enumerate(q_best):
if h_best[hi] == qi: # 互为最近邻
sim = float(sim_matrix[qi, hi])
matches.append((qi, hi, sim))
# 按相似度排序,取 top_k
matches.sort(key=lambda x: x[2], reverse=True)
matches = matches[:top_k]
return {
'cls_similarity': round(cls_sim, 4),
'matches': [(q, h, round(s, 4)) for q, h, s in matches],
'query_grid': feat1['patch_grid'],
'hist_grid': feat2['patch_grid'],
'query_image_size': feat1['image_size'],
'hist_image_size': feat2['image_size'],
'query_crop_offset': feat1.get('crop_offset', (0, 0)),
'hist_crop_offset': feat2.get('crop_offset', (0, 0)),
'query_resize_size': feat1.get('resize_size', (518, 518)),
'hist_resize_size': feat2.get('resize_size', (518, 518)),
}
|