ARTICLE DETAIL

资讯详情

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

AnimeGANv3手机端部署实战:ONNX压缩至5.6MB的完整链路

AnimeGANv3手机端部署实战:ONNX压缩至5.6MB的完整链路 上个月接了个挺有意思的活把AnimeGANv3模型搬到手机上跑动漫风格迁移。输入一张照片手机端直接出二次元风格化的结果整个模型ONNX格式压到5.6MB左右。这个体积放在今天遍地几百MB大模型的环境里算是非常能打了也正因为小普通手机纯CPU推理都能达到可用的帧率。整条链路是PyTorch训练权重转ONNX、再做图简化和精度压缩、最后接到Android端用ONNX Runtime推理。想把自己训练好的模型搬到移动端跑通的开发者或者刚入门模型部署、想完整走一遍流程的朋友这篇实战笔记应该能帮你少踩好几个坑。1. 部署思路先理清楚为什么是这套技术栈一个模型从训练到落地最难的不是训练本身而是“环境变了还能不能好好跑”。在服务器上PyTorch跑得好好的换到手机上问题一堆框架不支持、算子缺失、内存扛不住、体积太大用户不装。所以动手之前先花半天把思路捋顺后面会省很多事。1.1 5.6MB这个数字意味着什么先把这个数字讲透。AnimeGANv3的生成器网络结构本身不算夸张原始PyTorch权重是FP32精度存下来的一份下来大概十几MB。但PyTorch的权重文件里面除了参数还带了一些训练相关的元信息并不能直接塞进手机。真正给移动端用的是ONNX格式一个计算图描述文件把网络结构和权重打包在一起。我这次操作下来导出的FP32版ONNX大概11MB左右经过onnxsim做图层面的冗余消除再把权重转成FP16半精度最后定格在5.6MB。这个量级的实际意义非常直接对Android APK来说包体积增加不到6MB用户无感知模型加载进内存后峰值占用控制在几十MBCPU推理不需要额外申请GPU资源也能维持可用的速度。如果一个风格迁移模型做到这个体积还跑不动那问题基本不在模型而在集成方式上了。1.2 PyTorch到ONNX再到移动端链路怎么串整个部署链路就是一条单向流水线PyTorch模型 → ONNX → 优化压缩 → 移动端推理运行时。第一步把PyTorch模型用torch.onnx.export导出成标准ONNX图第二步对ONNX做两件事一是用onnxsim这类工具把计算图里冗余的节点、没用的分支清掉二是把FP32权重转成FP16甚至INT8来压缩体积第三步在Android/iOS工程里引入ONNX Runtime加载优化后的模型写好前后处理跑通推理。这套链路之所以是行业标配核心是ONNX这个中间格式的“翻译层”作用。模型训练时用的PyTorch也好、TensorFlow也好各有各的算子体系和运行环境移动端不可能全支持。ONNX相当于把模型翻译成一套相对统一的中间表示再由各平台上的推理引擎把它映射到本地的算子库。这个“先统一再分发”的思路比给每个框架单独写移动端推理实要靠谱得多。1.3 技术路线选型ONNX Runtime为主NCNN备选移动端推理引擎我最后选了ONNX Runtime但中间也认真对比过NCNN和MNN。简单说结论如果你的模型是从PyTorch导出的、又希望用一套模型文件同时覆盖Android和iOSONNX Runtime是投入产出比最高的选择——官方直接支持Android和iOSJava/Kotlin和Objective-C接口都齐文档完整遇到问题搜得到答案。NCNN是腾讯开源的老牌移动端推理库对ARM架构优化做得很深CPU推理速度往往比ONNX Runtime还快一点社区里做人脸、风格迁移的小项目也很多。但它有个绕不开的步骤需要用onnx2ncnn把ONNX再转一次NCNN格式转换过程遇到不支持的算子就要手动改网络结构这一下就把工作量抬上去了。MNN也是类似的情况。所以我的建议是先ONNX Runtime把整条链路跑通确认业务效果没问题再根据性能瓶颈决定要不要转到NCNN深度优化。部署这事的铁律是先能跑再谈跑得快。2. 环境准备与工具版本对齐模型部署最怕什么最怕环境版本对不上某个看似玄学的报错其实是版本冲突。我这次先花时间把所有工具的版本固定下来后面基本没遇到“不可描述”的问题。2.1 本机环境与依赖清单先说我的实验环境一台普通Windows笔记本CPU是Intel i7级别没有独立显卡所以导出和验证全程走CPUPython用的3.9。这里要提醒一下导出ONNX不需要GPU纯CPU环境完全够用别在这个环节被“一定要高端显卡”的误解劝退。依赖库版本如下torch1.13.1 torchvision0.14.1 onnx1.13.1 onnxruntime1.15.1 onnxsim0.4.33 numpy1.23.5 opencv-python4.8.0.74版本为什么要固定因为torch.onnx.export生成的图结构和算子版本跟PyTorch版本强相关而ONNX Runtime对ONNX中间表示的opset版本有支持范围。我见过太多人报错后一路升级依赖结果越升越乱最后都不知道是哪个版本出问题。2.2 模型文件与输入输出的确认部署前最重要的一件事把你手里模型的前向逻辑看明白。AnimeGANv3这类生成器网络输入是归一化到0到1区间的RGB图像张量形状为(1, 3, H, W)输出同样是(1, 3, H, W)的张量值域看最后一层激活函数而定。AnimeGAN系列一般用的Tanh输出范围在-1到1之间后处理时得把值重新映射回0到255再存成图片。这里有个容易被忽略的细节模型文件里存的是网络参数但并没有存“输入长什么样”。你必须回到训练代码里确认预处理方式包括缩放尺寸、归一化系数、通道顺序是RGB还是BGR。这决定了导出时的输入设计也决定了移动端写预处理时用什么参数。我前面这些信息就是翻了项目里的inference脚本确认的看清楚之后才动笔写代码。2.3 工具链选择背后的理由ONNX Runtime为什么是主力因为它兼顾了“少折腾”和“可用性能”。它对ONNX算子的支持覆盖面很广尤其图像类模型常见的卷积、归一化、插值、激活函数都有良好支持不需要像NCNN那样做一堆算子适配。而且它提供多种执行提供程序(Execution Provider)Android上可以用NNAPI也可以后续接XNNPACK性能调优有空间。onnxsim这个工具是必须要装的。PyTorch导出的ONNX图经常有大量冗余包括没用的Identity节点、可以合并的Reshape序列、常量折叠没做干净的节点。onnxsim做的事情就是把图结构自动简化和算子融合出来的图更干净对移动端推理引擎也更友好。我已经养成习惯了不管导出什么模型第一步永远是onnxsim。3. 从PyTorch导出ONNX每一步都有讲究导出ONNX是整个部署流程里技术含量最集中的一个环节报错率也最高。这里把每个关键步骤和参数选择都拆开讲。3.1 导出前的模型准备导出第一步不是写torch.onnx.export而是把模型调到eval模式并且把参数全部固定。PyTorch的BatchNorm和Dropout在训练和推理两种模式下的行为完全不同忘了写model.eval()会导致导出的图里带着训练态的逻辑移动端推理的结果直接不对。import torch import onnx from models.generator import Generator # 换成你自己的模型定义路径 # 1. 加载权重并固定到eval模式 model Generator() model.load_state_dict(torch.load(checkpoint/AnimeGANv3.pth, map_locationcpu)) model.eval() # 2. 确认输入尺寸这里以256x256为例 dummy_input torch.randn(1, 3, 256, 256) # 3. 导出ONNX torch.onnx.export( model, dummy_input, animeganv3_fp32.onnx, opset_version12, input_names[input], output_names[output], dynamic_axesNone, )最关键的是第三个参数dummy_input。torch.onnx.export默认走的是TorchScript的tracing机制拿一个真实的张量跑一遍前向把实际执行过的算子记录下来生成计算图。所以dummy_input的形状基本就定义了模型的输入形状。3.2 torch.onnx.export核心参数逐项拆解opset_version这个参数很多人不关心其实影响很大。ONNX Runtime的每个版本都只支持一定范围的opsetopset太高老版本Runtime不认opset太低又可能表达不了某些新算子。我这边用的opset 12覆盖了大部分常见算子同时兼容性也好。如果你的模型里有较新的算子导致导出失败再尝试往上调但移动端的Runtime版本也要跟着升级。input_names和output_names这俩参数是给输入输出起名字移动端加载模型后就是通过这个名字来绑定数据的建议起得直白一点。dynamic_axes这个参数我这次特意没有用原因后面细说。这里有一个经验如果你的模型在导出时因为某个不支持的算子报错最直接的办法是检查模型里用了什么特殊操作比如torch.where、动态shape的F.interpolate、某些索引操作。要么换等价算子要么把输入尺寸固定住。排查的通用思路就是把报错的op在代码里定位出来然后去查PyTorch的ONNX支持矩阵确认替代方案。3.3 导出后的正确性验证导出完不等于万事大吉必须立刻验证。验证方法很简单用同一张输入图分别过PyTorch模型和ONNX Runtime对比输出差异。这个步骤能拦住一大批问题比如算子映射错误、图结构不完整、值域对不上。import numpy as np import onnxruntime as ort import torch def validate_export(pytorch_model, onnx_path, input_tensor): # PyTorch推理 with torch.no_grad(): pt_output pytorch_model(input_tensor).numpy() # ONNX Runtime推理 ort_session ort.InferenceSession(onnx_path, providers[CPUExecutionProvider]) ort_input {ort_session.get_inputs()[0].name: input_tensor.numpy()} ort_output ort_session.run(None, ort_input)[0] diff np.abs(pt_output - ort_output).max() print(f最大绝对误差: {diff:.6f}) return diff最大绝对误差在1e-4量级属于正常范围因为PyTorch和ONNX Runtime底层算子实现有细微数值差异。如果误差到了0.1甚至1的量级说明图结构或者值域处理有问题赶紧回头查。这里顺带说一个动态shape的问题。AnimeGANv3里的F.interpolate在做上采样时卷积层和残差块对输入尺寸没有强制约束理论上支持任意分辨率输入。但动态shape会带来两个麻烦一是ONNX图里的Resize节点变为动态移动端引擎需要做更多shape推断性能受损二是部分量化工具对动态shape支持很差。所以我默认固定为256x256输入换来的是稳定和性能。如果你的业务强需求多尺寸输入建议固定2到3个档位分别导出模型而不是在推理时动态改shape。4. ONNX瘦身与精度压缩5.6MB是怎么压出来的导出拿到的是FP32版ONNX大概11MB。从11MB到5.6MB主要做了两步图结构简化和半精度转换。4.1 先用onnxsim做图简化onnxsim的使用非常简单一条命令就能完成python -m onnxsim animeganv3_fp32.onnx animeganv3_sim.onnx这条命令会做几件事常量折叠把不需要输入就能算出结果的节点提前算掉冗余节点删除去掉一堆对结果没有影响的Identity、Cast、Reshape算子融合把多个可以合并的操作合成一个。做完之后图的节点数量明显减少后续转FP16和移动端加载都有好处。需要注意的是onnxsim成功的前提是模型输入shape是确定的。如果你的模型保留了动态维度onnxsim会跳过很多优化效果大打折扣这也是前面坚持固定输入尺寸的原因之一。4.2 FP16半精度转换的实操FP16转换这块我没有用网上的杂牌脚本而是直接用onnx官方提供的工具链先转Float16再配合onnxruntime的模型处理函数做校验。python -m onnxconverter_common.float16 animeganv3_sim.onnx animeganv3_fp16.onnx不过要提醒一点纯用float16工具直接转有时候会出问题ONNX模型里并不是所有算子都支持FP16输入个别节点会有类型不匹配。如果转换后验证报错一个通用做法是给指定算子保留FP32只让卷积、矩阵乘这些大头权重用FP16。实际上我的经验是5.6MB这个体积目标只要主要权重转成FP16就能达成图里个别算子保持FP32完全不影响整体体积和速度。转换完一定要再跑一遍验证脚本对比FP32和FP16版的最大输出差异。图像生成模型对精度相对宽容FP16的差异肉眼基本看不出来但数值上还是能看到10的-3次方量级的误差这是正常的。4.3 INT8量化要不要做GAN类模型的特殊考量很多教程走到这一步会继续教你做INT8量化把体积再压到2MB以下推理速度还能再快一倍。但我在测试AnimeGANv3时发现INT8量化在这类图像生成模型上并不是免费午餐。用onnxruntime的静态量化工具需要准备一组校准图像from onnxruntime.quantization import quantize_static, QuantType, CalibrationDataReader class AnimeCalibReader(CalibrationDataReader): def __init__(self, image_paths, input_name): # 加载图像做预处理 pass calib_reader AnimeCalibReader(calib_images, input) quantize_static( animeganv3_sim.onnx, animeganv3_int8.onnx, calib_reader, quant_formatQuantType.QOperator, per_channelTrue, )但是静态量化之后输出图像的色彩过渡会出现肉眼可见的断层尤其天空和皮肤这些渐变区域banding效应非常明显。原因在于生成模型输出的是连续色调图像对每一层的激活值精度很敏感INT8的量化误差在多层累积后被放大最终反映在画质上。所以我的结论是如果你的目标是“手机能跑起来且画质可接受”FP16是最佳平衡点这也是5.6MB这个体积的由来。如果业务对画质要求没那么苛刻、更追求速度和体积可以再试INT8但一定要用一批有代表性的实际场景图做校准。另外量化后必须做主观画质对比不能只看指标差异图像生成任务的评价标准最终是人的眼睛。5. 移动端集成与推理手机真正跑起来模型文件准备完毕接下来是重头戏在Android工程里把它跑起来。这块我踩坑最多主要集中在前处理、后处理和Session配置三件事上。5.1 Android工程接入ONNX Runtime我用的Android StudioGradle里加一行依赖就能引入ONNX Runtimeimplementation com.microsoft.onnxruntime:onnxruntime-android:1.15.0强烈建议Android和上面的Python端onnxruntime版本保持同一大版本否则遇到语义不明的问题时很难排查是不是版本差异引起的。模型文件放在assets目录下加载方式如下val env OrtEnvironment.getEnvironment() val sessionOptions OrtSession.SessionOptions() sessionOptions.setOptimizationLevel(OptimizationLevel.ALL_OPT) sessionOptions.setNumThreads(4) val session env.createSession(assetManager.open(animeganv3_fp16.onnx), sessionOptions)setNumThreads这里值得多说一句。手机CPU通常多核但不是线程越多越快。AnimeGANv3这种卷积密集型的模型在4到8线程区间通常能跑出最好成绩超过之后线程调度开销反而拖慢速度。具体最优值建议在目标机型上多测几档不要拍脑袋定。5.2 图像预处理与后处理踩坑最多的部分移动端推理的难度不在调用接口而在张量数据怎么从Bitmap变成模型输入、又从模型输出变回Bitmap。这个环节我见过无数人栽跟头包括我自己。预处理要做的事情Bitmap转RGB数组、缩放到256x256、归一化到0到1、把HWC的布局转换成CHW。fun bitmapToFloatInput(bitmap: Bitmap): FloatArray { val scaled Bitmap.createScaledBitmap(bitmap, 256, 256, true) val intPixels IntArray(256 * 256) scaled.getPixels(intPixels, 0, 256, 0, 0, 256, 256) val input FloatArray(1 * 3 * 256 * 256) for (i in intPixels.indices) { val pixel intPixels[i] val r ((pixel shr 16) and 0xFF) / 255.0f val g ((pixel shr 8) and 0xFF) / 255.0f val b (pixel and 0xFF) / 255.0f input[i] r input[256 * 256 i] g input[2 * 256 * 256 i] b } return input }这里的坑在通道布局。PyTorch模型默认CHW而Android的Bitmap拿到的像素数据是HWC的必须显式拆开重排。我一开始图省事直接把RGB依次排成HWC送给模型出来的图像整个颜色和结构都是乱的。这个问题排查了很久最后是跟Python端验证脚本里的预处理逐行对比才发现的。推理调用本身很简洁val inputTensor OnnxTensor.createTensor(env, inputFloatArray, longArrayOf(1, 3, 256, 256)) val outputs session.run(mapOf(input to inputTensor)) val outputArray outputs[0].value as Array*后处理则是预处理的逆过程把输出张量从CHW拆回HWC值域映射回0到255再逐像素写入Bitmap。这里特别要注意输出值域。如果模型最后一层是Tanh输出在-1到1之间直接把负值强行转成无符号字节会出现黑色斑块必须先用公式(pixel 1.0f) / 2.0f映射到0到1再乘255。这一步错了输出图像会一片惨不忍睹。5.3 性能实测与调优方向我在一台骁龙7系列的中端机上做了测试固定256x256输入FP16模型4线程CPU推理单帧耗时在250到400毫秒之间。这个速度做实时视频流处理还不够但做“拍照→出图”这种交互是完全可用的用户等待感不明显。如果想让速度更进一步有几个方向可以试。一是开NNAPI执行提供程序把计算委托给手机的NPU或GPU代码改动很少但要注意不同厂商的驱动对算子支持不一致可能有兼容性问题。二是把输入分辨率从256降到192速度几乎线性提升画质损失在可接受范围内。三是回到NCNN路线做算子级优化这个工作量最大但针对ARM架构的卷积优化确实能榨出更多性能。我的实践顺序是先用ONNX Runtime默认CPU跑通再开NNAPI对比效果最后根据业务需求决定要不要转NCNN深调。实测顺序很重要别一上来就ALL-IN某个高端优化方案很容易被兼容性问题耗掉大量时间。6. 常见问题与排查技巧实录部署这种东西看别人写一万句经验不如自己踩一次坑。但有些坑的排查成本实在太高我把这轮遇到的典型问题和排查思路整理出来遇到类似情况可以直接照方抓药。6.1 导出阶段典型报错与解法导出时报错最常见的一类是“Unsupported operator”或者“Failed to export”。这类问题八九成出在模型内部用了PyTorch的某些动态操作上比如torch.where、torch.argmax、非固定维度的view/reshape。我的建议是不要在模型代码里兜圈子直接在代码里定位报错的行换成ONNX支持的等价实现。比如有些InstanceNorm相关的自定义逻辑可以用标准nn.InstanceNorm2d替代。第二个常见问题是导出成功但onnxsim报错“shape inference failed”。这种一般是图里存在shape不确定的节点尤其是生成模型里的Resize或者插值操作。解决办法是回到导出那一步确认dummy_input形状固定并把dynamic_axes去掉。如果还是报错试试opset_version改到13或14新版opset对shape推断的支持更好。6.2 画质异常排查如果模型能跑、输出也有内容但颜色明显不对排查思路按优先级来先检查值域映射是否做了Tanh的-1到1到0到255转换再检查通道顺序输出的是RGB还是BGR很多图片库默认BGR存储直接转Bitmap会红蓝互换最后检查归一化方式是除以255还是减去均值再除以方差。如果是整体模糊或出现奇怪的网格纹理多半是模型输入尺寸和训练时不匹配或者插值算子在转换时类型出了问题。把输入尺寸跟项目仓库里inference脚本对齐一般能解决。6.3 手机端延迟与内存问题延迟高先别急着上NNAPI先用多线程档位和输入尺寸做一轮简单调优。我见过同一个模型在同一台手机上从2线程调到4线程速度提升接近一倍再调8线程反而变慢。要做多档位对比别想当然。内存方面注意有些后处理方式会同时存在多个Bitmap副本256x256还好如果哪天换成高分辨率Bitmap内存和FloatArray内存会翻好几倍。我习惯在处理链路上尽早释放中间对象避免GC抖动导致的推理卡顿。下面是这轮遇到的重点问题速查问题现象可能原因排查与解决导出报Unsupported operator模型用了ONNX不支持的特殊算子定位报错行替换为等价算子onnxsim报shape推断失败图里存在动态shape节点固定输入shape去掉dynamic_axesONNX和PyTorch输出差距大算子映射错误或预处理不一致用同一输入逐层对比检查预处理输出图像红蓝互换通道顺序RGB/BGR搞混后处理时显式转换通道顺序输出有黑色斑块没处理Tanh的-1到1值域用(x1)/2映射回0到1再输出手机端单帧耗时过长线程数、精度、输入尺寸未调优对比多档线程考虑FP16或降低分辨率INT8量化后有色带生成模型对激活精度敏感回归FP16或改进校准数据集最后分享一点我的体会这次部署做完最大的感受是“模型小就是王道”。AnimeGANv3能够轻松跑到手机上归根结底是作者在设计网络结构时就把轻量化考虑进去了。作为部署方我们能做的优化——剪枝、量化、图简化——都是在既定结构上做文章上限摆在那里。选模型的时候如果注定要上移动端一开始就不要选几百MB的大网络能在源头省下的事不要拖到部署阶段用加班来还。另外强烈建议每一个做部署的人养成写验证脚本的习惯每一轮转换都跑一遍输出对比这个习惯能帮你拦住绝大多数低级错误也让你在排查问题时有据可依。
返回列表