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

PyTorch核心方法之state_dict()、parameters()參數(shù)打印與應(yīng)用案例

 更新時間:2025年12月13日 14:56:20   作者:木棉知行者  
PyTorch是一個流行的開源深度學(xué)習(xí)框架,提供了靈活且高效的方式來訓(xùn)練和部署神經(jīng)網(wǎng)絡(luò),這篇文章主要介紹了PyTorch核心方法之state_dict()、parameters()參數(shù)打印與應(yīng)用案例的相關(guān)資料,需要的朋友可以參考下

前言

本文以 LeNet-5 模型為案例,介紹了 PyTorch 中打印模型參數(shù)的相關(guān)方法。首先展示了 LeNet-5 模型的結(jié)構(gòu)定義及打印結(jié)果;隨后詳細說明了三種獲取模型參數(shù)的方式:

  • state_dict()方法返回有序字典形式的可學(xué)習(xí)參數(shù),包含參數(shù)名稱和對應(yīng)張量;
  • parameters()方法返回生成器,僅包含各層參數(shù)信息;
  • named_parameters()方法返回生成器,包含模型名稱和對應(yīng)參數(shù)信息;
    最后提供了利用named_parameters()進行模型結(jié)構(gòu)凍結(jié)的示例,可打印確認凍結(jié)的網(wǎng)絡(luò)名稱。

模型案例

本文以LeNet-5為基礎(chǔ)模型,快速驗證模型參數(shù)打印過程。

import os 
os.environ['CUDA_VISIBLE_DEVICES'] = '3'
import torch 
import torch.nn.functional as F 
import torch.nn as nn

class LeNet5(nn.Module):
    def __init__(self):
        super(LeNet5, self).__init__()
        # 1 input image channel, 6 output channels, 5x5 square convolution
        # kernel
        self.conv1 = nn.Conv2d(1, 6, 5)
        self.conv2 = nn.Conv2d(6, 16, 5)
        # an affine operation: y = Wx + b
        self.fc1 = nn.Linear(16 * 5 * 5, 120) # 這里論文上寫的是conv,官方教程用了線性層
        self.fc2 = nn.Linear(120, 84)
        self.fc3 = nn.Linear(84, 10)

    def forward(self, x):
        # Max pooling over a (2, 2) window
        x = F.max_pool2d(F.relu(self.conv1(x)), (2, 2))
        # If the size is a square you can only specify a single number
        x = F.max_pool2d(F.relu(self.conv2(x)), 2)
        x = x.view(-1, self.num_flat_features(x))
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = self.fc3(x)
        return x

    def num_flat_features(self, x):
        size = x.size()[1:]  # all dimensions except the batch dimension
        num_features = 1
        for s in size:
            num_features *= s
        return num_features

net = LeNet5()
print(net)

模型結(jié)構(gòu)打印如下。

A. state_dict()方法驗證

在 PyTorch 中,state_dict() 是核心方法之一,用于以有序字典(OrderedDict)的形式返回模型 / 優(yōu)化器等實例的可學(xué)習(xí)參數(shù)(或狀態(tài)),是模型保存、加載、遷移學(xué)習(xí)的基礎(chǔ)。

state_dict() 本質(zhì)是一個 Python 字典(PyTorch 中為 OrderedDict),鍵為參數(shù) / 狀態(tài)的名稱(字符串),值為對應(yīng)的張量(torch.Tensor)。

print(type(net.state_dict()))   # <class 'collections.OrderedDict'>
## 遍歷打印
for model_key in net.state_dict():      # 【字典格式】的遍歷,獲取的是模型的名稱
    print(f"{model_key}: {net.state_dict()[model_key].size()}")

對于Lenet-5模型進行打印,可以看到state_dict()的類型為 <class 'collections.OrderedDict'>,各層名稱及參數(shù)尺寸如下圖所示。

B. parameters()

parameters()方法也可以獲取到模型的參數(shù)??梢钥闯?,parameters()獲取到的是一個生成器,其中僅包含各層參數(shù)的信息。

params = net.parameters()   
print(type(params))   # <class 'generator'>  生成器  

for param in params:    
    print(param.size())   # 只包含參數(shù)信息:具體的參數(shù)尺寸

對Lenet-5進行模型參數(shù)打印。

如果也需要模型名稱信息,可以使用named_parameters()方法。該方法獲取的也是一個生成器,其中返回的是一個元組,包括模型名稱和對應(yīng)的參數(shù)。

named_params = net.named_parameters()   
print(type(named_params))   # <class 'generator'>  也是一個生成器

for name, param in named_params:
    print(f"{name}: {param.size()}")   # 同時獲取網(wǎng)絡(luò)名稱和網(wǎng)絡(luò)參數(shù)

對Lenet-5進行模型名稱及參數(shù)尺寸信息打印:

C. 模型結(jié)構(gòu)凍結(jié)示例

該方法可以在對模型結(jié)構(gòu)凍結(jié)時使用,如下述示例對模型結(jié)構(gòu)m的參數(shù)進行凍結(jié),同時打印確認凍結(jié)包含哪些網(wǎng)絡(luò)結(jié)構(gòu)。

# 示例
for name, param in m.named_parameters():
	param.requires_grad = False
	print(f"Freezing layer {name}")

總結(jié) 

到此這篇關(guān)于PyTorch核心方法之state_dict()、parameters()參數(shù)打印與應(yīng)用案例的文章就介紹到這了,更多相關(guān)PyTorch state_dict()、parameters()參數(shù)打印內(nèi)容請搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家!

相關(guān)文章

最新評論

新津县| 宣化县| 广宗县| 佳木斯市| 左云县| 开平市| 同仁县| 墨竹工卡县| 阜阳市| 龙川县| 彭阳县| 尖扎县| 兖州市| 根河市| 鄄城县| 无棣县| 小金县| 天镇县| 平乡县| 扶余县| 福安市| 体育| 杂多县| 宁蒗| 陈巴尔虎旗| 西丰县| 晋中市| 淮南市| 兴城市| 乌拉特中旗| 同仁县| 夏河县| 吕梁市| 洛隆县| 乌什县| 阿坝县| 全南县| 峨边| 二连浩特市| 正阳县| 广汉市|