PyTorch预训练模型转TFLite遇多余Transpose层问题求助
PyTorch预训练模型转TFLite时冗余Transpose层问题解决
我有PyTorch预训练模型(如ResNet50、MobileViT),需要将.pth格式转换为TFLite格式。当前转换流程为:先将.pth模型转为ONNX格式(此时无大量Transpose层),再通过onnx-tf转换为TensorFlow模型,最后转为TFLite时出现大量意外Transpose层,这会对模型在MCU上的部署产生不利影响。
最小复现代码
import torch import onnx from onnx_tf.backend import prepare import tensorflow as tf import torchvision import os import time import timm from pytorch_pretrained_vit import ViT def pth_to_onnx(output_path): torch_model = timm.create_model('resnet50', pretrained=True) x = torch.randn(1, 3, 224, 224) export_onnx_file = output_path torch.onnx.export(torch_model, x, export_onnx_file, opset_version=11, do_constant_folding=True, input_names=["input"], output_names=["output"], dynamic_axes={"input": {0: "batch_size"}, "output": {0: "batch_size"}} ) def onnx_to_pb(output_path): model = onnx.load(output_path) tf_rep = prepare(model) tf_rep.export_graph('resnet50') if __name__=='__main__': output_path = "resnet50.onnx" pth_to_onnx(output_path) onnx_to_pb(output_path) saved_model_dir = 'resnet50' converter = tf.lite.TFLiteConverter.from_saved_model(saved_model_dir) converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS] tflite_model = converter.convert() # Save the model. with open('resnet50.tflite', 'wb') as f: f.write(tflite_model)
问题根源
PyTorch默认采用NCHW(批量-通道-高度-宽度)的张量格式,而TensorFlow/TFLite默认使用NHWC格式。onnx-tf在转换过程中未做最优的格式适配,导致插入大量Transpose层来转换通道顺序,增加了MCU部署时的计算开销和延迟。
优化方案
1. 导出ONNX时适配NHWC格式
在PyTorch导出ONNX前,用包装类统一输入输出的通道格式,减少后续转换的格式转换操作:
def pth_to_onnx(output_path): torch_model = timm.create_model('resnet50', pretrained=True) torch_model.eval() # 包装模型,适配NHWC输入输出 class NHWCAdapter(torch.nn.Module): def __init__(self, model): super().__init__() self.model = model def forward(self, x): # NHWC转NCHW供PyTorch模型处理 x = x.permute(0, 3, 1, 2) x = self.model(x) return x wrapped_model = NHWCAdapter(torch_model) # 使用NHWC格式的输入张量导出 x = torch.randn(1, 224, 224, 3) torch.onnx.export(wrapped_model, x, output_path, opset_version=13, # 更高版本opset支持更多优化算子 do_constant_folding=True, input_names=["input"], output_names=["output"], dynamic_axes={"input": {0: "batch_size"}, "output": {0: "batch_size"}} )
2. 替换onnx-tf为tf2onnx工具
onnx-tf的转换优化能力有限,推荐使用tf2onnx(微软维护的ONNX与TensorFlow互转工具)来转换ONNX到TensorFlow SavedModel,能有效减少冗余Transpose层:
首先安装依赖:
pip install tf2onnx onnxruntime
然后执行转换命令:
# 从ONNX直接转换为TensorFlow SavedModel python -m tf2onnx.convert --input resnet50.onnx --output resnet50_tf.pb --saved-model resnet50_tf
再用TFLite转换器处理这个SavedModel即可。
3. TFLite转换时启用优化
开启TFLite的默认优化,让转换器自动合并冗余的Transpose层:
saved_model_dir = 'resnet50_tf' converter = tf.lite.TFLiteConverter.from_saved_model(saved_model_dir) # 启用默认优化,包含冗余算子合并 converter.optimizations = [tf.lite.Optimize.DEFAULT] converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS] tflite_model = converter.convert() with open('resnet50_opt.tflite', 'wb') as f: f.write(tflite_model)
4. 跳过ONNX,直接从PyTorch转TFLite
使用torch2tf工具链直接完成PyTorch到TensorFlow的转换,避免中间格式转换带来的问题:
import torch import tensorflow as tf from torch2tf import convert torch_model = timm.create_model('resnet50', pretrained=True) torch_model.eval() # 定义输入形状 input_shape = (1, 3, 224, 224) x = torch.randn(input_shape) # 转换为TensorFlow模型 tf_model = convert(torch_model, input_shape=input_shape) # 保存为SavedModel tf.saved_model.save(tf_model, 'resnet50_tf_direct') # 转换为TFLite converter = tf.lite.TFLiteConverter.from_saved_model('resnet50_tf_direct') converter.optimizations = [tf.lite.Optimize.DEFAULT] tflite_model = converter.convert() with open('resnet50_direct.tflite', 'wb') as f: f.write(tflite_model)
内容的提问来源于stack exchange,提问作者Lyn22
相关产品推荐
相关产品推荐

