同一个张量究竟该沿哪条轴标准化?BatchNorm 与 LayerNorm 的训练/推理差异
从批量大小变化引发的预测漂移出发,手算 BatchNorm 与 LayerNorm 的统计轴,拆解运行统计量、仿射参数和 PyTorch 2.13 训练/推理语义。
上一篇用 Momentum 与 Adam 重塑参数更新,但优化器只能处理反向传来的梯度。随着权重不断变化,中间激活的尺度仍可能漂移:同一学习率在不同层代表不同大小的函数变化,深层网络也更容易进入激活饱和或数值不稳定区域。
归一化层(Normalization Layer)试图在网络内部重新控制张量的尺度。真正困难的不是背下“减均值、除标准差”,而是回答:均值和方差究竟沿哪些轴计算?训练和推理时又使用谁的统计量? 本文只比较批归一化(Batch Normalization,BatchNorm)与层归一化(Layer Normalization,LayerNorm),把一个形状为 [N, C, L] 的张量逐轴拆开。
01 为什么“把整个张量标准化”会出错?#
设一批序列特征 x 的形状为 [N,C,L]:
- :batch 中的样本数;
- :通道或特征数;
- :每个样本的序列位置数。
若直接对全部 个数求一个均值和方差,通道 0 与通道 1 会互相改变尺度;一个样本也会影响另一个样本。这样虽然能让全局均值接近 0,却可能破坏“每个通道代表不同特征”的语义。
BatchNorm 与 LayerNorm 使用同一基础公式:
差别全在集合 :为了计算位置 的统计量,哪些元素被分到同一组? 是防止分母过小的数值稳定项; 是可学习的缩放和平移,使网络在需要时能够恢复或重塑归一化后的尺度。
02 BatchNorm 沿哪些轴计算?#
对 BatchNorm1d(C) 的三维输入 ,每个通道 单独计算:
也就是保留 轴,沿 轴聚合。 的形状都是 [C],通过广播作用到 [N,C,L]。
BatchNorm1d:同一通道跨样本、跨位置统计
X [N,C,L]
│
├─ channel 0: X[:,0,:] ─► mean₀,var₀ ─► γ₀,β₀
├─ channel 1: X[:,1,:] ─► mean₁,var₁ ─► γ₁,β₁
└─ channel c: X[:,c,:] ─► mean_c,var_c ─► γ_c,β_c
归约轴:(N,L) 保留轴:C 输出:Y [N,C,L]
一个样本的输出会受同一 batch 中其他样本影响。text对图像 BatchNorm2d(C) 的输入 [N,C,H,W],逻辑相同:保留通道 ,沿 (N,H,W) 统计。这里的 “1d/2d” 描述输入的空间结构,不是“只对一维向量求均值”。
03 LayerNorm 沿哪些轴计算?#
对序列模型常见输入 ,使用 LayerNorm(D) 时,每个样本、每个 token 都沿最后一个特征轴独立计算:
的可广播形状是 [N,L,1]; 的形状是 [D]。统计量不跨 batch,也不跨 token,所以改变其他样本不会改变当前 token 的输出。
LayerNorm(D):每个 token 在自己的 D 个特征内统计
X [N,L,D]
│
├─ X[0,0,:] ─► mean₀₀,var₀₀ ─► 同一组 γ[:],β[:]
├─ X[0,1,:] ─► mean₀₁,var₀₁ ─► 同一组 γ[:],β[:]
└─ X[n,l,:] ─► mean_nl,var_nl ─► 同一组 γ[:],β[:]
归约轴:最后的 D 保留轴:(N,L) 输出:Y [N,L,D]
每个 token 的统计量只由自身特征决定。textLayerNorm((C,H,W)) 则沿输入最后三个维度一起统计,且可学习参数形状也是 [C,H,W]。因此不能只说“LayerNorm 沿通道归一化”;必须同时写出 normalized_shape 和输入布局。
04 用四个数手算:同一输入为何得到两种答案?#
先忽略 ,令 。考虑两样本、两特征矩阵:
BatchNorm:逐列跨样本。 第一列 [1,5] 的均值为 3、方差为 4;第二列 [3,7] 的均值为 5、方差也为 4。因此:
LayerNorm:逐行跨特征。 第一行 [1,3] 的均值为 2、方差为 1;第二行 [5,7] 的均值为 6、方差为 1。因此:
现在把第二个样本改成 [105,107]。第一个样本的 LayerNorm 输出仍是 [-1,1];BatchNorm 的列统计量却被新样本改变,第一个样本的输出也随之变化。这就是 BatchNorm 的批间耦合(Batch Coupling)。
05 BatchNorm 为什么需要两套数据流?#
训练时,BatchNorm 用当前 mini-batch 的 归一化,同时更新运行均值(Running Mean)与运行方差(Running Variance):
其中 PyTorch 参数 momentum=m 默认是 0.1。它与上一篇优化器的 Momentum 定义不同:这里 越大,新 batch 的权重越高。
推理时,BatchNorm 默认不再依赖当前 batch,而使用训练期积累的 running_mean 和 running_var。这样单样本推理才不会因为“恰好和谁同批”而改变答案。
┌─ batch mean/var ─► 当前训练输出
train(): X [N,C,...] ────┤
└─ 更新 running_mean/running_var(buffer)
eval(): X [N,C,...] ───── running_mean/running_var ─► 推理输出
γ, β:Parameter,反向传播更新
running_mean, running_var:Buffer,前向时更新,不接收梯度textmodel.eval() 切换的是模块行为;torch.no_grad() 关闭的是梯度记录。二者不是同一件事:验证时通常两者都需要,少任何一个都可能造成错误或浪费。
若 track_running_stats=False,BatchNorm 不保存运行统计量,训练和评估都使用当前 batch 统计量。此时 eval() 也无法让输出脱离批组成,必须明确接受这一语义。
06 LayerNorm 为什么不区分训练和推理统计量?#
LayerNorm 每次都从当前样本的指定末尾维度计算统计量,不需要跨 batch 积累运行均值或方差。于是同一输入在 train() 与 eval() 下,单独看 LayerNorm 模块会使用相同的数据流。
这让 LayerNorm 适合批量大小变化大、逐样本推理以及序列长度动态变化的模型。但“没有运行统计量”不等于“没有可学习参数”:默认的逐元素 仍通过反向传播训练。
| 比较维度 | BatchNorm | LayerNorm |
|---|---|---|
| 典型输入 | CNN 的 [N,C,H,W] | Transformer 的 [N,L,D] |
| 统计轴 | 每通道沿 (N,H,W) | 每个 token 沿最后的 D |
| 统计量是否跨样本 | 是 | 否 |
| 训练/推理统计 | 当前 batch / running stats | 都来自当前输入 |
| 仿射参数 | 通常每通道 [C] | 每个归一化元素 [D] 或 normalized_shape |
| 小 batch 风险 | 统计噪声与训练/推理偏差 | 基本不受 batch 大小影响 |
| 主要代价 | 依赖批组成与状态同步 | 会消去每个样本归一化轴上的整体尺度信息 |
07 不调用归一化层,先写出 NumPy 本体#
下面让 axis 明确表达统计轴,并保留维度以便广播:
import numpy as np
def normalize(x, *, axis, gamma, beta, eps=1e-5):
"""x、gamma、beta 均可广播;axis 是被归约的轴。"""
mean = x.mean(axis=axis, keepdims=True)
variance = x.var(axis=axis, keepdims=True, ddof=0)
x_hat = (x - mean) / np.sqrt(variance + eps)
return gamma * x_hat + beta, mean, variance
x = np.array([[1.0, 3.0], [5.0, 7.0]]) # [N=2,D=2]
# BatchNorm:沿 N 统计,gamma/beta 对应 D
bn, bn_mean, bn_var = normalize(
x,
axis=(0,),
gamma=np.ones((1, 2)),
beta=np.zeros((1, 2)),
)
# LayerNorm:沿 D 统计,gamma/beta 仍对应 D
ln, ln_mean, ln_var = normalize(
x,
axis=(-1,),
gamma=np.ones((1, 2)),
beta=np.zeros((1, 2)),
)
assert bn.shape == ln.shape == x.shape
np.testing.assert_allclose(bn_mean, [[3.0, 5.0]])
np.testing.assert_allclose(ln_mean, [[2.0], [6.0]])python这个函数只复现单次前向,没有 BatchNorm 的运行统计量、可学习参数、反向传播和分布式同步。它的价值是让 axis、keepdims 和广播关系可以直接检查,而不是替代框架实现。
08 完整训练与推理伪代码#
训练 BatchNorm:
设置 model.train()
对每个 mini-batch X:
沿“除通道外”的指定轴计算 batch mean/variance
用 batch 统计量归一化并做 γ、β 仿射变换
用 momentum 更新 running mean/variance
完成后续前向、loss、backward、optimizer.step
验证 BatchNorm:
设置 model.eval()
进入 no_grad 上下文
用冻结的 running mean/variance 归一化
不更新 running stats,不更新参数
LayerNorm:
train/eval 都沿 normalized_shape 对应的末尾轴计算当前输入统计量
只有 γ、β 随训练更新;没有 running statstext09 用 PyTorch 2.13 正确落地#
PyTorch 2.13 官方 BatchNorm1d ↗ 接收 [N,C] 或 [N,C,L],num_features 必须等于 。官方 LayerNorm ↗ 则对 normalized_shape 指定的最后 个维度求统计量。
import torch
from torch import nn
# 卷积/通道布局:[N,C,L]
conv_block = nn.Sequential(
nn.Conv1d(in_channels=8, out_channels=16, kernel_size=3, padding=1),
nn.BatchNorm1d(num_features=16, eps=1e-5, momentum=0.1),
nn.ReLU(),
)
channels_first = torch.randn(32, 8, 50) # [N=32,C=8,L=50]
conv_output = conv_block(channels_first) # [32,16,50]
# 序列/特征布局:[N,L,D]
token_block = nn.Sequential(
nn.Linear(64, 64),
nn.LayerNorm(normalized_shape=64, eps=1e-5),
nn.GELU(),
)
tokens = torch.randn(32, 50, 64) # [N=32,L=50,D=64]
token_output = token_block(tokens) # [32,50,64]pythonLayerNorm(64) 只检查并归一化最后一维。若输入误写成 [N,D,L],最后一维是 ,要么尺寸不匹配直接报错,要么尺寸碰巧相同而静默归一化错轴。进入层前用断言固定布局:
assert tokens.ndim == 3 and tokens.shape[-1] == 64
token_output = token_block(tokens)python官方 API 中,BatchNorm1d(..., affine=True, bias=True) 默认学习每通道缩放和偏置;LayerNorm(..., elementwise_affine=True, bias=True) 默认学习 normalized_shape 大小的逐元素缩放和偏置。bias=False 只关闭加性偏置,不等于关闭全部仿射参数。
10 怎样验证训练/推理切换没有写错?#
用固定样本做三个检查:同批重复、换同伴、切模式。
import torch
from torch import nn
torch.manual_seed(7)
bn = nn.BatchNorm1d(2)
anchor = torch.tensor([[1.0, 3.0]])
near_partner = torch.tensor([[5.0, 7.0]])
far_partner = torch.tensor([[105.0, 107.0]])
bn.train()
with torch.no_grad():
train_near = bn(torch.cat([anchor, near_partner]))[0]
train_far = bn(torch.cat([anchor, far_partner]))[0]
assert not torch.allclose(train_near, train_far)
bn.eval()
with torch.no_grad():
eval_single = bn(anchor)
eval_with_partner = bn(torch.cat([anchor, far_partner]))[:1]
torch.testing.assert_close(eval_single, eval_with_partner)python训练阶段两次调用还会分别更新 running stats,所以这个例子用于验证“批组成会影响当前训练输出”,不是比较完全相同的模块状态。要做严格 A/B,应复制同一 state_dict 到两个模块后各前向一次。
部署前还应检查:
model.training与每个归一化子模块的.training是否符合预期;running_mean、running_var、num_batches_tracked是否有限且已更新;- 训练与部署的通道布局、预处理和输入尺度是否一致;
- 校准数据跑过后的指标是否比随机初始化 running stats 更合理;
- 分布式训练中,各设备局部 batch 是否小到需要
SyncBatchNorm或其他方案。
11 最常见的错误与最短调试路径#
- 忘记
model.eval()。 单样本服务仍使用当前 batch 统计量,输出随请求拼批方式漂移。 - 把
no_grad()当成eval()。 梯度虽不记录,BatchNorm 仍可能更新 running stats,Dropout 也仍随机。 - 轴与内存布局混淆。
BatchNorm1d(C)期待通道在第 2 维;LayerNorm(D)期待被归一化维度在末尾。 - batch 太小。 BatchNorm 的估计噪声大;
N=1且每通道只有一个值时,训练统计甚至无法成立。 - 把 BN 的
momentum当优化器动量。 其更新式权重方向相反;调参前先写出运行统计更新公式。 - 只加载权重、不加载 buffer。 不完整 checkpoint 会丢失 BatchNorm 的运行统计量;应加载完整
state_dict。 - 冻结参数却仍污染统计量。
requires_grad=False不会阻止 BN buffer 在训练模式更新;需单独管理模块模式。 - 在梯度累积中误判有效 batch。 BN 每次前向只看 micro-batch,不会因为累积多步梯度就自动得到大 batch 统计量。
最短路径:打印输入形状与归约轴 → 在归一化前后打印逐组均值/方差 → 固定 anchor 更换 batch 同伴 → 分别跑 train()/eval() → 检查 state_dict 中参数与 buffer → 最后再比较端到端指标。
12 它们各自会在哪里失败?#
BatchNorm 在小 batch、非独立同分布 batch、在线学习和单样本自回归推理中容易产生统计不稳或训练/部署错位。领域分布改变时,旧 running stats 也可能不再代表线上数据。
LayerNorm 不依赖 batch,但会消去每个样本在归一化轴上的共同平移与整体尺度;如果任务恰好需要这些绝对幅度信息,模型必须从旁路或其他特征重新获得。它也不是数值问题的万能修复:输入已有 NaN/Inf 时,增大 eps 通常只会掩盖而非解决根因。
邻近方法不能仅凭名字互换:
- 实例归一化(Instance Normalization)通常对每个样本、每通道沿空间轴统计,常见于风格迁移;
- 组归一化(Group Normalization)把通道分组,在每个样本内沿组内通道与空间轴统计,适合小 batch 视觉任务;
- RMSNorm 只按均方根缩放,通常不减均值,也没有 BatchNorm 的运行统计量;
SyncBatchNorm跨分布式进程同步 batch 统计量,能扩大统计样本,但增加通信成本。
选择归一化方法的第一步永远是写清楚输入布局、归约轴和部署时可用的信息,而不是先从模型名字推断。
13 今天真正需要记住什么?#
- 归一化的核心选择是统计集合:BatchNorm 保留通道、跨 batch 聚合;LayerNorm 保留样本位置、沿末尾特征聚合。
- BatchNorm 训练时使用 batch 统计并更新 buffer,推理时默认使用 running stats;LayerNorm 两种模式都从当前输入计算。
- 是可训练参数,running mean/variance 是 buffer;它们必须一起进入完整 checkpoint。
- 写出
[N,C,L]或[N,L,D]及每个中间张量形状,是发现错轴、错布局和广播错误最快的方法。 - BatchNorm 与 LayerNorm 改变了优化几何和网络可表达方式,但都不能替代正确初始化、学习率、数据处理与数值诊断。
14 思考题与小练习#
- 对 ,忽略 且令 ,分别手算 BatchNorm 与 LayerNorm 输出。将第三个样本改为
[40,60]后,哪些输出会变化? - 给定输入
[N=4,C=3,H=2,W=2],分别写出BatchNorm2d(3)与LayerNorm((3,2,2))的统计量形状、仿射参数形状和归约轴,并计算各自每个统计组包含多少个数。 - 训练一个含 BatchNorm 的微型网络,分别只保存
named_parameters()与保存完整state_dict();恢复后用同一输入比较eval()输出,并定位差异来自哪个 buffer。
相关工作#
- Ioffe & Szegedy (2015), Batch Normalization ↗:提出用 mini-batch 统计量归一化中间激活及运行统计推理。
- Ba, Kiros & Hinton (2016), Layer Normalization ↗:改为在单个样本内部沿特征统计,摆脱对 batch 的依赖。
- Ulyanov, Vedaldi & Lempitsky (2016), Instance Normalization ↗:展示逐实例、逐通道归一化对快速风格化的作用。
- Wu & He (2018), Group Normalization ↗:以通道分组替代 batch 统计,改善小批量视觉训练。
- Zhang & Sennrich (2019), Root Mean Square Layer Normalization ↗:讨论省略重中心化的 RMSNorm。
15 下一篇预告#
归一化能控制每个块内部的尺度,却没有解决“网络越深,信息必须穿过越多非线性变换”的路径问题。下一篇将研究残差连接如何建立恒等捷径,手算前向叠加与反向梯度分流,并比较普通残差块和预归一化结构。