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

如何将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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.07 09:24:01