双网络训练机制

DINOv3 双网络训练机制(Teacher-Student)笔记

配套笔记:

涉及代码:

1. 一句话定义

DINOv3 训练 = student(梯度下降)+ teacher(EMA of student)两条独立模型。 Student 看 global + local crops,通过 cross-entropy 拟合 teacher 给的"双随机 prototype 分布";teacher 自身不训练,每步用 student 参数的**指数移动平均(EMA)**更新自己。

2. 直观比喻

把 DINO 的 teacher-student 想象成"老师与学生“的考试机制:

角色行为训练方式
学生参加考试,答错挨打梯度下降(真学习)
老师不直接学习,给学生出"标准答案”看学生"慢动作回放"(EMA)
关键约束老师比学生学得慢teacher 的参数 = student 的滑动平均

为什么需要老师比学生慢? 如果老师立刻跟上学生,学生答什么老师就"认什么"——没法学。如果老师慢半拍,老师能给出一个稳定的目标(target),学生才有明确方向。


3. 架构:两份独立副本

DINOv3 在初始化时构造两份完全相同的网络(ssl_meta_arch.py:48-81):

student_model_dict = dict()
teacher_model_dict = dict()

# 同一个 cfg → 同结构,但权重独立
student_backbone, teacher_backbone, embed_dim = build_model_from_cfg(cfg)
student_model_dict["backbone"] = student_backbone
teacher_model_dict["backbone"] = teacher_backbone

# 两份独立的 head(各自有权重,不共享)
student_model_dict["dino_head"] = dino_head_class()
teacher_model_dict["dino_head"] = dino_head_class()
student_model_dict["ibot_head"] = ibot_head_class()
teacher_model_dict["ibot_head"] = ibot_head_class()

self.student = nn.ModuleDict(student_model_dict)
self.teacher = nn.ModuleDict(teacher_model_dict)

每份都包含:

  • backbone:ViT(patch_size=16,embed_dim=768/1024)
  • dino_head:CLS token 用的 head
  • ibot_head:patch token 用的 head

关键约束(ssl_meta_arch.py:137-138):

self.teacher.requires_grad_(False)   # teacher 不需要梯度!
self.model_ema.requires_grad_(False) # model_ema 就是 teacher

teacher 的参数永远冻结,不参与反向传播。


4. 关键设计:Multi-Crop(多裁剪)策略

DINOv3 用 multi-crop 增强策略(ssl_meta_arch.py:362-363):

Crop数量(DINOv3 默认)尺寸喂给谁
Global crops2224×224teacher 和 student 都看
Local crops896×96只给 student 看
            ┌───────────────────────────────────────────┐
            │             Multi-crop 增强              │
            │                                           │
            │   Global #1   Global #2                   │
            │   (224×224)   (224×224)                   │
            │   ↓             ↓                         │
            │   teacher   +  student (global 分支)      │
            │                                           │
            │   Local #1 ... #8                         │
            │   (96×96 × 8)                             │
            │   ↓                                       │
            │   student (local 分支)                    │
            └───────────────────────────────────────────┘

4.1 为什么要这么分?

  • teacher 看完整 2 个全局视角 → 输出稳定的"标准答案"(信息全、噪声少)
  • student 既看全局又看局部 → 学到"局部细节 + 全局语义"对齐
  • 形成不对称:teacher 的输入比 student 简单,但 student 要追 teacher

这是 DINO 论文的核心 trick 之一(Caron et al. 2021):让 student 看到"更难"的视角(小图、缺信息),teacher 看到"更容易"的视角(大图、全信息),迫使 student 学到"从小局部推断全局语义"的能力。

4.2 在 iBOT 里还有 Mask 增强

iBOT 在 multi-crop 之上额外对 patch 做 mask(随机遮掉部分 patch):

  • student 看:global 视图 + 部分 patch 被 mask 掉
  • teacher 看:global 视图 + 完整 patch

iBOT 损失的 target 只在"学生看到的被 mask 掉的那部分 patch“上算(ssl_meta_arch.py:450)——这迫使 student 学"从可见 patch 推断被 mask patch 的 prototype”。


5. 前向传播:teacher vs student 各自做什么

5.1 Teacher 前向(ssl_meta_arch.py:440-465)

with torch.no_grad():                                # teacher 不存梯度
    backbone_out = self.teacher.backbone(images, is_training=True)
    cls      = backbone_out["x_norm_clstoken"]        # [2*B, D]    (global crops)
    ibot_patch = backbone_out["x_norm_patchtokens"]  # [2*B, P, D]

    # iBOT head 只对"学生被 mask 的 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)

    # DINO head 对所有 CLS token 算
    cls_after_head = self.teacher.dino_head(cls)     # [2*B, K]

    # Sinkhorn-Knopp 投影到双随机 → 做 target
    cls_centered = self.dino_loss.sinkhorn_knopp_teacher(cls_after_head, teacher_temp=...)
    masked_patch_centered = self.ibot_patch_loss.sinkhorn_knopp_teacher(...)

teacher 的任务只是产生 target,不参与反传。

5.2 Student 前向(ssl_meta_arch.py:530-582)

global_out, local_out = self.student.backbone(
    [global_crops, local_crops.flatten(0, 1)],     # student 看 global + local
    masks=[masks, None],                            # student global 有 mask
    is_training=True,
)
g_cls, l_cls = global_out["x_norm_clstoken"], local_out["x_norm_clstoken"]

# DINO head:把 global + local CLS 串起来一起算
buffer = torch.cat([g_cls, l_cls], dim=0)         # [2*B + 8*B, D] = [10*B, D]
buffer = self.student.dino_head(buffer)           # [10*B, K]

student 的输出比 teacher 多(global + local),并被 cross-entropy 推着去匹配 teacher 的 target。

5.3 对比表

维度TeacherStudent
输入 crops仅 global(2 个 224×224)global(2 个 224×224)+ local(8 个 96×96)
是否看到 masked patch完整 patch部分 patch 被 mask
梯度冻结(requires_grad_(False))反向传播
输出用途经 Sinkhorn-Knopp 做 target经 cross-entropy 拟合 target
head 数dino_head + ibot_head同
forward 时的 torch.no_grad必须开不开(要梯度)

6. 三种"训练"机制

DINOv3 训练时实际有三种参数更新方式:

6.1 Student:标准梯度下降

跟普通训练一样(train.py:484-485):

optimizer.zero_grad(set_to_none=True)
total_loss, metrics_dict = model.forward_backward(data, ...)  # 反向传播
# ...optimizer.step() 在 train.py 后续代码里

student 是真正"学知识"的那个,由 AdamW / LAMB 等标准优化器更新。

6.2 Teacher:EMA(指数移动平均)——“看学生慢动作回放”

teacher 参数更新公式(ssl_meta_arch.py:713-726):

def update_ema(self, m):
    with torch.no_grad():
        # teacher ← m × teacher + (1 - m) × student
        torch._foreach_mul_(teacher_param_list, m)              # teacher *= m
        torch._foreach_add_(teacher_param_list, student_param_list, alpha=1 - m)  # teacher += (1-m) × student

展开为标量形式:

$$ \theta_{\text{teacher}}^{(t+1)} = m \cdot \theta_{\text{teacher}}^{(t)} + (1 - m) \cdot \theta_{\text{student}}^{(t)} $$

这就是 BYUA / MoCo / DINO 那一套的 “EMA update”——teacher 是 student 参数的时间加权平均。

等价的"半衰期"视角:

  • 假设 $m$ 是常数,$(1-m) \ll 1$(如 $1-m = 0.001$)
  • 当前 student 的参数"贡献"到 teacher 的"权重" = $(1-m)$
  • 当前 teacher 的"惯性" = $m$
  • $t$ 步前 student 的贡献权重 $\approx (1-m) \cdot m^{t-1}$
  • teacher 的"有效记忆窗口" = $\frac{-1}{\log m} \approx \frac{1}{1-m}$ 步

6.3 momentum 是动态 schedule 的

m 不是一个常数,而是按 cosine schedule 从小到大变(train.py:122-124):

momentum = dict(
    base_value=cfg.teacher["momentum_teacher"],      # 初始 m,如 0.994
    final_value=cfg.teacher["final_momentum_teacher"], # 终止 m,如 0.999
)
momentum_schedule = CosineScheduler(**momentum)

schedule 含义:

训练阶段m 大约物理效果
初期0.994(1-m = 0.006)teacher 更新较快,能跟住 student 学到东西
后期0.999(1-m = 0.001)teacher 更新极慢,target 极稳定,student 收敛

每一步调用(train.py:478, 532):

mom = momentum_schedule[it]      # 拿到当前 step 的 m
...
model.update_ema(mom)            # 更新 teacher 参数

6.4 Gram Teacher:DINOv3 引入的第三个模型

DINOv3 还引入了一个第三模型(ssl_meta_arch.py:180-181):

self.gram_teacher = nn.ModuleDict(gram_model_dict)
self.gram_teacher.requires_grad_(False)
  • 用于 Gram loss(一种特征空间相似度对齐的辅助任务)
  • 更新策略:update_gram(m=0) 表示直接复制 teacher 的参数(ssl_meta_arch.py:728-745)
  • 或者从预训练 checkpoint 加载(ssl_meta_arch.py:687-693)
  • 这块超出 teacher/student 主线,本文不展开

7. 完整时序:每一步训练循环

for each iteration it:
    │
    ├── 1. 取数据 → multi-crop 增强(2 global + 8 local)
    │
    ├── 2. student 前向
    │      backbone(global + local, with mask) → CLS + patch
    │      dino_head(CLS) → logits [10B, K]
    │      ibot_head(masked patch) → logits [n_masked, K]
    │
    ├── 3. teacher 前向(无梯度)
    │      backbone(global, no mask) → CLS + patch
    │      dino_head(CLS) → logits [2B, K]
    │      ibot_head(masked patch index) → logits [n_masked, K]
    │
    ├── 4. teacher 输出 → Sinkhorn-Knopp → 双随机 target
    │      cls_centered:           [2B, K]     (列和=1, 行和=B/K)
    │      masked_patch_centered:  [n_masked, K]
    │
    ├── 5. 算损失
    │      loss_dino_gs = DINO(student_global_CLS, teacher_CLS)     # student 学 teacher 全局
    │      loss_dino_ls = DINO(student_local_CLS,  teacher_CLS)     # student 局部→teacher 全局
    │      loss_ibot   = iBOT(student_masked_patch, teacher_masked_patch)  # 局部对齐
    │      loss_gram   = Gram(student_global_patch, gram_teacher_patch)    # 辅助
    │      loss_koleo  = KoLeo(student_CLS)                              # 均匀化
    │      total_loss = 加权和
    │
    ├── 6. 反向传播 + optimizer.step()    ← 更新 student
    │
    └── 7. EMA: teacher ← m*teacher + (1-m)*student   ← 更新 teacher(无梯度)

8. 为什么这样设计能"防 collapse"?

这是 DINO 的灵魂问题。如果让 student 自己跟自己学(teacher = student 副本),网络会迅速塌缩到"输出常数"(所有图像都给同一 prototype)。两个关键设计挡住了 collapse:

8.1 EMA 让 teacher 滞后且稳定

如果 teacher = student 自身,student 一更新,target 跟着变 → 学不到东西。 teacher 滞后 → target 稳定 → student 有明确学习目标。

8.2 Sinkhorn-Knopp 的"列和均匀"约束

teacher 输出的 logits 经过 SK 投影后强制让所有 prototype 被均匀使用(参见 双随机矩阵笔记 §5)。即使网络想 collapse 到某个 prototype,SK 也会"反向推回"均匀分布。

两者缺一不可:

  • 没有 EMA → student 跟 teacher 跑偏
  • 没有 SK → teacher 自己塌缩到常数
  • 两者结合 → 稳定训练

8.3 multi-crop 的不对称

teacher 只看 global(信息全),student 还要看 local(信息少):

  • 强迫 student 学"小图也能匹配大图的语义"
  • 阻止 student 通过"抄答案"作弊
  • 等价于一种信息瓶颈,迫使学到核心特征

9. 数据流总图

                  输入图像 (2 global + 8 local crops)
                              │
            ┌─────────────────┴─────────────────┐
            ↓                                   ↓
         Teacher                             Student
   (只看 2 global,                 (2 global + 8 local,
    无梯度)                          部分 patch 被 mask)
            │                                   │
        backbone                              backbone
            │                                   │
   ┌────────┴────────┐                  ┌───────┴───────┐
   ↓                 ↓                  ↓               ↓
x_norm_clstoken  x_norm_patch    x_norm_clstoken  x_norm_patch
   [2B, D]       [2B, P, D]       [10B, D]      [10B, P, D]
   │                 │                 │               │
dino_head        ibot_head         dino_head      ibot_head
   │             (取 mask             │          (取 mask
   │              索引子集)            │           索引子集)
   ↓                 ↓                 ↓               ↓
[2B, K]       [n_masked, K]       [10B, K]    [n_masked, K]
   │                 │                 │               │
Sinkhorn-Knopp   Sinkhorn-Knopp       │               │
   │                 │                 │               │
   ↓                 ↓                 ↓               ↓
[2B, K]       [n_masked, K]     cross-entropy    cross-entropy
(双随机 target)  (双随机 target)        ↓               ↓
                                  student 损失    student 损失
                                          ↓
                                     反向传播
                                          ↓
                                    更新 student
                                          ↓
                                    teacher ← EMA(student)

10. 训练完后

训练结束后,teacher 和 student 都可用,但通常用 teacher 作为推理用的预训练模型(更稳定):

  • 加载它的 backbone + dino_head
  • 丢弃 prototype 头,用 backbone 提取的特征做下游任务(分类、检测、分割)
  • DINOv3 的官方预训练权重默认就是 teacher 的参数

经验上:teacher 收敛轨迹更平滑(因为是 EMA 累积),在下游任务上常比 student 略好。


11. 几个易混点

  1. teacher 和 student 的架构必须完全一样:

    • 同 cfg.student.arch、同 embed_dim、同 head_n_prototypes
    • 加载蒸馏时尤其要 assert(ssl_meta_arch.py:271-275)
  2. model_ema 不一定等于 teacher:

    • 默认情况 self.model_ema = self.teacher(ssl_meta_arch.py:131)
    • 蒸馏时(cfg.distillation.enabled)model_ema 可能被覆盖为外部加载的 teacher
  3. EMA 更新在 optimizer.step() 之后:

    • 顺序:student 反向传播 → step → teacher EMA 更新
    • 不能颠倒(否则 student 用了"未来"的 teacher target)
  4. teacher 的 warm-up 阶段要小心:

    • 训练初期 student 随机,teacher 几乎 = student
    • 头几步的 target 没意义,要给 loss 一些 warmup 步数
    • 代码中通过 cfg.train.warmup_iterations 控制
  5. momentum schedule 的"为什么":

    • 训练初期 teacher 可以更新快一点(m 小)→ 跟上 student 早期学习
    • 训练后期 teacher 极慢(m 大)→ 给 student 极稳定目标
    • cosine schedule 是个相对"普适"的选择
  6. 多机分布式时的细节:

    • SK 的 all_reduce 让 target 是"全局 batch"上的双随机(参见 sinkhorn-knopp理解.md §5.1)
    • EMA 在每个 rank 内独立做(teacher 同步靠 DDP 的 state_dict 广播)

12. 一句话总结

DINOv3 的双网络训练 = student(梯度下降)+ teacher(EMA of student),teacher 用 multi-crop 的"global 视角"产生稳定的 Sinkhorn-Knopp 双随机 target,student 用 “global + local” 视角的输出拟合这个 target;EMA 滞后 + 双随机约束 + multi-crop 不对称三者协同防止 collapse、让 K=98304 的 prototype 字典被稳定学成。 训练完成后通常用 teacher 作为下游预训练模型。