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

pytorch加載訓練好的模型用來測試或者處理方式

 更新時間:2023年09月09日 16:32:17   作者:群星閃耀  
這篇文章主要介紹了pytorch加載訓練好的模型用來測試或者處理方式,具有很好的參考價值,希望對大家有所幫助,如有錯誤或未考慮完全的地方,望不吝賜教

1.直接加載預訓練模型

如果我們使用的模型和原模型完全一樣,

那么我們可以直接加載別人訓練好的模型:

import torchvision.models as models
resnet50 = models.resnet50(pretrained=True)

如果只需要網(wǎng)絡結(jié)構(gòu),不需要用預訓練模型的參數(shù)來初始化,

那么就是:

model =torchvision.models.resnet50(pretrained=False)

2.修改某一層

PyTorch中的torchvision里已經(jīng)有很多常用的模型了,

可以直接調(diào)用:

  • AlexNet
  • VGG
  • ResNet
  • SqueezeNet
  • DenseNet
import torchvision.models as models
resnet18 = models.resnet18()
alexnet = models.alexnet()
squeezenet = models.squeezenet1_0()
densenet = models.densenet_161()

但是對于我們的任務而言有些層并不是直接能用,需要我們微微改一下,

比如,resnet最后的全連接層是分1000類,而我們只有21類;

又比如,resnet第一層卷積接收的通道是3, 我們可能輸入圖片的通道是4,

那么可以通過以下方法修改:

resnet.conv1 = nn.Conv2d(4, 64, kernel_size=7, stride=2, padding=3, bias=False)
resnet.fc = nn.Linear(2048, 21)

3.加載部分預訓練模型

其實大多數(shù)時候我們需要根據(jù)我們的任務調(diào)節(jié)我們的模型,所以很難保證模型和公開的模型完全一樣,但是預訓練模型的參數(shù)確實有助于提高訓練的準確率,為了結(jié)合二者的優(yōu)點,就需要我們加載部分預訓練模型。

#加載model,model是自己定義好的模型
resnet50 = models.resnet50(pretrained=True) 
model =Net(...) 
#讀取參數(shù) 
pretrained_dict =resnet50.state_dict() 
model_dict = model.state_dict() 
#將pretrained_dict里不屬于model_dict的鍵剔除掉 
pretrained_dict =  {k: v for k, v in pretrained_dict.items() if k in model_dict} 
# 更新現(xiàn)有的model_dict 
model_dict.update(pretrained_dict) 
# 加載我們真正需要的state_dict
model.load_state_dict(model_dict)
# 加載我們真正需要的state_dict 
model.load_state_dict(model_dict)  

4. 保存和加載自己的模型

pytorch保存模型的方式有兩種:

  • 第一種:將整個網(wǎng)絡都都保存下來
  • 第二種:僅保存和加載模型參數(shù)(推薦使用這樣的方法)

4.1 保存和加載整個模型

# 保存
torch.save(model_object, Path)
# 加載
model = torch.load(Path)

4.2 僅保存和加載模型參數(shù)(推薦使用) 

# ----------------保存模型參數(shù)--------------------------
torch.save(model.state_dict(), PATH)
#example
torch.save(resnet50.state_dict(),'ckp/model.pth')    
# ----------------加載模型參數(shù)--------------------------
model = ModelClass(*args, **kwargs) # 這是你后來設(shè)置的模型
model.load_state_dict(torch.load(PATH)) # 加載參數(shù)
#example
resnet=resnet50(pretrained=True)
resnet.load_state_dict(torch.load('ckp/model.pth'))

4.3 每個epoch保存一個模型參數(shù)

for epoch in range(start_epoch, nEpochs + 1):
        train(training_data_loader, optimizer, model, criterion, epoch)
        save_checkpoint(model, epoch)
def save_checkpoint(model, epoch):
    model_out_path = "checkpoint/" + "model_epoch_{}.pth".format(epoch)
    state = {"epoch": epoch ,"model": model}
    if not os.path.exists("checkpoint/"):
        os.makedirs("checkpoint/")
    torch.save(state, model_out_path)
    print("Checkpoint saved to {}".format(model_out_path))

上面的代碼中start_epoch是開始保存模型的epoch,nEpochs是總共訓練的次數(shù)。

train()里面的參數(shù),是訓練的過程:一些訓練數(shù)據(jù),優(yōu)化器,模型和訓練標準。

總結(jié)

以上為個人經(jīng)驗,希望能給大家一個參考,也希望大家多多支持腳本之家。

相關(guān)文章

  • 一文搞懂Python中is和==的區(qū)別

    一文搞懂Python中is和==的區(qū)別

    is和==都是對對象進行比較判斷作用的,但對對象比較判斷的內(nèi)容并不相同,下面來看看具體區(qū)別在哪?對Python中is和==的區(qū)別感興趣的朋友跟隨小編一起看看吧
    2023-01-01
  • 使用python-Jenkins批量創(chuàng)建及修改jobs操作

    使用python-Jenkins批量創(chuàng)建及修改jobs操作

    這篇文章主要介紹了使用python-Jenkins批量創(chuàng)建及修改jobs操作,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2020-05-05
  • django實現(xiàn)用戶登陸功能詳解

    django實現(xiàn)用戶登陸功能詳解

    這篇文章主要介紹了django實現(xiàn)用戶登陸功能詳解,具有一定借鑒價值,需要的朋友可以參考下。
    2017-12-12
  • Pycharm配置導入torch報錯Traceback的問題及解決

    Pycharm配置導入torch報錯Traceback的問題及解決

    這篇文章主要介紹了Pycharm配置導入torch報錯Traceback的問題及解決方案,具有很好的參考價值,希望對大家有所幫助,如有錯誤或未考慮完全的地方,望不吝賜教
    2023-12-12
  • Python利用多進程將大量數(shù)據(jù)放入有限內(nèi)存的教程

    Python利用多進程將大量數(shù)據(jù)放入有限內(nèi)存的教程

    這篇文章主要介紹了Python利用多進程將大量數(shù)據(jù)放入有限內(nèi)存的教程,使用了multiprocessing和pandas來加速內(nèi)存中的操作,需要的朋友可以參考下
    2015-04-04
  • 如何原位刪除python字典的key簡單示例

    如何原位刪除python字典的key簡單示例

    這篇文章主要介紹了如何原位刪除python字典的key的相關(guān)資料,為了避免在遍歷字典時修改其結(jié)構(gòu)導致的運行時錯誤,可以使用list(dic.items())創(chuàng)建鍵的副本,然后遍歷這個副本進行刪除操作,需要的朋友可以參考下
    2025-05-05
  • Pycharm遠程解釋器配置方式(自用成功版)

    Pycharm遠程解釋器配置方式(自用成功版)

    文章介紹了在PyCharm中配置SSH解釋器并同步本地代碼到遠程服務器的方法,包括配置SSH、設(shè)置虛擬環(huán)境、本地代碼推送和瀏覽遠程主機
    2026-02-02
  • Python logging模塊用法示例

    Python logging模塊用法示例

    這篇文章主要介紹了Python logging模塊用法,結(jié)合實例形式分析了Python logging模塊相關(guān)配置、函數(shù)、組件等操作方法與注意事項,需要的朋友可以參考下
    2018-08-08
  • Python中使用hashlib模塊處理算法的教程

    Python中使用hashlib模塊處理算法的教程

    這篇文章主要介紹了Python中使用hashlib模塊處理算法的教程,代碼基于Python2.x版本,需要的朋友可以參考下
    2015-04-04
  • python把列表中的字符串轉(zhuǎn)成整型的3種方法詳解

    python把列表中的字符串轉(zhuǎn)成整型的3種方法詳解

    這篇文章主要介紹了python把列表中的字符串轉(zhuǎn)成整型的3種方法詳解,python中在不同類型數(shù)據(jù)轉(zhuǎn)換方面是有標準庫的,使用非常方便,但是在開發(fā)中,經(jīng)常在list中字符轉(zhuǎn)成整形的數(shù)據(jù)方便遇到問題,需要的朋友可以參考下
    2023-07-07

最新評論

石城县| 黎平县| 锦州市| 延安市| 龙里县| 建平县| 靖江市| 综艺| 乌拉特中旗| 名山县| 嘉善县| 松阳县| 互助| 尖扎县| 安多县| 陕西省| 尼勒克县| 田阳县| 永兴县| 南乐县| 黄大仙区| 马公市| 潜山县| 民丰县| 连城县| 潜山县| 佛冈县| 南开区| 太康县| 岳阳县| 和平区| 定边县| 永胜县| 双鸭山市| 阳新县| 渝北区| 栖霞市| 余江县| 渭源县| 博湖县| 丹巴县|