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

pytorch中的模型訓(xùn)練(以CIFAR10數(shù)據(jù)集為例)

 更新時(shí)間:2023年06月15日 09:23:50   作者:MarkAssassin  
這篇文章主要介紹了pytorch中的模型訓(xùn)練(以CIFAR10數(shù)據(jù)集為例),具有很好的參考價(jià)值,希望對大家有所幫助。如有錯(cuò)誤或未考慮完全的地方,望不吝賜教

在pytorch模型訓(xùn)練時(shí),基本的訓(xùn)練步驟可以大致地歸納為:

準(zhǔn)備數(shù)據(jù)集--->搭建神經(jīng)網(wǎng)絡(luò)--->創(chuàng)建網(wǎng)絡(luò)模型--->創(chuàng)建損失函數(shù)--->設(shè)置優(yōu)化器--->訓(xùn)練步驟開始--->測試步驟開始

本文以pytorch官網(wǎng)中torchvision中的CIFAR10數(shù)據(jù)集為例進(jìn)行講解。

需要用到的庫為(這里說一個(gè)小技巧,比如可以在沒有import對應(yīng)庫的情況下先輸入"torch",,之后將光標(biāo)移到torch處單擊,這時(shí)左邊就會(huì)出現(xiàn)一個(gè)紅色的小燈泡,點(diǎn)開它就可以import對應(yīng)的庫了):

import torch
import torchvision
from torch import nn
from torch.utils.data import DataLoader
from torch.utils.tensorboard import SummaryWriter

準(zhǔn)備數(shù)據(jù)集

數(shù)據(jù)集分為“訓(xùn)練數(shù)據(jù)集”+“測試數(shù)據(jù)集”。

CIFAR數(shù)據(jù)集是由50000訓(xùn)練集和10000測試集組成。

這里可以調(diào)用torchvision.datasets對CIFAR10數(shù)據(jù)集進(jìn)行獲?。?/p>

#訓(xùn)練數(shù)據(jù)集
train_data=torchvision.datasets.CIFAR10(root='./dataset',train=True,transform=torchvision.transforms.ToTensor(),download=True)
#測試數(shù)據(jù)集
test_data=torchvision.datasets.CIFAR10(root='./dataset',train=False,transform=torchvision.transforms.ToTensor(),download=True)
#利用dataloader來加載數(shù)據(jù)集
train_data_loader=DataLoader(train_data,batch_size=64)
test_data_loader=DataLoader(test_data,batch_size=64)

root是數(shù)據(jù)集保存的路徑,這里筆者使用的是相對路徑;對于訓(xùn)練集train=True,而測試集train=False;transform是將數(shù)據(jù)集的類型轉(zhuǎn)換為tensor類型;download一般設(shè)置為True。

之后利用DataLoader對數(shù)據(jù)集進(jìn)行加載即可,其中batch_size表示單次傳遞給程序用以訓(xùn)練的數(shù)據(jù)(樣本)個(gè)數(shù)。

(這里可以再獲取下測試集數(shù)據(jù)的長度,這樣后面在測試步驟時(shí)就可以通過計(jì)算得到整體測試集上的準(zhǔn)確率

搭建神經(jīng)網(wǎng)絡(luò)

pytorch官網(wǎng)提供了一個(gè)神經(jīng)網(wǎng)絡(luò)的簡單實(shí)例:

import torch.nn as nn
import torch.nn.functional as F
class Model(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 20, 5)
        self.conv2 = nn.Conv2d(20, 20, 5)
    def forward(self, x):
        x = F.relu(self.conv1(x))
        return F.relu(self.conv2(x))

這個(gè)實(shí)例展示了神經(jīng)網(wǎng)絡(luò)的基本架構(gòu)。

從百度上我們可以搜索到CIFAR10數(shù)據(jù)集的網(wǎng)絡(luò)基本架構(gòu):

 從原理圖中可以看到,該網(wǎng)絡(luò)從輸入(inputs)到輸出(outputs)先后經(jīng)過了

卷積(Convolution)--->最大池化(Max-pooling)--->卷積--->最大池化--->卷積--->最大池化--->展平(Flatten)--->2次線性層

由此可以開始搭建神經(jīng)網(wǎng)絡(luò),筆者將網(wǎng)絡(luò)命名為MXC:

class Flatten(nn.Module):
    def forward(self, input):
        return input.view(input.size(0), -1)
#搭建神經(jīng)網(wǎng)絡(luò)
class MXC(nn.Module):
    def __init__(self):
        super(MXC, self).__init__()
        self.model=nn.Sequential(
            nn.Conv2d(3,32,5,1,2),
            nn.MaxPool2d(2),
            nn.Conv2d(32,32,5,1,2),
            nn.MaxPool2d(2),
            nn.Conv2d(32,64,5,1,2),
            nn.MaxPool2d(2),
            Flatten(),
            nn.Linear(64*4*4,64),
            nn.Linear(64,10)
        )
    def forward(self,x):
        x=self.model(x)
        return x

注:這里Flatten類自己寫的原因是筆者使用的torch版本中沒有展平類,因此需要自己構(gòu)建,可以參考解決 ImportError: cannot import name ‘Flatten‘ from ‘torch.nn‘

創(chuàng)建網(wǎng)絡(luò)模型

mxc=MXC()

創(chuàng)建損失函數(shù)

loss_fn=nn.CrossEntropyLoss()

設(shè)置優(yōu)化器

learning_rate=1e-2
optimizer=torch.optim.SGD(params=mxc.parameters(),lr=learning_rate)

這里params是網(wǎng)絡(luò)模型;lr是學(xué)習(xí)速率,一般設(shè)置小一些(0.01)。

訓(xùn)練步驟+測試步驟開始

先設(shè)置一些參數(shù)

#設(shè)置訓(xùn)練網(wǎng)絡(luò)的一些參數(shù)
#記錄訓(xùn)練的次數(shù)
total_train_step=0
#記錄測試的次數(shù)
total_test_step=0
#記錄測試的準(zhǔn)確率
total_accuracy=0
#訓(xùn)練的輪數(shù)
epoch=10

將訓(xùn)練步驟和測試步驟放入一個(gè)大循環(huán)中,進(jìn)入循環(huán)開始訓(xùn)練:

for i in range(epoch):
    print("-----第{}輪訓(xùn)練開始------".format(i+1))
    #訓(xùn)練步驟開始
    for data in train_data_loader:
        imgs,targets=data
        output=mxc(imgs)
        loss=loss_fn(output,targets)
        #優(yōu)化器優(yōu)化模型
        optimizer.zero_grad()#梯度清零
        loss.backward()#反向傳播
        optimizer.step()#參數(shù)優(yōu)化
        total_train_step=total_train_step+1
        if total_train_step%100==0:
            print("訓(xùn)練次數(shù):{} , Loss:{}".format(total_train_step, loss))  # 更正規(guī)的可以寫成loss.item()
    #測試步驟開始
    total_test_loss=0
    with torch.no_grad():
        for data in test_data_loader:
            imgs,targets=data
            output=mxc(imgs)
            loss=loss_fn(output,targets)
            total_test_loss=total_test_loss+loss
            accuracy=(output.argmax(1)==targets).sum()
            total_accuracy=total_accuracy+accuracy
    print("整體測試集上的Loss:{}".format(total_test_loss))
    print("整體測試集上的準(zhǔn)確率:{}".format(total_accuracy/test_data_size))
    total_test_step=total_test_step+1#測試的次數(shù),其實(shí)就是第幾輪
    #保存模型
    torch.save(mxc,"mxc_cpu{}.pth".format(i+1))#mxc_1是cpu版的
    print("模型已保存")

這里每一個(gè)data中包含有圖片+標(biāo)簽,需要將圖片(imgs)放入之前搭建好的神經(jīng)網(wǎng)絡(luò)模型mxc中去。

筆者這里設(shè)置的是每訓(xùn)練100次打印1次,訓(xùn)練完一輪后會(huì)將模型進(jìn)行保存,注意文件類型是pth格式。

如果想讓訓(xùn)練后的結(jié)果可視化,有兩種方法:

1.在循環(huán)前調(diào)用SummaryWriter添加tensorboard:

#添加tensorboard
writer=SummaryWriter('./logs_train')

并在循環(huán)的適當(dāng)位置中插入writer.add_scalar

        if total_train_step%100==0:
            print("訓(xùn)練次數(shù):{} , Loss:{}".format(total_train_step, loss))  # 更正規(guī)的可以寫成loss.item()
            writer.add_scalar("train_loss",loss.item(),total_train_step)
print("整體測試集上的Loss:{}".format(total_test_loss))
print("整體測試集上的準(zhǔn)確率:{}".format(total_accuracy/test_data_size))
total_test_step=total_test_step+1#測試的次數(shù),其實(shí)就是第幾輪
writer.add_scalar("test_loss",total_test_loss,total_test_step)
writer.add_scalar("test_accuracy",total_accuracy/test_data_size,total_test_step)

在運(yùn)行結(jié)束后,打開Terminal輸入(注意,前面需要顯示pytorch,因?yàn)橹挥性趐ytorch環(huán)境下才可以,如果沒有顯示還要切換到pytorch才行,可以輸入activate pytorch):

tensorboard --logdir=logs_train --port=6007

這里的“logs_train”是在SummaryWriter中設(shè)置的保存路徑。運(yùn)行之后,就可以在tensorboard中查看隨著訓(xùn)練次數(shù)的增加測試集上的Loss和準(zhǔn)確率的趨勢圖像。

2.調(diào)用matlab庫自己進(jìn)行繪制

總結(jié)

本文簡要介紹了pytorch模型訓(xùn)練的一個(gè)基本流程,并以CIFAR10數(shù)據(jù)集進(jìn)行了演示。

但這種方法是在CPU(device=“cpu”)上進(jìn)行訓(xùn)練的,訓(xùn)練速度比較慢,如果數(shù)據(jù)集十分龐大不建議使用這種方法,應(yīng)該在GPU(device=“cuda”)上進(jìn)行訓(xùn)練。

以上為個(gè)人經(jīng)驗(yàn),希望能給大家一個(gè)參考,也希望大家多多支持腳本之家。

相關(guān)文章

  • Numpy?三維數(shù)組索引與切片的實(shí)現(xiàn)

    Numpy?三維數(shù)組索引與切片的實(shí)現(xiàn)

    本文主要介紹了Numpy?三維數(shù)組索引與切片,文中通過示例代碼介紹的非常詳細(xì),對大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來一起學(xué)習(xí)學(xué)習(xí)吧
    2023-03-03
  • Python批量實(shí)現(xiàn)文件查找、打包與修改名字

    Python批量實(shí)現(xiàn)文件查找、打包與修改名字

    這篇文章主要為大家詳細(xì)介紹了如何通過Python批量實(shí)現(xiàn)文件查找、打包與修改名字的功能,文中的示例代碼講解詳細(xì),有需要的小伙伴可以進(jìn)行參考
    2025-11-11
  • Python中代碼開發(fā)的調(diào)試技巧分享

    Python中代碼開發(fā)的調(diào)試技巧分享

    在軟件開發(fā)的世界中,調(diào)試是每個(gè)程序員都無法回避的核心技能,本文將深入探討Python程序員在面對復(fù)雜問題時(shí)應(yīng)具備的調(diào)試思維模式,并提供一套完整的調(diào)試方法和實(shí)踐工具,希望對大家有所幫助
    2025-11-11
  • python global的創(chuàng)建和修改實(shí)例講解

    python global的創(chuàng)建和修改實(shí)例講解

    在本篇文章里小編給大家整理了一篇關(guān)于python global的創(chuàng)建和修改實(shí)例講解內(nèi)容,有興趣的朋友們可以學(xué)習(xí)下。
    2021-09-09
  • 淺析Python3中遍歷目錄的三種方法

    淺析Python3中遍歷目錄的三種方法

    在學(xué)習(xí)中,工作中,我們經(jīng)常會(huì)說遍歷一下當(dāng)前目錄咯,那么Python3中遍歷目錄的方法具體都有哪些呢并且如何操作呢,下面小編就來和大家簡單聊聊吧
    2023-07-07
  • python3.6使用pymysql連接Mysql數(shù)據(jù)庫

    python3.6使用pymysql連接Mysql數(shù)據(jù)庫

    這篇文章主要為大家詳細(xì)介紹了python3.6使用pymysql連接Mysql數(shù)據(jù)庫,以及簡單的增刪改查操作,具有一定的參考價(jià)值,感興趣的小伙伴們可以參考一下
    2018-05-05
  • 在 Python 中接管鍵盤中斷信號(hào)的實(shí)現(xiàn)方法

    在 Python 中接管鍵盤中斷信號(hào)的實(shí)現(xiàn)方法

    要使用信號(hào),我們需用導(dǎo)入 Python 的signal庫。然后自定義一個(gè)信號(hào)回調(diào)函數(shù),當(dāng) Python 收到某個(gè)信號(hào)時(shí),調(diào)用這個(gè)函數(shù)。 ,下面通過實(shí)例代碼給大家介紹在 Python 中接管鍵盤中斷信號(hào),需要的朋友可以參考下
    2020-02-02
  • Python OpenCV特征檢測之特征匹配方式詳解

    Python OpenCV特征檢測之特征匹配方式詳解

    OpenCV中提供了兩種技術(shù)用于特征匹配,分別為Brute-Force匹配器和基于FLANN的匹配器。本文將為大家詳細(xì)介紹一下這兩種匹配方式,需要的可以參考一下
    2021-12-12
  • Python使用Django實(shí)現(xiàn)博客系統(tǒng)完整版

    Python使用Django實(shí)現(xiàn)博客系統(tǒng)完整版

    這篇文章主要為大家詳細(xì)介紹了Python利用Django完整的開發(fā)一個(gè)博客系統(tǒng),文中示例代碼介紹的非常詳細(xì),具有一定的參考價(jià)值,感興趣的小伙伴們可以參考一下
    2018-03-03
  • python中使用while循環(huán)的實(shí)例

    python中使用while循環(huán)的實(shí)例

    在本篇內(nèi)容里小編給各位整理的是關(guān)于python中使用while循環(huán)的實(shí)例以及相關(guān)知識(shí)點(diǎn),需要的朋友們學(xué)習(xí)下。
    2019-08-08

最新評論

涿鹿县| 五华县| 崇仁县| 大足县| 珠海市| 桦甸市| 曲阳县| 黎平县| 松江区| 耿马| 德保县| 潮安县| 中方县| 安仁县| 新营市| 永福县| 兰州市| 外汇| 逊克县| 郎溪县| 弋阳县| 长乐市| 营口市| 香河县| 秀山| 淮阳县| 三明市| 五大连池市| 淮北市| 柳河县| 丰台区| 宾川县| 龙山县| 竹山县| 蛟河市| 丁青县| 衡阳市| 孝义市| 饶平县| 仪征市| 临沧市|