PyTorch模型轻量化部署实战:从torch.fx剪枝到TensorRT引擎

发布时间:2026/10/11 10:50:08
PyTorch模型轻量化部署实战:从torch.fx剪枝到TensorRT引擎 简介本资源为山东大学软件学院《众智科学与网络化产业》2022年课程全套实验材料面向计算机与软件工程专业本科生及对集体智能、网络化系统实践感兴趣的学习者旨在通过动手实验深化对众智原理、数据采集、机器学习建模、分布式通信及综合项目开发的理解与应用。压缩包共17个文件含5份实验报告.docx、5个可执行程序.exe、5个核心源码.cpp、1份实验大纲.doc及1份说明文本.txt覆盖从Python/Java环境搭建、Web爬虫与Pandas数据分析到Scikit-learn算法实现、客户端-服务器通信再到推荐系统或社交网络分析等综合项目的完整链路。资源包仅1.7MB轻量易用目录按实验序号清晰组织便于逐级学习与复现。目前已有1720人下载学习提供即开即用的代码报告大纲三位一体支撑助读者高效掌握理论落地的关键环节与工程表达规范。1. 这不是一份普通实验材料它是一套可复用的“工业级轻量模型验证闭环”实战切片“山东大学软件学院众智2022年实验 代码及实验报告”——光看标题你可能以为这只是某门课的作业压缩包。但实际拆开你会发现它完整封装了一个在资源受限场景下如边缘设备、嵌入式终端、低配开发板验证模型轻量化效果的最小可行闭环从 PyTorch 模型剪枝 → ONNX 导出 → TensorRT 引擎构建 → C 推理接口封装 → 定点精度比对 → 实验报告自动生成模板。这不是教学演示而是某高校实验室为对接真实产线部署需求所设计的“带压测、带校验、带归档”的工程化训练靶场。适合正在做毕业设计、准备嵌入式AI项目落地、或需要快速验证模型压缩方案有效性的开发者。如果你正卡在“模型训好了但不知道怎么塞进树莓派/ Jetson Nano / STM32H7AI加速核里跑通”这份材料就是你缺的那块拼图——它不讲原理推导只给你能git clone make ./run_test的实操路径。2. 从 PyTorch 剪枝到 ONNX 导出为什么必须用torch.fx而非torch.nn.utils.prune2.1 为什么传统剪枝在部署时会“失真”很多同学直接用torch.nn.utils.prune.l1_unstructured对模型做剪枝训练完导出 ONNX 后发现推理结果和 PyTorch 不一致甚至报错Unsupported node kind: prune。根本原因在于nn.utils.prune是参数掩码mask-based剪枝它不修改模型结构只是把权重置零并加 mask 层而 ONNX 导出器无法识别PruningContainer这类运行时动态结构更不会自动剔除零权重通道——导致导出的 ONNX 图里还挂着一堆“死连接”TensorRT 编译时要么报错要么生成冗余计算节点最终推理耗时不降反升。提示剪枝 ≠ 压缩。只有结构可导出、通道可裁剪、权重可固化才算真正进入部署流水线。2.2 用torch.fx实现结构感知剪枝三步走通路该实验采用torch.fx图追踪 自定义TracerGraphModule重写的方式实现结构可导出剪枝。核心逻辑是先用fx.symbolic_trace获取模型计算图再遍历graph.nodes找到所有conv2d节点根据其输出通道的 L1 范数排序标记待裁剪通道索引最后用replace_node_module替换原Conv2d为通道数缩减后的新模块并同步更新后续BatchNorm2d和ReLU的输入维度。# prune_fx.py import torch import torch.fx as fx from torch import nn def trace_and_prune(model: nn.Module, example_input, ratio0.3): # 1. 符号追踪获取计算图 traced fx.symbolic_trace(model) # 2. 遍历图节点定位 conv2d 并统计通道重要性 conv_nodes [n for n in traced.graph.nodes if n.target torch.nn.functional.conv2d] channel_scores {} for node in conv_nodes: # 获取对应原始模块需提前注册 named_modules 映射 module_name node.args[0].target if hasattr(node.args[0], target) else None if module_name and hasattr(traced, module_name): conv_module getattr(traced, module_name) # 计算每个输出通道的 L1 范数按 output channel 维度求和 scores torch.norm(conv_module.weight.data, p1, dim(1,2,3)) channel_scores[node] scores # 3. 构建新 GraphModule裁剪 更新后续层 new_graph fx.Graph() env {} for node in traced.graph.nodes: if node in channel_scores: scores channel_scores[node] k int(len(scores) * ratio) keep_idx torch.topk(scores, len(scores)-k, largestTrue).indices # 创建裁剪后的新 Conv2d old_conv getattr(traced, node.target) new_conv nn.Conv2d( in_channelsold_conv.in_channels, out_channelslen(keep_idx), kernel_sizeold_conv.kernel_size, strideold_conv.stride, paddingold_conv.padding, biasold_conv.bias is not None ) # 权重与偏置复制仅保留重要通道 new_conv.weight.data old_conv.weight.data[keep_idx] if old_conv.bias is not None: new_conv.bias.data old_conv.bias.data[keep_idx] # 插入新模块 new_node new_graph.create_node(call_module, fpruned_{node.target}, argsnode.args, kwargsnode.kwargs) env[node] new_node setattr(new_graph, fpruned_{node.target}, new_conv) else: # 复制其他节点注意需同步更新依赖节点的 args new_node new_graph.node_copy(node, lambda x: env[x] if x in env else x) env[node] new_node new_graph.output(env[traced.graph.nodes[-1]]) return fx.GraphModule(traced, new_graph) # 使用示例 model resnet18(pretrainedTrue).eval() pruned_model trace_and_prune(model, torch.randn(1,3,224,224), ratio0.3) pruned_model(torch.randn(1,3,224,224)) # ✅ 可正常前向关键参数说明ratio0.3表示裁剪掉 30% 的输出通道实际值需结合val_acc_drop 1.5%和latency_gain 25%双指标调优torch.norm(..., p1, dim(1,2,3))L1 范数对通道敏感比 L2 更利于稀疏性且避免了 BN 层 scale 影响new_graph.node_copy(...)中的 lambda 函数确保下游节点如 ReLU、Add的输入自动指向裁剪后的新 conv 输出这是结构可导出的核心保障。3. ONNX 导出与 TensorRT 引擎构建绕过dynamic_axes陷阱的静态 shape 策略3.1 为什么dynamic_axes在边缘端是“伪需求”实验报告中明确要求所有 ONNX 导出必须使用固定 shape如--input_shape 1,3,224,224禁用dynamic_axes。理由很现实Jetson Nano 的 TensorRT 7.1.3 不支持Resizedynamic_axes混合编译STM32Cube.AI 工具链完全不认动态 batch连最新版onnx-simplifier在含If控制流的模型上也会因 dynamic shape 失效。所谓“动态适配”在真实嵌入式场景里90% 是靠预设几组常用 shape1x3x224x224、4x3x224x224、1x3x320x320分别编译多个引擎运行时查表加载——这才是工业界真实做法。3.2 用--opset-version 13--do_constant_folding导出稳定 ONNX该实验统一使用 ONNX opset 13兼容 TensorRT 7.x ~ 8.6并强制开启常量折叠。关键命令如下python -m torch.onnx.export \ --opset-version 13 \ --do_constant_folding \ --input-names input \ --output-names output \ --dynamic_axes {} \ # 空字典显式禁用 pruned_model.pth \ model.onnx \ --example-inputs torch.randn(1,3,224,224) \ --verbose # 查看导出日志中的 warning导出后必检三项onnx.shape_inference.infer_shapes_path(model.onnx)确认所有 tensor 具有完整 static shapeonnx.checker.check_model(model.onnx)无 error 即通过onnxsim.simplify(model.onnx, perform_optimizationTrue)简化后模型体积应缩小 15%~30%且onnxruntime.InferenceSession加载后session.get_inputs()[0].shape [1,3,224,224]。注意若onnxsim报Unsupported operator: Resize说明模型中存在F.interpolate且未被torch.fx正确替换——需回退到 2.2 节检查interpolate节点是否被转为aten::upsample_bilinear2d并手动替换为nn.Upsample模块。3.3 TensorRT 引擎构建用trtexec命令行而非 Python API 更可靠实验采用trtexecTensorRT 自带命令行工具构建引擎因其错误提示更明确、环境依赖更少、且能直接输出latency和VRAM usage统计。典型命令如下trtexec \ --onnxmodel.onnx \ --saveEnginemodel.engine \ --fp16 \ --workspace2048 \ --minShapesinput:1x3x224x224 \ --optShapesinput:1x3x224x224 \ --maxShapesinput:1x3x224x224 \ --shapesinput:1x3x224x224 \ --avgRuns100 \ --duration10 \ --separateProfileRun \ --exportProfileprofile.json参数深解--fp16启用半精度Jetson Nano 上提速约 2.1x精度损失 0.8%实测 ResNet18 Top1--workspace2048指定最大 GPU 显存工作区MB低于 1024 会导致某些 layer 无法使用优化 kernel--min/opt/maxShapes全部设为相同值强制静态 shape规避 profile 不匹配问题--separateProfileRun先 warmup 再测速避免首次运行抖动污染 latency 数据profile.json包含各 layer 的耗时占比可用于定位瓶颈 layer如某Conv2d占 63% 时间则需考虑分组卷积或 depthwise 替换。4. C 推理接口封装与定点精度比对用std::vectorfloat而非cv::Mat传参4.1 为什么不用 OpenCV 封装输入/输出实验报告特别强调禁止在推理核心路径中引入cv::Mat。原因有三①cv::Mat默认内存布局为HWCheight-width-channel而 TensorRT 引擎输入要求NCHW转换需cv::transposecv::reshape引入额外 memcpy②cv::Mat的 ROI、step、refcount 机制在多线程 infer 场景下易引发 dangling pointer③cv::Mat依赖 OpenCV 动态库增加部署包体积12MB和符号冲突风险。正确做法是用std::vectorfloat直接管理NCHW格式数据通过memcpy填充 device buffer。// infer_engine.h class TRTInfer { public: TRTInfer(const std::string engine_file); void infer(const std::vectorfloat input, std::vectorfloat output); private: void* m_context; // IExecutionContext* void* m_buffers[2]; // input output device ptr size_t m_input_size; // bytes size_t m_output_size; }; // infer_engine.cpp void TRTInfer::infer(const std::vectorfloat input, std::vectorfloat output) { // 1. memcpy host - device cudaMemcpy(m_buffers[0], input.data(), m_input_size, cudaMemcpyHostToDevice); // 2. execute async auto stream /* get CUDA stream */; ((IExecutionContext*)m_context)-enqueueV2(m_buffers, stream, nullptr); // 3. memcpy device - host cudaMemcpy(output.data(), m_buffers[1], m_output_size, cudaMemcpyDeviceToHost); cudaStreamSynchronize(stream); }关键细节input.data()必须是float*类型、连续内存、已按NCHW排列预处理阶段完成HWC→NCHWnormalizem_input_size 1 * 3 * 224 * 224 * sizeof(float)不可硬编码应从 engine 的 binding info 动态读取cudaStreamSynchronize不可省略否则output可能读到旧数据玄学翻车高发点。4.2 定点精度比对用abs(a-b)/max(|a|,|b|)替代 MSE实验报告要求精度比对必须使用相对误差而非 MSE 或 MAE因为 float32 的绝对误差在不同量级输出上无意义如分类 logits 为 [-5, 12]回归输出为 [0.001, 0.999]。该实验采用工业界通用公式$$ \text{rel_err} \frac{|a - b|}{\max(|a|, |b|) \varepsilon},\quad \varepsilon 1e^{-6} $$并设定阈值rel_err 0.0151.5%视为通过。C 实现如下bool compare_outputs(const std::vectorfloat fp32_out, const std::vectorfloat int8_out) { const float eps 1e-6f; size_t n fp32_out.size(); size_t fail_cnt 0; for (size_t i 0; i n; i) { float a fabs(fp32_out[i]); float b fabs(int8_out[i]); float max_val fmaxf(a, b) eps; float rel_err fabs(fp32_out[i] - int8_out[i]) / max_val; if (rel_err 0.015f) { fail_cnt; if (fail_cnt 5) { // 打印前5个失败项 printf(idx %zu: fp32%.6f, int8%.6f, rel_err%.4f\n, i, fp32_out[i], int8_out[i], rel_err); } } } printf(Total fails: %zu / %zu (%.2f%%)\n, fail_cnt, n, 100.0f * fail_cnt / n); return fail_cnt 0; }血泪经验若fail_cnt 0优先检查int8_out是否被错误地 cast 为uint8_t后再 reinterpret_cast 为float常见于 TensorRT INT8 calibration 后忘记 dequantize。5. 实验报告自动生成与避坑指南3 个让导师当场签字的细节5.1 报告自动生成用pandocjinja2模板替代 Word 手动填写该实验提供report_gen.py输入为results.json含各阶段 latency、acc、size 数据输出为 PDF 格式实验报告。核心是jinja2模板report.md.j2# 实验报告{{ model_name }} 轻量化验证 ## 性能对比 | 指标 | 原始模型 | 剪枝后 | TensorRT FP16 | TensorRT INT8 | |--------------|----------|--------|----------------|----------------| | 参数量(M) | {{ orig_params }} | {{ pruned_params }} | — | — | | 推理延迟(ms) | {{ orig_lat }} | {{ pruned_lat }} | {{ trt_fp16_lat }} | {{ trt_int8_lat }} | | Top1 准确率 | {{ orig_acc }} | {{ pruned_acc }} | {{ trt_fp16_acc }} | {{ trt_int8_acc }} | ## 关键结论 - 剪枝使参数量下降 {{ %.1f|format((orig_params-pruned_params)/orig_params*100) }}%但 Top1 仅下降 {{ %.2f|format(orig_acc-pruned_acc) }}% - TensorRT FP16 引擎在 Jetson Nano 上提速 {{ %.1f|format(orig_lat/trt_fp16_lat) }}x满足实时性要求 50ms执行命令python report_gen.py --results results.json --template report.md.j2 --output report.pdf提示pandoc生成 PDF 需系统安装texlive-latex-recommended和texlive-fonts-recommendedUbuntu 下一行解决sudo apt install texlive-latex-recommended texlive-fonts-recommended5.2 避坑指南这 4 个错误让 73% 的同学重跑超 3 小时现象 1trtexec报错Assertion failed: dims.nbDims 4 || dims.nbDims 5原因ONNX 模型输出 tensor 的 shape 为[1,1000]2D但 TensorRT 要求至少 4DNCHW。解决导出 ONNX 前在模型末尾加unsqueeze(2).unsqueeze(3)或用onnx-graphsurgeon修改 output shapeimport onnx_graphsurgeon as gs graph gs.import_onnx(onnx.load(model.onnx)) graph.outputs[0] gs.Variable(nameoutput, dtypenp.float32, shape(1,1000,1,1)) onnx.save(gs.export_onnx(graph), model_fixed.onnx)现象 2C infer 结果全为nan原因CUDA stream 创建后未显式cudaStreamCreate(stream)或enqueueV2第四参数传了nullptr但 stream 未初始化。解决严格按TRTInfer构造函数中创建 stream并在infer()中传入cudaStream_t stream; cudaStreamCreate(stream); ((IExecutionContext*)m_context)-enqueueV2(m_buffers, stream, nullptr);现象 3INT8 校准后精度暴跌 5%原因校准 dataset 仅用 10 张图且未覆盖光照/遮挡/尺度变化。解决必须用 ≥ 500 张图且来自验证集非训练集并确保每张图都经过与训练时完全一致的预处理包括mean[0.485,0.456,0.406],std[0.229,0.224,0.225]。现象 4onnxsim后模型无法被 TensorRT 加载原因onnxsim默认开启optimize_model会将Gemm层合并但某些老版本 TensorRT 不支持合并后的MatMulAdd结构。解决关闭优化仅做 shape inferenceonnxsim --skip-optimization model.onnx model_sim.onnx6. 进阶技巧用nvprof定位 TensorRT 层级瓶颈与一个我坚持了 3 年的习惯6.1 用nvprof替代trtexec --exportProfile查看 kernel 级耗时trtexec的--exportProfile只能给出 layer 级耗时如conv_1: 12.4ms但无法告诉你这个conv_1是被拆成了多少个 CUDA kernel、哪个 kernel 占用最多 SM 资源。这时必须上nvprofnvprof --unified-memory-profiling off \ --profile-from-start off \ --events sms__inst_executed,sms__sass_thread_inst_executed_op_fadd_pred_on,sms__sass_thread_inst_executed_op_fmul_pred_on \ --metrics sms__inst_executed,sms__sass_thread_inst_executed_op_fadd_pred_on \ ./trt_infer --engine model.engine --input input.bin关键输出解读sms__inst_executedSM 执行的总指令数越高说明计算密度大sms__sass_thread_inst_executed_op_fadd_pred_on浮点加法指令数若此值远高于fmul说明 kernel 是 memory-bound访存瓶颈而非 compute-bound若某 kernel 的achieved__inst_per_warp 20说明 warp 利用率低需检查是否因 bank conflict 或 divergent branch 导致。提示nvprof在 JetPack 4.6 已被nsys取代但nsys profile -t cuda,nvtx ./trt_infer输出更复杂对新手不友好nvprof虽 deprecated但对本实验的 Jetson Nanocompute capability 5.3仍最稳定。6.2 一个表格TensorRT 7.1.3 在 Jetson Nano 上各优化策略实测收益优化项启用方式延迟降低精度损失Top1备注FP16--fp162.1x 0.3%必选无脑开INT8--int8 --calibdata.calib3.8x 1.2%校准得当需 500 校准图否则崩DLA--useDLA01.4xvs GPU0%DLA core 仅支持有限 opResNet18 全链路不支持Layer fusion自动——trtexec默认开启无需配置6.3 我坚持了 3 年的习惯每次git commit前必跑./validate.sh该实验附带validate.sh它不是一个 fancy 的 CI 脚本而是极简的 5 行 shell却帮我避开 90% 的低级错误#!/bin/bash # validate.sh python -c import torch; print(torch.__version__) | grep -q 1.10 || { echo PyTorch version mismatch; exit 1; } onnx-checker model.onnx /dev/null 21 || { echo ONNX invalid; exit 1; } trtexec --onnxmodel.onnx --dryRun 2/dev/null || { echo TRT dry run failed; exit 1; } ./trt_infer --enginemodel.engine --warmup10 --runs100 /dev/null 21 || { echo Infer binary broken; exit 1; } echo ✅ All checks passed它不追求覆盖率只守三个底线环境版本对、模型格式对、引擎可加载、二进制可跑。每次改完代码git add . ./validate.sh git commit -m fix: xxx已成肌肉记忆。这种“小步快跑即时反馈”的节奏比写 100 行 unit test 但半年不跑一次更能守住交付质量。希望帮到你。本文还有配套的精品资源点击获取