最近在公司做一个小项目,需要把训练好的BERT模型部署到生产环境。一开始想用PyTorch自带的JIT trace,结果动态shape直接给我整不会了。后来转ONNX,又遇到LayerNorm和GELU算子不支持,各种改代码绕路,最后导出的模型精度还掉了0.3%。我心态有点崩,感觉自己在瞎折腾。想问问各位老哥,现在实际工业落地,如果不想用TensorRT(公司没N卡)也不想上服务端框架(太复杂),有没有比较稳的部署方案?或者是我用ONNX的方式不对?求指条明路,谢谢了!
PyTorch转JIT被坑惨了,ONNX导出也报算子错误,大佬们现在部署都用啥?
全部回复
共 68 条ONNX这个坑我也踩过,GELU和LayerNorm确实是重灾区,后来我们直接换成了ONNX Runtime的优化算子集,精度掉0.3%大概率是某些op被替换成低精度实现导致的,你可以检查下导出时的opset版本和优化级别。如果不想折腾,其实可以试试把模型转成ONNX后直接用onnxruntime的C++接口,动态shape用symbolic shape配合dynamic axes能解决大部分问题,别用JIT那条路了。你现在的部署环境是CPU还是自研芯片?如果对延迟不敏感,直接用PyTorch的torchscript加上torch.set_grad_enabled(False)加eval模式,配合torch.compile可能都比绕ONNX省心。
跟你情况差不多,之前也被JIT的动态shape坑过,最后直接换了思路用ONNX Runtime配合动态轴,算子问题靠升级版本和改导出参数绕过去了,精度没掉那么多。你要是没N卡,其实CPU上ONNX Runtime已经挺能打了,关键是把opset调到13以上,很多新算子都支持了。要是还卡在LayerNorm,试试把模型里那些自定义实现改成PyTorch原生的,导出会顺很多。
说实话ONNX这坑我也踩过,BERT的LayerNorm报错基本是opset版本和优化器没对齐,试试把opset升到17以上然后关掉graph优化,精度掉0.3%大概率是动态shape导致某些层被折叠了。你要是图省事,直接用CTranslate2吧,专门优化transformer的,CPU上跑得飞快还支持动态shape,导出也简单。不过得注意它只支持fp32和int8,精度需要自己验证下。
试试转成TorchScript时把dynamic_axes配全,或者直接用CTranslate2,BERT支持贼稳,精度也不掉。
试试把动态shape固定成几个档位再导出,精度掉了可能是前面改算子改出问题,建议分步验证。
ONNX绕不过去就试试ONNXRuntime直接上,自带优化比瞎折腾算子强,我上次就这么救回来的。
试试把动态轴固定到最大长度,加padding和mask,精度损失能小很多,ONNX对静态shape友好。
试试ONNX Runtime的DNNL/OpenVINO后端吧,动态shape和算子兼容比直接导出稳多了。
看到这个经历太真实了,BERT转ONNX那堆算子问题我当初也踩过,GELU还好说自己拼个近似,LayerNorm才是真折磨。不过精度掉0.3%大概率不是算子问题,是动态shape导致某些维度被固定后数值路径变了,建议你试试把输入padding到固定长度再trace,虽然浪费点显存但能省一堆麻烦。另外如果你只是CPU部署,其实不用死磕ONNX,直接上Intel的OpenVINO,它对Transformer系列优化得很透,而且自带LayerNorm融合,转换脚本写个十几行就能跑。要是连OpenVINO都嫌重,还有个野路子:把模型导出成TorchScript后关掉shape校验,用torch.jit.optimize_for_inference配合固定长度假输入,虽然丑但能用。最后提醒下,如果精度敏感,转完一定要用同一批测试集对比每层输出,别只看最终指标,误差可能早就在中间层积累了。