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

pytorch保存和加載模型的方法及如何load部分參數(shù)

 更新時間:2024年03月11日 14:55:32   作者:BigerBang  
本文總結(jié)了pytorch中保存和加載模型的方法,以及在保存的模型文件與新定義的模型的參數(shù)不一一對應(yīng)時,我們該如何加載模型參數(shù),對pytorch保存和加載模型相關(guān)知識感興趣的朋友一起看看吧

本文總結(jié)了pytorch中保存和加載模型的方法,以及在保存的模型文件與新定義的模型的參數(shù)不一一對應(yīng)時,我們該如何加載模型參數(shù)。

1. 模型保存和加載的基本方式

在PyTorch中,模型可以通過兩種方式保存和加載:保存整個模型(包括模型架構(gòu)和參數(shù))或僅保存模型的參數(shù)(state_dict)。

保存整個模型: 保存模型的架構(gòu)和所有的權(quán)重參數(shù)。這樣做的好處是可以直接加載使用,無需再定義模型架構(gòu),但是無法再對模型做出調(diào)整,不夠靈活。

python
import torch
import torchvision.models as models
# 實例化一個預(yù)訓(xùn)練的resnet模型
model = models.resnet18(pretrained=True)
# 保存整個模型
torch.save(model, 'model.pth')
# 加載整個模型
model = torch.load('model.pth')

僅保存模型參數(shù)
通常推薦此方式,因為它僅保存權(quán)重參數(shù),體積更小,更靈活,需要時可用新定義的模型結(jié)構(gòu)加載參數(shù)。
保存的參數(shù)通過model.state_dict()獲取,得到一個有序字典類型:collections.OrderedDict,其中key是參數(shù)名稱,value是保存了參數(shù)數(shù)值的tensor類型。

OrderedDict是 Python 標(biāo)準(zhǔn)庫 collections 模塊中的一種字典(dict)類的子類。和普通的字典相比,OrderedDict 繳存了元素插入的順序,所以當(dāng)對其進(jìn)行迭代時,鍵值對會按照添加的先后次序返回,而不是基于鍵的散列值。

保存模型參數(shù)示例:

# 保存模型的state_dict
torch.save(model.state_dict(), 'model_state_dict.pth')

加載模型參數(shù)示例:

# 首先需要重新定義模型的結(jié)構(gòu),這里假設(shè)我們已經(jīng)有了一模一樣的模型定義
model = models.resnet18(pretrained=False) # 取消預(yù)訓(xùn)練權(quán)重
# 加載模型參數(shù)
model.load_state_dict(torch.load('model_state_dict.pth'))

2. 保存的模型文件和當(dāng)前定義的模型參數(shù)不完全一致時

有時候我們會對一個pretrained model的若干層進(jìn)行一些修改,涉及到層的添加和減少,同時未改變的那些層想要load pretrained model的參數(shù)。
假設(shè)新定義的模型是new_net, pretrained模型是old_net, 以下兩種方式適用于以下所有場景:
1. old_net的參數(shù)是new_net的子集
2. new_net的參數(shù)是old_net的子集
3. new_net和old_net的參數(shù)有交集

strict=False
一個直接的方式是在load_state_dict時strict=False,這樣在load參數(shù)時pytorch會匹配兩個模型中參數(shù)名字相同的參數(shù)進(jìn)行導(dǎo)入。

net_2.load_state_dict(torch.load("net_1.pth"), strict=False)

一種更靈活的方式,可自行添加更多的規(guī)則

def load(save_path, model):
    pretraind_dict = torch.load(save_path)
    model_dict =  model.state_dict()
    # 只將pretraind_dict中那些在model_dict中的參數(shù),提取出來
    state_dict = {k:v for k,v in pretraind_dict.items() if k in model_dict.keys()}
    # 將提取出來的參數(shù)更新到model_dict中,而model_dict有的,而state_dict沒有的參數(shù),不會被更新
    model_dict.update(state_dict)
    model.load_state_dict(model_dict)

可利用上面的代碼自行設(shè)計一些規(guī)則,比如如果不要laod某個參數(shù),就可以在上面的代碼中修改:

state_dict = {k:v for k,v in pretraind_dict.items() if k in model_dict.keys() and k != 'conv1.weight'}

3. 驗證代碼

import torch
from torch import nn as nn
class model_2_convs(nn.Module):
    def __init__(self) -> None:
        super().__init__()
        self.conv1 = nn.Conv2d(3, 64, 3)
        self.relu = nn.ReLU()
        self.conv2 = nn.Conv2d(64, 32, 3)
        self.mlp = nn.Linear(32, 10)
    def forward(self, x):
        x = self.conv1(x)
        x = self.relu(x)
        x = self.conv2(x)
        x = self.relu(x)
        return x
class model_3_convs(nn.Module):
    def __init__(self) -> None:
        super().__init__()
        self.conv1 = nn.Conv2d(3, 64, 3)
        self.relu = nn.ReLU()
        self.conv2 = nn.Conv2d(64, 32, 3)
        self.conv3 = nn.Conv2d(32, 64, 3)
    def forward(self, x):
        x = self.conv1(x)
        x = self.relu(x)
        x = self.conv2(x)
        x = self.relu(x)
        return x
def load(save_path, model):
    pretraind_dict = torch.load(save_path)
    model_dict =  model.state_dict()
    # 只將pretraind_dict中那些在model_dict中的參數(shù),提取出來
    state_dict = {k:v for k,v in pretraind_dict.items() if k in model_dict.keys()}
    # print(state_dict.keys())
    # 將提取出來的參數(shù)更新到model_dict中,而model_dict有的,而state_dict沒有的參數(shù),不會被更新
    model_dict.update(state_dict)
    model.load_state_dict(model_dict)
def load_weight_from_3_conv_to_2_conv(use_strict=False):
    net_1 = model_3_convs()
    net_2 = model_2_convs()
    torch.save(net_1.state_dict(), "net_1.pth")
    if use_strict:
        net_2.load_state_dict(torch.load("net_1.pth"), strict=False)
    else:
        load("net_1.pth", net_2)
    for key, para in net_2.state_dict().items():
        print(key)
        if key in net_1.state_dict().keys():
            print(torch.equal(para, net_1.state_dict()[key]))
def load_weight_from_2_conv_to_3_conv(use_strict=False):
    net_1 = model_3_convs()
    net_2 = model_2_convs()
    torch.save(net_2.state_dict(), "net_2.pth")
    if use_strict:
        net_1.load_state_dict(torch.load("net_2.pth"), strict=False)
    else:
        load("net_2.pth", net_1)
    for key, para in net_1.state_dict().items():
        print(key)
        if key in net_2.state_dict().keys():
            print(torch.equal(para, net_2.state_dict()[key]))
if __name__ == "__main__":
    load_weight_from_3_conv_to_2_conv(use_strict=True)
    load_weight_from_3_conv_to_2_conv(use_strict=False)
    load_weight_from_2_conv_to_3_conv(use_strict=True)
    load_weight_from_2_conv_to_3_conv(use_strict=False)

到此這篇關(guān)于pytorch保存和加載模型以及如何load部分參數(shù)的文章就介紹到這了,更多相關(guān)pytorch保存和加載模型內(nèi)容請搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家!

相關(guān)文章

  • 解決Tensorboard 不顯示計算圖graph的問題

    解決Tensorboard 不顯示計算圖graph的問題

    今天小編就為大家分享一篇解決Tensorboard 不顯示計算圖graph的問題,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2020-02-02
  • python: line=f.readlines()消除line中\(zhòng)n的方法

    python: line=f.readlines()消除line中\(zhòng)n的方法

    這篇文章主要介紹了python: line=f.readlines()消除line中\(zhòng)n的方法,需要的朋友可以參考下
    2018-03-03
  • python爬取微信公眾號文章的方法

    python爬取微信公眾號文章的方法

    這篇文章主要為大家詳細(xì)介紹了python爬取微信公眾號文章的方法,具有一定的參考價值,感興趣的小伙伴們可以參考一下
    2019-02-02
  • 一文秒懂pandas中iloc()函數(shù)

    一文秒懂pandas中iloc()函數(shù)

    iloc[]函數(shù)屬于pandas庫全稱為index?location,即對數(shù)據(jù)進(jìn)行位置索引,從而在數(shù)據(jù)表中提取出相應(yīng)的數(shù)據(jù),本文通過實例代碼介紹pandas中iloc()函數(shù),感興趣的朋友一起看看吧
    2023-04-04
  • 在PYQT5中QscrollArea(滾動條)的使用方法

    在PYQT5中QscrollArea(滾動條)的使用方法

    今天小編就為大家分享一篇在PYQT5中QscrollArea(滾動條)的使用方法,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2019-06-06
  • Python中__new__()方法適應(yīng)及注意事項詳解

    Python中__new__()方法適應(yīng)及注意事項詳解

    這篇文章主要介紹了Python中__new__()方法適應(yīng)及注意事項的相關(guān)資料,new()方法是Python中的一個特殊構(gòu)造方法,用于在創(chuàng)建對象之前調(diào)用,并負(fù)責(zé)返回類的新實例,它與init()方法不同,需要的朋友可以參考下
    2025-03-03
  • pandas的唯一值、值計數(shù)以及成員資格的示例

    pandas的唯一值、值計數(shù)以及成員資格的示例

    今天小編就為大家分享一篇pandas的唯一值、值計數(shù)以及成員資格的示例,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2018-07-07
  • Python中defaultdict與dict的差異詳情

    Python中defaultdict與dict的差異詳情

    這篇文章主要介紹了Python中defaultdict與dict的差異,在collections模塊中的defauldict使用時與dict有何不同,為何我們用dict中的key值不存在時會報錯,而defaudict不會報錯,下面文章做出解答,需要的朋友可以參考一下
    2021-11-11
  • python 時間處理之月份加減問題

    python 時間處理之月份加減問題

    這篇文章主要介紹了python 時間處理之月份加減問題,具有很好的參考價值,希望對大家有所幫助。如有錯誤或未考慮完全的地方,望不吝賜教
    2022-11-11
  • Pycharm使用Database?Navigator連接mysql數(shù)據(jù)庫全過程

    Pycharm使用Database?Navigator連接mysql數(shù)據(jù)庫全過程

    這篇文章主要介紹了Pycharm使用Database?Navigator連接mysql數(shù)據(jù)庫全過程,具有很好的參考價值,希望對大家有所幫助。如有錯誤或未考慮完全的地方,望不吝賜教
    2022-07-07

最新評論

安庆市| 丽水市| 昆山市| 台江县| 胶南市| 沁阳市| 康定县| 延长县| 临潭县| 迁西县| 江孜县| 山阳县| 双柏县| 湟中县| 贵州省| 辽阳县| 萨迦县| 北安市| 四子王旗| 南郑县| 开远市| 青州市| 闽清县| 比如县| 洪洞县| 山西省| 邵阳县| 五家渠市| 弋阳县| 镇沅| 大竹县| 西城区| 诸城市| 成武县| 布尔津县| 新干县| 大姚县| 清河县| 庆元县| 青州市| 东山县|