layerscale理解

DINOv3 中的 LayerScale 笔记

对应代码:

1. 一句话定义

LayerScale = 残差分支上一个 per-channel 的可学习缩放 $\gamma$,从 $10^{-5}$ 这种接近零的值初始化。 它让 Transformer 在初始化时近似恒等映射,训练中再"拧大"每个分支的实际贡献。

来源:Touvron et al. 2021, “Going Deeper with Image Transformers”(LayerScale 原始论文)。

2. 公式:就是逐通道的一个可学习缩放向量

layer_scale.py:12-29 整个模块 8 行有效代码:

class LayerScale(nn.Module):
    def __init__(self, dim, init_values=1e-5, inplace=False, device=None):
        super().__init__()
        self.inplace = inplace
        self.gamma = nn.Parameter(torch.empty(dim, device=device))
        self.init_values = init_values

    def reset_parameters(self):
        nn.init.constant_(self.gamma, self.init_values)

    def forward(self, x):
        return x.mul_(self.gamma) if self.inplace else x * self.gamma

数学上等价于给残差分支的输出左乘一个对角矩阵 $\text{diag}(\gamma_1, \ldots, \gamma_D)$:

$$ \mathbf{x}_{\text{out}} = \gamma \odot \mathbf{x}_{\text{residual}} $$

注意是 per-channel(每个特征维度一个 $\gamma_i$),不是 per-token 也不是 per-scalar。形状 [dim],对所有 token、所有 batch、所有空间位置共享同一组 $\gamma$。

3. 在 SelfAttentionBlock 里放在哪?

block.py:54 和 block.py:66:

self.ls1 = LayerScale(dim, init_values=init_values, device=device) if init_values else nn.Identity()  # attention 分支
self.ls2 = LayerScale(dim, init_values=init_values, device=device) if init_values else nn.Identity()  # mlp 分支

对应的前向(block.py:121-122):

x_attn = x        + self.ls1(self.attn(self.norm1(x), rope=rope))   # 注意力残差
x_ffn  = x_attn   + self.ls2(self.mlp(self.norm2(x_attn)))         # MLP 残差

LayerScale 夹在分支输出和残差相加之间——它是"残差分支上的可调音量旋钮"。

关键 if 分支:当 init_values=None 时,LayerScale 整层被替换为 nn.Identity(),完全关闭该机制(不做缩放)。DINOv3 默认在视觉骨干里是开着的;消融实验里通常会关掉它做对比。

CausalSelfAttentionBlock(block.py:231, 244)用同样的模式:ls1 在 attention 分支前,ls2 在 FFN 分支前。

4. 为什么 $\gamma$ 初始化成 1e-5 那么小?

这是 LayerScale 的精髓——让模型从"恒等映射"开始学。

4.1 初始化时的行为

$\gamma$ 初始化为 $\epsilon$(默认 $10^{-5}$),所以初始化那一瞬间:

$$ x + \epsilon \cdot f(x) \approx x $$

整个 Transformer block 在第 0 步的输入输出几乎一样,等价于一个接近恒等映射的浅层。

这为什么重要?

  • 没加 LayerScale 时,深层 Transformer 在初始化阶段就有"乱来的残差贡献",叠加 N 次后信号方差爆炸或梯度消失
  • 加了之后,模型是"从浅到深"逐步"打开"每个 block 的容量

4.2 训练过程中的行为

$\gamma$ 是可学习的,所以训练中网络会自适应地:

  • 对重要分支(比如浅层 attention、深层 MLP)把对应的 $\gamma$ 学大
  • 对不重要分支让 $\gamma$ 保持小或继续收缩

等价于一种"软剪枝 / 软深度选择"——网络可以自己"决定"每个 block 实际贡献多少非线性。

5. 与其他残差修饰机制的对比

机制作用位置操作粒度训练/推理
残差连接 $x + f(x)$分支出口加法per-element一致
DropPath (Stochastic Depth)分支出口随机把整分支置 0per-sample(整个分支)训练时随机,推理关闭
LayerScale分支出口乘一个可学习 $\gamma$per-channel训练 / 推理都生效

LayerScale 的相对优势:

  • 比 DropPath 更"温和":DropPath 是 0/1 二值离散;LayerScale 是连续可微的,学习更平滑
  • 比固定缩放 $\lambda$ 更灵活:每个维度独立 $\gamma_i$,给网络更细的控制
  • 推理时也工作:DropPath 在推理时关闭,而 LayerScale 的 $\gamma$ 训完固定,部署时不需要切换模式

6. 一个数值直觉

设 $\gamma$ 初始化 $10^{-5}$,分支输出 $f(x)$ 的某元素 $f_j = 100$(未缩放时残差贡献很大):

训练阶段$\gamma_j$残差贡献 $f_j \cdot \gamma_j$效果
初始化$10^{-5}$$10^{-3}$几乎不动,近似恒等
训练中(重要分支)学到 $0.5$$50$开始显著贡献
训练后(无用分支)维持 $10^{-5}$$10^{-3}$自动"关掉"

7. 整条链路:代码 ↔ 公式对照

代码片段公式含义
self.gamma = nn.Parameter(torch.empty(dim))$\gamma \in \mathbb{R}^{D}$,per-channel 可学习向量
nn.init.constant_(self.gamma, init_values)$\gamma \leftarrow \epsilon$,默认 $\epsilon = 10^{-5}$
x * self.gamma残差分支逐元素乘 $\gamma$
if init_values else nn.Identity()init_values=None → 整层关闭,不缩放
ls1(attention 分支)控制 $\text{attn}(\cdot)$ 输出对残差的贡献
ls2(mlp 分支)控制 $\text{mlp}(\cdot)$ 输出对残差的贡献
x + ls(f(norm(x)))“缩放后的残差分支 + 恒等分支”——LayerScale 的完整残差块结构

8. DINOv3 里的两个特殊点

8.1 inplace=False(默认)

DINOv3 不使用 in-place 乘法。这让 gamma 的梯度可以正常回传(in-place 操作会破坏 autograd 的版本计数),也方便和 torch.compile 配合——block.py:17 设置 torch._dynamo.config.automatic_dynamic_shapes = False,整体对图编译友好。

8.2 配合 sample dropping 的训练

DINOv3 的 SelfAttentionBlock 还用了样本级 stochastic depth(block.py:90-119):对 batch 内每个样本独立地以 sample_drop_ratio 概率"跳过"分支。LayerScale 和 DropPath 在这里是互补的:

  • DropPath 控制"哪些样本走分支"(粗粒度,per-sample 二值)
  • LayerScale 控制"走分支时贡献多大"(细粒度,per-channel 连续)

两者叠加,训练时信号既不会爆(LayerScale 缩 + DropPath 跳),又能在推理时完整利用全部容量(LayerScale 已训好 + DropPath 关闭)。

9. 小结

LayerScale 的核心思想就一句话:让残差分支"从静音开始"。$\gamma$ 从 $10^{-5}$ 起步,网络在训练中自己"调音量",既稳定了超深 Transformer 的初始化,又给了每条分支 per-channel 的精细控制。它和 DropPath 不是替代关系——前者是"音量旋钮",后者是"开关"——在 DINOv3 里两者一起用,构成了训练超深 ViT 骨干的稳定器。