最近在部署一个BERT-like的小模型(大概110M参数),用PyTorch动态图推理单条样本大概12ms,想着转成ONNX用TensorRT加速一下。结果导出很顺利,但用onnxruntime-gpu跑下来居然要18ms,反而慢了50%?我确认了输入输出都是动态shape,也开了graph optimization level,甚至试了固定seq_len到128,依然没改善。查了下网上说可能是算子不支持导致fallback到CPU,但看profiler又显示在CUDA上。有没有大佬遇到过类似情况?是TensorRT不擅长处理Transformer结构,还是我动态shape的设置有问题?或者量化INT8才是正解?求指点,谢谢!
PyTorch转ONNX后推理速度反而变慢,是姿势不对还是框架问题?
全部回复
共 85 条说实话你这个情况我太熟了,之前调一个RoBERTa也踩过一模一样的坑。ONNX导出后大概率是embedding层或者attention里的某些算子被拆碎了,TensorRT虽然显示在跑CUDA,但实际可能走了插件或者低效的融合路径,尤其是动态shape下,显存分配和kernel launch开销会吃掉不少性能。我建议你先用onnxruntime的CUDA EP单独测一下,排除TensorRT的干扰,看看是不是onnx模型本身就有问题。另外你试过用onnxsim简化图结构吗?有时候官方导出的图冗余节点特别多,优化器根本来不及融合。还有个思路是直接用TensorRT的Python API从PyTorch导出engine,跳过ONNX这层,很多支持不好的算子能自动处理掉。不过说实话,110M的模型单条12ms已经很快了,要是生产环境对延迟没那么敏感,不如直接上PyTorch的torch.compile或者用C++ libtorch推理,省得折腾转换。你要实在想加速,试试把seq_len固定成你实际用的最大值,然后关掉动态shape,有时候动态shape的开销比算子不优化还致命。
这问题我上个月刚踩过一遍,110M的BERT转TRT反而变慢太正常了。你注意看下onnxruntime的execution provider是不是真的用上了TRT,有时候CUDA EP和TRT EP混着跑,算子图被拆得稀碎,反而比纯PyTorch的算子融合差。我那次是发现LayerNorm和Gelu被拆成十几个小kernel,每个都有kernel launch开销,12ms变18ms基本就是这浪费的。
动态shape确实是另一个坑,你固定到128没改善的话,建议看看是不是attention的score矩阵被当成非batch维度去优化了。我后来是直接转成TensorRT的engine文件,不走onnxruntime,用trtexec调了层级的精度和算法选择,才压到8ms。不过你这模型规模,我怀疑是不是显存带宽卡住了,PyTorch的eager模式有些操作反而能走cudnn的融合路径,ONNX导出的图结构不一定能触发同样的底层优化。
还有个骚操作你可以试试,把模型里所有reshape和transpose手动合并掉,有时候onnx的shape推理会生成一堆冗余的拷贝节点。另外确认下你onnxruntime版本,1.16和1.17对BERT的优化差异很大,我换了个版本直接快了30%。如果还不行,干脆用FasterTransformer那套推理逻辑,专门吃Transformer结构,比TRT稳多了。
这情况我也踩过坑,110M的BERT转TRT反而变慢,大概率不是动态shape的锅,而是算子融合没吃透。你试试把onnxruntime的execution_mode设成ORT_SEQUENTIAL,然后开enable_cpu_mem_arena=False,有时候内存池分配策略会影响小batch的延迟。另外你说profiler显示在CUDA上,但有没有留意过是不是有算子被拆成了多个小kernel?比如LayerNorm在ONNX里可能被拆成多个ReduceMean和Sub,TRT对这类小算子反而有调度开销。我自己的经验是,对于BERT这种结构,直接用TensorRT的官方BERT插件或者换成FasterTransformer会好很多,ONNX中间层优化本来就是玄学。还有个思路,你可以试试把模型量化到FP16再转TRT,有时候精度掉一点点但速度能翻倍,特别是你这种110M的规模,显存带宽才是瓶颈。最后想问下,你测速的时候有没有做warmup?TRT第一次推理会做engine build和cudnn调优,不排除你测的是包含初始化时间在内的数据。
110M的bert转trt反而变慢,这情况我碰到过好几次,大概率不是框架问题,是优化没吃到点上。onnxruntime-gpu走的是它自己的cuda kernel,跟tensorrt完全是两码事,你现在这个对比其实是在拿pytorch的eager模式跟ort比,中间还隔了一层onnx的图优化,本身就有损耗。我建议你先别急着固定seq_len,把dynamic axes去掉,用onnx-simplifier过一遍图,再看下是不是有gather、where这类算子被拆成了多个小算子,bert里这种情况特别多,每个小算子都启动一次kernel,延迟就上去了。另外你提到profiler显示在cuda上,但有没有看具体每个节点的时间?很可能大部分时间都耗在reshape和transpose这类内存拷贝上,而不是matmul。我之前有个类似模型,把onnx的opset版本调到13以上,再配合trt的fp16,直接冲到了6ms,但纯ort怎么调都压不进10ms。所以你要是真追求低延迟,建议直接跳过onnx,用torchscript转trt,或者干脆用faster-transformer那套,专门优化bert的。还有个细节,你测速的时候有没有做warmup?onnxruntime第一次推理会初始化cuda context,那一下能占到5-10ms,不排除你测的就是这个。
我之前也踩过类似的坑,110M的BERT转TRT反而变慢大概率是动态shape或者算子融合没吃透。建议先试试把attention mask和token type ids这些输入全部固定成具体值,用trtexec的minShapes/optShapes/maxShapes三档配上看看,很多时候是显存分配和kernel选择在动态维度下太保守。另外你profiler显示在CUDA上不代表没走CPU fallback,可以加个ENABLE_MS_CUDA_LAZY_LOADING环境变量再测,或者直接看onnxruntime的session options里有没有把cudnn_conv_algo_search设为HEURISTIC。还有一个思路是干脆绕开TensorRT,试下onnxruntime的transformers优化工具,它针对BERT类模型有专门的fusion,有时候比TRT还快。