双网络训练机制
DINOv3 双网络训练机制(Teacher-Student)笔记
配套笔记:
sinkhorn-knopp理解.md— teacher 的 target 怎么用 Sinkhorn-Knopp 投影到双随机双随机矩阵.md— 双随机矩阵的几何意义ssl_loss理解.md— student 端的 cross-entropy 损失细节涉及代码:
dinov3/train/ssl_meta_arch.py:48-138— 构造 student/teacher 两份独立模型 + 冻结 teacherdinov3/train/ssl_meta_arch.py:355-475— 前向、反向、forward_backwarddinov3/train/ssl_meta_arch.py:530-582— student 前向(global + local)dinov3/train/ssl_meta_arch.py:713-726—update_ema(teacher 移动平均更新)dinov3/train/train.py:122-124, 478, 532— momentum schedule 与每步调用
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 用的 headibot_head:patch token 用的 head
关键约束(ssl_meta_arch.py:137-138):
self.teacher.requires_grad_(False) # teacher 不需要梯度!
self.model_ema.requires_grad_(False) # model_ema 就是 teacherteacher 的参数永远冻结,不参与反向传播。
4. 关键设计:Multi-Crop(多裁剪)策略
DINOv3 用 multi-crop 增强策略(ssl_meta_arch.py:362-363):
| Crop | 数量(DINOv3 默认) | 尺寸 | 喂给谁 |
|---|---|---|---|
| Global crops | 2 | 224×224 | teacher 和 student 都看 |
| Local crops | 8 | 96×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 对比表
| 维度 | Teacher | Student |
|---|---|---|
| 输入 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. 几个易混点
teacher 和 student 的架构必须完全一样:
- 同
cfg.student.arch、同embed_dim、同head_n_prototypes - 加载蒸馏时尤其要 assert(ssl_meta_arch.py:271-275)
- 同
model_ema不一定等于teacher:- 默认情况
self.model_ema = self.teacher(ssl_meta_arch.py:131) - 蒸馏时(
cfg.distillation.enabled)model_ema可能被覆盖为外部加载的 teacher
- 默认情况
EMA 更新在
optimizer.step()之后:- 顺序:student 反向传播 → step → teacher EMA 更新
- 不能颠倒(否则 student 用了"未来"的 teacher target)
teacher 的 warm-up 阶段要小心:
- 训练初期 student 随机,teacher 几乎 = student
- 头几步的 target 没意义,要给 loss 一些 warmup 步数
- 代码中通过
cfg.train.warmup_iterations控制
momentum schedule 的"为什么":
- 训练初期 teacher 可以更新快一点(m 小)→ 跟上 student 早期学习
- 训练后期 teacher 极慢(m 大)→ 给 student 极稳定目标
- cosine schedule 是个相对"普适"的选择
多机分布式时的细节:
- SK 的
all_reduce让 target 是"全局 batch"上的双随机(参见 sinkhorn-knopp理解.md §5.1) - EMA 在每个 rank 内独立做(teacher 同步靠 DDP 的
state_dict广播)
- SK 的
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 作为下游预训练模型。