sinkhorn-knopp理解
Sinkhorn-Knopp 笔记
对应代码:
dinov3/loss/dino_clstoken_loss.py:42-70—DINOLoss.sinkhorn_knopp_teacher(DINO 损失用)dinov3/loss/ibot_patch_loss.py:20-58—SinkhornKnoppTeacher(iBOT 损失用,独立 module 以便torch.compile)- 调用点:
ssl_meta_arch.py:617-633(DINO)、ssl_meta_arch.py:449-465(iBOT)
1. 一句话定义
Sinkhorn-Knopp 算法 = 对一个正矩阵做"交替行归一 / 列归一"的迭代,最终让它收敛到"双随机矩阵"(每行和 = 每列和 = 1)。 在 DINO/iBOT 中,它把 teacher 的 logits 矩阵变成"列和=1、行和均匀“的"软分配概率”,作为 student 的对齐目标。
2. 为什么要用 Sinkhorn-Knopp?
2.1 teacher 输出是什么?
teacher 把每个 sample 的特征投到 K 维(K = 65536 / 98304)“原型"空间,输出一个 logits 矩阵
$$ Z \in \mathbb{R}^{B \times K} $$其中 $B$ 是 batch size,$K$ 是 prototype 数量。
如果我们直接做 $\mathrm{softmax}(Z/\tau)$:
- 每行(每个 sample)确实和为 1
- 但每列(每个 prototype)没有任何保证——有的 prototype 可能从来不被激活(prototype collapse),有的被所有 sample 同时激活(dimension collapse)
2.2 我们想要什么 teacher 目标?
DINO 论文里希望 teacher 给的"软标签"满足:
$$ \boxed{\;P \in \mathbb{R}^{B \times K},\quad P \ge 0,\quad P\mathbf{1} = \mathbf{1},\quad P^\top \mathbf{1} = \tfrac{B}{K}\mathbf{1}\;} $$意思是:
- 每行(sample)= 1:是个概率分布
- 每列(prototype)的均值 = $B/K$:所有 prototype 均匀被使用(避免 collapse)
这正是”双随机矩阵的"按比例缩放""——把 $B \times K$ 矩阵的列和规整为均匀分布。
2.3 Sinkhorn-Knopp 就是干这个事
对任意正矩阵 $Q \in \mathbb{R}^{B \times K}_+$,如果它支持集(哪些位置 $Q_{ij} > 0$)足够"连通",就唯一存在对角矩阵 $U \in \mathbb{R}^{B \times B}$、$V \in \mathbb{R}^{K \times K}$(正对角元)使 $UQV$ 是双随机的。Sinkhorn-Knopp 就是求 $U, V$ 的迭代算法。
3. 数学基础
3.1 双随机矩阵(doubly stochastic matrix)
$$ D \in \mathbb{R}^{n \times m}_+,\quad D\mathbf{1} = \mathbf{1},\quad D^\top \mathbf{1} = \mathbf{1} $$行和=1 且 列和=1。在 DINO 里我们并不要求严格"列和=1",而是要求"列和均匀",最后再 $\times B$ 让列和=B——本质相同。
3.2 Sinkhorn 定理(矩阵缩放定理)
定理(Sinkhorn 1964):设 $A \in \mathbb{R}^{n \times m}_+$。当且仅当 $A$ 的"支撑图"(即 $\{ (i,j) : A_{ij} > 0 \}$)包含一个完美匹配(即 $A$ 的"零模式"有完美匹配),存在对角正定矩阵 $D_1, D_2$ 使 $D_1 A D_2$ 是双随机的。
而且这个 $D_1, D_2$ 唯一(差一个标量倍数)。
应用到 DINO:
- $A = \exp(Z/\tau)$ 是严格正的(无零元素)
- 所以 $D_1, D_2$ 存在且唯一
- 算法就是求它们
3.3 缩放后的形式
等价地,迭代地求行缩放因子 $\mathbf{u} \in \mathbb{R}^B_+$、列缩放因子 $\mathbf{v} \in \mathbb{R}^K_+$,使
$$ \tilde{Q}_{ij} = u_i \cdot Q_{ij} \cdot v_j $$是双随机的。
4. 完整算法 + 三个等价视角
4.1 算法伪代码
输入: Q = exp(Z / τ) # 形状 [B, K],每行/每列均 > 0
参数: n_iterations
Q ← Q / sum(Q) # 整体缩放到总和=1(数值稳定用)
repeat n_iterations times:
# 行归一
Q ← Q / (Q · 1) [B, 1] # 每行 sum → 1
# 列归一
Q ← Q / (1^T · Q) [1, K] # 每列 sum → 1
return Q注意 DINOv3 在每步多除一个 K 和 B(见 §5.3 的细节)。
4.2 视角 A:迭代缩放(Sinkhorn 1964 原始)
直接由 §3.3 推导而来。对 $Q$ 转置(成 $K \times B$):
$$ Q \leftarrow \mathrm{diag}(\mathbf{u})^{-1}\, Q\, \mathrm{diag}(\mathbf{v})^{-1} $$- $\mathbf{u} = $ 每行和(=1)→ 推出 $\mathbf{u} = 1 / (Q\mathbf{1})$
- $\mathbf{v} = $ 每列和(=1)→ 推出 $\mathbf{v} = 1 / (Q^\top \mathbf{1})$
写开就是"先除行和再除列和"。
4.3 视角 B:带熵正则的 Optimal Transport(Cuturi 2013)
这是 Sinkhorn-Knopp 重新"翻红"的原因。考虑 OT 问题:
$$ \min_{P \in \Pi(\mathbf{a}, \mathbf{b})} \langle C, P \rangle - \varepsilon H(P) $$其中 $H(P) = -\sum_{ij} P_{ij} \log P_{ij}$ 是 Shannon 熵,$\mathbf{a}, \mathbf{b}$ 是给定的边际分布。
定理:最优解必有形式 $P^\star_{ij} = u_i \cdot K_{ij} \cdot v_j$(其中 $K_{ij} = e^{-C_{ij}/\varepsilon}$),并且 $\mathbf{u}, \mathbf{v}$ 由 Sinkhorn 迭代给出:
$$ \mathbf{u} \leftarrow \mathbf{a} \,/\, (K \mathbf{v}),\quad \mathbf{v} \leftarrow \mathbf{b} \,/\, (K^\top \mathbf{u}) $$当 $\mathbf{a} = \mathbf{1}/B$、$\mathbf{b} = \mathbf{1}/K$ 时,这就是把 $K$ 缩到双随机(再 $\times$ 一个标量)。
4.4 视角 C:指数族投影(DINO 论文视角)
teacher logits $Z$ 是某个"原始能量"的负值。teacher 目标 $P$ 是:
- 是 $Z$ 的"温度 softmax"
- 在 $K$ 个 prototype 上"均匀"
用 Sinkhorn-Knopp 在 $\exp(Z/\tau)$ 上做"双随机投影",等价于解
$$ P^\star = \arg\min_{P \ge 0,\, P\mathbf{1}=\mathbf{1},\, P^\top\mathbf{1}=B/K\mathbf{1}} \mathrm{KL}(P \,\|\, \mathrm{softmax}(Z/\tau)) $$也就是"在双随机约束下,找最接近 softmax 的分布"。这就是 DINO 论文 Eq. (2) 的来源。
5. DINOv3 代码逐行拆解
@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第一阶段:归一化到"总和=1"
| 代码 | 含义 |
|---|---|
Q = exp(z/τ).t() | 转成 $[K, B]$,“行=原型,列=样本"视角 |
sum_Q + all_reduce | 跨 rank 求总和(每 rank 只看到自己 batch,全局必须加起来) |
Q /= sum_Q | 让 $\sum_{ij} Q_{ij} = 1$(数值稳定:防止后续 Q / row_sum 时除以小量) |
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第二阶段:3 次交替缩放
| 步骤 | 维度 | 目标 | 数值 |
|---|---|---|---|
| 行归一 | dim=1(K 维) | 每行 sum → 1/K | Q /= row_sum; Q /= K |
| 列归一 | dim=0(B 维) | 每列 sum → 1/B | Q /= col_sum; Q /= B |
| 整体 | — | 最终 Q *= B | 还原为每列和=1 |
为什么每步除 K / B? 单纯"除以 row_sum"只能让每行=1;这里想最终让列和=1(每样本概率分布),所以要在每步把"基准量"也按 $1/K$ 或 $1/B$ 缩。两者组合后:每行 $1/K$、每列 $1/B$、总和 $1$。最后
Q *= B把"每列 1/B"变成"每列 1”。数学上等价于"先做 3 次标准 Sinkhorn(行 1、列 1),再乘 $B/K$"。
5.1 “all_reduce 在 process_subgroup 上”
DINO 用的是数据并行子组(ssl_meta_arch.py:788-805),与 gram teacher 的 default group 不同:
- DINO 的
world_size= 参与同一 SSL 训练的 rank 数(数据并行) - 跨这些 rank 的
all_reduce让"行 sum"、“列 sum”、“全局 sum"是全训练集的统计量
# 跨 rank 同步 row_sum 后再除 → 每行 sum 反映"全训练 batch"
sum_of_rows = torch.sum(Q, dim=1, keepdim=True)
dist.all_reduce(sum_of_rows, group=get_process_subgroup())
Q /= sum_of_rows5.2 iBOT 版的差异
iBOT SinkhornKnoppTeacher 几乎一样,唯一差别:
B = n_masked_patches_tensor # 直接从 dataloader 拿不用 Q.shape[1] * world_size,因为:
- iBOT 只对被掩码的 patch 算 SK
- 每个 sample 被掩的 patch 数不固定
- 所以用一个预先
all_reduce过的n_masked_patches_tensor(来自data["n_masked_patches"])当全局 batch size
iBOT 把 SinkhornKnoppTeacher 写成独立 nn.Module 是为了 torch.compile(注释里说"a single function decorator is bad”,见 ibot_patch_loss.py:23-26)。
5.3 n_iterations=3 够不够?
理论上 Sinkhorn 收敛速度 $O(1/n)$,对 K=65536 量级的矩阵,3 次迭代误差已经 < 1%(DINO 论文实测)。代码注释里也写"using very few iterations"。
论文中给的经验:3-5 次迭代对 prototype 空间训练基本无差异;太多反而把分布"打平"成完全均匀,丢掉 teacher 的判别性信号。
6. 收敛性:什么情况下不收敛?
6.1 数学要求
Sinkhorn 收敛要求输入矩阵 $Q$ 的支撑图(哪些位置非零)有完美匹配。如果 $Q$ 中存在全零行 / 全零列,算法除以 0 会爆。
6.2 在 DINO 里为什么安全?
- $Q = \exp(Z/\tau)$:严格正(无零元素)
- $\tau > 0$ 有限:不会因温度太低而数值下溢到 0
- 整体缩放到 sum=1 防止 row_sum ≈ 0
6.3 极端情况
如果某个 sample 的 logits $z$ 在某个维度非常负($-50$ 量级),$\exp(-50/\tau) \approx 0$,可能导致浮点下溢 → 某行实际为 0 → 除法爆 nan。所以代码里用 float() 升精度 + 整体归一化兜底。
7. 与"softmax-centering"的对比(DINOv1 vs DINOv2 起)
DINO 原始论文(v1)用的是 softmax + EMA center:
# 伪代码(DINOv1)
center = 0.9 * center + 0.1 * mean(teacher_output, dim=0) # 累积均值
teacher_target = softmax((teacher_output - center) / teacher_temp)| 维度 | softmax-center(DINOv1) | Sinkhorn-Knopp(DINOv2+) |
|---|---|---|
| 机制 | EMA 减均值,显式“去偏” | 迭代投影到双随机约束,隐式均匀 |
| 维护 | 需要 center buffer 持续更新 | 无状态——不需要维护额外参数 |
| 目标 | 让 logits 减去均值,prototype “用得平均” | 强制列和=1,比"减均值"更"硬" |
| 抗 collapse | 弱:可能只是"减均值后 softmax 趋平" | 强:直接保证 prototype 使用率 |
| batch 依赖 | 弱(center 是 EMA) | 强(每次 Sinkhorn 都看当前 batch) |
| 与温度关系 | $\tau$ 越低分布越尖 | $\tau$ + Sinkhorn 迭代次数共同决定 |
为什么 DINOv2 之后改成 SK?
- 更鲁棒:避免
center累积速度与学习率耦合的玄学 - 更简单:不需要 freeze last layer 等 trick 来稳定 center 更新
- 更通用:直接对"目标分布"做约束,对 batch size 不那么敏感
8. 整条链路:代码 ↔ 公式对照
| 代码片段 | 公式含义 |
|---|---|
Q = exp(z/τ).t() | 软分配矩阵 $Q_{ij} = \exp(z_{ij}/\tau)$,转置为 $[K, B]$ |
Q /= sum(Q) + all_reduce | $\sum_{ij} Q_{ij} = 1$(跨 rank 同步) |
Q /= sum_of_rows; Q /= K | 每行 sum → $1/K$ |
Q /= sum_of_cols; Q /= B | 每列 sum → $1/B$ |
for _ in range(3) | 3 次 Sinkhorn 迭代 |
Q *= B | 还原"列和=1"形态(行和=$1/K$) |
return Q.t() | 转回 $[B, K]$,给交叉熵当 target |
n_iterations=3 | 论文经验值 |
process_subgroup | 数据并行 rank 通信组 |
9. 几个易混点
“列和=1” vs “行和=1”:
- 矩阵
Q.t()形状是 $[K, B]$:所以"行"指 $K$ 维度(prototype),“列"指 $B$ 维度(sample) - 最终
Q *= B还原后,每列(sample)和=1——这是 student 交叉熵的合法 target
- 矩阵
为什么
world_size不参与每步的缩放? 因为sum_of_rows已经all_reduce过了,它的值就是"全局行和”;不需要再用world_size缩放。顺序敏感吗? Sinkhorn 迭代是"先 row 再 col"还是"先 col 再 row",理论上极限一样(都是双随机),但迭代次数有限时数值会有微差。DINO 论文固定 row → col 顺序,代码也保持一致。
能否不开
torch.no_grad? 一定要开(代码里也写了)。teacher 不参与反传,开了省内存。teacher_temp(DINOv3 默认 0.04-0.07)和n_iterations=3是耦合的:温度越低 → 分布越尖 → 双随机投影"打平"越多 → 越接近均匀分布。经验上调低温度要同时减少迭代次数(否则目标太均匀)。
10. 一句话总结
Sinkhorn-Knopp = “交替行/列归一"的迭代算法,把任意正矩阵缩放到"双随机"形式。 在 DINOv3 里它把 teacher logits 变成”列和=1(合法概率分布)且行和均匀(防 prototype collapse)“的软目标,替代了 DINOv1 的 EMA-center,做法更鲁棒、免维护中心 buffer;3 次迭代 + 数据并行 all_reduce 即可获得全局一致的伪标签,是 SSL 多视图对齐的关键一环。