ssl_loss理解
DINOv3 自监督损失笔记
对应代码:
dinov3/layers/dino_head.py—DINOHead(DINO 和 iBOT 共用的投影头)dinov3/loss/dino_clstoken_loss.py—DINOLoss+ Sinkhorn-Knoppdinov3/loss/ibot_patch_loss.py—iBOTPatchLoss+SinkhornKnoppTeacherdinov3/loss/koleo_loss.py—KoLeoLoss/KoLeoLossDistributeddinov3/loss/gram_loss.py—GramLoss(DINOv3 新增)dinov3/train/ssl_meta_arch.py— 损失组装 + 调度(ssl_meta_arch.py:584-684)dinov3/configs/ssl_default_config.yaml— 默认超参dinov3/configs/train/dinov3_vit7b16_gram_anchor.yaml— DINOv3 7B Gram Anchoring 配置
0. 整体架构:四把锁怎么拧
DINOv3 的自监督学习沿用 DINOv2 的多任务配方。一个 SSL step 跑 4 个损失,总和反传:
| 损失 | 作用对象 | teacher 目标 | 核心思路 | DINOv2 | DINOv3 |
|---|---|---|---|---|---|
| DINO loss | student 全部 crop 的 [CLS] | teacher global crop 的 [CLS](Sinkhorn 中心化) | 把所有 view 的 [CLS] 拉到统一原型分布 | ✓ | ✓ |
| iBOT loss | student 被掩码 的 patch | teacher 同位置 patch(Sinkhorn 中心化) | 掩码 patch 重建式对齐 | ✓ | ✓ |
| KoLeo loss | student global crop 的 [CLS] pre-head | 无(自正则) | 在 L2 球面均匀分散 | ✓ | ✓ |
| Gram loss | student patch 特征 | 另一份更高分辨率的 gram teacher | patch 间相关矩阵对齐 | ✗ | 新增 |
注:教师自身在每步后做 EMA(teacher = m·teacher + (1-m)·student),并对教师预测做 centering / sharpening(DINO 部分用 Sinkhorn,iBOT 也用 Sinkhorn,但 centering buffer 不同)。
1. 投影头 DINOHead:把 backbone 特征映射到"原型"空间
1.1 结构
DINOHead(
in_dim, # = backbone.embed_dim(D)
out_dim, # = 65536 (DINO) / 98304 (iBOT 7B) — 原型数 K
nlayers=3, # MLP 层数
hidden_dim=2048, # 中间层宽度
bottleneck_dim=256, # 最后一层输入宽度(注意:先 3 层 MLP → bottleneck → L2 norm → 线性到 K)
use_bn=False,
)结构示意(nlayers=3):
in_dim → Linear → hidden_dim → GELU → Linear → hidden_dim → GELU
→ Linear → bottleneck_dim → L2Norm → Linear(bias=False) → out_dim要点:
last_layer无 bias(dino_head.py:32)— 论文里的 trick,让 head 的梯度范数易控。- forward 走两步法(dino_head.py:43-50):
no_last_layer=True:只跑 MLP+bottleneck+L2 norm,用于拿"归一化后特征"(如做 KoLeo 之前)。only_last_layer=True:只跑最后一层last_layer(用于"unbind last layer"梯度解耦,论文里常用)。
1.2 3 个 head 在 SSLMetaArch 里各自的位置
student.dino_head/teacher.dino_head:输入 [CLS] pre-head,输出K=65536维的"原型 logits",用于 DINO 损失。student.ibot_head/teacher.ibot_head:输入被掩码 patch 的 pre-head 特征,输出K=98304维(7B)原型 logits,用于 iBOT 损失。- 论文里说"weight normalization on the last layer"对应
force_weight_norm(ssl_default_config.yaml:24,本仓库默认false)。
2. DINO 损失:CLS token 的多视图对齐
dinov3/loss/dino_clstoken_loss.py— 124 行 核心方法:sinkhorn_knopp_teacher(第 42-70 行)+forward(第 72-99 行)
2.1 一句话定义
DINO 损失 = student 各 crop 的 [CLS] 软标签,与 teacher global crop [CLS] 中心化分布做交叉熵,teacher 中心用 Sinkhorn-Knopp 估。
2.2 三个核心部件
A. center 缓冲(第 25-30 行)
self.register_buffer("center", torch.full((1, out_dim), math.nan))
self.updated = True # 是否还有未应用的反向累加
self.reduce_handle = None # 异步 all-reduce 句柄
self.async_batch_center = None # 累加中的"本步中心估计"center形状[1, K],用 nan 初始化,强制init_weights()必须显式清零。- 异步流水线:先发起异步
all_reduce,等下一次softmax_center_teacher调用时再wait()应用。update_centers控制本次是否更新(首步和被冻结 last-layer 阶段不更新,见下面 §2.4)。
B. softmax_center_teacher(第 36-40 行)
@torch.no_grad()
def softmax_center_teacher(self, teacher_output, teacher_temp, update_centers=True):
if update_centers:
self.apply_center_update()
return F.softmax((teacher_output - self.center) / teacher_temp, dim=-1)- centering:减
center(EMA 累计的均值),避免某一两个 prototype 主导。 - sharpening:除以
teacher_temp(默认 0.04→0.07,warmup 后),温度越低分布越尖。 - 注意:这个函数在 DINOv3 流程里其实没用上,因为 teacher 走的是
sinkhorn_knopp_teacher(DINOv2 起就用 SK 替代了简单的 softmax-center)。
C. sinkhorn_knopp_teacher(第 42-70 行)
DINOv2 起默认用 Sinkhorn-Knopp 做 teacher 中心化,比 softmax-center 更稳定。它把 teacher 输出归一化为一个双随机矩阵(行列和都为 1),并按真实 batch 大小
B缩放,使每列和=1(分配概率)。
@torch.no_grad()
def sinkhorn_knopp_teacher(self, teacher_output, teacher_temp, n_iterations=3):
teacher_output = teacher_output.float()
world_size = get_subgroup_size() if dist.is_initialized() else 1
Q = torch.exp(teacher_output / teacher_temp).t() # K × B
B = Q.shape[1] * world_size # 跨 rank 总样本数
K = Q.shape[0] # 原型数
sum_Q = torch.sum(Q)
if dist.is_initialized():
dist.all_reduce(sum_Q, group=get_process_subgroup()) # 跨 rank 总和
Q /= sum_Q # 矩阵总 = 1
for _ in range(n_iterations):
# 1) 行归一:每原型总权重 = 1/K
sum_of_rows = torch.sum(Q, dim=1, keepdim=True)
if dist.is_initialized():
dist.all_reduce(sum_of_rows, group=get_process_subgroup())
Q /= sum_of_rows
Q /= K
# 2) 列归一:每样本总权重 = 1/B
Q /= torch.sum(Q, dim=0, keepdim=True)
Q /= B
Q *= B
return Q.t() # B × K,每列和=1关键点:
Q = exp(t/τ):先把 teacher logits 变成正的非归一化"软分配"矩阵。- 3 次迭代交替做行 / 列归一。
all_reduce在process_subgroup上做(DINO 用数据并行子组,区别于 gram teacher 用的 default group,详见 ssl_meta_arch.py:788-805)。- 返回的张量每列和=1,对应每个样本在 K 个 prototype 上的概率分布。
D. forward(第 72-99 行)——学生-教师交叉熵
def forward(self, student_logits, teacher_probs, ignore_diagonal=False):
student_crops, B, K = student_logits.shape # [S, B, K]
teacher_crops, _, _ = teacher_probs.shape # [T, B, K]
student_logits = F.log_softmax(student_logits.float() / self.student_temp, dim=-1)
if not ignore_diagonal:
loss = -torch.einsum("s b k, t b k -> ", student_logits, teacher_probs)
return loss / (B * student_crops * teacher_crops)
else:
loss = -torch.einsum("s b k, t b k -> s t", student_logits, teacher_probs)
min_st = min(student_crops, teacher_crops)
loss = torch.diagonal_scatter(loss, loss.new_zeros(min_st))
return loss.sum() / (B * student_crops * teacher_crops - B * min_st)实现细节:
student_temp默认 0.1,比 teacher temp 更大(student 更软,target 更尖)。- 矩阵
student_logits[s, b]与teacher_probs[t, b]点积 = “student crop s 用 teacher crop t 当 target” 的交叉熵。einsum把所有 (s, t, b) 三轴一起求和。 ignore_diagonal=True时(默认 ssl_default_config.yaml:12):把同 crop 对 (s==t) 的项置零。直觉是同一张图被增强成两个 global view,自对齐(A↔A)信息泄露,应被排除。diagonal_scatter是 PyTorch 2.x 的高效 API,比手工构造 mask 快。forward的"按 (s,t,b) 三轴 sum / count"对应"所有 student-teacher crop 对求平均"。
2.3 多视图配对:哪几张 view 互相对齐?
# 1) student(local) vs teacher(global)
dino_local_crops_loss = self.dino_loss(
student_logits=student_local["cls_after_head"],
teacher_probs=teacher_global["cls_centered"],
)
# 2) student(global) vs teacher(global)
dino_global_crops_loss = self.dino_loss(
student_logits=student_global["cls_after_head"],
teacher_probs=teacher_global["cls_centered"],
ignore_diagonal=self.dino_global_ignore_diagonal, # 默认 True
)再算 dino_global_scale / dino_local_scale(ssl_meta_arch.py:602-607):
dino_global_terms = n_global_crops * (n_global_crops - 1) # 2 * 1 = 2
dino_local_terms = n_global_crops * n_local_crops # 2 * 8 = 16
dino_global_scale = 2 / 18
dino_local_scale = 16 / 18含义:DINO 损失合计是 1.0 * (2/18 * dino_global + 16/18 * dino_local),让 local 和 global 项的总权重平衡(数量多的小 view 不至于主导梯度)。
2.4 教师 EMA + center 异步流水线
update_center → reduce_center_update → apply_center_update 三步法(第 101-124 行):
reduce_center_update:发起异步all_reduce求 teacher 输出在dim=0(= crops×B)上的和。- 下一轮前向
softmax_center_teacher(update_centers=True)时再wait(),然后做 EMA:center = center * 0.9 + (sum / len / world_size) * 0.1 - 配合
dinov3/train/train.py在每 step 末尾调用update_ema(m)(ssl_meta_arch.py:713-726):用torch._foreach_mul_(teacher_param_list, m) # teacher *= m torch._foreach_add_(teacher_param_list, student_param_list, alpha=1 - m)_foreach原地更新,省内存带宽。m从 0.992 cosine 到 1.0。
冻结最后层:DINO 论文里为稳定训练,前几 epoch 冻结 head 的 last_layer 权重更新(freeze_last_layer_epochs=1,见 ssl_default_config.yaml:139)。
3. iBOT 损失:被掩码 patch token 的对齐
dinov3/loss/ibot_patch_loss.py— 142 行 核心方法:SinkhornKnoppTeacher(第 20-58 行)+forward_masked(第 96-117 行)
3.1 一句话定义
iBOT 损失 = 把 student 看到的"被掩码的 patch 位置"的特征,经过 ibot_head 投到 K 维原型空间,与 teacher 同位置 patch 的 Sinkhorn 软标签做交叉熵。
3.2 与 DINO 损失的差异
| 维度 | DINO | iBOT |
|---|---|---|
| 对象 | 全 crop 的 [CLS] | 只有被掩码 的 patch |
| teacher head 头 | 共享 dino_head | 独立的 ibot_head(separate_head=True) |
| teacher 中心 | 用 [B, K] 中心 | 用 [1, 1, K] 中心(per-dim 全局) |
| center 累加 | sum(teacher_output, dim=0) | sum(teacher_patch.mean(1), dim=0)(先按 N patch 取平均) |
| sinkhorn B | Q.shape[1] * world_size | 直接拿 n_masked_patches_tensor(来自 data["n_masked_patches"]) |
| 归一化 | softmax(s · t / temp) | 一样的 lossfunc 模块函数(第 16-17 行) |
3.3 SinkhornKnoppTeacher 单独写成 module(第 20-58 行)
注释里写了原因:
This is because we want to torch.compile it, and torch.compil-ing a single function with the @torch.compile decorator is bad. It’s better to module.compile() it.
所以在 iBOTPatchLoss.init 里:
self.sinkhorn_knopp_teacher = SinkhornKnoppTeacher()
self.sinkhorn_knopp_teacher.compile()核心流程和 DINO 的几乎一样,唯一差别是 B = n_masked_patches_tensor(含跨 rank all-reduce),因为不是每个 sample 都有相同数量的被掩码 patch,必须用一个外部张量来表达"全局被掩码 patch 总数"。
3.4 forward_masked(第 96-117 行)
def forward_masked(self, student_patch_tokens_masked, teacher_patch_tokens_masked,
student_masks_flat, n_masked_patches=None, masks_weight=None):
t = teacher_patch_tokens_masked # 已是 SK 软分布
s = student_patch_tokens_masked # 原始 logits
loss = lossfunc(t, s, self.student_temp) # = sum(t * log_softmax(s/tau))
if masks_weight is None:
masks_weight = (
(1 / student_masks_flat.sum(-1).clamp(min=1.0))
.unsqueeze(-1)
.expand_as(student_masks_flat)[student_masks_flat]
)
if n_masked_patches is not None:
loss = loss[:n_masked_patches]
loss = loss * masks_weight
return -loss.sum() / student_masks_flat.shape[0]要点:
student_masks_flat形状[n_crops * B, P],是 bool mask(哪些 patch 被掩了)。masks_weight(来自 dataloader,ssl_meta_arch.py:373)保留了在采样掩码时按 block 面积的归一化权重,让所有 patch 的"被采样概率"是均匀的(每个 patch 的平均权重为 1/P)。n_masked_patches截断到 dataloader 算好的实际数(避免一些切分带来的尾段噪声)。- 最终除以 B(不是被掩码 patch 数)做 batch 级归一:不同 step 掩码数不同也没关系。
3.5 teacher / student 输出取掩码 patch 的实现
在 ssl_meta_arch.py:449-465(teacher 侧)和 ssl_meta_arch.py:553-555(student 侧):
# teacher: ibot_head 只对 student 掩码的那些 patch 跑
buffer = torch.index_select(ibot_patch.flatten(0, 1), dim=0, index=mask_indices_list)
masked_patch_after_head = self.teacher.ibot_head(buffer) # [n_masked_patches, K]
# SK
masked_patch_centered = self.ibot_patch_loss.sinkhorn_knopp_teacher(
masked_patch_after_head, teacher_temp, n_masked_patches_tensor
) # [n_masked_patches, K]
# student
masked_patches_pre_head = torch.index_select(g_patch.flatten(0, 1), dim=0, index=mask_indices_list)
global_masked_patch_after_head = self.student.ibot_head(masked_patches_pre_head)mask_indices_list 是 dataloader 算好的"所有 global crop 摊平后,被掩码 patch 的全局 index 列表"。
3.6 teacher 中心 [1, 1, K] 与 mean(1) 求和
reduce_center_update(第 123-129 行):
self.async_batch_center = torch.sum(teacher_patch_tokens.mean(1), dim=0, keepdim=True)为什么先 mean(1) 再 sum(0)?因为不同样本的 patch 数 N 不固定,先对每张图的 N 个 patch 取平均(变成 [B, K]),再对 B 求和,得到一个 per-dim 全局中心([1, 1, K])。这与 DINO 用 [1, K] 形状不同。
4. KoLeo 损失:让 CLS 特征在 L2 球面均匀分散
dinov3/loss/koleo_loss.py— 113 行 引用:Sablayrolles et al. 2018, “Spreading vectors for similarity search”
4.1 一句话定义
KoLeo 损失 = 对每个 L2 归一化后的 CLS 特征,找其 batch 内最近邻,鼓励二者 L2 距离的对数更大。
直观上,这是一个几何均匀化正则——把 [CLS] 推到单位球面上均匀覆盖的位置,避免 collapse 到几个聚类。
4.2 KoLeoLoss.forward(第 33-43 行)
def forward(self, student_output, eps=1e-8):
with torch.autocast("cuda", enabled=False): # 强制 fp32,防 normalize 数值问题
student_output = F.normalize(student_output, eps=eps, p=2, dim=-1)
indices = self.pairwise_NNs_inner(student_output)
distances = self.pdist(student_output, student_output[indices]) # [B]
loss = -torch.log(distances + eps).mean()
return losspairwise_NNs_inner(第 21-31 行):
def pairwise_NNs_inner(self, x):
dots = torch.mm(x, x.t()) # [B, B]
n = x.shape[0]
dots.view(-1)[:: (n + 1)].fill_(-1) # 把对角线置 -1(不算自己)
_, indices = torch.max(dots, dim=1) # 最大内积 = 最小 L2 距离(因为已归一化)
return indices数学形式(对每个样本 $i$):
$$ \mathcal{L}_{\text{KoLeo}} = -\frac{1}{B}\sum_{i=1}^{B} \log\bigl(\| \mathbf{x}_i - \mathbf{x}_{n(i)}\|_2 + \epsilon\bigr) $$其中 $n(i) = \arg\max_{j \ne i} \langle \mathbf{x}_i, \mathbf{x}_j\rangle$。
4.3 KoLeoLossDistributed(第 46-113 行)
单卡版本只用 local batch 找最近邻;分布式版本要 all-gather 跨 rank 的特征一起找。
关键点:
all_gather把所有 rank 的student_output拼成[global_B, D]。- 把 global batch 分成
n_groups = global_B / loss_group_size个组(每组loss_group_size个样本)。 - 把 local rank 所在的那一组当做"最近邻候选集",算 NN。
- 要求
loss_group_size是 local batch 的整数倍,且global_B能被loss_group_size整除(第 89-96 行)。 - 用
topk > 1时把 topk 个 NN 的平均距离也算进来(重复展开student_output→[local_B * topk, D])。
实际配置默认走非分布式版(ssl_default_config.yaml:19 koleo_loss_distributed: false),所以 batch 内的 [CLS] 数量等于 2 * B(2 个 global crop)。
4.4 怎么用
koleo_loss = sum(self.koleo_loss(x) for x in student_global["cls_pre_head"]) / n_global_crops
loss_dict["koleo_loss"] = koleo_loss
loss_accumulator += self.dino_koleo_loss_weight * koleo_scale * koleo_losskoleo_scale = n_global_crops(=2)—— 简单地把koleo_loss_weight(默认 0.1)按 crop 数缩放。- 输入是
student_global["cls_pre_head"],即 student 的 [CLS] 特征 过 head 之前(backbone 出来的x_norm_clstoken),形状[n_global_crops, B, D]。 - 在
koleo_loss内部DINOHead不参与;这里用的是原始 backbone 特征(x_norm_clstoken)的 L2 归一化空间。
5. Gram 损失:DINOv3 新增的"特征间相关矩阵对齐"
dinov3/loss/gram_loss.py— 84 行 这是 DINOv3 与 DINOv2 唯一不同的损失。 论文核心:让 student 的 patch 特征两两内积矩阵(G = X·Xᵀ)去拟合另一份冻结 gram teacher的 patch 内积矩阵。
5.1 一句话定义
Gram 损失 = MSE(student patch 相关矩阵 ‖ teacher patch 相关矩阵),常在 L2 归一化 + 截断负值后做。
5.2 GramLoss.forward(第 34-84 行)
def forward(self, output_feats, target_feats, img_level=True):
output_feats = output_feats.float()
target_feats = target_feats.float()
# 1) 目标侧:可选 L2 norm,按 img_level 决定是否 flatten
if self.apply_norm:
target_feats = F.normalize(target_feats, dim=-1)
if not img_level and len(target_feats.shape) == 3:
target_feats = target_feats.flatten(0, 1) # (B*N, D)
target_sim = torch.matmul(target_feats, target_feats.transpose(-1, -2)) # (B, N, N) 或 (BN, BN)
# 2) student 侧同上
if self.apply_norm:
output_feats = F.normalize(output_feats, dim=-1)
if not img_level and len(output_feats.shape) == 3:
output_feats = output_feats.flatten(0, 1)
student_sim = torch.matmul(output_feats, output_feats.transpose(-1, -2))
# 3) 截断负值
if self.remove_neg:
target_sim[target_sim < 0] = 0.0
student_sim[student_sim < 0] = 0.0
elif self.remove_only_teacher_neg:
target_sim[target_sim < 0] = 0.0
student_sim[(student_sim < 0) & (target_sim < 0)] = 0.0 # 只在 teacher 也是负时才把 student 截断
return self.mse_loss(student_sim, target_sim)5.3 关键参数
| 参数 | 默认(DINOv3 7B) | 含义 |
|---|---|---|
apply_norm | True | 算 Gram 前对特征做 L2 归一化 |
img_level | True | 图像内 算 N×N Gram;为 False 时把 (B, N, D) 展成 (BN, D) 算一个全局 Gram |
remove_neg | False(DINOv3 7B) | 把 teacher 和 student 两侧的负值都置 0(DINOv2 默认这个,DINOv3 改成 remove_only_teacher_neg) |
remove_only_teacher_neg | False(DINOv3 7B) | 只在 teacher sim<0 的位置把 student 同步置 0 |
7B Gram Anchor 的真实配置(gram_anchor.yaml:57-60):
normalized: true
img_level: true
remove_neg: false
remove_only_teacher_neg: false而 ssl_default_config.yaml:56-59 默认是:
normalized: true
img_level: false # 默认 batch 级
remove_neg: false
remove_only_teacher_neg: falseDINOv3 7B 的两个关键改动:
img_level: true—— 把 Gram 限制在单张图内部,避免跨图 patch 比对的语义噪声。- 不截断负值 —— 让 student 学会"哪些 patch 互斥",比只学"哪些 patch 相近"信息更丰富。
5.4 Gram teacher 的来源与更新
见 ssl_meta_arch.py:476-528 的 get_gram_teacher_output:
两个来源二选一(cfg.gram.ema_teacher):
| 模式 | 含义 | 配置 |
|---|---|---|
ema_teacher: true | 直接用 EMA 的 teacher 作为 gram teacher | gram.ckpt=null(强制) |
ema_teacher: false | 维护一个独立的 gram teacher backbone | 用 gram.ckpt 指定初始权重,训练中每 N 步 fork 一次 |
第二种模式在 7B 配置里实际是:
it_load_ema_teacher: -1(gram_anchor.yaml:52)—— 本次训练不主动 fork。ckpt: ignore(gram_anchor.yaml:51)—— 走占位逻辑(由脚本里更外层给真实 ckpt 路径)。rep_update: true+update_frequency: 10000+it_first_update: 1010000(gram_anchor.yaml:53-55)—— 第 1.01M 步后每 10K 步 fork 一次,最多 3 次(max_updates: 3,gram_anchor.yaml:56)。gram_teacher_crops_size: 512(gram_anchor.yaml:165)—— gram teacher 看到的图是 512×512,而 student/teacher 是 256×256。
高分辨率 trick:DINOv3 论文里,gram teacher 用 4× 的分辨率(512 vs 256),所以 patch 数多 16×。在 ssl_meta_arch.py:494-510 里会把 gram teacher 的 patch 特征 reshape 到 (D, N, N) → interpolate → flatten 回 (B, P_student, D),把两份对齐到同一空间分辨率。
fork 实现(update_gram,ssl_meta_arch.py:728-745):
def update_gram(self, m=0):
if not self.has_gram_teacher:
return
# ... 拼出 (gramteacher_params, teacher_params) 列表 ...
with torch.no_grad():
torch._foreach_mul_(gramteacher_param_list, m)
torch._foreach_add_(gramteacher_param_list, teacher_param_list, alpha=1 - m)m=0 表示 gram teacher 直接被 teacher 当前权重覆盖(不是 EMA),这与 EMA teacher 那种 m≈0.999 不同。
5.5 gram teacher 的 crops
- 数据端在 dataloader 里专门加了一组
collated_gram_teacher_crops(ssl_meta_arch.py:380),大小由cfg.crops.gram_teacher_crops_size决定。 gram_teacher_no_distortions: true(gram_anchor.yaml:169)—— gram teacher 的图不做颜色扰动(ColorJitter、灰度、Gaussian Blur),保持"原始信号";而 student/teacher 的两张 global view 正常加扰动。
5.6 tokens_used 选项(ssl_meta_arch.py:515-520)
if self.gram_tokens_used == "masked":
student_patches = student_patches[masks]
teacher_patches = teacher_patches[masks]
elif self.gram_tokens_used == "unmasked":
student_patches = student_patches[~masks]
teacher_patches = teacher_patches[~masks]all(DINOv3 7B 默认):所有 patch 都参与。masked/unmasked:只算被掩码 / 没被掩码 的 patch 子集的 Gram。这两个模式必须配img_level: false(ssl_meta_arch.py:230-231),因为单图内的"被掩码 patch 数"通常 10-50 个,不够算可靠的 N×N。
5.7 怎么用(ssl_meta_arch.py:651-682)
if self.gram_use_loss:
gram_loss = self.gram_loss(
gram_global["student_patches"],
gram_global["teacher_patches"],
img_level=self.gram_img_level,
)
# 支持 loss_weight schedule(warmup + cosine)
if self.gram_loss_schedule is not None:
gram_loss_weight = self.gram_loss_schedule[iteration]
else:
gram_loss_weight = self.gram_loss_weight
loss_dict["gram_loss_weight"] = gram_loss_weight
loss_accumulator += gram_loss * gram_loss_weight
loss_dict["gram_loss"] = gram_lossDINOv3 7B 的 loss_weight_schedule(gram_anchor.yaml:64-69):
loss_weight_schedule:
start: 0.0 # 开头 1000 epoch warmup 内一直是 0(Gram teacher 还没准备好)
peak: 0.0
end: 2.0 # 1 epoch cosine 升到 2.0
warmup_epochs: 1000
cosine_epochs: 1即"先正常训 1M+ 步让 gram teacher 稳定下来,再快速把 Gram 损失权重拉到 2.0"——这是个 curriculum:早期别用 Gram(teacher 噪声大),后期主推 Gram。
6. 总损失 = 各路加权和
dinov3/train/ssl_meta_arch.py:584-684—compute_losses
loss_accumulator = 0.0
# (1) DINO local: student(local crop CLS) vs teacher(global CLS), 无 diagonal mask
dino_local_crops_loss = self.dino_loss(
student_logits=student_local["cls_after_head"],
teacher_probs=teacher_global["cls_centered"],
)
loss_accumulator += dino_loss_weight * dino_local_scale * local_weight * dino_local_crops_loss
# (2) DINO global: student(global crop CLS) vs teacher(global CLS), diagonal 被屏蔽
dino_global_crops_loss = self.dino_loss(
student_logits=student_global["cls_after_head"],
teacher_probs=teacher_global["cls_centered"],
ignore_diagonal=self.dino_global_ignore_diagonal,
)
loss_accumulator += dino_loss_weight * dino_global_scale * dino_global_crops_loss
# (3) KoLeo: student global crop 的 [CLS] pre-head 在 L2 球面均匀化
koleo_loss = sum(self.koleo_loss(x) for x in student_global["cls_pre_head"]) / n_global_crops
loss_accumulator += dino_koleo_loss_weight * koleo_scale * koleo_loss
# (4) iBOT: student 被掩码 patch vs teacher 同样位置 patch (均经过 SK)
ibot_patch_loss = self.ibot_patch_loss.forward_masked(
student_global["masked_patch_after_head"],
teacher_global["masked_patch_centered"],
student_masks_flat=masks,
n_masked_patches=mask_indices_list.shape[0],
masks_weight=masks_weight,
)
loss_accumulator += ibot_loss_weight * ibot_patch_loss
# (5) Gram (DINOv3): MSE(Student Gram, Gram-Teacher Gram)
if self.gram_use_loss:
gram_loss = self.gram_loss(gram_global["student_patches"], gram_global["teacher_patches"],
img_level=self.gram_img_level)
gram_loss_weight = self.gram_loss_schedule[iteration] if self.gram_loss_schedule is not None \
else self.gram_loss_weight
loss_accumulator += gram_loss * gram_loss_weight6.1 默认权重的数学校验(ssl_default_config.yaml)
| 损失 | 配置 | 乘以的 scale | 总系数 |
|---|---|---|---|
| dino_global | loss_weight: 1.0 | 2/18 | 0.111 |
| dino_local | loss_weight: 1.0 | 16/18 | 0.889 |
| koleo | koleo_loss_weight: 0.1 | n_global_crops=2 | 0.2 |
| ibot | loss_weight: 1.0 | 1.0 | 1.0 |
| gram | loss_weight: 1.0 | 1.0(除非用 schedule) | 1.0 |
直观:DINO 局部 8 张 vs 全局 2 张的"配对总数"平衡;KoLeo 0.1 是 gentle 正则;iBOT / Gram 是 dense patch 对齐主信号。
6.2 真正的反传:FSDP 友好的 backprop_loss
def backprop_loss(self, loss):
loss.backward() # 一切梯度靠 PyTorch autograd;FSDP/compile 接管ssl_meta_arch.py:710-712 — 简单透传。
7. 端到端:4 个损失的"数据流"对照
把数据流画成图(每条箭头 = 一次前向 → 一次后向):
[2 global crops + 8 local crops + masks]
│
├─→ student.backbone (global+local 一次过)
│ │
│ ├─ [CLS] ── student.dino_head ─→ dino_logits (global+local)
│ │ │
│ │ ├─→ dino_local_loss vs teacher(global) [CLS]_SK
│ │ └─→ dino_global_loss vs teacher(global) [CLS]_SK (skip A-A,B-B)
│ │
│ └─ patches ─→ index_select(masked) → student.ibot_head ─→ ibot_logits (masked)
│ │
│ └─→ ibot_loss vs teacher(masked patch)_SK
│
├─→ student [CLS] pre-head (global only)
│ └─→ koleo_loss (球面均匀化正则)
│
├─→ teacher.backbone (global only, no grad, EMA)
│ ├─ [CLS] ── teacher.dino_head ─→ [CLS]_SK (给 DINO 用)
│ └─ patches ─→ teacher.ibot_head(masked indices) ─→ masked_patch_SK (给 iBOT 用)
│
└─→ (DINOv3 only) gram_teacher.backbone (高分辨 crops, no grad)
└─ patches (resized to student 分辨率) ─→ 与 student patches 算 Gram → gram_loss8. 数学速查表
把上面所有损失写成形式化公式:
设 $V_g$ = global crop 数(=2),$V_l$ = local crop 数(=8),$B$ = batch size。
8.1 DINO loss
- Teacher 侧:$\mathbf{p}_t = \mathrm{SinkhornKnopp}(\mathbf{z}_t / \tau_t) \in \mathbb{R}^{K}$,列和=1。
- Student 侧:$\log \mathbf{p}_s = \log\mathrm{softmax}(\mathbf{z}_s / \tau_s)$。
- 损失(包含 diagonal mask 选项):
- 总和项:$\text{global\_terms} = V_g(V_g - 1) = 2$,$\text{local\_terms} = V_g \cdot V_l = 16$。
- 损失项:$\mathcal{L}_{\mathrm{DINO}} = \frac{2}{18} \mathcal{L}_{\mathrm{DINO, g}} + \frac{16}{18} \mathcal{L}_{\mathrm{DINO, l}}$。
8.2 iBOT loss
- 掩码集合 $\mathcal{M} = \{(c, p) : \text{patch } p \text{ in crop } c \text{ 被掩码}\}$,大小 $M = |\mathcal{M}|$。
- Teacher:$\mathbf{p}_t^{(c,p)} = \mathrm{SinkhornKnopp}(\mathbf{z}_t^{(c,p)} / \tau_t)$,列和=1;归一化常数 $B' = M$。
- Student:$\log\mathbf{p}_s^{(c,p)} = \log\mathrm{softmax}(\mathbf{z}_s^{(c,p)} / \tau_s)$。
- 损失(按 dataloader 给的
masks_weight加权):
8.3 KoLeo loss
- L2 归一化:$\tilde{\mathbf{x}}_i = \mathbf{x}_i / \|\mathbf{x}_i\|_2$。
- 最近邻:$n(i) = \arg\max_{j \ne i} \langle \tilde{\mathbf{x}}_i, \tilde{\mathbf{x}}_j\rangle$。
8.4 Gram loss(DINOv3 新增)
- patch 特征 $X_s, X_t \in \mathbb{R}^{N \times D}$,先做 L2 归一化。
- Gram 矩阵:$G_s = X_s X_s^\top$,$G_t = X_t X_t^\top$。
- 截断(按配置):$\hat{G} = \max(G, 0)$(或仅在 teacher 负值处截断 student)。
img_level=true时 $G$ 的形状是 $[B, N, N]$(每图独立算);img_level=false时把 $(B, N, D)$ 展成 $(BN, D)$,算一个 $[BN, BN]$ 的全局 Gram。
9. 几个易混点 / 实战 Tips
- teacher head 是否参与反传? 不参与:
self.teacher.requires_grad_(False)(ssl_meta_arch.py:137)。teacher.dino_head的权重只是被 EMA 复制到 student。 - Sinkhorn 的
B在 DINO 和 iBOT 不同:DINO 用Q.shape[1] * world_size(每个 rank 的样本数 × rank 数 = 总样本数);iBOT 用外部传进来的n_masked_patches_tensor(data["n_masked_patches"])。iBOT 用外数是因为 patch 数不固定。 - center 缓冲初始值:注册时是
nan,必须显式init_weights()清零。FSDP checkpoint 加载时跳过dino_loss.center/ibot_patch_loss.center(ssl_meta_arch.py:317, 333)。 - EMA 用
_foreach_mul_/_foreach_add_:原地操作省内存,PyTorch 内部会用 fused kernel。 - Gram teacher 的 head 跳过加载(ssl_meta_arch.py:689-696):
skip_load_prefixes = ["dino_head.", "ibot_head."],因为 gram teacher 只用 backbone 拿 patch 特征。 mask_k_bias/qkv.bias_mask不被 sharded(ssl_meta_arch.py:320):FSDP 在 7B 用的qkv.bias_mask与rope_embed.periods形状都是 [head_dim//2] 之类小张量,shard 不划算。- 为什么 KoLeo 用
torch.autocast("cuda", enabled=False)? L2 归一化 +log(dist)在 fp16/bf16 下数值不稳(log(eps)会爆),强制 fp32 算 NN 距离更稳。 - Gram 损失在 7B 里为什么是
img_level: true? 论文里的经验:跨图算 Gram 会把"语义上无关的 patch"混在一起算相关,引入噪声;单图内算 Gram 只关心"图内部结构",更适合做 dense 自监督。 force_masking_even_with_zero_weight:当 IBOT 损失权重为 0 时仍生成 mask(保留数据流),方便复用 dataloader。
10. 论文 ↔ 代码一页纸对照
| 论文概念 | 代码位置 |
|---|---|
| DINO 头 (3-layer MLP + L2 norm + unbind last layer) | dinov3/layers/dino_head.py |
| CLS 多视图对比 | dino_clstoken_loss.py:forward |
| Sinkhorn-Knopp centering | dino_clstoken_loss.py:42-70 |
| iBOT masking + 局部 patch 对齐 | ibot_patch_loss.py + vision_transformer.py mask_token + data augmentor |
| KoLeo 正则 | koleo_loss.py |
| Gram anchoring(DINOv3 新增) | gram_loss.py + ssl_meta_arch.py:476-528 + dinov3_vit7b16_gram_anchor.yaml |
| Multi-crop 配对 + scale balance | ssl_meta_arch.py:602-633 |
| 教师 EMA + 冻结 last layer | ssl_meta_arch.py:713-726 + freeze_last_layer_epochs=1 |
| 高分辨率 gram teacher 路径 | ssl_meta_arch.py:494-510 |
| Gram teacher fork | ssl_meta_arch.py:686-745 |
| 分布式通信组(DINO 用 process_subgroup,gram teacher 用 default group) | ssl_meta_arch.py:788-805 |