
DINO 中的中心化与锐化:为什么它们能避免模型坍塌?理解 DINO 中的中心化(centering)和锐化(sharpening),最重要的不是先背公式,而是先回答三个问题:自监督学习为什么会发生模型坍塌(model collapse)?中心化到底“减掉”了什么?锐化为什么只需要调一个温度参数,就能让输出变得更明确?整个逻辑可以先浓缩成一句话:中心化(centering)防止所有样本都跑到同一个输出维度;锐化(sharpening)防止模型对所有维度都给出差不多的概率。两者作用方向相反,因此可以互相制衡。一、先建立场景:DINO 到底在学习什么?假设我们有一张图片xxx。对它进行两种不同的数据增强,例如:第一张裁掉左边;第二张裁掉右边;或者一张改变颜色,另一张改变亮度。得到:x1=Aug1(x)x_1=\operatorname{Aug}_1(x)x1=Aug1(x)x2=Aug2(x)x_2=\operatorname{Aug}_2(x)x2=Aug2(x)DINO 中存在两个网络:学生网络(student network)教师网络(teacher network)教师看到x1x_1x1,学生看到x2x_2x2。我们的目标是:虽然两张图片长得不完全一样,但是它们来自同一张原图,因此教师和学生应该产生相似的表示。可以简单写成:pt(x1)≈ps(x2)p_t(x_1)\approx p_s(x_2)pt(x1)≈ps(x2)其中:ptp_tpt表示教师的输出概率;psp_sps表示学生的输出概率。学生通过训练不断模仿教师。看起来非常合理,但这里隐藏着一个严重问题。二、为什么会出现模型坍塌?1. 模型发现了一种“作弊方法”假设教师和学生都输出三个数字。理想情况下,不同图片应该产生不同结果。例如图片 A:[0.8,0.1,0.1][0.8,0.1,0.1][0.8,0.1,0.1]图片 B:[0.1,0.8,0.1][0.1,0.8,0.1][0.1,0.8,0.1]图片 C:[0.1,0.1,0.8][0.1,0.1,0.8][0.1,0.1,0.8]说明模型能够区分不同图片。但是模型可能发现:如果不管输入什么图片,我永远输出同样的答案,那么教师和学生不就永远一致了吗?例如所有图片都输出:[1,0,0][1,0,0][1,0,0]那么对于任意图片:pt(x)=ps(x)=[1,0,0]p_t(x)=p_s(x)=[1,0,0]pt(x)=ps(x)=[1,0,0]教师和学生确实完全一致。训练损失可能也很低。但是问题在于:模型虽然完成了“教师和学生保持一致”这个任务,却完全没有学习到图片之间的区别。这就是模型坍塌(model collapse)。三、坍塌其实有两个不同的极端理解这一点,是理解中心化和锐化的关键。第一种:单一维度坍塌假设无论输入什么图片,模型都输出:[0.98,0.01,0.01][0.98,0.01,0.01][0.98,0.01,0.01]第一维永远特别大。也就是说:所有图片都被认为属于同一个方向。可以把这种情况称为单一维度主导(single-dimension domination)。第二种:均匀坍塌还有另一种极端。无论输入什么图片,模型都输出:[1/3,1/3,1/3][1/3,1/3,1/3][1/3,1/3,1/3]看起来似乎没有某个维度统治其他维度。但是所有图片依然完全相同。所以这仍然是坍塌。这种情况可以称为均匀坍塌(uniform collapse)。因此,我们面对两个相反的危险:[1,0,0]和[1/3,1/3,1/3][1,0,0]\qquad\text{和}\qquad[1/3,1/3,1/3][1,0,0]和[1/3,1/3,1/3]DINO 的设计非常巧妙:中心化(centering)主要阻止第一种情况;锐化(sharpening)主要阻止第二种情况。四、模型的输出是怎样变成概率的?假设教师网络最后产生KKK个数字:z=[z1,z2,…,zK]z=[z_1,z_2,\dots,z_K]z=[z1,z2,…,zK]这些数字叫做logits(未归一化分数)。例如:z=[2,1,0]z=[2,1,0]z=[2,1,0]它们本身不是概率。为了变成概率,需要经过Softmax 归一化(softmax normalization):pi=exp(zi/τ)∑j=1Kexp(zj/τ)p_i=\frac{\exp(z_i/\tau)}{\sum_{j=1}^{K}\exp(z_j/\tau)}pi=∑j=1Kexp(zj/τ)exp(zi/τ)其中τ\tauτ叫做温度(temperature)。最终得到:p=[p1,p2,…,pK]p=[p_1,p_2,\dots,p_K]p=[p1,p2,…,pK]并且:∑i=1Kpi=1\sum_{i=1}^{K}p_i=1i=1∑Kpi=1为了方便理解,高中阶段可以暂时把pip_ipi想象成:模型对“第iii个潜在类别”的偏好程度。需要注意:DINO 并没有真正人工定义“类别 1 是猫、类别 2 是狗”,这里的“类别”只是帮助理解。五、中心化到底是什么?先记住一个最简单的统计学操作。假设三个人的成绩为:80,90,10080,\ 90,\ 10080,90,100平均值是:xˉ=90\bar{x}=90xˉ=90每个人减去平均数:80−90=−1080-90=-1080−90=−1090−90=090-90=090−90=0100−90=10100-90=10100−90=10于是:[80,90,100]→[−10,0,10][80,90,100]\rightarrow[-10,0,10][80,90,100]→[−10,0,10]原来的数据围绕 90 分布。减去平均值以后,数据围绕 0 分布。这就是最基本的“中心化”。六、DINO 中为什么也需要中心化?假设教师网络有三个输出维度。观察大量图片以后发现:c=[10,2,1]c=[10,2,1]c=[10,2,1]也就是说:第一维平均输出大约是 10;第二维平均输出大约是 2;第三维平均输出大约是 1。这说明第一维存在一个非常严重的“天然偏高”。现在三张图片分别得到:zA=[12,3,2]z_A=[12,3,2]zA=[12,3,2]zB=[11,4,1]z_B=[11,4,1]zB=[11,4,1]zC=[13,2,2]z_C=[13,2,2]zC=[13,2,2]只看原始数字,会发现:每一张图片都是第一维最大。如果直接经过 Softmax,很可能所有图片都变成:[1,0,0][1,0,0][1,0,0]于是发生单一维度坍塌。但注意:第一维很大,不一定说明这张图片真的特别属于第一维。也可能只是:第一维平时本来就特别大。因此我们应该去掉这种“长期偏置”。七、中心化具体做什么操作?最直观的形式是:z′=z−cz'=z-cz′=z−c也就是:当前输出减去长期平均输出。对于图片 A:zA′=[12,3,2]−[10,2,1]z'_A=[12,3,2]-[10,2,1]zA′=[12,3,2]−[10,2,1]得到:zA′=[2,1,1]z'_A=[2,1,1]zA′=[2,1,1]对于图片 B:zB′=[11,4,1]−[10,2,1]z'_B=[11,4,1]-[10,2,1]zB′=[11,4,1]−[10,2,1]得到:zB′=[1,2,0]z'_B=[1,2,0]zB′=[1,2,0]对于图片 C:zC′=[13,2,2]−[10,2,1]z'_C=[13,2,2]-[10,2,1]zC′=[13,2,2]−[10,2,1]得到:zC′=[3,0,1]z'_C=[3,0,1]zC′=[3,0,1]现在情况发生了变化。原来:三张图片第一维都最大。中心化以后:A 更偏向第一维;B 更偏向第二维;C 更偏向第一维。这时网络开始表现出样本之间的区别。八、中心化真正减掉的是什么?可以把网络输出想象成两部分:当前输出=长期偏置+当前图片真正带来的变化\text{当前输出}=\text{长期偏置}+\text{当前图片真正带来的变化}当前输出=长期偏置+当前图片真正带来的变化例如:z=[10+a,2+b,1+c]z=[10+a,2+b,1+c]z=[10+a,2+b,1+c]其中:[10,2,1][10,2,1][10,2,1]是各个维度长期存在的平均偏置。真正由当前图片决定的是:[a,b,c][a,b,c][a,b,c]如果中心为:ccenter≈[10,2,1]c_{\text{center}}\approx[10,2,1]ccenter≈[10,2,1]那么:z−ccenter≈[a,b,c]z-c_{\text{center}}\approx[a,b,c]z−ccenter≈[a,b,c]因此,中心化真正做的事情是:去掉“这个神经元平时就很大”的影响,只看“这张图片让它比平时高了多少”。这是理解中心化最重要的一句话。九、中心ccc是怎样计算出来的?DINO 使用下面的更新公式:c←mc+(1−m)1B∑i=1Bgθt(xi),(4)c\leftarrow mc+(1-m)\frac{1}{B}\sum_{i=1}^{B}g_{\theta_t}(x_i),\tag{4}c←mc+(1−m)B1i=1∑B