DINOv2 多头注意力:3 步看懂视觉聚焦机制
【免费下载链接】dinov2PyTorch code and models for the DINOv2 self-supervised learning method.项目地址: https://gitcode.com/GitHub_Trending/di/dinov2
DINOv2 的视觉 Transformer(vision Transformer,用 Transformer 架构处理图像的模型)里,多头注意力(multi-head attention,把一次注意力拆成多个并行视角)把一张 512×512 的荧光细胞图处理成一组能直接用于分类、分割、深度估计的视觉特征,全程不需要任何标注。
无标签自监督预训练,多头注意力自己学会看图
🔍 注意力头的并行特征提取:从图像块到聚焦特征
图像先被切成 token
语言模型读文章要先分词,视觉模型同理。dinov2/layers/patch_embed.py里的PatchEmbed用一个"卷积核大小=步长"的卷积,把 224×224 的图直接切成 14×14 的小块并投影成向量,得到 256 个 patch token(每个小块对应的向量标记)。
最前面再拼上一个 CLS token(classification token,全局信息的汇聚点),所有 token 加上位置编码(positional encoding,记录每个 token 在图像里的空间位置),后续每一层都在这串 token 上工作。
QKV 与缩放点积注意力在算什么
注意力头就像一组各有偏好的审稿人,每个头独立打分、各盯一种视觉模式。
Attention类只有一个线性层:nn.Linear(dim, dim * 3)同时产出 Q(query,查询)、K(key,键)、V(value,值)三组矩阵,再按头数拆开并行计算,完整实现在 dinov2/layers/attention.py:
# dinov2/layers/attention.py def forward(self, x: Tensor, is_causal: bool = False) -> Tensor: B, N, C = x.shape qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads) # ← 关键:一个线性层产出 Q/K/V q, k, v = torch.unbind(qkv, 2) q, k, v = [t.transpose(1, 2) for t in [q, k, v]] # ← 关键:把头维提到批次维,各头并行 x = nn.functional.scaled_dot_product_attention( q, k, v, attn_mask=None, dropout_p=self.attn_drop if self.training else 0, is_causal=is_causal ) x = x.transpose(1, 2).contiguous().view(B, N, C) x = self.proj_drop(self.proj(x)) # ← 关键:各头输出拼回 C 维,再投影一次 return x核心是缩放点积注意力(Scaled Dot-Product Attention):Q 与 K 点积得到相似度分数,按头维取 -0.5 次方缩放(防止数值过大导致 softmax 输出饱和),softmax 归一化后对 V 加权求和。以 ViT-L 为例,1024 维拆成 16 个头、每头 64 维,各自学着关注不同的视觉模式。
Transformer 块如何逐层叠加
每一层块像一轮审校:注意力改写 token 之间的关系,前馈网络补充每个 token 自身的信息,残差连接(residual connection,把输入原样加回输出)保证训练稳定。
dinov2/layers/block.py的Block骨架是LayerNorm → Attention → 残差 → LayerNorm → MLP → 残差,外面再套 LayerScale(可学习的逐维缩放)和 DropPath(随机深度,训练时按概率丢弃整条残差路径做正则)。dinov2/models/vision_transformer.py 的DinoVisionTransformer把 12~40 个这样的块叠起来,注意力在前、MLP 在后,逐层把"聚焦"细化。
⚡ 显存友好的注意力实现:xFormers 与降级路径
注意力分数矩阵规模是 N×N:DINOv2 默认输入 518×518,切 14×14 块后有 1369 个 token,一次要算上千万个注意力分数,全精度展开非常吃显存。
所有vit_*工厂函数都默认挂载MemEffAttention:优先调用 xFormers 库的分块注意力(memory efficient attention,分块计算、不落地 N² 大矩阵);没装 xFormers 时自动降级,功能不受影响:
# dinov2/layers/attention.py class MemEffAttention(Attention): def forward(self, x: Tensor, attn_bias=None) -> Tensor: if not XFORMERS_AVAILABLE: return super().forward(x) # ← 关键:未装 xFormers 时回退到原生 SDPA B, N, C = x.shape qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads) q, k, v = unbind(qkv, 2) x = memory_efficient_attention(q, k, v, attn_bias=attn_bias) # ← 关键:分块计算,避开 N² 注意力矩阵 return self.proj_drop(self.proj(x.reshape([B, N, C])))想用环境变量强制关闭时,设置XFORMERS_DISABLED即可,加载逻辑会自动回退。
最小可运行路径:4 步拿到 DINOv2 特征输出
- 环境:加载预训练模型只依赖 PyTorch,
pip install torch(建议带 CUDA);可选pip install xformers启用显存优化。 - 拉代码:
git clone https://gitcode.com/GitHub_Trending/di/dinov2 - 加载并调用(入口在
hubconf.py):
# 模型入口见 hubconf.py import torch model = torch.hub.load('<本地仓库路径>', 'dinov2_vitb14', source='local') # ← 关键:一行加载 ViT-B/14 预训练权重 x = torch.randn(1, 3, 224, 224) feats, cls = model.get_intermediate_layers(x, n=1, return_class_token=True) # ← 关键:n=1 表示取末尾 1 层特征 print(feats[0].shape, cls[0].shape)- 看输出:
feats[0]是[1, 256, 768],即 256 个 patch token 的 768 维特征,可直接接下游头;cls[0]是[768]的全局特征,适合接分类器。
延伸方向:细胞显微成像、语义分割与深度估计
- 细胞显微镜场景:仓库自带 Channel-Adaptive DINO,把多通道图像逐通道过 backbone 再做跨通道聚合,CHAMMI 数据集上线性评测(linear evaluation,特征冻结只训一层线性头)结果见 docs/README_CHANNEL_ADAPTIVE_DINO.md:
| 方法 | WTC - Task 1 | HPA - Task 2 | CP - Task 4 |
|---|---|---|---|
| KNN(复现) | 80.3 | 61.4 | 18.4 |
| KNN(论文) | 79.4 | 59.3 | 18.5 |
| 线性评测(复现) | 89.9 | 87.2 | 32.5 |
| 线性评测(论文) | 90.5 | 84.7 | 32.7 |
- 像素级任务:语义分割头与深度估计头(如 DPT head,把 patch 特征重采样回像素分辨率再回归深度)直接消费前面拿到的特征,入口在
dinov2/eval/segmentation/与dinov2/eval/depth/,配套示例见notebooks/depth_estimation.ipynb。 - 训练侧机制:teacher-student 自蒸馏(self-distillation,教师网络把自己学到的东西教给学生网络)的完整流程在
dinov2/train/ssl_meta_arch.py,想复现自监督预训练可以从这里读起。
那张 512×512 的细胞图,从切块、256 个 token 到一组可直接迁移的特征,全程没有一张标签。想继续深挖,直接翻dinov2/layers/目录的源码,或到项目 issue 区提问。
【免费下载链接】dinov2PyTorch code and models for the DINOv2 self-supervised learning method.项目地址: https://gitcode.com/GitHub_Trending/di/dinov2
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考