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

Pytorch參數注冊和nn.ModuleList nn.ModuleDict的問題

 更新時間:2023年01月03日 09:22:39   作者:luputo  
這篇文章主要介紹了Pytorch參數注冊和nn.ModuleList nn.ModuleDict的問題,具有很好的參考價值,希望對大家有所幫助。如有錯誤或未考慮完全的地方,望不吝賜教

參考自官方文檔

參數注冊

嘗試自己寫GoogLeNet時碰到的問題,放在字典中的參數無法自動注冊,所謂的注冊,就是當參數注冊到這個網絡上時,它會隨著你在外部調用net.cuda()后自動遷移到GPU上,而沒有注冊的參數則不會隨著網絡遷到GPU上,這就可能導致輸入在GPU上而參數不在GPU上,從而出現錯誤,為了說明這個現象。

舉一個有點鐵憨憨的例子:

import torch
import torch.nn as nn
import torch.nn.functional as F

class Net(nn.Module):
?? ?def __init__(self):
?? ??? ?super(Net,self).__init__()
?? ??? ?self.weight = torch.rand((3,4)) # 這里其實可以直接用nn.Linear,但為了舉例這里先憨憨一下
?? ?
?? ?def forward(self,x):
?? ??? ?return F.linear(x,self.weight)

if __name__ == "__main__":
?? ?batch_size = 10
?? ?dummy = torch.rand((batch_size,4))
?? ?net = Net()
?? ?print(net(dummy))

上面的代碼可以成功運行,因為所有的數值都是放在CPU上的,但是,一旦我們要把模型移到GPU上時

import torch
import torch.nn as nn
import torch.nn.functional as F

class Net(nn.Module):
?? ?def __init__(self):
?? ??? ?super(Net,self).__init__()
?? ??? ?self.weight = torch.rand((3,4))
?? ?
?? ?def forward(self,x):
?? ??? ?return F.linear(x,self.weight)

if __name__ == "__main__":
?? ?batch_size = 10
?? ?dummy = torch.rand((batch_size,4)).cuda()
?? ?net = Net().cuda()
?? ?print(net(dummy))

運行后就會出現

...
RuntimeError: Expected object of backend CUDA but got backend CPU for argument #2 'mat2'

這就是因為self.weight沒有隨著模型一起移到GPU上的原因,此時我們查看模型的參數,會發(fā)現并沒有self.weight

import torch
import torch.nn as nn
import torch.nn.functional as F

class Net(nn.Module):
?? ?def __init__(self):
?? ??? ?super(Net,self).__init__()
?? ??? ?self.weight = torch.rand((3,4))
?? ?
?? ?def forward(self,x):
?? ??? ?return F.linear(x,self.weight)

if __name__ == "__main__":
?? ?net = Net()
?? ?for parameter in net.parameters():
?? ??? ?print(parameter)

上面的代碼沒有輸出,因為net根本沒有參數

那么為了讓net有參數,我們需要手動地將self.weight注冊到網絡上

import torch
import torch.nn as nn
import torch.nn.functional as F

class Net(nn.Module):
?? ?def __init__(self):
?? ??? ?super(Net,self).__init__()
?? ??? ?self.weight = nn.Parameter(torch.rand((3,4))) # 被注冊的參數必須是nn.Parameter類型
?? ??? ?self.register_parameter('weight',self.weight) # 手動注冊參數
?? ??? ?
?? ?
?? ?def forward(self,x):
?? ??? ?return F.linear(x,self.weight)

if __name__ == "__main__":
?? ?net = Net()
?? ?for parameter in net.parameters():
?? ??? ?print(parameter)

?? ?batch_size = 10
?? ?net = net.cuda()
?? ?dummy = torch.rand((batch_size,4)).cuda()
?? ?print(net(dummy))

此時網絡的參數就有了輸出,同時會隨著一起遷到GPU上,輸出就類似這樣

Parameter containing:
tensor([...])
tensor([...])

不過后來我實驗了以下,好像只寫nn.Parameter不寫register也可以被默認注冊

nn.ModuleList和nn.ModuleDict

有時候我們?yōu)榱藞D省事,可能會這樣寫網絡

import torch
import torch.nn as nn
import torch.nn.functional as F

class Net(nn.Module):
?? ?def __init__(self):
?? ??? ?super(Net,self).__init__()
?? ??? ?self.linears = [nn.Linear(4,4),nn.Linear(4,4),nn.Linear(4,2)]
?? ?
?? ?def forward(self,x):
?? ??? ?for linear in self.linears:
?? ??? ??? ?x = linear(x)
?? ??? ??? ?x = F.relu(x)
?? ??? ?return x

if __name__ == '__main__':
?? ?net = Net()
?? ?for parameter in net.parameters():
?? ??? ?print(parameter)??

     

同樣,輸出網絡的參數啥也沒有,這意味著當調用net.cuda時,self.linears里面的參數不會一起走到GPU上去

此時我們可以在__init__方法中手動對self.parameters()迭代然后把每個參數注冊,但更好的方法是,pytorch已經為我們提供了nn.ModuleList,用來代替python內置的list,放在nn.ModuleList中的參數將會自動被正確注冊

import torch
import torch.nn as nn
import torch.nn.functional as F

class Net(nn.Module):
?? ?def __init__(self):
?? ??? ?super(Net,self).__init__()
?? ??? ?self.linears = nn.ModuleList([nn.Linear(4,4),nn.Linear(4,4),nn.Linear(4,2)])
?? ?
?? ?def forward(self,x):
?? ??? ?for linear in self.linears:
?? ??? ??? ?x = linear(x)
?? ??? ??? ?x = F.relu(x)
?? ??? ?return x

if __name__ == '__main__':
?? ?net = Net()
?? ?for parameter in net.parameters():
?? ??? ?print(parameter)?? ??? ?

此時就有輸出了

Parameter containing:
tensor(...)
Parameter containing:
tensor(...)
...

nn.ModuleDict也是類似,當我們需要把參數放在一個字典里的時候,能夠用的上,這里直接給一個官方的例子看一看就OK

class MyModule(nn.Module):
? ? def __init__(self):
? ? ? ? super(MyModule, self).__init__()
? ? ? ? self.choices = nn.ModuleDict({
? ? ? ? ? ? ? ? 'conv': nn.Conv2d(10, 10, 3),
? ? ? ? ? ? ? ? 'pool': nn.MaxPool2d(3)
? ? ? ? })
? ? ? ? self.activations = nn.ModuleDict([
? ? ? ? ? ? ? ? ['lrelu', nn.LeakyReLU()],
? ? ? ? ? ? ? ? ['prelu', nn.PReLU()]
? ? ? ? ])

? ? def forward(self, x, choice, act):
? ? ? ? x = self.choices[choice](x)
? ? ? ? x = self.activations[act](x)
? ? ? ? return x

需要注意的是,雖然直接放在python list中的參數不會自動注冊,但如果只是暫時放在list里,隨后又調用了nn.Sequential把整個list整合起來,參數仍然是會自動注冊的

另外一點要注意的是ModuleList和ModuleDict里面只能放Module的子類,也就是nn.Conv,nn.Linear這樣的,但不能放nn.Parameter,如果要放nn.Parameter,用nn.ParameterList即可,用法和nn.ModuleList一樣

總結

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

相關文章

  • 詳解Django中CSRF和CORS的區(qū)別

    詳解Django中CSRF和CORS的區(qū)別

    本文主要介紹了詳解Django中CSRF和CORS的區(qū)別,文中通過示例代碼介紹的非常詳細,對大家的學習或者工作具有一定的參考學習價值,需要的朋友們下面隨著小編來一起學習學習吧
    2022-08-08
  • 對python中url參數編碼與解碼的實例詳解

    對python中url參數編碼與解碼的實例詳解

    今天小編就為大家分享一篇對python中url參數編碼與解碼的實例詳解,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2019-07-07
  • Python回溯法(Backtracking)的具體使用

    Python回溯法(Backtracking)的具體使用

    在Python中,我們可以應用回溯法解決各種問題,如八皇后問題、子集問題等,本文就來介紹一下Python回溯法(Backtracking)的具體使用,感興趣的可以了解一下
    2023-12-12
  • pymongo中聚合查詢的使用方法

    pymongo中聚合查詢的使用方法

    這篇文章主要給大家介紹了關于pymongo中聚合查詢的使用方法,文中通過示例代碼介紹的非常詳細,對大家學習或者使用pymongo具有一定的參考學習價值,需要的朋友們下面來一起學習學習吧
    2019-03-03
  • 用Python的Django框架完成視頻處理任務的教程

    用Python的Django框架完成視頻處理任務的教程

    這篇文章主要介紹了用Python的Django框架完成視頻處理任務的教程,包括用戶的視頻上傳和播放以及下載功能的實現,需要的朋友可以參考下
    2015-04-04
  • Python 分享10個PyCharm技巧

    Python 分享10個PyCharm技巧

    這篇文章主要介紹了Python 分享10個PyCharm技巧,今天要跟大家分享幾個PyCharm小技巧,幫助大家提升工作效率!,需要的朋友可以參考下
    2019-07-07
  • Python+PuLP實現線性規(guī)劃的求解

    Python+PuLP實現線性規(guī)劃的求解

    線性規(guī)劃(Linear?programming),在線性等式或不等式約束條件下求解線性目標函數的極值問題,常用于解決資源分配、生產調度和混合問題。本文將利用PuLP實現線性規(guī)劃的求解,需要的可以參考一下
    2022-04-04
  • 如何卸載python插件

    如何卸載python插件

    在本篇文章里小編給大家分享了關于python插件如何卸載的相關文章,需要的朋友們可以參考下。
    2020-07-07
  • python實現矩陣打印

    python實現矩陣打印

    這篇文章主要為大家詳細介紹了python實現矩陣打印的相關代碼,具有一定的參考價值,感興趣的小伙伴們可以參考一下
    2019-03-03
  • Python實現將字典(列表按列)存入csv文件

    Python實現將字典(列表按列)存入csv文件

    這篇文章主要介紹了Python實現將字典(列表按列)存入csv文件方式,具有很好的參考價值,希望對大家有所幫助,如有錯誤或未考慮完全的地方,望不吝賜教
    2024-06-06

最新評論

三台县| 建阳市| 郯城县| 桐乡市| 成武县| 福泉市| 红安县| 怀仁县| 甘泉县| 微博| 合作市| 阿尔山市| 邯郸市| 屯昌县| 资源县| 大石桥市| 太保市| 中超| 来凤县| 福安市| 琼结县| 马山县| 行唐县| 会同县| 磐石市| 禹城市| 册亨县| 伊春市| 桓台县| 常德市| 沁水县| 射阳县| 曲靖市| 庆安县| 贵溪市| 漳平市| 宜黄县| 绥滨县| 双鸭山市| 清丰县| 上饶市|