ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Pytorch导出ONNX文件教程

Pytorch导出ONNX文件教程 新建一个工程文件夹用于存放python文件和生成的onnx文件开启anaconda创建环境输入conda create -n mpu6050task python3.10 -ympu6050task为环境名可自定义启动创建的环境conda activate mpu6050task切换文件路径到创建的工程文件夹在当前窗口中运行以下命令一次性安装 PyTorch、ONNX 及其解析支持库固定使用兼容性最佳的numpy2pip install torch torchvision onnx onnxscript numpy2 -i https://pypi.tuna.tsinghua.edu.cn/simple创建python工程选择已经生成的环境解释器加入代码import torch import torch.nn as nn import torch.onnx # 1. 定义适配 MCU 算力的超轻量 1D 卷积神经网络 class TinyVibrationCNN(nn.Module): def __init__(self, num_classes3): super(TinyVibrationCNN, self).__init__() # 输入: [Batch1, Channel3(X/Y/Z), TimeSteps32] self.conv1 nn.Conv1d( in_channels3, out_channels8, kernel_size3, padding1 ) self.relu nn.ReLU() self.pool nn.MaxPool1d(kernel_size2) # 时间步: 32 - 16 self.conv2 nn.Conv1d( in_channels8, out_channels16, kernel_size3, padding1 ) # 时间步: 16 - 8 self.fc nn.Linear(16 * 8, num_classes) def forward(self, x): x self.pool(self.relu(self.conv1(x))) x self.pool(self.relu(self.conv2(x))) x x.view(x.size(0), -1) # 展平为向量 x self.fc(x) return x if __name__ __main__: # 2. 实例化模型并设为评估模式 model TinyVibrationCNN(num_classes3) model.eval() # 3. 构造虚拟输入 (Batch1, 3轴通道, 32个采样点) dummy_input torch.randn(1, 3, 32, dtypetorch.float32) onnx_filename vibration_model.onnx # 4. 导出为兼容 STM32Cube.AI 的单文件 ONNX torch.onnx.export( model, dummy_input, onnx_filename, export_paramsTrue, opset_version11, # STM32Cube.AI 支持最稳定的算子集 dynamoFalse, # 关键禁用图拆分将权重直接打包进单个 .onnx 文件 do_constant_foldingTrue, input_names[sensor_input], output_names[class_probabilities], ) print(fONNX 模型导出成功: {onnx_filename})在conda终端输入命令进行生成onnx文件python export_model.py
返回列表