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] |
其中 D 是 hidden_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 | 输入 |
它的关键思想不是简单地“做一次激活”,而是:
让一条分支产生内容,另一条分支决定这些内容应该通过多少。
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_projup_projdown_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 | gate = gate_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 | gate: [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 | 输入 x |
核心代码其实只有:
1 | gate_up = self.gate_up_proj(x) |
总结
RMSNorm 和 SwiGLU 都是现代 LLM 中非常常见的基础组件,但两者解决的问题不同。
RMSNorm 主要解决的是 数值尺度稳定性。学习时最需要关注最后一个维度的归一化、keepdim=True、FP32 计算以及可学习 weight。
SwiGLU 主要解决的是 MLP 中的信息选择和特征变换。学习时最需要理解 gate 和 up 的作用、8 / 3 的参数量来源、multiple_of 的对齐意义,以及最后为什么必须经过 down_proj 回到 hidden size。
这两个模块看起来代码都不长,但真正重要的是理解每一步为什么存在。只有把 shape、数值精度和参数量这些细节想清楚,后面阅读完整 Transformer 或 LLM 代码时才不会只停留在 API 层面。