sinkhorn-knopp理解

Sinkhorn-Knopp 笔记

对应代码:

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 代码逐行拆解

dino_clstoken_loss.py:42-70:

@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/KQ /= row_sum; Q /= K
列归一dim=0(B 维)每列 sum → 1/BQ /= 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_rows

5.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. “列和=1” vs “行和=1”:

    • 矩阵 Q.t() 形状是 $[K, B]$:所以"行"指 $K$ 维度(prototype),“列"指 $B$ 维度(sample)
    • 最终 Q *= B 还原后,每列(sample)和=1——这是 student 交叉熵的合法 target
  2. 为什么 world_size 不参与每步的缩放? 因为 sum_of_rows 已经 all_reduce 过了,它的值就是"全局行和”;不需要再用 world_size 缩放。

  3. 顺序敏感吗? Sinkhorn 迭代是"先 row 再 col"还是"先 col 再 row",理论上极限一样(都是双随机),但迭代次数有限时数值会有微差。DINO 论文固定 row → col 顺序,代码也保持一致。

  4. 能否不开 torch.no_grad? 一定要开(代码里也写了)。teacher 不参与反传,开了省内存。

  5. 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 多视图对齐的关键一环。