Group Normalization 分组归一化:从数学定义到 PyTorch 实现的完整剖析
【免费下载链接】annotated_deep_learning_paper_implementations🧑🏫 60+ Implementations/tutorials of deep learning papers with side-by-side notes 📝; including transformers (original, xl, switch, feedback, vit, ...), optimizers (adam, adabelief, sophia, ...), gans(cyclegan, stylegan2, ...), 🎮 reinforcement learning (ppo, dqn), capsnet, distillation, ... 🧠项目地址: https://gitcode.com/gh_mirrors/an/annotated_deep_learning_paper_implementations
本文以labml_nn/normalization/group_norm/模块为主体,完整讲解 Group Normalization(分组归一化)的数学定义、四种归一化层的统一形式、仓库中手写GroupNorm层的逐步实现,以及配套的 CIFAR-10 训练实验。读完后,你将理解"为什么 BatchNorm 在小 batch 下失效、GroupNorm 如何按通道分组消除对 batch 大小的依赖",并能独立编写与调参一个使用 GroupNorm 的卷积分类网络。
1. 背景:为什么需要 Group Normalization
Group Norm 文档开篇给出了引入动机:Batch Normalization 在 batch 足够大时表现良好,但它是对 batch 维度做归一化,因此对很小的 batch size 效果很差;而受设备显存限制,大模型往往无法使用大 batch 训练。
Group Normalization 论文的核心思想是:把特征通道分成若干组,对每个组内的所有特征一起做归一化。这一设计借鉴了 SIFT、HOG 等传统视觉特征"分组式"的组织方式。由于统计量只取自单个样本内的通道与空间维度,GroupNorm 完全不依赖 batch 中其他样本,这正是它对小 batch 训练、推理一致性友好(训练与推理行为一致,无需维护 running statistics)的根本原因。
在仓库的归一化层总览中,Group Normalization 与 Batch Norm、Layer Norm、Instance Norm、Weight Standardization、Batch-Channel Norm、DeepNorm 并列,是labml_nn.normalization包的核心实现之一。
2. 数学定义:所有归一化层的统一形式
模块文档中首先给出了归一化层的统一计算形式:
$$\hat{x}_i = \frac{1}{\sigma_i}(x_i - \mu_i)$$
其中 $x$ 是表示整个 batch 的张量,$i$ 是单个数值的索引。以 2D 图像为例,$i = (i_N, i_C, i_H, i_W)$ 分别是 batch 内图像索引、特征通道索引、垂直坐标和水平坐标。均值与方差在索引集合 $\mathcal{S}_i$ 上计算:
$$\mu_i = \frac{1}{m}\sum_{k \in \mathcal{S}i} x_k, \qquad \sigma_i = \sqrt{\frac{1}{m}\sum{k \in \mathcal{S}_i}(x_k - \mu_i)^2 + \epsilon}$$
$\mathcal{S}_i$ 就是"与索引 $i$ 一起参与统计的那些索引的集合",$m = |\mathcal{S}_i|$ 对所有 $i$ 相同。不同归一化层的区别仅在于 $\mathcal{S}_i$ 的定义:
- Batch Normalization:$\mathcal{S}_i = {k \mid k_C = i_C}$,所有共享同一特征通道的值一起归一化(跨 batch、跨空间位置);
- Layer Normalization:$\mathcal{S}_i = {k \mid k_N = i_N}$,同一 batch 内同一样本的所有值一起归一化;
- Instance Normalization:$\mathcal{S}_i = {k \mid k_N = i_N, k_C = i_C}$,同一样本、同一通道内一起归一化;
- Group Normalization:
$$\mathcal{S}_i = \left{k ;\middle|; k_N = i_N,; \left\lfloor \frac{k_C}{C/G} \right\rfloor = \left\lfloor \frac{i_C}{C/G} \right\rfloor \right}$$
其中 $G$ 是分组数,$C$ 是通道数。即:同一样本内、同一组通道内的所有值一起归一化。每个组包含 $C/G$ 个相邻通道,组的划分由通道号除以每组的通道数 $C/G$ 取整得到。
直观对比四种归一化在 $(N, C, H, W)$ 张量上的统计范围:BN 沿 $N$、$H$、$W$(固定 $C$),LN 沿 $C$、$H$、$W$(固定 $N$),IN 沿 $H$、$W$(固定 $N$、$C$),GN 沿 $C_{\text{group}}$、$H$、$W$(固定 $N$ 与组号)。GN 介于 LN 与 IN 之间,且 $G=1$ 时退化为 LN,$G=C$ 时退化为 IN——这一"分组数插值"的性质从上述公式可以直接推出。
3. PyTorch 实现逐行剖析
仓库在 GroupNorm 实现 中手写了一个nn.Module,下面按构造函数与 forward 两段拆解。
3.1 构造函数:参数与约束
class GroupNorm(nn.Module): def __init__(self, groups: int, channels: int, *, eps: float = 1e-5, affine: bool = True): super().__init__() assert channels % groups == 0, \ "Number of channels should be evenly divisible by the number of groups" self.groups = groups self.channels = channels self.eps = eps self.affine = affine if self.affine: self.scale = nn.Parameter(torch.ones(channels)) self.shift = nn.Parameter(torch.zeros(channels))关键设计点(见init.py#L94-L113):
| 参数 | 说明 | 约束/默认值 |
|---|---|---|
groups | 通道被划分成的组数 $G$ | 必须整除channels,否则assert报错 |
channels | 输入通道数 $C$ | 必须与 forward 中x.shape[1]一致 |
eps | 数值稳定项 $\epsilon$,进入 $\sqrt{\mathrm{Var} + \epsilon}$ | 默认1e-5 |
affine | 是否对归一化结果做仿射缩放/平移($\gamma$、$\beta$) | 默认True;为True时创建scale(初始全 1,形状[channels])与shift(初始全 0)两个可学习参数 |
注意仿射参数是逐通道的(每通道一个 $\gamma_{i_C}$、$\beta_{i_C}$),而非逐组的——这是与数学公式 $y_{i_C} = \gamma_{i_C}\hat{x}{i_C} + \beta{i_C}$ 一致的细节。
3.2 forward:reshape → 统计 → 归一化 → 仿射 → 还原
def forward(self, x: torch.Tensor): # x 的形状是 [batch_size, channels, *],例如卷积特征图 # [batch_size, channels, height, width] x_shape = x.shape # 保留原始形状,最后要还原 batch_size = x_shape[0] assert self.channels == x.shape[1] # 关键一步:重排为 [batch_size, groups, -1] x = x.view(batch_size, self.groups, -1) # 沿最后一维(组内所有通道+空间位置)计算组均值与平方的均值 mean = x.mean(dim=[-1], keepdim=True) mean_x2 = (x ** 2).mean(dim=[-1], keepdim=True) # Var[x] = E[x^2] - E[x]^2 var = mean_x2 - mean ** 2 # 归一化:(x - E[x]) / sqrt(Var[x] + eps) x_norm = (x - mean) / torch.sqrt(var + self.eps) # 逐通道仿射缩放/平移 if self.affine: x_norm = x_norm.view(batch_size, self.channels, -1) x_norm = self.scale.view(1, -1, 1) * x_norm + self.shift.view(1, -1, 1) # 还原原始形状后返回 return x_norm.view(x_shape)几个值得注意的实现细节(见init.py#L115-L154):
- reshape 是整个技巧的核心。输入张量形状为
[batch_size, channels, *](*可以是任意个维度,2D 卷积时即[N, C, H, W])。由于 PyTorch 中通道维是第 2 维,且组由连续的 $C/G$ 个通道构成,直接view(batch_size, groups, -1)就能把张量切成[N, G, (C/G) * H * W]——最后一维恰好是公式中 $\mathcal{S}_i$ 所定义的那组"同一样本、同一组通道"的所有值。因此一次mean(dim=[-1])就完成了第 2 节公式中的 $\mu_i$ 计算,且keepdim=True保证形状可广播回原张量。 - 方差用 E[x²] − E[x]² 计算,即
var = mean_x2 - mean ** 2,与公式 $\mathrm{Var}[x] = \mathbb{E}[x^2] - \mathbb{E}[x]^2$ 一致,只需两次沿同一维的 reduce,比"先减均值再平方"少一次减法。 - 输入是任意维度后缀:由于 reshape 用
-1吃掉剩余维度,同一实现可用于[N, C]、[N, C, H, W]、[N, C, D]等张量,对 1D/2D/3D 卷积乃至纯全连接输入都成立。 - 训练与推理行为完全一致:实现中没有任何
running_mean/running_var,也未按self.training分支——这正是 GroupNorm 相对 BatchNorm 的结构优势。 - 仿射参数通过
.view(1, -1, 1)广播为[1, C, 1],作用在[N, C, -1]张量上,实现逐通道缩放/平移。
文件末尾的_test()提供了一个最小自测:输入形状[2, 6, 2, 4](batch=2、channels=6、2×4 空间),GroupNorm(2, 6)分成 2 组、每组 3 通道,验证输出形状不变。可以将其作为接入自己项目时的最小冒烟测试模板。
4. CIFAR-10 实验:把 GroupNorm 放进 VGG 风格网络
文档配套的实验在 experiment.py(对应 experiment.ipynb 可交互版本),目标是用 GroupNorm 训练一个 CIFAR-10 图像分类卷积网络。
4.1 模型结构:VGG 风格 + GroupNorm 卷积块
实验基于 CIFAR10VGGModel 这个通用 VGG 风格架构,子类只需覆写conv_block就能替换归一化策略:
class Model(CIFAR10VGGModel): def conv_block(self, in_channels, out_channels) -> nn.Module: return nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1), fnorm.GroupNorm(self.groups, out_channels), # 用 GroupNorm 替换默认归一化 nn.ReLU(inplace=True), ) def __init__(self, groups: int = 32): self.groups = groups super().__init__([[64, 64], [128, 128], [256, 256, 256], [512, 512, 512], [512, 512, 512]])(fnorm是同目录下 group_norm 模块的别名,独立复现时可写为from labml_nn.normalization import group_norm as fnorm。)
父类CIFAR10VGGModel的行为(见 cifar10.py#L68-L111):5 个卷积块,每块内含若干"3×3 卷积(padding=1,保持空间尺寸)+ 归一化 + ReLU"层,块尾接一个MaxPool2d(2, 2);5 次下采样把 32×32 的 CIFAR-10 图像压到 1×1,随后展平接一个 10 类线性层。
一个容易忽略的通道-分组整除约束:由于GroupNorm构造时有channels % groups == 0断言,各卷积输出通道数为 64/128/256/512,所以groups必须同时整除这些数,可选值为 1、2、4、8、16、32。实验的超参取groups=16(见下文Configs),而Model的默认参数是 32,两者均在合法范围内。
4.2 训练配置与运行
实验的配置类与入口(见 experiment.py#L38-L68):
class Configs(CIFAR10Configs): # 分组数 groups: int = 16 @option(Configs.model) def model(c: Configs): """创建模型""" return Model(c.groups).to(c.device) def main(): experiment.create(name='cifar10', comment='group norm') conf = Configs() experiment.configs(conf, { 'optimizer.optimizer': 'Adam', 'optimizer.learning_rate': 2.5e-4, }) with experiment.start(): conf.run()要点:
- 实验使用
labml框架(experiment.create/experiment.configs/experiment.start)组织,configs覆盖项把优化器指定为Adam,学习率 2.5e-4;groups默认16。 - 数据集侧继承自 CIFAR10Configs:训练集做
RandomCrop(32, padding=4)+RandomHorizontalFlip增强,归一化均值/标准差均为0.5;验证集不增强。数据加载器参数继承自 datasets.py 的CIFAR10Configs:训练 batch size 默认 64、验证 1024、训练集默认 shuffle。 - 训练循环继承自 MNISTConfigs(
CIFAR10Configs同时继承数据集配置与MNISTConfigs的训练骨架):交叉熵损失、Accuracy指标、默认 10 个 epoch,每个 epoch 内按inner_iterations=10交替做训练与验证;step中依次执行前向、算损失、backward、优化器步进,并在每个 epoch 最后一个 batch 记录模型参数与梯度。
运行方式:按仓库 readme 的方式安装labml相关依赖后,直接执行python -m labml_nn.normalization.group_norm.experiment,或用 Jupyter 打开 experiment.ipynb 逐步运行;首次运行会自动下载 CIFAR-10 数据集(download=True)。
5. 仓库内 GroupNorm 的其他落点
从源码结构看,GroupNorm 在本仓库中并不止服务于分类实验——图像生成方向的扩散模型实现里大量使用了 PyTorch 内置的nn.GroupNorm,这印证了 GN"小 batch 下依然稳定"的定位(扩散模型训练常以小 batch 进行):
- DDPM 的 UNet:残差块内
nn.GroupNorm(n_groups, in_channels)与nn.GroupNorm(n_groups, out_channels)分别放在两个卷积之后; - Stable Diffusion 的 AutoEncoder 与 UNet 注意力模块:均使用
nn.GroupNorm(num_groups=32, num_channels=channels, eps=1e-6)。
对比可以发现一个实用细节:仓库手写实现默认eps=1e-5,而扩散模型中的用法取eps=1e-6、分组数普遍取32。从源码结构看,32 组是视觉生成模型的常见选择——分组越细,每组的统计量越接近 InstanceNorm(保留组间信息),分组越粗则越接近 LayerNorm,可按任务在 1/8/16/32 之间权衡,唯一硬性约束是"组数整除通道数"。
6. 小结与适用建议
- 什么时候选 GroupNorm:batch size 很小(显存受限的大模型、生成模型训练)、需要训练/推理行为严格一致(不维护 running statistics)、以及扩散模型 U-Net / 自编码器等场景;
- 与相邻归一化层的关系:BN 跨 batch 统计(大 batch 首选)、LN 跨整个样本统计(NLP 常用)、IN 单通道统计(风格化、图像合成)、GN 按通道分组统计(小 batch 视觉任务),四者共享第 2 节的统一公式,差异只在 $\mathcal{S}_i$ 的定义;
- 实现要点回顾:
[N, C, *] → [N, G, (C/G)*H*W]的一次view完成分组;方差用E[x²] − E[x]²;逐通道可学习scale/shift;eps防止除零;channels % groups == 0是必须满足的约束。
参考仓库文件:文档、实现、实验、实验 Notebook、CIFAR-10 训练器、数据集配置、训练骨架。
【免费下载链接】annotated_deep_learning_paper_implementations🧑🏫 60+ Implementations/tutorials of deep learning papers with side-by-side notes 📝; including transformers (original, xl, switch, feedback, vit, ...), optimizers (adam, adabelief, sophia, ...), gans(cyclegan, stylegan2, ...), 🎮 reinforcement learning (ppo, dqn), capsnet, distillation, ... 🧠项目地址: https://gitcode.com/gh_mirrors/an/annotated_deep_learning_paper_implementations
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考