从 BatchNorm 到 RMSNorm:如何理解大模型中的归一化

从 BatchNorm 到 RMSNorm:如何理解大模型中的归一化
刚接触归一化时我只记住了一句话减去均值再除以标准差。后来学习 Transformer我发现它使用 LayerNorm到了 LLaMA又变成了 RMSNorm同时还经常出现 Pre-Norm、Post-Norm。这些概念看起来相似实际上讨论的是两个不同问题• BatchNorm、LayerNorm、RMSNorm 决定对哪些数值进行归一化以及具体怎样归一化• Pre-Norm、Post-Norm 决定把归一化放在残差结构的什么位置。理解这两条主线才能真正明白为什么现代大模型经常采用 Pre-RMSNorm。一、归一化到底在控制什么假设某一层接收到隐藏向量 (x)经过线性变换得到如果网络很深每一层都会改变隐藏状态的数值尺度。某些层可能不断放大激活某些层可能不断缩小激活最终导致不同层之间的数值范围差异很大。这对 Attention 尤其敏感。注意力分数为假设隐藏状态整体放大两倍经过线性投影后(Q) 和 (K) 通常也会近似放大两倍于是 (QK^\top) 会放大约四倍。例如原本两个注意力分数为经过 Softmax 后如果 (QK^\top) 整体放大四倍那么注意力分布会突然变得非常尖锐。模型可能过早地将注意力集中到少数位置梯度形态也会明显改变。因此我更愿意把归一化理解成一种“数值尺度控制器”归一化让 Attention 和 MLP 面对尺度相对稳定的输入从而降低深层网络的优化难度。它并不是单纯为了让数据在统计意义上“更漂亮”而是在减少网络对绝对数值尺度的敏感性。二、BatchNorm依赖一批样本共同计算Batch Normalization 的核心特征是利用一批样本共同计算统计量。假设一个批次中第 (j) 个特征有 (m) 个取值BN 首先计算该特征在当前批次中的均值然后计算方差接下来完成归一化最后通过可学习参数进行缩放和平移其中(\gamma_j) 和 (\beta_j) 并不是多余的。BN 先把数值调整到稳定范围再允许模型根据任务需要重新学习合适的尺度和中心位置。一个直观例子假设某个特征在一个 batch 中的取值为均值为方差为忽略 (\epsilon)标准差为 1因此如果设置那么最终输出为BN 的本质样本之间存在耦合BN 的统计量来自整个 batch所以一个样本发生变化其他样本的归一化结果也可能跟着变化。例如将第二个样本从 3 改成 5此时均值和方差都会改变第一个样本 1 的归一化结果也会随之改变。这就是 BN 与后面两种方法最根本的区别BN 不是让每个样本独立完成归一化而是让一批样本共同决定归一化尺度。在忽略 (\epsilon) 且缩放系数为正数的情况下如果整个批次做统一的缩放和平移那么归一化结果基本不变但 BN 并不适合作为大语言模型的主流归一化方案主要有三个原因• 统计结果依赖 batchmicro-batch 较小时容易波动• 不同样本相互耦合不适合变长序列和逐 Token 自回归生成• 训练时使用当前批次统计量推理时通常使用累计的 running mean 和 running variance训练与推理的计算方式不完全一致。需要说明的是BN 并不天然要求跨 GPU 通信。只有使用同步 BatchNorm希望多张 GPU 共同计算统计量时才会额外引入设备间通信。【图片1BN、LN、RMSNorm 的归一化维度对比】三、LayerNorm让每个 Token 独立归一化LayerNorm 不依赖 batch而是对单个样本内部的特征维度进行统计。在 Transformer 中隐藏状态通常表示为其中• (B) 表示批次大小• (L) 表示序列长度• (d) 表示隐藏维度。对于其中某个 Token它的隐藏向量可以写成LayerNorm 只在这个 Token 自己的 (d) 个隐藏维度上计算均值然后计算方差最终输出为一个 Token 如何完成 LayerNorm假设某个 Token 的隐藏向量为均值为方差为标准差约为暂时令 (\gamma1,\beta0)则如果把输入统一放大两倍再整体加上 10新的均值为 14标准差约为 1.633归一化结果仍然约为因此在忽略 (\epsilon) 且 (a0) 时LayerNorm 同时消除了统一平移和正比例缩放的影响。更重要的是每个 Token 都可以独立完成 LayerNorm。同一批次中其他句子发生变化不会影响当前 Token 的归一化结果。这让它具备三个明显优势• 不依赖 batch size• 适合变长序列• 训练和推理采用完全相同的计算规则。因此LayerNorm 自然成为 Transformer 中最经典的归一化方式。四、RMSNorm只控制尺度不减均值RMSNorm 的计算更加直接。它不再计算均值也不要求输出以 0 为中心而是只利用均方根控制隐藏向量的整体尺度。对于隐藏向量它的均方根为RMSNorm 的输出为标准 RMSNorm 通常只保留缩放参数 (\gamma)不使用平移参数 (\beta)。用四个数字走完 RMSNorm假设先计算平方的平均值所以令 (\gamma1)归一化结果为可以看到RMSNorm 没有把输出均值调整为 0而只是用同一个 RMS 值除以所有维度。现在把输入整体放大 10 倍它的均方根也会放大 10 倍所以结果基本不变。因此在忽略 (\epsilon) 且 (a0) 时我认为理解 RMSNorm 最关键的不是记住公式而是理解它完成了两件事• 消除隐藏向量整体尺度的影响• 保留各维度之间的相对比例和符号关系。RMSNorm 将送入 Attention 和 MLP 的输入控制在较稳定的尺度上可以减少 (Q)、(K)、(V) 投影以及激活函数输入的剧烈波动。【图片2不同尺度的隐藏向量经过 RMSNorm 后被拉回稳定范围】五、三种归一化的复杂度究竟有什么差别这一部分最容易出现两个极端要么简单地说“RMSNorm 复杂度更低”要么看到三者都是 (O(N))就认为它们没有任何性能差别。这两种说法都不够准确。1. 渐进复杂度确实相同假设 Transformer 的隐藏张量大小为无论是 BN、LN 还是 RMSNorm都至少需要读取并处理张量中的每个元素因此理论时间复杂度都是三者的差异主要来自常数开销、统计范围和底层 Kernel 实现。2. LayerNorm 和 RMSNorm 实际少算了什么以单个 Token 的 (d) 维向量为例LayerNorm 在概念上需要完成读取 (d) 个元素并计算均值对每个元素减去均值计算方差计算倒数标准差乘以 (\gamma)再加上 (\beta)。RMSNorm 则需要计算每个元素的平方和得到均方根用均方根缩放输入乘以 (\gamma)。因此RMSNorm 主要少了• 均值统计• 逐元素减均值• 平移参数 (\beta) 及相应加法。如果隐藏维度为 (d)LayerNorm 通常具有 (2d) 个可学习参数而标准 RMSNorm 通常只有 (d) 个参数不过相对于大模型中数十亿个矩阵参数这部分参数量差异通常很小。RMSNorm 真正有价值的地方是计算路径更短、更容易高效实现。3. 为什么同为 (O(N))实际速度仍然不同因为渐进复杂度只说明数据规模增大时计算量如何增长并不反映每个元素具体执行了多少次运算、访存和同步。归一化通常不是像矩阵乘法那样的计算密集型算子而更接近显存带宽受限算子• 需要从显存中读取大量激活• 每个元素只进行少量计算• 计算结束后还要把结果重新写回显存。在这种情况下少一次归约、少一次逐元素操作以及减少中间结果读写都可能降低 Kernel 延迟。RMSNorm 计算流程更简单也更容易与前后算子融合成 fused kernel从而减少• Kernel 启动次数• 中间张量写回显存的次数• 重复读取激活的开销。这才是“RMSNorm 通常比 LayerNorm 更轻”的工程含义。4. BN 的开销为什么不能直接和 LN、RMSNorm 横向排序BN、LN 和 RMSNorm 虽然都是 (O(N))但它们统计的维度不同。• BN 需要沿 batch 等维度聚合相同特征• LN 和 RMSNorm 只需要在每个 Token 的隐藏维度内独立归约。单卡训练时BN 不一定比 LN 更慢在 CNN 中高度优化的 BN Kernel 甚至可能非常高效。但如果使用多卡同步 BN为了获得跨设备的全局统计量就需要额外通信。因此不能简单写成BN 并行性差LN 并行性好RMSNorm 并行性最好。更准确的结论是三者的理论复杂度相同RMSNorm 的单次计算最简单但最终性能取决于张量形状、硬件、数值精度以及是否使用融合 Kernel。它的优势通常是降低归一化算子的常数开销而不是显著改变整个大模型的总 FLOPs。大模型的大部分计算仍然来自线性层、MLP 和 Attention。RMSNorm 可以带来工程收益但不会让整个模型按相同比例加速。六、Pre-Norm 与 Post-NormNorm 放在哪里同样重要BN、LN、RMSNorm回答的是“怎样归一化”Pre-Norm 和 Post-Norm回答的是“归一化放在哪里”。设某一层的 Attention 或 MLP 子层为1. Post-Norm先计算残差再做归一化原始 Transformer 使用的 Post-Norm 可以写成它的计算顺序是输入 (x_l) 进入 Attention 或 MLP得到分支输出 (F(x_l))与残差主干 (x_l) 相加对相加后的结果进行归一化。如果暂时忽略 Norm那么残差结构本来具有一条非常直接的路径但在 Post-Norm 中残差相加之后还要再经过一次 Norm。因此从这一层输出向前反向传播时所有梯度路径都必须经过 Norm 的雅可比矩阵。它的局部梯度可以写成这里• (I) 是残差连接产生的恒等项• (J_F) 是子层 (F) 的雅可比矩阵• (J_{\operatorname{Norm}}) 是归一化操作的雅可比矩阵。关键点在于虽然残差连接产生了 (I)但整个结果仍然会被左侧的 (J_{\operatorname{Norm}}) 调制。当网络堆叠很多层时这些雅可比矩阵会连续相乘使梯度对初始化、学习率和 warmup 更加敏感。这并不意味着 Post-Norm 必然梯度消失也不意味着它无法训练深层模型。更准确地说Post-Norm 缺少一条完全绕开各层变换的恒等梯度通路因此深层训练通常更难稳定。2. Pre-Norm先归一化再进入计算分支Pre-Norm 的结构为它的计算顺序是对输入 (x_l) 做归一化将归一化结果送入 Attention 或 MLP将分支输出与未经归一化的 (x_l) 直接相加。此时局部梯度为与 Post-Norm 相比最大的区别是Pre-Norm 的恒等项 (I) 位于整个表达式最外层不再被 Norm 的雅可比矩阵包裹。也就是说即使计算分支中的梯度很小残差主干仍然保留这就是所谓的恒等梯度通路。3. 用一个简化数字理解差别为了直观理解可以暂时把某个梯度方向上的雅可比近似看成标量。假设那么 Pre-Norm 的局部梯度近似为而 Post-Norm 的局部梯度近似为这个例子不是对真实 Transformer 梯度的完整计算只是为了展示结构差异• Pre-Norm 保留独立的恒等项 1• Post-Norm 的整个梯度都会受到 Norm 的调制。当网络只有一两层时这种区别可能不明显但堆叠几十层后连续相乘的差异会逐渐累积。4. Pre-Norm 为什么更适合深层 Transformer我认为可以从两个方向理解。第一计算分支的输入尺度稳定。每次进入 Attention 或 MLP 之前输入都会先经过 Norm因此不管残差主干的数值如何累积计算分支看到的输入尺度始终相对稳定。第二残差主干保持通畅。原始隐藏状态可以沿残差连接逐层向前传递反向传播时也始终存在恒等梯度项。这让深层网络对初始化和学习率更加宽容。不过Pre-Norm 也不是没有代价。因为 Norm 只作用于计算分支残差主干本身会不断累加所以经过很多层后残差流的尺度仍可能逐渐增长。这也是 Pre-Norm 模型通常会在所有 Transformer Block 之后再增加一次最终归一化的原因然后才送入语言模型输出层因此Pre-Norm 的本质不是“每一层输出都被归一化”而是每个 Attention 和 MLP 分支的输入被归一化同时残差主干保持直接连通。【图片3Pre-Norm 与 Post-Norm 的残差及梯度通路对比】七、为什么现代大模型经常采用 Pre-RMSNorm现在可以把前面的两条主线合起来• RMSNorm 决定如何处理隐藏状态• Pre-Norm 决定在残差结构的哪个位置处理隐藏状态。对于一个典型的 Decoder Block首先执行 Attention然后执行 MLP对应的伪代码为u h attention(attention_norm(h))h u mlp(ffn_norm(u))这里必须注意MLP 接收的是 Attention 残差相加后的 (u_l)而不是再次使用原始的 (h_l)。经过全部 Decoder Block 后通常还会执行最终归一化h final_norm(h)logits lm_head(h)Pre-RMSNorm 的收益可以准确地概括为三点。1. 稳定计算分支的输入尺度RMSNorm 让 Attention 和 MLP 不必适应残差流不断变化的绝对尺度从而降低 (QK^\top) 和非线性激活发生剧烈波动的风险。2. 保留直接的残差梯度通路Pre-Norm 将 Norm 放在计算分支内部使残差主干不必经过归一化操作深层网络因此更容易优化。3. 降低归一化算子的常数开销RMSNorm 不计算均值也不执行中心化计算路径比 LayerNorm 更简单并且更容易与其他算子进行 Kernel 融合。不过我不会把 Pre-RMSNorm 描述成绝对最优方案。它不能单独保证模型一定可以稳定训练也不能代替合理的初始化、学习率、warmup、残差缩放和数值精度控制。RMSNorm 也不一定在所有任务中都优于 LayerNorm。更准确的说法是对新设计的深层 Decoder-only TransformerPre-RMSNorm 是一个经过充分验证、计算简单且训练稳定的默认起点但对于已有预训练模型不应脱离原始架构随意替换归一化方式。八、总结我最终如何区分这些概念对比项BatchNormLayerNormRMSNorm统计对象一批样本中的同类特征单个 Token 的隐藏维度单个 Token 的隐藏维度是否计算均值是是否尺度统计量方差方差均方根是否依赖 Batch是否否样本之间是否耦合是否否训练与推理规则不完全一致一致一致常见可学习参数(\gamma,\beta)(\gamma,\beta)通常只有 (\gamma)渐进复杂度(O(N))(O(N))(O(N))典型应用CNNTransformer现代大语言模型我最终形成的理解是• BN 通过一批样本共同确定统计尺度• LayerNorm 让每个 Token 独立完成中心化和尺度归一化• RMSNorm 只控制每个 Token 隐藏向量的整体尺度• Pre-Norm 和 Post-Norm 不决定如何归一化而是决定 Norm 在残差结构中的位置。现代大模型中常见的结构可以概括为其中• RMSNorm 负责稳定计算分支的输入尺度• Pre-Norm 负责保留残差主干的恒等梯度通路• 二者结合使深层 Transformer 同时获得相对稳定的前向数值和更加顺畅的反向传播。一句话总结BN、LayerNorm 和 RMSNorm 回答对谁归一化、怎样归一化Pre-Norm 和 Post-Norm 回答在哪里归一化。现代大模型经常采用 Pre-RMSNorm本质上是为了同时控制分支输入尺度、保护残差梯度通路并降低归一化的工程开销。学AI大模型的正确顺序千万不要搞错了2026年AI风口已来各行各业的AI渗透肉眼可见超多公司要么转型做AI相关产品要么高薪挖AI技术人才机遇直接摆在眼前有往AI方向发展或者本身有后端编程基础的朋友直接冲AI大模型应用开发转岗超合适就算暂时不打算转岗了解大模型、RAG、Prompt、Agent这些热门概念能上手做简单项目也绝对是求职加分王给大家整理了超全最新的AI大模型应用开发学习清单和资料手把手帮你快速入门学习路线:✅大模型基础认知—大模型核心原理、发展历程、主流模型GPT、文心一言等特点解析✅核心技术模块—RAG检索增强生成、Prompt工程实战、Agent智能体开发逻辑✅开发基础能力—Python进阶、API接口调用、大模型开发框架LangChain等实操✅应用场景开发—智能问答系统、企业知识库、AIGC内容生成工具、行业定制化大模型应用✅项目落地流程—需求拆解、技术选型、模型调优、测试上线、运维迭代✅面试求职冲刺—岗位JD解析、简历AI项目包装、高频面试题汇总、模拟面经以上6大模块看似清晰好上手实则每个部分都有扎实的核心内容需要吃透我把大模型的学习全流程已经整理好了抓住AI时代风口轻松解锁职业新可能希望大家都能把握机遇实现薪资/职业跃迁这份完整版的大模型 AI 学习资料已经上传CSDN朋友们如果需要可以微信扫描下方CSDN官方认证二维码免费领取【保证100%免费】