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

Pytorch中實現(xiàn)只導(dǎo)入部分模型參數(shù)的方式

 更新時間:2020年01月02日 16:57:09   作者:咆哮的阿杰  
今天小編就為大家分享一篇Pytorch中實現(xiàn)只導(dǎo)入部分模型參數(shù)的方式,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧

我們在做遷移學(xué)習(xí),或者在分割,檢測等任務(wù)想使用預(yù)訓(xùn)練好的模型,同時又有自己修改之后的結(jié)構(gòu),使得模型文件保存的參數(shù),有一部分是不需要的(don't expected)。我們搭建的網(wǎng)絡(luò)對保存文件來說,有一部分參數(shù)也是沒有的(missed)。如果依舊使用torch.load(model.state_dict())的辦法,就會出現(xiàn) xxx expected,xxx missed類似的錯誤。那么在這種情況下,該如何導(dǎo)入模型呢?

好在Pytorch中的模型參數(shù)使用字典保存的,鍵是參數(shù)的名稱,值是參數(shù)的具體數(shù)值。我們使用model.state_dict()獲得這個字典,之后就能利用參數(shù)名稱來實現(xiàn)導(dǎo)入。

請看下面的一個例子。

我們先搭建一個小小的網(wǎng)絡(luò)。

import torch as t
from torch.nn import Module
from torch import nn
from torch.nn import functional as F
class Net(Module):
  def __init__(self):
    super(Net,self).__init__()
    self.conv1 = nn.Conv2d(3,32,3,1)
    self.conv2 = nn.Conv2d(32,3,3,1)
    self.w = nn.Parameter(t.randn(3,10))
    for p in self.children():
      nn.init.xavier_normal_(p.weight.data)
      nn.init.constant_(p.bias.data, 0)
  def forward(self, x):
    out = self.conv1(x)
    out = self.conv2(x)
 
    out = F.avg_pool2d(out,(out.shape[2],out.shape[3]))
    out = F.linear(out,weight=self.w)
    return out

然后我們保存這個網(wǎng)絡(luò)的初始值。

model = Net()
t.save(model.state_dict(),'xxx.pth')

現(xiàn)在我們將Net修改一下,多加幾個卷積層,但并不加入到forward中,僅僅出于少些幾行的目的。

import torch as t
from torch.nn import Module
from torch import nn
from torch.nn import functional as F
 
 
class Net(Module):
  def __init__(self):
    super(Net, self).__init__()
    self.conv1 = nn.Conv2d(3, 32, 3, 1)
    self.conv2 = nn.Conv2d(32, 3, 3, 1)
    self.conv3 = nn.Conv2d(3,64,3,1)
    self.conv4 = nn.Conv2d(64,32,3,1)
    for p in self.children():
      nn.init.xavier_normal_(p.weight.data)
      nn.init.constant_(p.bias.data, 0)
 
    self.w = nn.Parameter(t.randn(3, 10))
  def forward(self, x):
    out = self.conv1(x)
    out = self.conv2(x)
 
    out = F.avg_pool2d(out, (out.shape[2], out.shape[3]))
    out = F.linear(out, weight=self.w)
    return out

我們現(xiàn)在試著導(dǎo)入之前保存的模型參數(shù)。

path = 'xxx.pth'
model = Net()
model.load_state_dict(t.load(path))
 
'''
RuntimeError: Error(s) in loading state_dict for Net:
 Missing key(s) in state_dict: "conv3.weight", "conv3.bias", "conv4.weight", "conv4.bias". 
'''

出現(xiàn)了沒有在模型文件中找到error中的關(guān)鍵字的錯誤。

現(xiàn)在我們這樣導(dǎo)入模型

path = 'xxx.pth'
model = Net()
save_model = t.load(path)
model_dict = model.state_dict()
state_dict = {k:v for k,v in save_model.items() if k in model_dict.keys()}
print(state_dict.keys()) # dict_keys(['w', 'conv1.weight', 'conv1.bias', 'conv2.weight', 'conv2.bias'])
model_dict.update(state_dict)
model.load_state_dict(model_dict)

看看上面的代碼,很容易弄明白。其中model_dict.update的作用是更新代碼中搭建的模型參數(shù)字典。為啥更新我其實并不清楚,但這一步驟是必須的,否則還會報錯。

為了弄清楚為什么要更新model_dict,我們不妨分別輸出state_dict和model_dict的關(guān)鍵值看一看。

for k in state_dict.keys():
  print(k)
 
'''
w
conv1.weight
conv1.bias
conv2.weight
conv2.bias
'''
for k in model_dict.keys():
  print(k)
 
'''
w
conv1.weight
conv1.bias
conv2.weight
conv2.bias
conv3.weight
conv3.bias
conv4.weight
conv4.bias
'''

這個結(jié)果也是預(yù)料之中的,所以我猜測,update之后,model_dict和state_dict中具有相同鍵的值已經(jīng)同步了。updata的目的就是使model_dict帶有state_dict中都具有的那一部分參數(shù)的值,對于model_dict中有的,但是save_dict中沒有的參數(shù),值不改變,參數(shù)仍然使用初始值。

以上這篇Pytorch中實現(xiàn)只導(dǎo)入部分模型參數(shù)的方式就是小編分享給大家的全部內(nèi)容了,希望能給大家一個參考,也希望大家多多支持腳本之家。

相關(guān)文章

  • python安裝pytorch方式

    python安裝pytorch方式

    這篇文章主要介紹了python安裝pytorch方式,具有很好的參考價值,希望對大家有所幫助,如有錯誤或未考慮完全的地方,望不吝賜教
    2024-01-01
  • Python yield的使用詳解

    Python yield的使用詳解

    您可能聽說過,帶有 yield 的函數(shù)在 Python 中被稱之為、generator(生成器),何謂 generator ?我們先拋開 generator,以一個常見的編程題目來展示 yield 的概念
    2021-10-10
  • python和php哪個更適合寫爬蟲

    python和php哪個更適合寫爬蟲

    這篇文章主要介紹了python和php哪個更適合寫爬蟲的相關(guān)對比知識點,需要的朋友們可以學(xué)習(xí)下。
    2020-06-06
  • Python3監(jiān)控疫情的完整代碼

    Python3監(jiān)控疫情的完整代碼

    這篇文章主要介紹了Python3監(jiān)控疫情的完整代碼,代碼簡單易懂,非常不錯具有一定的參考借鑒價值,需要的朋友可以參考下
    2020-02-02
  • python抓取網(wǎng)頁內(nèi)容示例分享

    python抓取網(wǎng)頁內(nèi)容示例分享

    這篇文章主要介紹了python抓取網(wǎng)頁內(nèi)容示例,在抓取的時候?qū)τ趃bk編碼網(wǎng)頁還需要轉(zhuǎn)化一下,具體看下面的示例吧
    2014-02-02
  • python Pygal庫生成SVG(可縮放矢量圖形)圖表示例

    python Pygal庫生成SVG(可縮放矢量圖形)圖表示例

    這篇文章主要為大家介紹了python Pygal庫生成SVG(可縮放矢量圖形)圖表示例,有需要的朋友可以借鑒參考下,希望能夠有所幫助,祝大家多多進(jìn)步,早日升職加薪
    2024-01-01
  • python使用pymongo與MongoDB基本交互操作示例

    python使用pymongo與MongoDB基本交互操作示例

    這篇文章主要介紹了python使用pymongo與MongoDB基本交互操作,結(jié)合實例形式詳細(xì)分析了python基于pymongo庫實現(xiàn)與MongoDB基本交互相關(guān)操作技巧與注意事項,需要的朋友可以參考下
    2020-04-04
  • 如何利用PyQt5美化你的GUI界面

    如何利用PyQt5美化你的GUI界面

    python的腳本開發(fā)簡單,有時候只需幾行代碼就能實現(xiàn)豐富的功能,而且python本身是跨平臺的,所以深受程序員的喜愛,下面這篇文章主要給大家介紹了關(guān)于如何利用PyQt5美化你的GUI界面的相關(guān)資料,需要的朋友可以參考下
    2022-01-01
  • 使用Python神器對付12306變態(tài)驗證碼

    使用Python神器對付12306變態(tài)驗證碼

    這篇文章主要介紹了使用Python神器對付12306變態(tài)驗證碼的相關(guān)資料,需要的朋友可以參考下
    2016-01-01
  • Python使用pandas處理CSV文件的實例講解

    Python使用pandas處理CSV文件的實例講解

    今天小編就為大家分享一篇Python使用pandas處理CSV文件的實例講解,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2018-06-06

最新評論

德惠市| 临海市| 驻马店市| 固始县| 大港区| 肥东县| 庆阳市| 教育| 黄浦区| 尼木县| 台中县| 深泽县| 彩票| 延川县| 和静县| 日喀则市| 南宫市| 丰台区| 富蕴县| 遂平县| 淄博市| 庄河市| 长白| 独山县| 芜湖县| 土默特左旗| 扎囊县| 灵山县| 青浦区| 若尔盖县| 邹平县| 陈巴尔虎旗| 五台县| 德格县| 平湖市| 洛扎县| 南和县| 武鸣县| 峨边| 根河市| 罗定市|