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

pytorch模型保存與加載中的一些問(wèn)題實(shí)戰(zhàn)記錄

 更新時(shí)間:2022年10月28日 12:37:20   作者:colourmind  
一般來(lái)說(shuō),保存模型是把參數(shù)全部用model.cpu().state_dict(),然后加載模型時(shí)一般用model.load_state_dict(torch.load(model_path)),下面這篇文章主要給大家介紹了關(guān)于pytorch模型保存與加載中的一些問(wèn)題實(shí)戰(zhàn)記錄,需要的朋友可以參考下

前言

最近使用pytorch訓(xùn)練模型,保存模型后再次加載使用出現(xiàn)了一些問(wèn)題。記錄一下解決方案!

一、torch中模型保存和加載的方式

1、模型參數(shù)和模型結(jié)構(gòu)保存和加載

torch.save(model,path)
torch.load(path)

2、只保存模型的參數(shù)和加載——這種方式比較安全,但是比較稍微麻煩一點(diǎn)點(diǎn)

torch.save(model.state_dict(),path)
model_state_dic = torch.load(path)
model.load_state_dic(model_state_dic)

二、torch中模型保存和加載出現(xiàn)的問(wèn)題

1、單卡模型下保存模型結(jié)構(gòu)和參數(shù)后加載出現(xiàn)的問(wèn)題

模型保存的時(shí)候會(huì)把模型結(jié)構(gòu)定義文件路徑記錄下來(lái),加載的時(shí)候就會(huì)根據(jù)路徑解析它然后裝載參數(shù);當(dāng)把模型定義文件路徑修改以后,使用torch.load(path)就會(huì)報(bào)錯(cuò)。

把model文件夾修改為models后,再加載就會(huì)報(bào)錯(cuò)。

import torch
from model.TextRNN import TextRNN
 
load_model = torch.load('experiment_model_save/textRNN.bin')
print('load_model',load_model)

這種保存完整模型結(jié)構(gòu)和參數(shù)的方式,一定不要改動(dòng)模型定義文件路徑

2、多卡機(jī)器單卡訓(xùn)練模型保存后在單卡機(jī)器上加載會(huì)報(bào)錯(cuò)

在多卡機(jī)器上有多張顯卡0號(hào)開始,現(xiàn)在模型在n>=1上的顯卡訓(xùn)練保存后,拷貝在單卡機(jī)器上加載

import torch
from model.TextRNN import TextRNN
 
load_model = torch.load('experiment_model_save/textRNN_cuda_1.bin')
print('load_model',load_model)

會(huì)出現(xiàn)cuda device不匹配的問(wèn)題——你保存的模代碼段 小部件型是使用的cuda1,那么采用torch.load()打開的時(shí)候,會(huì)默認(rèn)的去尋找cuda1,然后把模型加載到該設(shè)備上。這個(gè)時(shí)候可以直接使用map_location來(lái)解決,把模型加載到CPU上即可。

load_model = torch.load('experiment_model_save/textRNN_cuda_1.bin',map_location=torch.device('cpu'))

3、多卡訓(xùn)練模型保存模型結(jié)構(gòu)和參數(shù)后加載出現(xiàn)的問(wèn)題

當(dāng)用多GPU同時(shí)訓(xùn)練模型之后,不管是采用模型結(jié)構(gòu)和參數(shù)一起保存還是單獨(dú)保存模型參數(shù),然后在單卡下加載都會(huì)出現(xiàn)問(wèn)題

a、模型結(jié)構(gòu)和參數(shù)一起保然后在加載

torch.distributed.init_process_group(backend='nccl')

模型訓(xùn)練的時(shí)候采用上述多進(jìn)程的方式,所以你在加載的時(shí)候也要聲明,不然就會(huì)報(bào)錯(cuò)。

b、單獨(dú)保存模型參數(shù)

model = Transformer(num_encoder_layers=6,num_decoder_layers=6)
state_dict = torch.load('train_model/clip/experiment.pt')
model.load_state_dict(state_dict)

同樣會(huì)出現(xiàn)問(wèn)題,不過(guò)這里出現(xiàn)的問(wèn)題是參數(shù)字典的key和模型定義的key不一樣

原因是多GPU訓(xùn)練下,使用分布式訓(xùn)練的時(shí)候會(huì)給模型進(jìn)行一個(gè)包裝,代碼如下:

model = torch.load('train_model/clip/Vtransformers_bert_6_layers_encoder_clip.bin')
print(model)
model.cuda(args.local_rank)
。。。。。。
model = nn.parallel.DistributedDataParallel(model,device_ids=[args.local_rank],find_unused_parameters=True)
print('model',model)

包裝前的模型結(jié)構(gòu):

包裝后的模型

在外層多了DistributedDataParallel以及module,所以才會(huì)導(dǎo)致在單卡環(huán)境下加載模型權(quán)重的時(shí)候出現(xiàn)權(quán)重的keys不一致。

三、正確的保存模型和加載的方法

    if gpu_count > 1:
        torch.save(model.module.state_dict(),save_path)
    else:
        torch.save(model.state_dict(),save_path)
    model = Transformer(num_encoder_layers=6,num_decoder_layers=6)
    state_dict = torch.load(save_path)
    model.load_state_dict(state_dict)

這樣就是比較好的范式,加載不會(huì)出錯(cuò)。

總結(jié)

到此這篇關(guān)于pytorch模型保存與加載中的一些問(wèn)題的文章就介紹到這了,更多相關(guān)pytorch模型保存與加載內(nèi)容請(qǐng)搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家!

相關(guān)文章

  • Python文件如何讀取read()函數(shù)

    Python文件如何讀取read()函數(shù)

    這篇文章主要介紹了Python文件如何讀取read()函數(shù)問(wèn)題,具有很好的參考價(jià)值,希望對(duì)大家有所幫助,如有錯(cuò)誤或未考慮完全的地方,望不吝賜教
    2024-02-02
  • Python模擬簡(jiǎn)易版淘寶客服機(jī)器人的示例代碼

    Python模擬簡(jiǎn)易版淘寶客服機(jī)器人的示例代碼

    這篇文章主要介紹了Python模擬簡(jiǎn)易版淘寶客服機(jī)器人的示例代碼,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來(lái)一起學(xué)習(xí)學(xué)習(xí)吧
    2021-03-03
  • Python自動(dòng)化開發(fā)學(xué)習(xí)之三級(jí)菜單制作

    Python自動(dòng)化開發(fā)學(xué)習(xí)之三級(jí)菜單制作

    這篇文章主要為大家詳細(xì)介紹了Python自動(dòng)化開發(fā)學(xué)習(xí)之三級(jí)菜單的制作方法,具有一定的參考價(jià)值,感興趣的小伙伴們可以參考一下
    2017-07-07
  • Python內(nèi)置的字符串處理函數(shù)整理

    Python內(nèi)置的字符串處理函數(shù)整理

    Python內(nèi)置的字符串處理函數(shù)整理,收集常用的Python 內(nèi)置的各種字符串處理 函數(shù)的使用方法
    2013-01-01
  • 詳解python爬蟲系列之初識(shí)爬蟲

    詳解python爬蟲系列之初識(shí)爬蟲

    這篇文章主要介紹了python爬蟲系列之初識(shí)爬蟲,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來(lái)一起學(xué)習(xí)學(xué)習(xí)吧
    2019-04-04
  • 基于python中的TCP及UDP(詳解)

    基于python中的TCP及UDP(詳解)

    下面小編就為大家?guī)?lái)一篇基于python中的TCP及UDP(詳解)。小編覺(jué)得挺不錯(cuò)的,現(xiàn)在就分享給大家,也給大家做個(gè)參考。一起跟隨小編過(guò)來(lái)看看吧,希望對(duì)大家有所幫助
    2017-11-11
  • python 中xpath爬蟲實(shí)例詳解

    python 中xpath爬蟲實(shí)例詳解

    這篇文章主要介紹了python實(shí)例:xpath爬蟲實(shí)例,本文通過(guò)實(shí)例代碼給大家介紹的非常詳細(xì),具有一定的參考借鑒價(jià)值,需要的朋友可以參考下
    2019-08-08
  • Python?字符替換的四方法

    Python?字符替換的四方法

    本文主要介紹了Python?字符替換的四方法,主要包括replace、translate、maketrans?和正則這是四種方法,具有一定的參考價(jià)值,感興趣的可以了解一下
    2024-01-01
  • python orm 框架中sqlalchemy用法實(shí)例詳解

    python orm 框架中sqlalchemy用法實(shí)例詳解

    這篇文章主要介紹了python orm 框架中sqlalchemy用法,結(jié)合實(shí)例形式詳細(xì)分析了Python orm 框架基本概念、原理及sqlalchemy相關(guān)使用技巧,需要的朋友可以參考下
    2020-02-02
  • 使用pandas庫(kù)對(duì)csv文件進(jìn)行篩選保存

    使用pandas庫(kù)對(duì)csv文件進(jìn)行篩選保存

    這篇文章主要介紹了使用pandas庫(kù)對(duì)csv文件進(jìn)行篩選保存,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友可以參考下
    2020-05-05

最新評(píng)論

紫阳县| 临武县| 桦川县| 江口县| 玉山县| 罗江县| 元阳县| 哈尔滨市| 扎囊县| 达州市| 苍梧县| 枣阳市| 始兴县| 茌平县| 离岛区| 平遥县| 金堂县| 左贡县| 四川省| 安多县| 稻城县| 彭水| 浮梁县| 丁青县| 达州市| 河津市| 灵川县| 甘德县| 镇平县| 渭南市| 渭源县| 山西省| 南靖县| 新干县| 甘泉县| 镇雄县| 时尚| 吉安市| 磴口县| 丰城市| 清新县|