You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.21 22:18:29