RMSNorm

RMSNorm 的核心目的,是让每个 token 的 hidden states 保持在一个相对稳定的数值尺度上。

它和 LayerNorm 最大的区别,是 不减均值。RMSNorm 不关心特征整体有没有偏移,而主要关注向量的大小,也就是 Root Mean Square。

对于一个 hidden vector,核心思想可以理解成:

先计算所有 hidden 维度的平方平均值,再用它对原始向量做缩放。

公式为:

$$
\operatorname{RMS}(x)

\sqrt{\frac{1}{d}\sum_{i=1}^{d}x_i^2+\epsilon}
$$

归一化以后,再乘上一个可学习的 weight


学习重点

RMSNorm 最重要的不是背公式,而是理解它到底在 哪个维度 上做计算。

假设输入 shape 是:

1
[B, T, D]

其中 Dhidden_size

RMSNorm 是对每一个 token 单独处理,所以平方、求平均都是沿最后一个维度进行:

1
x.pow(2).mean(dim=-1, keepdim=True)

最终 variance 的 shape 是:

1
[B, T, 1]

这里的 1 非常关键,因为后面还需要和 [B, T, D] 的原始输入进行广播计算。

因此,看到:

1
keepdim=True

要立刻想到:

保留最后这个维度,是为了后面的 broadcasting。


为什么要先转成 float32

RMSNorm 中会涉及平方:

1
x.pow(2)

如果模型本身使用 FP16 或 BF16,平方、求平均、开根号这些操作更容易出现数值精度问题。

所以实际实现中通常会先:

1
x_fp32 = x.float()

在 FP32 中完成归一化,再把结果转换回输入原来的 dtype。

这个细节很重要,因为训练大模型时经常会使用混合精度。

可以记住一个经验:

归一化、方差、平方、指数等数值敏感操作,需要特别关注计算精度。


weight 的作用

RMSNorm 虽然进行了归一化,但并不是把结果永远固定在某个尺度。

模型还需要自己学习:

每一个 hidden 维度最终应该放大还是缩小。

因此会有:

1
self.weight = nn.Parameter(torch.ones(hidden_size))

初始化为 1,意味着模型刚开始时不会额外改变归一化结果。

之后训练过程中,这个参数会被不断更新。


容易踩的坑

第一个坑是 求均值的维度写错

如果直接写:

1
x.mean()

那么 batch、sequence、hidden 全部混在一起计算了,这就完全不是 RMSNorm 需要的行为。

正确的是沿最后一个 hidden 维度:

1
mean(dim=-1, keepdim=True)

第二个坑是忘记 keepdim=True

如果去掉以后,shape 会从:

1
[B, T, D]

变成:

1
[B, T]

后续广播会变得不直观,甚至可能产生错误。

第三个坑是忽略 dtype。

归一化过程可以临时使用 FP32,但最终一般应该恢复到输入 dtype,否则可能导致后续模块出现不必要的类型变化。


一句话理解 RMSNorm

RMSNorm 可以理解成:

对每个 token 的 hidden vector 根据自身的整体大小进行重新缩放,使数值尺度更加稳定,但不做均值中心化。


SwiGLU

SwiGLU 是 Transformer 中 MLP 的一种门控结构。

普通 MLP 可以简单理解为:

1
升维 → 激活函数 → 降维

而 SwiGLU 多了一条 gate 分支:

1
2
3
4
5
6
7
输入
├─ gate
└─ up

SiLU(gate) * up

down projection

它的关键思想不是简单地“做一次激活”,而是:

让一条分支产生内容,另一条分支决定这些内容应该通过多少。


gate 和 up 怎么理解

up 分支可以理解为真正需要进行特征变换的信息。

gate 分支则更像一个门控信号。

计算过程是:

1
F.silu(gate) * up

这里是逐元素乘法。

因此 gate 和 up 的 shape 必须完全一致。

直觉上可以理解成:

up 提供候选信息,gate 决定哪些信息重要,以及保留多少。

这也是 SwiGLU 相比普通激活函数更灵活的地方。


为什么会出现 8 / 3

这是 SwiGLU 学习中最容易死记硬背的地方。

其实 8 / 3 不是随便选出来的。

传统 Transformer MLP 经常使用:

1
hidden_size → 4 × hidden_size → hidden_size

它主要有两个大的线性层,所以参数量大约是:

$$
8d^2
$$

而 SwiGLU 有三个投影:

  • gate_proj
  • up_proj
  • down_proj

如果 intermediate size 还是 4d,参数量会明显增加。

因此为了让 SwiGLU 的参数规模和普通 4d MLP 大致接近,需要缩小 intermediate size。

最后得到:

$$
\text{intermediate size} \approx \frac{8}{3}d
$$

所以:

1
intermediate_size = int(hidden_size * 8 / 3)

理解这一点以后,就不需要把 8 / 3 当成魔法数字。

它本质上是在做:

因为多了一个线性层,所以把中间维度缩小,从而控制整体参数量。


multiple_of 是做什么的

理论上算出来的 intermediate_size 不一定是一个硬件友好的数字。

例如:

1
10922

实际实现中可能会把它向上调整到:

1
11008

也就是某个固定数的倍数。

这样做主要是为了让矩阵计算更加规整,更适合 GPU 上的高效计算。

所以 multiple_of 本质上是在做:

把理论上的 intermediate size 调整成更适合实际计算的尺寸。

这里要注意,它通常是 向上对齐,而不是简单取整或者向下取整。


为什么 gate 和 up 可以合并成一个 Linear

最直观的写法是:

1
2
gate = gate_proj(x)
up = up_proj(x)

但这两个操作:

  • 输入完全一样;
  • 输出维度一样;
  • 本质上都是一次矩阵乘法。

因此可以把它们合并:

1
gate_up_proj(x)

一次直接输出:

1
2 × intermediate_size

然后再切成两半:

1
gate, up = torch.chunk(gate_up, 2, dim=-1)

这样做的核心意义是:

减少重复的投影调用,让计算更加集中和高效。


torch.chunk 的坑

假设:

1
gate_up.shape = [B, T, 2I]

那么:

1
torch.chunk(gate_up, 2, dim=-1)

会返回两个 tensor:

1
2
gate: [B, T, I]
up: [B, T, I]

这里比较容易写成:

1
torch.chunk(...)[0]

这样实际上只拿到了第一部分,也就是 gate。

正确的思路应该是:

1
gate, up = torch.chunk(...)

因为 SwiGLU 两条分支都需要。


为什么最后还需要 down_proj

经过 gate 和 up 计算以后,tensor 仍然处于 intermediate size:

1
[B, T, I]

但 Transformer 主干的 hidden size 是:

1
[B, T, D]

所以必须通过:

1
down_proj

把 intermediate size 再映射回 hidden size。

也就是说,SwiGLU 内部可以升维做复杂的特征变换,但模块最终输出的 shape 必须和输入保持一致,才能继续参与残差连接和后续 Transformer Block 计算。


SwiGLU 最值得记住的流程

整个过程可以压缩成:

1
2
3
4
5
6
7
8
9
10
11
输入 x

gate_up_proj

拆成 gate 和 up

SiLU(gate) * up

down_proj

回到 hidden_size

核心代码其实只有:

1
2
3
gate_up = self.gate_up_proj(x)
gate, up = torch.chunk(gate_up, 2, dim=-1)
output = self.down_proj(F.silu(gate) * up)

总结

RMSNorm 和 SwiGLU 都是现代 LLM 中非常常见的基础组件,但两者解决的问题不同。

RMSNorm 主要解决的是 数值尺度稳定性。学习时最需要关注最后一个维度的归一化、keepdim=True、FP32 计算以及可学习 weight

SwiGLU 主要解决的是 MLP 中的信息选择和特征变换。学习时最需要理解 gate 和 up 的作用、8 / 3 的参数量来源、multiple_of 的对齐意义,以及最后为什么必须经过 down_proj 回到 hidden size。

这两个模块看起来代码都不长,但真正重要的是理解每一步为什么存在。只有把 shape、数值精度和参数量这些细节想清楚,后面阅读完整 Transformer 或 LLM 代码时才不会只停留在 API 层面。