乐于分享
好东西不私藏

AI 视觉模型导出 ONNX 后,BatchNorm 为什么不见了?

AI 视觉模型导出 ONNX 后,BatchNorm 为什么不见了?
·

AI INFRA · CHAPTER 6

AI 视觉模型导出 ONNX 后,BatchNorm 为什么不见了?

本章会把两种归一化算子放在一起认识:BatchNormalization 和 LayerNormalization。它们都会整理输入数值,但统计方向和常见应用不同。由于本章使用 AI 视觉模型作为应用场景,后续代码和 ONNX 实验以 BatchNorm 为主;LayerNorm 只用于建立对比,不重复展开完整计算过程。

SECTION 01

先回顾旧算子:Conv、Relu 和 Add

前面几章已经介绍过本章会再次出现的三个算子:

算子
已经认识的作用
在本章中的位置
Conv
用卷积核读取局部区域,生成新的特征图
残差块的主支路
Relu
把负数变成 0,保留正数
Conv 后和残差相加后
Add
把两个 Shape 相同的张量逐位置相加
主支路与直连支路的汇合点

本章不会重复这三个算子的内部计算,而是观察新加入的归一化算子怎样连接在 Conv 后面,以及导出 ONNX 后节点怎样变化。

ResNet 常被用作图像分类、目标检测和图像分割模型的骨干网络。是最常见的视觉模型的基础结构。它可以连续堆叠很多卷积层,其中最有辨识度的结构就是残差连接。

残差连接本身很简单:输入 x 一路进入卷积支路,另一路直接送到后面,两路在 Add 处相加。

输出 = Relu(F(x) + x)

F(x) 负责学习这一次需要增加或修改的内容,直连支路负责保留输入。本章不展开梯度推导,只用两个最小残差块观察一个更直接的问题:PyTorch 代码里明明有 BatchNorm,为什么导出的 ONNX 图中有时看得见,有时却不见了?

图 6-1(残差结构):输入 x 同时进入卷积支路和直连支路,最后在 Add 节点汇合。

SECTION 02

本章的新重点:BatchNorm 与 LayerNorm

BatchNorm 是 Batch Normalization(批归一化)的简称,常接在卷积层后面。它会把同一通道中偏大、偏小的数值重新整理到较稳定的范围,再用可学习的缩放和偏移参数调整结果。

训练时,BatchNorm 根据当前一批样本计算均值和方差,同时记录供推理使用的统计量。推理时,它通常不再统计当前输入,而是直接使用训练阶段保存下来的均值和方差,按照固定规则调整每个通道。

LayerNorm 是 Layer Normalization(层归一化)的简称,常用于 Transformer 等以特征向量为主的模型。它不依赖同一批中的其他样本,而是在单个样本自己的特征内部计算均值和方差。

两者最关键的区别是“沿哪个方向统计”:

对比项
BatchNorm / BatchNormalization
LayerNorm / LayerNormalization
统计范围
一批样本中同一通道的数值
单个样本自己的特征
常见场景
Conv 后的视觉特征图
Transformer 中的特征向量
本章安排
TinyResNet、ONNX 导出与折叠实验的主角
用于对比两类归一化算子

前面的“从零学习大模型 Transformer”系列第 40 章《LayerNorm:把每个向量拉回稳定的分布》已经讲过 LayerNorm 的计算过程,以及它和 BatchNorm 的统计方向差异。本章只保留这段必要对比,后续重点观察 BatchNorm 在视觉模型中的应用和 ONNX 折叠现象。

SECTION 03

用两个 BasicBlock 搭建 TinyResNet

由于本章的应用场景是视觉模型,下面使用 BatchNorm 构建最简单的 BasicBlock;LayerNorm 不进入这个 TinyResNet。主支路依次经过:

Conv -> BatchNormalization -> Relu -> Conv -> BatchNormalization

输入沿另一条路直接到达 Add,相加后再经过 Relu

class BasicBlock(torch.nn.Module):    def __init__(self, channels: int):        super().__init__()        self.conv1 = 

              torch.nn.Conv2d(                channels,

                channels,

                kernel_size=3, 

                padding=1,

                bias=False        )        self.bn1 = 

            torch.nn.BatchNorm2d(channels)        self.conv2 = 

            torch.nn.Conv2d(              channels, 

              channels, 

              kernel_size=3, 

              padding=1, 

              bias=False        )        self.bn2 = 

           torch.nn.BatchNorm2d(channels)    def forward(self, x):        residual = x        out = 

          torch.relu(

            self.bn1(self.conv1(x))

          )        out = 

          self.bn2(self.conv2(out))        return torch.relu(out + residual)

教学用的 TinyResNet 连续调用两个 BasicBlock。输入和输出 Shape 都是 [1,2,4,4]

features -> BasicBlock 1 -> BasicBlock 2 -> output

图 6-2(TinyResNet):两个残差块都包含两层 Conv、两层 BatchNorm 和一次残差 Add。

SECTION 04

固定 BatchNorm 参数,再导出两份 ONNX

需要先认识常量折叠可以先按一句话理解:推理开始前就能确定结果的计算,提前算完并写进模型,不必等每次推理再重复计算。 例如 Constant(2) 与 Constant(3) 相加,可以在导出时直接变成 Constant(5)

Conv 与 BatchNorm 的合并稍有不同:它不是提前算出模型输出,而是把 BatchNorm 的固定参数换算进 Conv 的权重和 Bias,更准确地说属于参数折叠。两者的严格区别留到后面的图优化章节再讲。

在本文使用的传统 ONNX 导出器中,do_constant_folding 控制是否执行常量折叠这一步。设为 True 时,BatchNorm 会折叠进 Conv;设为 False 时,BatchNormalization 会作为独立节点保留。下面保持模型和输入完全相同,只改变这个参数:

model = TinyResNet().eval()features = 

    torch.arange(32, 

                 dtype=torch.float32)

features = 

    features.reshape(1, 2, 4, 4) / 10torch.onnx.export(    model,    (features,),    "06_tiny_resnet_folded.onnx",    input_names=["features"],    output_names=["output"],    opset_version=18,    external_data=False,    dynamo=False,    do_constant_folding=True,)torch.onnx.export(    model,    (features,),    "06_tiny_resnet_unfolded.onnx",    input_names=["features"],    output_names=["output"],    opset_version=18,    external_data=False,    dynamo=False,    do_constant_folding=False,)

本例只用它观察折叠前后的节点变化。

图 6-3(折叠对比):不折叠时 BatchNormalization 是独立节点;折叠后它的效果进入 Conv 参数,残差 Add 不受影响。

SECTION 05

不折叠:ONNX 保留四个 BatchNormalization

设置 do_constant_folding=False 后,每个 BasicBlock 的主支路与 PyTorch 代码一致:

Conv -> BatchNormalization -> Relu     -> Conv -> BatchNormalization -> Add -> Relu

两个 BasicBlock 一共导出 14 个节点,其中包含 4 个 Conv、4 个 BatchNormalization、4 个 Relu 和 2 个 Add

图 6-4(不折叠 ONNX):四个绿色 BatchNormalization 节点完整保留,两条直连支路仍分别进入两个 Add

这种图更接近 PyTorch Module 的写法,适合学习算子连接或确认 BatchNorm 位于哪一层之后。代价是图更长、节点更多。

SECTION 06

折叠后:BatchNormalization 消失,Add 保留

设置 do_constant_folding=True 后,每个 BasicBlock 变成:

Conv -> Relu -> Conv -> Add -> Relu

节点总数从 14 个减少到 10 个,4 个独立的 BatchNormalization 全部消失。Netron 中的 Conv 现在带有 Bias,虽然 PyTorch 代码里的 Conv2d 设置了 bias=False;这些新参数已经包含原来 BatchNorm 的缩放和平移效果。

图 6-5(折叠后 ONNX):BatchNormalization 不再显示,但两个 Add 和两条残差直连都没有变化。

两份图可以直接对比:

对比项
不折叠
折叠后
ONNX 节点总数
14
10
BatchNormalization
4
0
Add
2
2

使用同一输入比较两份 ONNX,ONNX Runtime 输出之间的最大绝对误差是 0.0。这说明 BatchNorm 并没有被删除,它的计算效果只是搬进了 Conv 参数。图变短也不等于已经证明推理一定更快;真实性能还需要后续使用性能工具测量。

EPILOGUE

小结

本章先回顾了 Conv、Relu 和 Add,再对比了 BatchNormalization 与 LayerNormalization:BatchNorm 常用于视觉模型的卷积特征,LayerNorm 常用于单个样本的特征向量。

后续实验以 BatchNorm 为主,同一份 PyTorch TinyResNet 在不折叠时保留 Conv -> BatchNormalization,折叠后只看到参数已经更新的 Conv。

无论是否折叠,残差支路和 Add 都保持不变。

看到 ONNX 中缺少 BatchNormalization 时,先检查导出和图优化设置,不要直接判断模型少了一层。

如果你对 AI Infra 系列感兴趣,或者身边有朋友正在学习 AI,欢迎关注本账号,并转发给需要的朋友。