ssl_loss理解

DINOv3 自监督损失笔记

对应代码:

0. 整体架构:四把锁怎么拧

DINOv3 的自监督学习沿用 DINOv2 的多任务配方。一个 SSL step 跑 4 个损失,总和反传:

损失作用对象teacher 目标核心思路DINOv2DINOv3
DINO lossstudent 全部 crop 的 [CLS]teacher global crop 的 [CLS](Sinkhorn 中心化)把所有 view 的 [CLS] 拉到统一原型分布✓✓
iBOT lossstudent 被掩码 的 patchteacher 同位置 patch(Sinkhorn 中心化)掩码 patch 重建式对齐✓✓
KoLeo lossstudent global crop 的 [CLS] pre-head无(自正则)在 L2 球面均匀分散✓✓
Gram lossstudent patch 特征另一份更高分辨率的 gram teacherpatch 间相关矩阵对齐✗新增

注:教师自身在每步后做 EMA(teacher = m·teacher + (1-m)·student),并对教师预测做 centering / sharpening(DINO 部分用 Sinkhorn,iBOT 也用 Sinkhorn,但 centering buffer 不同)。


1. 投影头 DINOHead:把 backbone 特征映射到"原型"空间

dinov3/layers/dino_head.py:11-67

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 里各自的位置

ssl_meta_arch.py:80-126:

  • 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 互相对齐?

见 ssl_meta_arch.py:611-633:

# 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 损失的差异

维度DINOiBOT
对象全 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 BQ.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 loss

pairwise_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 的特征一起找。

关键点:

  1. all_gather 把所有 rank 的 student_output 拼成 [global_B, D]。
  2. 把 global batch 分成 n_groups = global_B / loss_group_size 个组(每组 loss_group_size 个样本)。
  3. 把 local rank 所在的那一组当做"最近邻候选集",算 NN。
  4. 要求 loss_group_size 是 local batch 的整数倍,且 global_B 能被 loss_group_size 整除(第 89-96 行)。
  5. 用 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 怎么用

见 ssl_meta_arch.py:636-638:

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_loss
  • koleo_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_normTrue算 Gram 前对特征做 L2 归一化
img_levelTrue图像内 算 N×N Gram;为 False 时把 (B, N, D) 展成 (BN, D) 算一个全局 Gram
remove_negFalse(DINOv3 7B)把 teacher 和 student 两侧的负值都置 0(DINOv2 默认这个,DINOv3 改成 remove_only_teacher_neg)
remove_only_teacher_negFalse(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: false

DINOv3 7B 的两个关键改动:

  1. img_level: true —— 把 Gram 限制在单张图内部,避免跨图 patch 比对的语义噪声。
  2. 不截断负值 —— 让 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 teachergram.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_loss

DINOv3 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_weight

6.1 默认权重的数学校验(ssl_default_config.yaml)

损失配置乘以的 scale总系数
dino_globalloss_weight: 1.02/180.111
dino_localloss_weight: 1.016/180.889
koleokoleo_loss_weight: 0.1n_global_crops=20.2
ibotloss_weight: 1.01.01.0
gramloss_weight: 1.01.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_loss

8. 数学速查表

把上面所有损失写成形式化公式:

设 $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 选项):
$$ \mathcal{L}_{\mathrm{DINO}} = - \frac{1}{B\,S\,T} \sum_{b,s,t} [\![s\ne t]\!] \cdot \langle \log\mathbf{p}_s^{(b,s)},\; \mathbf{p}_t^{(b,t)}\rangle $$
  • 总和项:$\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 加权):
$$ \mathcal{L}_{\mathrm{iBOT}} = - \frac{1}{B} \sum_{(c,p) \in \mathcal{M}} w_{c,p} \cdot \langle \mathbf{p}_t^{(c,p)},\; \log\mathbf{p}_s^{(c,p)}\rangle $$

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$。
$$ \mathcal{L}_{\mathrm{KoLeo}} = -\frac{1}{B}\sum_{i=1}^{B} \log\bigl(\|\tilde{\mathbf{x}}_i - \tilde{\mathbf{x}}_{n(i)}\|_2 + \epsilon\bigr) $$

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)。
$$ \mathcal{L}_{\mathrm{Gram}} = \mathrm{MSE}(\hat{G}_s,\; \hat{G}_t) $$
  • img_level=true 时 $G$ 的形状是 $[B, N, N]$(每图独立算);img_level=false 时把 $(B, N, D)$ 展成 $(BN, D)$,算一个 $[BN, BN]$ 的全局 Gram。

9. 几个易混点 / 实战 Tips

  1. teacher head 是否参与反传? 不参与:self.teacher.requires_grad_(False)(ssl_meta_arch.py:137)。teacher.dino_head 的权重只是被 EMA 复制到 student。
  2. Sinkhorn 的 B 在 DINO 和 iBOT 不同:DINO 用 Q.shape[1] * world_size(每个 rank 的样本数 × rank 数 = 总样本数);iBOT 用外部传进来的 n_masked_patches_tensor(data["n_masked_patches"])。iBOT 用外数是因为 patch 数不固定。
  3. center 缓冲初始值:注册时是 nan,必须显式 init_weights() 清零。FSDP checkpoint 加载时跳过 dino_loss.center / ibot_patch_loss.center(ssl_meta_arch.py:317, 333)。
  4. EMA 用 _foreach_mul_/_foreach_add_:原地操作省内存,PyTorch 内部会用 fused kernel。
  5. Gram teacher 的 head 跳过加载(ssl_meta_arch.py:689-696):skip_load_prefixes = ["dino_head.", "ibot_head."],因为 gram teacher 只用 backbone 拿 patch 特征。
  6. mask_k_bias / qkv.bias_mask 不被 sharded(ssl_meta_arch.py:320):FSDP 在 7B 用的 qkv.bias_mask 与 rope_embed.periods 形状都是 [head_dim//2] 之类小张量,shard 不划算。
  7. 为什么 KoLeo 用 torch.autocast("cuda", enabled=False)? L2 归一化 + log(dist) 在 fp16/bf16 下数值不稳(log(eps) 会爆),强制 fp32 算 NN 距离更稳。
  8. Gram 损失在 7B 里为什么是 img_level: true? 论文里的经验:跨图算 Gram 会把"语义上无关的 patch"混在一起算相关,引入噪声;单图内算 Gram 只关心"图内部结构",更适合做 dense 自监督。
  9. 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 centeringdino_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 balancessl_meta_arch.py:602-633
教师 EMA + 冻结 last layerssl_meta_arch.py:713-726 + freeze_last_layer_epochs=1
高分辨率 gram teacher 路径ssl_meta_arch.py:494-510
Gram teacher forkssl_meta_arch.py:686-745
分布式通信组(DINO 用 process_subgroup,gram teacher 用 default group)ssl_meta_arch.py:788-805