You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在ONNX格式中保留PyTorch模型的Batch Normalization层?

解决PyTorch导出ONNX时BatchNorm层被合并及TrainingMode未定义问题

问题原因

  • BatchNorm层被合并到卷积层:默认导出时模型处于eval模式,且do_constant_folding=True会触发算子融合,把BatchNorm和卷积合并为单个节点,导致BN层在ONNX中无单独记录。
  • NameError错误:TrainingMode是torch.onnx模块下的枚举类,原代码未导入该模块或未指定完整路径,导致无法识别。

修正方案

  1. 补充必要导入:导入torch.onnx模块以使用TrainingMode枚举,或直接用字符串参数替代枚举值。
  2. 调整导出参数:关闭常量折叠(do_constant_folding=False),同时指定训练模式,强制保留BatchNorm层的独立结构。
  3. 修复输入缺失:原代码未定义dummy_input,需补充与模型输入维度匹配的测试张量。

修正后的完整代码

import torch
from torch import nn
import torch.onnx  # 导入torch.onnx以使用TrainingMode枚举

device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Using {device} device")

class NeuralNetwork(nn.Module):
    def __init__(self):
        super(NeuralNetwork, self).__init__()
        self.Network = nn.Sequential(
            nn.Conv2d(1,1,3),
            nn.BatchNorm2d(1, track_running_stats=False),
            nn.ReLU(),
        )

    def forward(self,x):
        output = self.Network(x)
        return output

model = NeuralNetwork().to(device)
print(model)

# 定义与模型输入维度匹配的测试张量
dummy_input = torch.randn(1, 1, 28, 28)  # batch=1, channel=1, 28x28

# 导出ONNX模型
torch.onnx.export(
    model, 
    dummy_input, 
    "test.onnx", 
    verbose=True, 
    export_params=True, 
    opset_version=12,
    do_constant_folding=False,  # 关闭常量折叠,避免BN与卷积融合
    input_names=['input'], 
    output_names=['output'],
    training=torch.onnx.TrainingMode.TRAINING  # 指定训练模式
)

# 验证ONNX模型合法性
import onnx
onnx_model = onnx.load("test.onnx")
onnx.checker.check_model(onnx_model)
print("ONNX模型验证通过")

补充说明

  • 若不想导入torch.onnx,可直接用字符串training='training'替代torch.onnx.TrainingMode.TRAINING,效果完全一致。
  • track_running_stats=False确保BatchNorm在训练模式下不使用滑动均值,导出时会保留完整的BN计算逻辑。
  • dummy_input的维度必须与模型实际输入维度匹配,否则会触发导出错误。

内容的提问来源于stack exchange,提问作者Eddymage

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.17 08:25:31