PyTorch如何利用parameters()獲取模型參數(shù)
利用parameters()獲取模型參數(shù)
在PyTorch中,可以使用parameters函數(shù)來(lái)獲取模型中的所有可學(xué)習(xí)參數(shù)。
以下是一個(gè)示例:
import torch.nn as nn class MyModel(nn.Module): ? ? def __init__(self): ? ? ? ? super(MyModel, self).__init__() ? ? ? ? self.fc1 = nn.Linear(10, 5) ? ? ? ? self.fc2 = nn.Linear(5, 1) ? ? def forward(self, x): ? ? ? ? x = self.fc1(x) ? ? ? ? x = self.fc2(x) ? ? ? ? return x model = MyModel() params = list(model.parameters())
在這個(gè)示例中,我們首先定義了一個(gè)包含兩個(gè)線性層的神經(jīng)網(wǎng)絡(luò),然后通過(guò)list(model.parameters())獲取了模型中的所有可學(xué)習(xí)參數(shù)。
這些參數(shù)存儲(chǔ)在一個(gè)Python列表中,可以用于進(jìn)行優(yōu)化器的初始化和模型的保存和加載。
PyTorch中模型的parameters()方法
首先先定義一個(gè)模型:
import torch as t import torch.nn as nn class A(nn.Module): ? ? def __init__(self): ? ? ? ? super().__init__() ? ? ? ? self.conv1 = nn.Conv2d(2, 2, 3) ? ? ? ? self.conv2 = nn.Conv2d(2, 2, 3) ? ? ? ? self.conv3 = nn.Conv2d(2, 2, 3) ? ? def forward(self, x): ? ? ? ? x = self.conv1(x) ? ? ? ? x = self.conv2(x) ? ? ? ? x = self.conv3(x) ? ? ? ? return x
然后打印出該模型的參數(shù):
pythona = A() print(a.parameters()) #<generator object Module.parameters at 0x7f7b740d2360>
以上代碼說(shuō)明parameters()會(huì)返回一個(gè)生成器(迭代器)
然后將其迭代打印出來(lái):
print(list(a.parameters())):#將迭代器轉(zhuǎn)換成列表 Parameter containing: tensor([[[[-0.0299, ?0.0891, ?0.0303], ? ? ? ? ? [ 0.0869, -0.0230, -0.1760], ? ? ? ? ? [ 0.1408, ?0.0348, ?0.1795]], ? ? ? ? ?[[ 0.2001, ?0.0023, -0.1775], ? ? ? ? ? [ 0.0947, -0.0231, -0.1756], ? ? ? ? ? [ 0.1201, -0.0997, -0.0303]]], ? ? ? ? [[[-0.0425, ?0.0748, -0.1754], ? ? ? ? ? [-0.1191, -0.1203, -0.1219], ? ? ? ? ? [-0.0794, ?0.0895, -0.1719]], ? ? ? ? ?[[ 0.1968, -0.0463, ?0.0550], ? ? ? ? ? [-0.0386, ?0.1594, ?0.1282], ? ? ? ? ? [-0.0009, ?0.2167, -0.1783]]]], requires_grad=True) Parameter containing: tensor([ 0.0147, -0.0406], requires_grad=True) Parameter containing: tensor([[[[-0.0578, -0.1114, -0.1194], ? ? ? ? ? [-0.1469, -0.1175, -0.1616], ? ? ? ? ? [-0.2289, -0.0975, -0.1700]], ? ? ? ? ?[[-0.0894, ?0.0074, ?0.1222], ? ? ? ? ? [-0.0176, -0.0509, ?0.1622], ? ? ? ? ? [-0.0405, -0.1349, ?0.1782]]], ? ? ? ? [[[-0.0739, ?0.2167, ?0.1864], ? ? ? ? ? [ 0.0956, -0.1761, ?0.0464], ? ? ? ? ? [ 0.0062, -0.0685, ?0.0748]], ? ? ? ? ?[[ 0.1085, ?0.1481, ?0.1334], ? ? ? ? ? [ 0.2236, -0.0706, -0.0224], ? ? ? ? ? [ 0.0079, -0.1835, -0.0407]]]], requires_grad=True) Parameter containing: tensor([-8.0720e-05, ?1.6026e-01], requires_grad=True) Parameter containing: tensor([[[[-0.0702, ?0.1846, ?0.0419], ? ? ? ? ? [-0.1891, -0.0893, -0.0024], ? ? ? ? ? [-0.0349, -0.0213, ?0.0936]], ? ? ? ? ?[[-0.1062, ?0.1242, ?0.0391], ? ? ? ? ? [-0.1924, ?0.0535, -0.1480], ? ? ? ? ? [ 0.0400, -0.0487, -0.2317]]], ? ? ? ? [[[ 0.1202, ?0.0961, ?0.2336], ? ? ? ? ? [ 0.2225, -0.2294, -0.2283], ? ? ? ? ? [-0.0963, -0.0311, -0.2354]], ? ? ? ? ?[[ 0.0676, -0.0439, -0.0962], ? ? ? ? ? [-0.2316, -0.0639, -0.0671], ? ? ? ? ? [ 0.1737, -0.1169, -0.1751]]]], requires_grad=True) Parameter containing: tensor([-0.1939, -0.0959], requires_grad=True)
從以上結(jié)果可以看出列表中有6個(gè)元素,由于nn.Conv2d()的參數(shù)包括self.weight和self.bias兩部分,所以每個(gè)2D卷積層包括兩部分的參數(shù).注意self.bias是加在每個(gè)通道上的,所以self.bias的長(zhǎng)度與output_channl相同
心得:
parameters()會(huì)返回一個(gè)生成器(迭代器),生成器每次生成的是Tensor類(lèi)型的數(shù)據(jù).
總結(jié)
以上為個(gè)人經(jīng)驗(yàn),希望能給大家一個(gè)參考,也希望大家多多支持腳本之家。
相關(guān)文章
python做聲音識(shí)別的實(shí)現(xiàn)示例
本文主要介紹了python做聲音識(shí)別的實(shí)現(xiàn)示例,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來(lái)一起學(xué)習(xí)學(xué)習(xí)吧2025-10-10
Python實(shí)現(xiàn)隨機(jī)游走的示例代碼
隨機(jī)游走是一個(gè)數(shù)學(xué)對(duì)象,稱為隨機(jī)或隨機(jī)過(guò)程,它描述了一條路徑,該路徑由一些數(shù)學(xué)空間上的一系列隨機(jī)步驟組成,下面我們就來(lái)學(xué)習(xí)一下Python如何實(shí)現(xiàn)隨機(jī)游走的吧2023-12-12
python 計(jì)算方位角實(shí)例(根據(jù)兩點(diǎn)的坐標(biāo)計(jì)算)
今天小編就為大家分享一篇python 計(jì)算方位角實(shí)例(根據(jù)兩點(diǎn)的坐標(biāo)計(jì)算),具有很好的參考價(jià)值,希望對(duì)大家有所幫助。一起跟隨小編過(guò)來(lái)看看吧2020-01-01
Python DNS查詢放大攻擊實(shí)現(xiàn)原理解析
這篇文章主要介紹了Python DNS查詢放大攻擊實(shí)現(xiàn),文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來(lái)一起學(xué)習(xí)吧2022-10-10
Python反射用法實(shí)戰(zhàn)完整學(xué)習(xí)筆記
反射就是程序在運(yùn)行時(shí)能夠觀察自己,獲取、檢查和修改自身狀態(tài)或行為的一種能力,這篇文章主要介紹了Python反射用法的相關(guān)資料,文中通過(guò)代碼介紹的非常詳細(xì),需要的朋友可以參考下2026-01-01
Python Excel vlookup函數(shù)實(shí)現(xiàn)過(guò)程解析
這篇文章主要介紹了Python Excel vlookup函數(shù)實(shí)現(xiàn)過(guò)程解析,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友可以參考下2020-06-06
一文學(xué)會(huì)使用OpenCV構(gòu)建文檔掃描儀
本文將使用 OpenCV,創(chuàng)建一個(gè)簡(jiǎn)單的文檔掃描儀,就像常用的攝像頭掃描儀應(yīng)用程序一樣,這篇文章主要給大家介紹了關(guān)于使用OpenCV構(gòu)建文檔掃描儀的相關(guān)資料,需要的朋友可以參考下2022-11-11

