如何将PyTorch pth权重文件转ONNX模型后再转换为TensorFlow模型
错误根因
你当前报错的核心原因是模型结构与权重文件不匹配:你用torchvision.models.vgg16()初始化的VGG卷积神经网络结构,和pth文件存储的ViTSTR(基于视觉Transformer的文本识别模型)结构完全不一致,权重张量形状无法对应,因此加载失败。
最简转换方案
前置依赖安装
先安装所需工具库:
pip install torch timm onnx onnx-tf
步骤1:加载匹配的ViTSTR模型结构并导出ONNX
直接用timm库的预定义ViT结构匹配ViTSTR权重,无需手动搭建模型:
import torch import timm import onnx # 1. 初始化匹配的ViT结构:vit_base_patch16_224 和你的权重版本对应 # num_classes设置为ViTSTR对应的字符集大小,官方默认是96(包含数字、字母、符号),如果你的权重自定义过字符集需要对应修改 model = timm.create_model('vit_base_patch16_224', pretrained=False, num_classes=96) # 2. 加载权重 state_dict = torch.load('/content/drive/MyDrive/VitSTR/vitstr_base_patch16_224_aug.pth') # 可选:如果pth是DataParallel训练保存的,需要先去除权重key的`module.`前缀 # new_state_dict = {} # for k, v in state_dict.items(): # name = k.replace('module.', '') # new_state_dict[name] = v # state_dict = new_state_dict # 若仍有key不匹配问题,打印state_dict和model.state_dict()的key列表对照调整即可 model.load_state_dict(state_dict) model.eval() # 3. 导出ONNX,opset选14以上支持Transformer算子 dummy_input = torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy_input, "vitstr.onnx", opset_version=14, do_constant_folding=True, input_names=['input'], output_names=['output'] ) # 验证ONNX文件合法性 onnx_model = onnx.load("vitstr.onnx") onnx.checker.check_model(onnx_model)
步骤2:ONNX转TensorFlow模型
用onnx-tf直接转成TensorFlow支持的SavedModel格式,可直接用于后续微调:
import onnx from onnx_tf.backend import prepare onnx_model = onnx.load("vitstr.onnx") tf_rep = prepare(onnx_model) # 导出SavedModel格式 tf_rep.export_graph("vitstr_tf")
后续微调说明
转换完成后直接用TensorFlow加载vitstr_tf文件夹的SavedModel即可,和普通TF模型的微调方式完全一致。
内容的提问来源于stack exchange,提问作者Roua Rouatbi
相关产品推荐
相关产品推荐

