ARTICLE DETAIL

资讯详情

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

PyTorch训练MobileNetV2并转TensorRT部署全流程

PyTorch训练MobileNetV2并转TensorRT部署全流程 简介面向PyTorch学习者的图像分类实战资源包基于植物幼苗数据集部分样本构建12类别分类任务完整演示MobilenetV2从PyTorch训练、ONNX转换到TensorRT部署的全流程。资源包共2000个文件以大量png图像数据为主体辅以8个Python脚本和2个pyc文件压缩包约936MB目录结构清晰适合直接对照运行。内容覆盖torchvision.models中MobilenetV2模型的调用方式、自定义数据集加载、Cutout与Mixup数据增强、余弦退火学习率调整、训练与验证流程以及训练后模型加载预测。数据集中12类植物幼苗图像涵盖不同生长状态与背景结合增强方式可有效提升模型泛化能力脚本按功能模块划分训练、验证、转换、推理各步骤均可独立运行与修改。部署部分则从PyTorch导出ONNX开始依次完成ONNX推理、ONNX转TensorRT并给出TensorRT推理实现。通过这份资源可以系统掌握图像分类从实验训练到工程落地的完整链条。目前已有1854人学习下载资源适用于希望打通算法与部署环节的中高级深度学习开发者。无论是用于课程设计、竞赛实践还是生产部署预研这份包含数据、脚本与方案说明的压缩包都有较高参考价值。1. 为什么 MobileNetV2 和 TensorRT 是图像分类落地的最稳组合手里只有一张老卡比如 GTX 1070还想把图像分类模型部署到生产环境跑实时推理直接用 PyTorch 的原始权重单张 224 的图推理要 30 到 50 毫秒显存还动不动就飙到 1GB 以上。但换成 MobileNetV2 做 backbone训练出来转成 TensorRT 引擎同样的卡能把单张延迟压到 5 毫秒以内显存占用掉一半还多。这就是这篇要讲的事在 PyTorch 下把 MobileNetV2 从零训练到分类部署先导出 ONNX 再转成 TensorRT 引擎走完整条落地链路。这篇适合两类人一是刚把 PyTorch 环境搭好、想跑通一个真实分类项目的初学者二是已经在用 PyTorch 训练模型、但被线上推理速度卡住、想给老显卡找一条性能出路的工程师。全文按训练、ONNX 导出、TensorRT 转换、踩坑排查、推理验证的顺序展开每个操作都有可直接复现的脚本和参数说明。2. MobileNetV2 用 PyTorch 训练从数据集到权重的完整流程2.1 先搞懂 MobileNetV2 的两个设计点再决定训练策略MobileNetV2 能在嵌入式设备上长期被选作分类主干核心是倒残差结构Inverted Residual Block和线性瓶颈Linear Bottleneck。普通残差网络是“先压缩、再卷积、再扩张”MobileNetV2 反过来先把通道扩展 6 倍用深度可分离卷积做特征提取再压缩回低维输出。这样设计的好处是低维信息不需要经过 ReLU 激活避免信息丢失计算量也集中在高维但计算便宜的深度可分离卷积上。这两个结构特点直接影响你的训练决策。第一扩展系数是 6中间层的通道数是输入通道的 6 倍显存占用会比参数数量显得更高训练时 batch_size 要根据显存调整。第二网络内部大量使用 ReLU6导出的 ONNX 里会出现 ReLU6 算子TensorRT 老版本或低精度模式对它支持不稳定这一点到后面转换环节再说。从工程角度看我不建议你从头随机初始化训练整个网络。图像分类任务里torchvision 提供了在 ImageNet 上预训练好的 mobilenet_v2 权重直接加载它做迁移学习通常几百张图训练几十个 epoch 就能收敛到可用的精度。除非你的数据集和 ImageNet 分布差异极大才需要考虑从头训练。2.2 环境准备与数据组织PyTorch 版本、CUDA 和数据集目录结构训练环境以 PyTorch 为主建议安装 GPU 版。常见做法是先用 conda 建一个干净环境再安装与你的 CUDA 驱动匹配的 PyTorch。比如驱动支持 CUDA 11.8就装对应的 PyTorch 版本不要直接pip install torch装成 CPU 版否则后面转 TensorRT 时你会发现模型推理慢得没法看。我一般这样组织数据集目录ImageFolder 可以直接读取data/ train/ cat/ 001.jpg 002.jpg ... dog/ 001.jpg 002.jpg ... val/ cat/ ... dog/ ...这样用torchvision.datasets.ImageFolder一条代码就能同时拿到数据和标签映射。森林图像分类这类场景也适用只要把类别换成具体树种即可。2.3 训练脚本完整可复现的最小实现下面是一套可以直接跑的训练脚本。数据增强、模型加载、训练循环和验证都包含在内你在自己的数据集上只要改data_root和类别数即可。import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms, models # 数据增强训练集做随机裁剪和翻转验证集只做缩放和中心裁剪 train_trans transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ColorJitter(0.2, 0.2, 0.2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) val_trans transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) train_data datasets.ImageFolder(data/train, transformtrain_trans) val_data datasets.ImageFolder(data/val, transformval_trans) train_loader DataLoader(train_data, batch_size64, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_data, batch_size64, shuffleFalse, num_workers4, pin_memoryTrue) # 加载预训练权重替换最后一层分类头 model models.mobilenet_v2(weightsmodels.MobileNet_V2_Weights.DEFAULT) num_classes len(train_data.classes) model.classifier[1] nn.Linear(model.last_channel, num_classes) model model.cuda() criterion nn.CrossEntropyLoss() optimizer optim.SGD(model.parameters(), lr0.01, momentum0.9, weight_decay1e-4) scheduler optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max20) def train_one_epoch(epoch): model.train() running_loss 0.0 for images, labels in train_loader: images, labels images.cuda(), labels.cuda() optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * images.size(0) print(fEpoch {epoch1}, Loss: {running_loss / len(train_data):.4f}) def evaluate(): model.eval() correct 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.cuda(), labels.cuda() outputs model(images) _, preds torch.max(outputs, 1) correct (preds labels).sum().item() print(fVal Acc: {correct / len(val_data):.4f}) for epoch in range(20): train_one_epoch(epoch) evaluate() scheduler.step() torch.save(model.state_dict(), mobilenetv2_classification.pth)这段代码的设计逻辑是先用 ImageNet 权重初始化模型替换最后分类头然后用较小的学习率微调全部参数。SGD 加动量在分类任务上比 Adam 更容易收敛到平坦的最优点配合余弦退火调度器20 个 epoch 一般就够。batch_size 用 64 是因为 MobileNetV2 在 224 分辨率下显存开销不大如果你只有 6GB 显存可以降到 32学习率也按比例降到 0.005。验证集评估必须保持和训练时相同的归一化参数即 mean 和 std 固定为 0.485/0.456/0.406 和 0.229/0.224/0.225。后面转 TensorRT 时最容易翻车的点就是推理端忘记做同样的归一化模型精度直接从 95% 掉到 20%看起来像模型坏了实际是预处理没对齐。3. PyTorch 到 TensorRT 的部署链路先转 ONNX 再过一遍 TRT3.1 为什么中间层必须加 ONNX而不是把 PyTorch 权重直接给 TensorRTTensorRT 无法直接读取 PyTorch 的 .pth 权重文件。PyTorch 的权重只是张量字典里没有完整的网络计算图结构TensorRT 拿不到它做层融合和算子替换所需的信息。行业的标准路径是先把 PyTorch 模型导出成 ONNX 格式ONNX 是中间表示层描述完整的数据流图再交给 TensorRT 做图优化和引擎生成。ONNX 这一步也承担了算子归一化的角色。PyTorch 里有几百种自定义算子TensorRT 不可能全部支持。导出 ONNX 后你可以用onnxsim和polygraphy检查哪些算子不受支持提前替换掉。尤其是 MobileNetV2 这类结构化非常规整的网络导出的 ONNX 一般都能直接被 TensorRT 完整解析不会像 Transformer 那一类网络出现大量自定义算子需要手工兜底。3.2 导出 ONNX 的注意事项opset、动态维度、输出名导出 ONNX 这一段是整条链路里最容易出问题的。常见做法是用torch.onnx.export但有几个参数必须显式配置否则后面转 TensorRT 会踩到固定 batch 的坑。import torch from torchvision import models model models.mobilenet_v2(weightsmodels.MobileNet_V2_Weights.DEFAULT) in_features model.classifier[1].in_features model.classifier[1] torch.nn.Linear(in_features, 2) # 按你的类别数改 model.load_state_dict(torch.load(mobilenetv2_classification.pth)) model.eval().cuda() dummy_input torch.randn(1, 3, 224, 224).cuda() torch.onnx.export( model, dummy_input, mobilenetv2.onnx, opset_version13, do_constant_foldingTrue, input_names[input], output_names[output], dynamic_axes{ input: {0: batch}, output: {0: batch} } ) print(ONNX exported.)这里最关键的配置是dynamic_axes把 batch 维度标记为动态。如果不加这一项导出的 ONNX 会把输入形状固定为 (1, 3, 224, 224)转出来的 TensorRT 引擎一次只能推理一张图无法在批量推理场景下提升吞吐。opset_version设成 13 是因为它和 TensorRT 的兼容性较好ReLU6、深度可分离卷积这些算子在这个版本下都能正确映射。如果你装了更高版本 PyTorch默认 opset 可能是 17 或 18也不要直接用默认值除非你能确认当前 TensorRT 版本支持该 opset。do_constant_foldingTrue会把权重常量折叠进图里减少推理时的计算节点。导出后建议先用 Polygraphy 或 ONNX Runtime 做一次前向对比确保 ONNX 输出和 PyTorch 输出在数值误差范围内一致再进行下一步。3.3 两种方式把 ONNX 转成 TensorRT 引擎trtexec 命令行与 Python API转引擎最常见的方式是用 TensorRT 自带的trtexec命令行工具它适合快速验证和批量打引擎。命令如下trtexec --onnxmobilenetv2.onnx --saveEnginemobilenetv2.engine --fp16 --workspace1024--fp16让引擎以半精度推理对 MobileNetV2 这种深可分离卷积结构精度损失通常在 0.5% 以内速度提升 1.5 到 2 倍。--workspace1024指定工作空间大小单位是 MB老显卡显存紧张时这个参数要重点调。如果转换过程中报错先去掉--fp16再试排除精度模式导致的算子支持问题。第二种方式是用 Python API。适合需要把引擎构建流程集成进训练或部署脚本的场景可以动态设置 profile、动态 batch 范围。import tensorrt as trt logger trt.Logger(trt.Logger.WARNING) builder trt.Builder(logger) network builder.create_network( 1 int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH) ) parser trt.OnnxParser(network, logger) with open(mobilenetv2.onnx, rb) as f: if not parser.parse(f.read()): for i in range(parser.num_errors): print(parser.get_error(i)) config builder.create_builder_config() config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, 1 30) config.set_flag(trt.BuilderFlag.FP16) engine builder.build_serialized_network(network, config) with open(mobilenetv2.engine, wb) as f: f.write(engine) print(Engine saved.)Python API 里EXPLICIT_BATCH标志必须开启否则网络定义是隐式 batch 模式配合动态维度导出会出现维度不匹配错误。create_builder_config是新版 TensorRT 的接口老版本用的是builder.max_workspace_size如果你用的是 TensorRT 8.x 之前的版本接口写法不同要以你环境里的版本为准。转出来的 .engine 文件是一个二进制序列化引擎里面已经做了算子融合和显存优化这个文件直接部署到目标机器即可不需要再依赖 PyTorch 或 ONNX。4. TensorRT 部署避坑清单五个翻车现场与解决办法4.1 推理结果全错但训练精度很高预处理不一致现象是同一张测试图PyTorch 推理输出正确的猫TensorRT 推理输出的概率分布完全不对甚至 argmax 直接错成别的类别。原因基本可以锁定在预处理不一致。PyTorch 训练时做了Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225])这个操作是把像素值从 0-255 先除以 255 得到 0-1 区间再按公式(x - mean) / std做标准化。TensorRT 部署时如果只用cv2.resize和除以 255或者忘了减 mean 除 std输入分布就和训练时完全不同。解决方法是把预处理写成和训练完全一致并且最好在部署代码中显式写出来不要依赖任何隐式处理。我一般把归一化直接固化到 ONNX 图里这样部署端只需要做resize和ToTensor。4.2 转出的引擎只能跑 batch1dynamic_axes 没设置现象是引擎构建成功推理也正确但输入必须固定为 (1, 3, 224, 224)想一次喂 8 张图就报维度错误。原因是在 3.2 节导出 ONNX 时没有设置dynamic_axes。TensorRT 构建引擎时会根据 ONNX 的静态输入形状做显存优化如果输入维度是静态的引擎就锁死为单张推理。很多新手在这里翻车后都会以为是 TensorRT 不支持批量其实是导出环节缺了一步。解决方法是重新导出 ONNX加上dynamic_axes{input: {0: batch}, output: {0: batch}}转引擎时再用 Python API 设置动态 batch 范围最小 batch 为 1最大为 16 或 32优化 batch 设为 8。这样引擎支持可变 batch吞吐测试时能明显看到优势。4.3 转换时报 ReLU6 算子不支持opset 版本和 onnx-simplifier 双管齐下现象是trtexec转换到某个层时报Node ... type ReLU6 unsupported或者转换成功但该层被回退到 CPU 执行速度掉一大截。原因是低版本 TensorRT 对 ReLU6 的支持不完善或者导出的 ONNX 中 ReLU6 展开方式与 TensorRT 期望不一致。MobileNetV2 内部大量使用 ReLU6这个问题逃不掉。解决分两步。第一步把 opset 版本调到 13 或更高新版 ONNX 为 ReLU6 定义了标准算子。第二步跑一次python -m onnxsim mobilenetv2.onnx mobilenetv2_sim.onnx把图简化很多 ONNX 导出时的多余节点会被消除。如果依然报错就在转换时加--precisionOverrideReLU6:FP32让这一个算子保持 FP32 精度。4.4 老显卡放不下引擎或推断卡顿workspace 和精度模式一起调现象是 GTX 1070 这种 8GB 卡在转引擎时提示显存不足或者引擎构建成功但推理时 CPU 占用异常飙升。原因有两个方向一是 workspace 设得太大TensorRT 会为所有可能融合的层预留显存超过显卡容量直接被拒二是某些算子不支持 FP16被回退到 CPU 执行造成 CPU 占用高、速度反而更慢。解决方法是先限制workspace为 512MB降低显存压力再打开 TensorRT 的 layer 信息输出检查哪些层没有跑在 GPU 上。对于必须回退的层用--layerPrecisions单独指定为 FP16 无法解决时就考虑整体转 FP32MobileNetV2 在 FP32 下速度也已经比 PyTorch 快很多。4.5 Python API 换版本后接口全变以官方 sample 为准不要硬背 API现象是把网上老代码粘进来builder.max_workspace_size报错或者create_network少传参数再或者build_engine返回 None 却没有报错信息。原因是 TensorRT 各版本 API 变化很大7 到 8 再到 9核心接口流程都有调整。网络上的教程往往对应特定版本直接照抄必然出问题。解决方法是先确定你实际安装的 TensorRT 版本进官方 Python sample 目录按它的写法调整。核心流程展开后本质没变创建 logger、创建 builder、解析 ONNX、配 config、build engine只是接口名字和参数在变。把这五个步骤作为主干按当前版本查文档补参数比死记代码可靠得多。5. TensorRT 推理端的验证与调优从精度对接到 batch 吞吐5.1 用 Python 写一套 TensorRT 推理代码并对比 PyTorch 精度拿到 .engine 文件后第一件事是验证精度不是先测速度。写一个最小推理脚本喂同一张图对比 PyTorch 和 TensorRT 的 softmax 输出。import numpy as np import tensorrt as trt import torch import torchvision.transforms as transforms from PIL import Image TRT_LOGGER trt.Logger(trt.Logger.WARNING) def load_engine(engine_path): with open(engine_path, rb) as f: runtime trt.Runtime(TRT_LOGGER) return runtime.deserialize_cuda_engine(f.read()) def preprocess(image_path): img Image.open(image_path).convert(RGB) trans transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) return trans(img).unsqueeze(0).numpy() engine load_engine(mobilenetv2.engine) context engine.create_execution_context() input_data preprocess(test.jpg) output_shape (1, 2) # 改成你的类别数 # 动态 batch 场景下必须显式设置输入形状 context.set_input_shape(input, input_data.shape) output_data np.empty(output_shape, dtypenp.float32) # 分配显存并执行推理 d_input trt.DeviceAllocator(input_data.nbytes) d_output trt.DeviceAllocator(output_data.nbytes) d_input.copy_host_to_device(np.ascontiguousarray(input_data)) context.set_tensor_address(input, d_input.address) context.set_tensor_address(output, d_output.address) context.execute_async_v3(0) d_output.copy_device_to_host(output_data) print(TRT output:, output_data)这段逻辑里最需要注意的是set_input_shape只要引擎是动态 batch 构建的执行上下文必须在推理前声明当前输入形状。execute_async_v3是新版接口老版本是context.execute_async(batch_size, bindings, stream)写法差异很大。精度对比时不要直接比较 argmax要看 softmax 概率分布的差异。一般 FP16 引擎的 top-1 结果与 FP32 PyTorch 完全一致但概率值可能有 0.01 以内的偏差。超过这个范围就要排查预处理或归一化。5.2 latency 和吞吐测试一套简单的计时流程推理速度测试要区分两个指标单张延迟和批量吞吐。单张延迟测试直接用 5.1 的代码循环 100 次计时取平均注意要把 warmup 的前 10 次排除因为 TensorRT 首次推理要做 cuDNN 和 cublas 的初始化会明显偏慢。批量吞吐测试更有参考价值因为生产场景通常希望一次处理多张图。把输入 shape 改成 (8, 3, 224, 224)用同一个 context 连续推理 50 次统计总耗时算8 * 50 / 总耗时得到每秒处理张数。GTX 1070 上 MobileNetV2 FP16 引擎单张延迟大约 3 到 6 毫秒batch8 时吞吐能到每秒 100 张以上。如果发现批量吞吐没有随 batch 线性增长先看是不是显存带宽瓶颈再看 CPU 预处理是否阻塞主线程。常见做法是把图片解码、resize 和推理拆成三个线程跑流水线吞吐会有明显提升。5.3 动态 batch 的配置实践min-opt-max 的选择动态 batch 配置直接影响显存占用和性能。opt值设得越大TensorRT 为它做的层融合优化越激进但显存预留也越多。我的经验是max batch16opt batch8在一个 8GB 卡上是比较均衡的配置min batch 设 1 保证单帧请求也能处理。engine 文件打包时需要记录好 min/opt/max 的 batch 值因为这是生成时定的运行时无法修改。如果你的线上流量波动大要么把 max 设大一点求稳要么生成多个不同 max batch 的引擎按流量切换。实际项目里我一般直接设 max32engine 文件会比静态 batch 大一些但换来了灵活性。6. 更进一步把预处理固化进 ONNX顺便做 INT8 量化前面所有步骤跑通之后还有两个方向值得做一是把归一化直接写进 ONNX 图里让部署端代码更短、出错概率更低二是尝试 INT8 量化把老显卡的算力再榨一轮。预处理固化的思路很简单归一化本质是减均值除方差这些计算可以表示成 ONNX 里的Sub和Div节点加在模型输入之后即可。你可以先按常规方式导出一版 ONNX再用onnx库手动注入这几个节点或者直接用torch.onnx.export时在模型外包一层nn.Sequential前面挂归一化模块。后端推理代码就不用再管 mean 和 std 了。INT8 量化是重头戏需要用到 TensorRT 的 calibrator。先准备一批代表性的校准图片数量 500 到 1000 张即可太少会导致量化后精度暴跌太多会拉长校准时间。校准过程会统计每层激活值的分布决定 INT8 的缩放因子。MobileNetV2 对 INT8 比较友好通常精度损失在 1% 到 2%但如果你数据集类别数很少比如只有两类可能会因为特征区分度集中而出现精度掉得更多的情况这时先看看是哪一层被量化得最狠单独设回 FP16。我自己的习惯是走到 INT8 之前先确保 FP16 引擎已经在线上稳定跑一段因为量化是个玄学活模型结构、校准数据分布都会影响最终精度。把这条链路完整吃透之后再碰 INT8遇到问题也容易排查。前阵子做森林图像分类项目就是栽在了预处理对齐上部署端少乘了一次 1/255精度从 94% 掉到 30%排查了两天才发现是它。希望这些教训能让你少走这段弯路。本文还有配套的精品资源点击获取
返回列表