Gemma RMSNorm:从纯归一化开始学习缩放
🎯 先猜一猜
self.weight 初始化为全 0,但 forward 仍执行了 RMS 归一化。
此时 output 最准确的描述是什么?
先补的知识
- x 的 shape 是 [batch_size, seq_len, hidden_size];RMSNorm 只沿最后的 hidden_size 维计算均方。
- variance = x.float().pow(2).mean(-1, keepdim=True) 得到 [B,S,1],再通过广播作用到 [B,S,H]。
- x_norm 已经是 FP32 的归一化结果;self.weight 是长度 H 的可训练参数,会沿 B、S 两维广播。
- Gemma 的 weight 初始化为 0,所以真正的缩放因子 1 + weight 初始为 1;return 时必须用 type_as(x) 恢复输入 dtype。
图解原理
标准缩放参数若直接乘 weight,就需要把 weight 初始化为 1。Gemma 换一种参数化:保存从 1 的偏移量,实际乘 (1+w)。当 w=0 时,层先只做 RMS 归一化;训练再学习每个隐藏特征应该从 1 向上或向下调整多少。
shape 与广播
variance[B,S,1]
x_norm[B,S,H],FP32
weight[H],每个隐藏特征一个缩放偏移
output[B,S,H],返回时恢复 x.dtype
初始化时发生什么
这不是恒等映射 output=x;它仍然做了 RMS 归一化,只是没有额外改变归一化后的特征缩放。
| 测试 | 期望 | 它证明什么 |
|---|---|---|
| weight=0 | out 等于手算的 x_norm | 正确使用了 1+w |
| weight 非零 | out2 与 out 不同 | 可训练缩放真正生效 |
| FP16 输入 | FP16 输出 | 最后恢复了原 dtype |
语法热身:从基准亮度 1 学习每个通道的偏移
normalized_pixels = torch.tensor([
[0.5, -1.0, 0.25],
])
channel_offset = torch.tensor([0.0, 0.2, -0.1])
scale = 1 + channel_offset
adjusted = normalized_pixels * scale
print(scale) # [1.0, 1.2, 0.9]
print(adjusted.shape)独立例子如何迁移回 Notebook
normalized_pixels对应已经算好的x_norm。channel_offset对应可训练的self.weight。adjusted对应 TODO 要创建的output。- Notebook 的 return 已负责
type_as(x),因此 TODO 保持 FP32 运算即可。
巩固一下
为什么 variance 使用 keepdim=True?
学完这一段,试着做
用一个小动作确认自己理解了;最后再进入官方题目。