深度学习进阶文档深度页
Batch Normalization
按通道跨 batch/空间统计并维护运行均值方差。
学习目标
- 识别统计轴
- 管理 running stats
- 处理小 batch
为什么需要这个概念
中间激活尺度变化会让优化不稳定。
- CNN 常用
- train/eval 行为不同
学习前需要掌握
均值方差内容待补充train/eval内容待补充
背景与问题
中间激活尺度变化会让优化不稳定。
概念定义
BatchNorm 在训练用当前 batch 统计,评估用 running stats,再学习 γ/β。
直观理解
每个通道用当前训练群体校准,再允许模型重新缩放和平移。
核心原理
BatchNorm2d 对 [N,C,H,W] 每通道跨 N/H/W 统计;小 batch 统计噪声大。
冻结 γ/β 不等于冻结 running stats;每次 model.train 后可能需重新设 BN eval。
数学表达
核心公式
每个通道独立统计。
变量含义
γ,β可学习缩放和平移running_mean评估使用的 buffer
计算步骤
- 1每通道 32×16×16
代码对应位置
示例 1Batch Normalization · PyTorch 示例:验证Batch Normalization的输入、计算和输出。
变量解释
| 变量 | 含义 |
|---|---|
γ,β | 可学习缩放和平移 |
running_mean | 评估使用的 buffer |
完整数值示例
BN2d 统计数
已知条件
- N=32,H=W=16
- 1
每通道 32×16×16
8192 个值
处理前后对比
计算与实现步骤
- 1
按通道统计
- 2
归一化
- 3
仿射变换
- 4
更新 running
- 5
eval 使用 running
代码实现
示例 1
Batch Normalization · PyTorch 示例
Batch Normalization · PyTorch 示例
PyTorch
example_01.py用途:验证Batch Normalization的输入、计算和输出。
bn=nn.BatchNorm2d(64)
x=torch.randn(32,64,16,16)
print(bn(x).shape,bn.weight.shape)代码解析
解析始终位于完整代码下方,并按实际代码段逐项对应。
输入数据与任务
验证Batch Normalization的输入、计算和输出。
Step 1 · 1–3 行
输出中间参数、形状或最终指标,用于核对代码是否符合预期。
bn=nn.BatchNorm2d(64)
x=torch.randn(32,64,16,16)
print(bn(x).shape,bn.weight.shape)预期输出或运行结果
[32,64,16,16] 与 weight [64]。
常见错误 · 5 条
- 验证忘 eval
- num_features 填 H
- 梯度累积当大 BN batch
- 只冻参数
- 确认样本轴、特征轴以及 X 与 y 的第一维完全对应。
实际应用
- CNN
- 稳定 batch MLP
常见错误
验证忘 eval
num_features 填 H
梯度累积当大 BN batch
只冻参数
文档来源
14_归一化方法_BatchNorm与LayerNorm原始文档
学习资料/deeplearning/14_归一化方法_BatchNorm与LayerNorm.md- §2 通用公式
- §3–7 BatchNorm
- §8 冻结
- §9 CNN 位置
常见问题
no_grad 会切换 BN 吗?
不会。
N=1 一定可用吗?
取决于空间统计与任务,仍可能不稳定。
推荐学习资料
暂未收录相关资料
参考库不会生成虚假资源或无效外部链接。