最新国产好看的视频,伊人天堂AV在线,国产Aaaaaa视频,蜜臀视频在线观看一区,人妻av色图,密臀久久久精品影片,青青视频免费观看毛片,久草在线观看视,国产三级精品色情在线

pytorch模型部署 pth轉(zhuǎn)onnx的方法

 更新時(shí)間:2023年05月18日 09:37:57   作者:aoyou19  
這篇文章主要介紹了pytorch模型部署 pth轉(zhuǎn)onnx的相關(guān)知識(shí),本文通過(guò)實(shí)例代碼給大家介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或工作具有一定的參考借鑒價(jià)值,需要的朋友可以參考下

Pytorch轉(zhuǎn)ONNX的意義

一般來(lái)說(shuō)轉(zhuǎn)ONNX只是一個(gè)手段,在之后得到ONNX模型后還需要再將它做轉(zhuǎn)換,比如轉(zhuǎn)換到TensorRT上完成部署,或者有的人多加一步,從ONNX先轉(zhuǎn)換到caffe,再?gòu)腸affe到tensorRT。Pytorch自帶的torch.onnx.export轉(zhuǎn)換得到的ONNX,ONNXRuntime需要的ONNX,TensorRT需要的ONNX都是不同的。

將pytorch訓(xùn)練保存的pth文件轉(zhuǎn)為onnx文件,為后續(xù)模型部署做準(zhǔn)備。

一、分類模型

import torch
import os
import timm
import argparse
from utils_net import Resnet
parser = argparse.ArgumentParser()
parser.add_argument("--pth_path", default='classify_model.pth')
parser.add_argument("--save_onnx_path", default='classify_model.onnx')
parser.add_argument("--input_width", default=416)
parser.add_argument("--input_height", default=416)
parser.add_argument("--input_channel", default=1)
parser.add_argument("--num_classes", default=6)
args = parser.parse_args()
def pth_to_onnx(pth_path, onnx_path, in_hig, in_wid, in_chal, num_cls):
    if not onnx_path.endswith('.onnx'):
        print('Warning! The onnx model name is not correct,\
              please give a name that ends with \'.onnx\'!')
        return 0
    model = Resnet(num_classes=num_cls)
    model.load_state_dict(torch.load(pth_path))
    model.eval()
    print(f'{pth_path} model loaded')
    input_names = ['input']
    output_names = ['output']
    im = torch.rand(1, in_chal, in_hig, in_wid)
    torch.onnx.export(model, im, onnx_path,
                      verbose=False,
                      input_names=input_names,
                      output_names=output_names)
    print("Exporting .pth model to onnx model has been successful!")
    print(f"Onnx model save as {onnx_path}")
if __name__ == '__main__':
    pth_to_onnx(pth_path=args.pth_path,
                onnx_path=args.save_onnx_path,
                in_hig=args.input_height,
                in_wid=args.input_width,
                in_chal=args.input_channel,
                num_cls=args.num_classes)

運(yùn)行結(jié)果:

classify_model.pth model loaded
Exporting .pth model to onnx model has been successful!
Onnx model save as classify_model.onnx

Process finished with exit code 0

二、分割模型

import torch
import os
import argparse
from utils_net import seg_net
parser = argparse.ArgumentParser()
parser.add_argument("--pth_path", default='segment_model.pth')
parser.add_argument("--save_onnx_path", default='segment_model.onnx')
parser.add_argument("--input_width", default=416)
parser.add_argument("--input_height", default=416)
parser.add_argument("--input_channel", default=1)
parser.add_argument("--num_classes", default=4)
args = parser.parse_args()
def pth_to_onnx(pth_path, onnx_path, in_hig, in_wid, in_channel, num_cls):
    if not onnx_path.endswith('.onnx'):
        print('Warning! The onnx model name is not correct,\
              please give a name that ends with \'.onnx\'!')
        return 0
    model = seg_net(in_channel=in_channel, num_cls=num_cls)
    model.load_state_dict(torch.load(pth_path))
    model.eval()
    print(f'{pth_path} model loaded')
    input_names = ['input']
    output_names = ['output']
    im = torch.rand(1, in_channel, in_hig, in_wid)
    torch.onnx.export(model, im, onnx_path,
                      verbose=False,
                      input_names=input_names,
                      output_names=output_names,
                      opset_version=11)
    print("Exporting .pth model to onnx model has been successful!")
    print(f"Onnx model save as {onnx_path}")
if __name__ == '__main__':
    pth_to_onnx(pth_path=args.pth_path,
                onnx_path=args.save_onnx_path,
                in_hig=args.input_height,
                in_wid=args.input_width,
                in_channel=args.input_channel,
                num_cls=args.num_classes)

運(yùn)行結(jié)果:

segment_model.pth model loaded
Exporting .pth model to onnx model has been successful!
Onnx model save as segment_model.onnx

Process finished with exit code 0

三、目標(biāo)檢測(cè)模型

在這里插入代碼片
import torch
import onnx
import argparse
from utils_net import YoloBody
parser = argparse.ArgumentParser()
parser.add_argument("--pth_path", default='yolo.pth')
parser.add_argument("--save_onnx_path", default='yolo.onnx')
parser.add_argument("--input_width", default=416)
parser.add_argument("--input_height", default=416)
parser.add_argument("--num_classes", default=2)
parser.add_argument("--anchors_mask", default=[[6, 7, 8], [3, 4, 5], [0, 1, 2]])
args = parser.parse_args()
def pth_to_onnx(pth_path: str, save_onnx_path: str, num_cls: int,
                in_hig: int, in_wid: int, anchor_mask: list,
                opset_version: int = 12, simplify: bool = False):
    """
    :param pth_path: pth文件文件
    :param save_onnx_path: 準(zhǔn)備保存的onnx路徑
    :param num_cls: 檢測(cè)目標(biāo)類別數(shù)
    :param in_hig: 網(wǎng)絡(luò)輸入高度
    :param in_wid: 網(wǎng)絡(luò)輸入寬度
    :param anchor_mask: anchor寬高索引
    :param opset_version: onnx算子集版本
    :param simplify: 是否對(duì)模型進(jìn)行簡(jiǎn)化
    :return:保存onnx到指定路徑
    """
    # Build model, load weights
    net = YoloBody(anchors_mask=anchor_mask,
                   num_classes=num_cls)
    # device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    # net.load_state_dict(torch.load(pth_path, map_location=device))
    net.load_state_dict(torch.load(pth_path))
    # print(next(net.parameters()).device)
    net = net.eval()
    print(f'{pth_path} model loaded')
    im = torch.zeros(1, 3, in_hig, in_wid).to('cpu')
    input_layer_names = ['images']
    output_layer_names = ['output']
    # Export the model
    print(f'Starting export with onnx {onnx.__version__}.')
    torch.onnx.export(net,
                      im,
                      f=save_onnx_path,
                      verbose=False,
                      opset_version=opset_version,
                      training=torch.onnx.TrainingMode.EVAL,
                      do_constant_folding=True,
                      input_names=input_layer_names,
                      output_names=output_layer_names,
                      dynamic_axes=None)
    # Checks
    model_onnx = onnx.load(save_onnx_path)  # load onnx model
    onnx.checker.check_model(model_onnx)  # check onnx model
    # Simplify onnx
    if simplify:
        import onnxsim
        print(f'Simplifying with onnx-simplifier {onnxsim.__version__}.')
        model_onnx, check = onnxsim.simplify(
            model_onnx,
            dynamic_input_shape=False,
            input_shapes=None)
        assert check, 'assert check failed'
        onnx.save(model_onnx, save_onnx_path)
    print('Onnx model save as {}'.format(save_onnx_path))
if __name__ == '__main__':
    pth_to_onnx(pth_path=args.pth_path,
                save_onnx_path=args.save_onnx_path,
                num_cls=args.num_classes,
                in_hig=args.input_height,
                in_wid=args.input_width,
                anchor_mask=args.anchors_mask)

運(yùn)行結(jié)果:

yolo.pth model loaded
Starting export with onnx 1.11.0.
Onnx model save as yolo.onnx

Process finished with exit code 0

參考鏈接:

1.yolo
2.模型部署翻車記:pytorch轉(zhuǎn)onnx踩坑實(shí)錄

到此這篇關(guān)于pytorch模型部署 pth轉(zhuǎn)onnx的文章就介紹到這了,更多相關(guān)pytorch模型部署內(nèi)容請(qǐng)搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家!

相關(guān)文章

  • Python賦值邏輯的實(shí)現(xiàn)

    Python賦值邏輯的實(shí)現(xiàn)

    本文主要介紹了 Python賦值邏輯的實(shí)現(xiàn),文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來(lái)一起學(xué)習(xí)學(xué)習(xí)吧
    2023-02-02
  • 利用Python爬蟲爬取金融期貨數(shù)據(jù)的案例分析

    利用Python爬蟲爬取金融期貨數(shù)據(jù)的案例分析

    從技術(shù)角度來(lái)看,經(jīng)過(guò)一步步解析,任務(wù)是簡(jiǎn)單的,入門requests爬蟲及入門pandas數(shù)據(jù)分析就可以完成,本文重點(diǎn)給大家介紹Python爬蟲爬取金融期貨數(shù)據(jù)的案例分析,感興趣的朋友一起看看吧
    2022-06-06
  • Win8下python3.5.1安裝教程

    Win8下python3.5.1安裝教程

    這篇文章主要為大家詳細(xì)介紹了Win8下python3.5.1安裝教程,文中安裝步驟介紹的非常詳細(xì),具有一定的參考價(jià)值,感興趣的小伙伴們可以參考一下
    2018-07-07
  • 簡(jiǎn)單介紹Python中的filter和lambda函數(shù)的使用

    簡(jiǎn)單介紹Python中的filter和lambda函數(shù)的使用

    這篇文章主要簡(jiǎn)單介紹了Python中的filter和lambda函數(shù)的使用,是Python學(xué)習(xí)中的基礎(chǔ),同時(shí)lambda匿名函數(shù)的使用也是經(jīng)常被用來(lái)對(duì)比各種編程語(yǔ)的重要特性,言需要的朋友可以參考下
    2015-04-04
  • Python3 pyecharts生成Html文件柱狀圖及折線圖代碼實(shí)例

    Python3 pyecharts生成Html文件柱狀圖及折線圖代碼實(shí)例

    這篇文章主要介紹了Python3 pyecharts生成Html文件柱狀圖及折線圖代碼實(shí)例,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友可以參考下
    2020-09-09
  • Python django導(dǎo)出excel詳解

    Python django導(dǎo)出excel詳解

    這篇文章主要介紹了Python django導(dǎo)出excel的方法 ,分享了相關(guān)代碼示例,小編覺(jué)得還是挺不錯(cuò)的,具有一定借鑒價(jià)值,需要的朋友可以參考下
    2021-11-11
  • python用quad、dblquad實(shí)現(xiàn)一維二維積分的實(shí)例詳解

    python用quad、dblquad實(shí)現(xiàn)一維二維積分的實(shí)例詳解

    今天小編大家分享一篇python用quad、dblquad實(shí)現(xiàn)一維二維積分的實(shí)例詳解,具有很好的參考價(jià)值,希望對(duì)大家有所幫助。一起跟隨小編過(guò)來(lái)看看吧
    2019-11-11
  • Python實(shí)現(xiàn)多線程并發(fā)請(qǐng)求測(cè)試的腳本

    Python實(shí)現(xiàn)多線程并發(fā)請(qǐng)求測(cè)試的腳本

    這篇文章主要為大家分享了一個(gè)Python實(shí)現(xiàn)多線程并發(fā)請(qǐng)求測(cè)試的腳本,文中的示例代碼簡(jiǎn)潔易懂,具有一定的借鑒價(jià)值,需要的小伙伴可以了解一下
    2023-06-06
  • 深入理解?python?虛擬機(jī)

    深入理解?python?虛擬機(jī)

    這篇文章主要介紹了深入理解?python?虛擬機(jī)的相關(guān)資料,需要的朋友可以參考下
    2023-04-04
  • python寫日志封裝類實(shí)例

    python寫日志封裝類實(shí)例

    這篇文章主要介紹了python寫日志封裝類,實(shí)例分析了Python操作日志的相關(guān)技巧,需要的朋友可以參考下
    2015-06-06

最新評(píng)論

桐城市| 元朗区| 囊谦县| 南涧| 安远县| 辽阳县| 庄浪县| 吉木萨尔县| 石门县| 望都县| 英德市| 盐山县| 平陆县| 甘德县| 福州市| 神农架林区| 河北区| 柯坪县| 梁平县| 商丘市| 巫山县| 宜昌市| 贵阳市| 莱西市| 防城港市| 固原市| 澄迈县| 梁山县| 岢岚县| 嘉义县| 项城市| 论坛| 荆州市| 额尔古纳市| 皋兰县| 红河县| 革吉县| 临海市| 德清县| 朝阳市| 内黄县|