ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

PyTorch模型转PaddleLite集成Android端侧推理全流程实战

PyTorch模型转PaddleLite集成Android端侧推理全流程实战 最近这段时间一直在折腾一件事把 PyTorch 训练好的模型转成 PaddleLite 模型再集成到 Android 应用里跑起来。整个链路从模型导出、格式转换到 Android 端 JNI 调用每一步都藏着不少坑。尤其是 PyTorch 转 PaddleLite 并没有官方的一条龙方案中间基本绕不开 ONNX转换完还得靠 opt 工具做一轮优化才能得到真正适合端侧加载的 .nb 文件。如果你也正准备把自训练的小模型塞进手机 App这篇全流程记录应该能帮你少走很多弯路。这套方案比较适合以下情况模型结构不算特别冷门、需要离线推理、想在 Android 上保持较低 CPU 占用和可控包体。我会按实际操作顺序把整体设计方案、环境准备、ONNX 导出、X2Paddle 转 Paddle、opt 优化、Android 工程集成以及最后的高频问题和性能调优全部展开。所有命令和代码我都尽量按能直接复制的标准来写版本建议也会标注清楚照做即可。1. 整体方案设计为什么选择 PaddleLite 这条链路1.1 比来比去PaddleLite 更适合端侧落地Android 端跑深度学习模型可选方案其实不少PyTorch Mobile、ONNX Runtime Mobile、NCNN、MNN 都有人用。我自己一开始也考虑过 PyTorch Mobile毕竟模型本来就是 PyTorch 训练出来的理论上最省事只要导出 TorchScript 就能加载。但实际用下来PyTorch Mobile 在 Android 上对算子支持、包体大小、推理性能的平衡并不理想尤其是某些自定义前处理逻辑转 TorchScript 时很容易报“无法跟踪”的错。ONNX Runtime Mobile 也不错但需要自己对算子做裁剪和定制集成成本和后期维护成本都不低。NCNN 在性能和轻量性上很突出可如果模型里有比较特殊的算子转换工具链会让人头大。PaddleLite 吸引我的点主要在三点第一Paddle 官方提供了 X2Paddle 转换工具对 ONNX 的支持比较完善PyTorch 导出的模型只要不是特别离谱的算子基本都能顺利转过去。第二PaddleLite 针对 ARM CPU、Android 平台做了大量优化支持多线程、功耗模式设置、CPU/GPU 异构还内置了量化方案。第三部署包相对干净一个 .nb 模型文件加几个 .so就能在 App 里把模型跑起来适合不想引入一堆运行时依赖的项目。当然选型不能只看优点。PaddleLite 社区更新节奏确实不算快文档有时候写得比较散遇到问题经常要靠翻源码和 GitHub issue。但对我这种需要快速在 Android 端做验证的场景PaddleLite 的转换链路和集成体验在现有框架里算是很顺的。1.2 转换链路设计PyTorch - ONNX - Paddle - PaddleLitePaddleLite 官方没有提供“PyTorch 直接转 PaddleLite”的完整工具实际可行的链路是分步转换PyTorch 模型先用torch.onnx.export导出为 ONNX 格式再用 X2Paddle 把 ONNX 转成 PaddlePaddle 模型格式得到model.pdmodel和model.pdiparams接着用 PaddleLite 的 opt 工具把 Paddle 模型优化、量化成端侧推理文件.nb最后在 Android 工程里通过 PaddleLite 的 Java API 加载.nb文件执行推理为什么中间非要加 ONNX 这一层因为没有直接的 PyTorch 转 Paddle 路线X2Paddle 虽然号称支持 PyTorch但走 ONNX 中转的兼容性和报错可控性最好。ONNX 本身就是一种标准中间表示导出后能先用 onnx.checker 验证一遍模型结构就算后面转换出问题也容易定位是在哪个算子、哪个维度上出的问题。另外要提醒一点既然最终目标是 Android 端模型输入尺寸最好在一开始就确定下来。如果训练时用的是动态输入导出 ONNX 时就算强行固定 batch 维度后续 opt 优化也可能因为 dynamic shape 信息残留而失败。我建议直接把 batch 固定为 1输入尺寸固定成部署时预期的大小这样能省掉很多 shape 相关的坑。2. 环境准备与模型导出2.1 环境清单与安装步骤转换过程不需要 GPUCPU 环境下就能完成。为了避免把本机 Python 环境搞乱我习惯用 conda 单独建一个环境专门给 PyTorch、Paddle 和 X2Paddle 用。组件版本建议作用Python3.8 或 3.9转换脚本运行环境PyTorch1.10 - 2.0训练/加载模型并导出 ONNXPaddlePaddle2.4 及以上 CPU 版运行 X2Paddle 转换后的模型X2Paddle1.4.0 及以上把 ONNX 转成 Paddle 模型PaddleLite2.12 及以上提供 opt 工具生成 .nbAndroid Studio当前稳定版编写和编译 Android 工程安装命令可以直接复制conda create -n torch2paddle python3.8 conda activate torch2paddle pip install torch1.13.1 torchvision0.14.1 pip install paddlepaddle2.4.2 pip install x2paddle onnx onnxruntime onnx-simplifier这里特别说一下版本搭配。PyTorch 版本别追太新2.0 之后导出 ONNX 的默认算子集版本偏高X2Paddle 不一定能完全吃下如果导出后转换报“Unsupported onnx op”先检查是不是 opset 版本太高。PaddlePaddle 只需要 CPU 版因为转换过程只是做一次网络结构解析和权重转换不需要反向传播CPU 版足够。2.2 PyTorch 导出 ONNX 的实操细节我以一个常见的分类模型为例假设你已经把训练好的权重存成了model.pth接下来要做的是先加载模型切到 eval 模式再构造一个 dummy input。import torch import torch.nn as nn # 1. 加载模型 model MyModel() model.load_state_dict(torch.load(model.pth, map_locationcpu)) model.eval() # 2. 构造固定尺寸的输入 dummy_input torch.randn(1, 3, 224, 224) # 3. 导出 ONNX with torch.no_grad(): torch.onnx.export( model, dummy_input, model.onnx, opset_version11, input_names[input], output_names[output], dynamic_axesNone )这段代码里几个参数要解释一下。opset_version我建议用 11ONNX opset 11 对大部分常规算子支持成熟Paddle 侧转换也比较友好。input_names和output_names不是随便写的后面如果用 onnxruntime 做输出对比会靠这两个名字获取输入输出张量。dynamic_axes我直接设成None也就是完全不启用动态维度这样导出的模型天然是固定 shape后续 opt 优化最省心。导出完成后强烈建议立刻做一次完整性检查import onnx model onnx.load(model.onnx) onnx.checker.check_model(model) print(ONNX model is valid)如果检查时弹出 warning比如某个节点的输出没被使用这种通常不影响转换但如果是 error就要回去检查模型结构。还有一个常用的小技巧是用 onnx-simplifier 做一次简化它能合并一些冗余的 Transpose、Reshape让 ONNX 结构更干净也能降低 X2Paddle 报错的概率。python -m onnxsim model.onnx model_sim.onnx之后用简化后的model_sim.onnx继续后续转换。3. 模型转换与优化3.1 用 X2Paddle 把 ONNX 转成 Paddle 模型X2Paddle 提供了命令行和 Python API 两种方式。我个人习惯用命令行因为输出信息更直观。进入刚才创建的 conda 环境执行x2paddle --frameworkonnx \ --modelmodel_sim.onnx \ --save_dirpd_model成功后pd_model目录下会出现inference_model文件夹里面包含两个核心文件model.pdmodel存储网络结构model.pdiparams存储权重参数在转换前最好先确认这两个文件的生成时间如果时间不对或者文件大小是 0说明转换没有真正完成。命令行通常会打印转换日志看到结束字样才算成功。如果你希望在自己的 Python 脚本里调用也可以这么写from x2paddle.convert import onnx2paddle onnx2paddle(model_sim.onnx, pd_model)转换成 Paddle 模型后我强烈建议先用 onnxruntime 和 PaddlePaddle 各跑一遍对比输出结果是否一致。这一步虽然多花几分钟但能提前发现精度问题不至于到 Android 端才抓瞎。对比方法比较简单给同一个输入分别在 onnxruntime 和 PaddlePaddle 里前向推理计算输出的最大绝对误差。如果误差小于 1e-4基本可以认为转换没问题。如果误差较大优先检查 BN 层、Normalize 层是否被转换器正确映射以及权重是否加载成功。3.2 用 opt 工具生成 PaddleLite 端侧模型拿到model.pdmodel和model.pdiparams后还需要经过 PaddleLite 的 opt 工具做一次优化才能得到适合 Android 端加载的.nb文件。opt 工具需要从 PaddleLite release 页面下载对应平台的预编译程序我是在 Linux 下执行的./opt --model_dirpd_model/inference_model \ --model_filemodel.pdmodel \ --param_filemodel.pdiparams \ --optimize_outandroid_model \ --optimize_out_typenaive_buffer \ --valid_targetsarm参数说明整理成表格参数含义建议--model_dirPaddle 模型所在目录使用 X2Paddle 输出的 inference_model--model_file网络结构文件名model.pdmodel--param_file权重文件名model.pdiparams--optimize_out输出文件前缀选一个见名知意的名字--optimize_out_type输出格式用 naive_buffer最终只有一个 .nb 文件--valid_targets目标平台Android 端填 arm成功后会在当前目录生成android_model.nb。.nb文件就是最终要放进 Android Assets 的推理模型文件。opt 工具做的工作包括算子融合、常量折叠、内存布局优化、算子映射等类似编译器的优化层能让模型在 ARM 上跑得更快。这里有个常见陷阱如果在导出 ONNX 时启用了dynamic_axes转换出的 Paddle 模型里可能存在动态 shape 信息opt 阶段很容易报 shape 相关的错误。遇到这种情况最好的办法不是硬调 opt而是回到 PyTorch 导出步骤把dynamic_axes去掉重新导出、转换。另外opt 工具是跨平台的在 x86 的 Linux 机器上也可以生成 arm 目标跟最终运行的硬件无关。4. 集成到 Android 工程4.1 Android 工程配置与依赖导入PaddleLite 在 Android 端以 JAR 包 动态库的形式提供。去 PaddleLite release 页面下载 Android 推理库解压后能看到PaddlePredictor.jar和jniLibs目录。把PaddlePredictor.jar拷贝到app/libs把 jniLibs 里的arm64-v8a等目录拷贝到app/src/main/jniLibs。然后在app/build.gradle里做几件事android { defaultConfig { ndk { abiFilters arm64-v8a } } sourceSets { main { jniLibs.srcDirs [src/main/jniLibs] } } } dependencies { implementation files(libs/PaddlePredictor.jar) }我建议只保留arm64-v8a现在市面上的 Android 手机基本都是 64 位 ARM保留 armv7 会白白增加包体。如果产品还要覆盖老设备再加一层armeabi-v7a对应的.nb文件是通用的不用重新生成但 .so 必须对应补齐。模型文件android_model.nb放入app/src/main/assets。assets 里的文件不能直接以绝对路径的方式传给 PaddleLite需要先复制到 App 私有目录才能拿到可读的文件路径。private String copyModelToCache(Context context, String assetName) throws IOException { File modelFile new File(context.getCacheDir(), assetName); try (InputStream is context.getAssets().open(assetName); OutputStream os new FileOutputStream(modelFile)) { byte[] buffer new byte[8192]; int len; while ((len is.read(buffer)) ! -1) { os.write(buffer, 0, len); } } return modelFile.getAbsolutePath(); }4.2 加载模型并执行推理PaddleLite 的 Java API 使用起来很直接。先初始化MobileConfig再创建PaddlePredictor之后每次推理只需要往输入 Tensor 里塞数据调用run()再从输出 Tensor 里取结果。import com.baidu.paddle.lite.MobileConfig; import com.baidu.paddle.lite.PaddlePredictor; import com.baidu.paddle.lite.Tensor; public class PaddleLiteEngine { private PaddlePredictor predictor; public PaddleLiteEngine(String modelPath, int threads) { MobileConfig config new MobileConfig(); config.setModelFromFile(modelPath); config.setThreads(threads); predictor PaddlePredictor.createPaddlePredictor(config); } public float[] predict(float[] inputData, long[] inputShape) { Tensor input predictor.getInput(0); input.resize(inputShape); input.setData(inputData); predictor.run(); Tensor output predictor.getOutput(0); return output.getFloatData(); } }有一个细节如果你不确定当前模型的输入输出名字可以用predictor.getInputNames()和predictor.getOutputNames()先打印出来。大多数情况下顺序跟我们导出 ONNX 时设置的input_names、output_names保持一致但经过 opt 优化后名字不一定原样保留打印一下最稳。图像预处理是最容易出错的地方。模型训练时如果用了 ImageNet 的均值和方差Android 端做推理前也必须做同样的归一化。Bitmap 默认是 RGBA 格式要转成 NCHW 的 float 数组通道顺序还要从 RGBA 调整成 RGB。Kotlin 代码大概是这样的val pixels IntArray(width * height) bitmap.getPixels(pixels, 0, width, 0, 0, width, height) val inputData FloatArray(3 * width * height) for (i in pixels.indices) { val p pixels[i] val r ((p shr 16) and 0xFF) / 255.0f val g ((p shr 8) and 0xFF) / 255.0f val b (p and 0xFF) / 255.0f inputData[i] (r - 0.485f) / 0.229f inputData[width * height i] (g - 0.456f) / 0.224f inputData[2 * width * height i] (b - 0.406f) / 0.225f }还要记住PaddlePredictor不要每次预测都重新创建初始化一次后复用即可。否则每次加载模型都会重新做内存分配和参数初始化App 很容易卡顿甚至内存溢出。我一般会把引擎包成单例App 启动后异步初始化第一次推理前再预热一次。5. 高频问题与性能调优5.1 高频坑排查速查表现象可能原因处理方式x2paddle 转换时报 Unsupported opONNX opset 版本太高或算子过于冷门降低 opset 到 11或替换成标准算子opt 生成 .nb 时报 shape 相关错误模型存在动态 shape去掉 dynamic_axes固定 batch 和输入尺寸重新导出Android 端 predictor 为 null.nb 文件目标平台不对、so 缺失、路径错误确认 valid_targets 是 arm检查 jniLibs 和 modelPath运行崩溃 UnsatisfiedLinkErrorjar 和 so 版本不一致统一用同一版本 PaddleLite 的 jar 和 so推理结果与 PyTorch 差异巨大预处理不一致或模型转换精度异常先用全 1 输入对比输出再逐层定位App 启动后首次推理很慢模型加载和初始化耗时启动时预创建 predictor并做一次预热推理这里我特别想强调的是版本一致性。很多人下载了最新版的PaddlePredictor.jar却用了老版本的libpaddle_lite_jni.so最后崩在UnsatisfiedLinkError上。排查起来很浪费时间。建议从 release 页面下载某个固定版本后把 jar 和 so 全部放在同一个工程目录里标注好版本号避免随手乱换。5.2 性能调优与实测心得模型跑通之后就是调性能。我的经验是按照这个顺序来固定输入 shape减少动态 shape 推导。设置合理的线程数。线程数太多会带来调度开销太少又吃不满多核。我在测试机上试下来双核到四核之间往往最优。选择合适的功耗模式。PaddleLite 的MobileConfig支持功耗模式设置比如LITE_POWER_HIGH偏性能LITE_POWER_LOW偏省电应用场景对功耗敏感的话可以调一下。再往下是量化。先把 FP32 模型整条链路跑通确认精度没问题再用 PaddleLite 的 INT8 量化方案压缩模型大小和加速推理。量化会损失一点精度所以量化后要在真实样本上重新验证。关于预处理Bitmap 转 float 数组的操作对性能影响也很大尤其是大图。建议在后台线程做避免阻塞 UI。另外如果模型输入分辨率不是很大可以提前把 Bitmap resize 到固定尺寸避免每次推理都在 resize 上花时间。站在实际项目的角度端侧部署最大的成本不是推理而是模型转换和后处理调试。先用一个极小的 demo 模型把全链路跑通再换真实模型会顺畅很多。我在实际项目中踩得最多的坑是版本搭配问题不同版本的 PyTorch、X2Paddle、PaddleLite 之间偶尔会有兼容性问题。建议把转换环境用 conda 单独隔离同时把每个工具的版本记到项目 README 里方便以后回查。还有一些模型量化压缩的进阶玩法等 FP32 版本稳定运行后可以再慢慢尝试效果往往比盲目优化代码更明显。
返回列表