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

PyTorch基于MNIST的手寫數(shù)字識(shí)別

 更新時(shí)間:2026年01月19日 08:54:56   作者:子夜江寒  
本文介紹了使用PyTorch框架構(gòu)建深度學(xué)習(xí)模型處理MNIST手寫數(shù)字識(shí)別的完整流程,文中通過示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來一起學(xué)習(xí)學(xué)習(xí)吧

1. 深度學(xué)習(xí)與PyTorch簡(jiǎn)介

深度學(xué)習(xí)作為機(jī)器學(xué)習(xí)的重要分支,已在計(jì)算機(jī)視覺、自然語(yǔ)言處理等領(lǐng)域取得了顯著成果。PyTorch是由Facebook開源的深度學(xué)習(xí)框架,以其動(dòng)態(tài)計(jì)算圖和直觀的API設(shè)計(jì)而廣受歡迎。本文以經(jīng)典的MNIST手寫數(shù)字?jǐn)?shù)據(jù)集為例,展示如何利用PyTorch框架構(gòu)建并訓(xùn)練深度學(xué)習(xí)模型。

2. 環(huán)境配置與數(shù)據(jù)準(zhǔn)備

2.1 環(huán)境檢查

首先檢查PyTorch及相關(guān)庫(kù)的版本,確保環(huán)境配置正確:

import torch
import torchvision
import torchaudio
from torch import nn
from torch.utils.data import DataLoader
from torchvision import datasets
from torchvision.transforms import ToTensor
from matplotlib import pyplot as plt

print(torch.__version__)
print(torchaudio.__version__)
print(torchvision.__version__)

2.2 數(shù)據(jù)加載與預(yù)處理

MNIST數(shù)據(jù)集包含60,000個(gè)訓(xùn)練樣本和10,000個(gè)測(cè)試樣本,每個(gè)樣本為28×28像素的灰度手寫數(shù)字圖像。

training_data = datasets.MNIST(
    root="data",
    train=True,
    download=True,
    transform=ToTensor(),
)

test_data = datasets.MNIST(
    root="data",
    train=False,
    download=True,
    transform=ToTensor(),
)

參數(shù)

  • root:數(shù)據(jù)存儲(chǔ)路徑
  • train:是否為訓(xùn)練集
  • download:是否自動(dòng)下載
  • transform:數(shù)據(jù)預(yù)處理轉(zhuǎn)換,ToTensor()將PIL圖像轉(zhuǎn)換為張量并歸一化到[0,1]

2.3 數(shù)據(jù)可視化

我們可以查看數(shù)據(jù)集的樣本分布:

print(len(training_data))

figure = plt.figure()
for i in range(9):
    img, label = training_data[i + 59000]
    figure.add_subplot(3, 3, i + 1)
    plt.title(label)
    plt.axis("off")
    plt.imshow(img.squeeze(), cmap="gray")
plt.show()

2.4 數(shù)據(jù)批量加載

使用DataLoader實(shí)現(xiàn)數(shù)據(jù)的批量加載和隨機(jī)打亂:

# 增加批次大小
train_dataloader = DataLoader(training_data, batch_size=128)  # 增大batch size
test_dataloader = DataLoader(test_data, batch_size=128)

for X, y in test_dataloader:
    print(f"Shape of X[N,C,H,W]:{X.shape}")
    print(f"Shape of y:{y.shape} {y.dtype}")
    break

3. 神經(jīng)網(wǎng)絡(luò)模型設(shè)計(jì)

3.1 設(shè)備選擇

根據(jù)可用硬件選擇計(jì)算設(shè)備:

device = "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"
print(f"Using {device} device")

3.2 神經(jīng)網(wǎng)絡(luò)架構(gòu)

設(shè)計(jì)一個(gè)包含多個(gè)全連接層的深度神經(jīng)網(wǎng)絡(luò):

class NeuralNetwork(nn.Module):
    def __init__(self):
        super().__init__()
        self.a = 10
        self.flatten = nn.Flatten()
        原始架構(gòu)
        self.hidden1 = nn.Linear(28 * 28, 128)
        self.hidden2 = nn.Linear(128, 256)
        self.out = nn.Linear(256, 10)
        
    
    def forward(self, x):
        # 原始前向傳播
        x = self.flatten(x)
        x = self.hidden1(x)
        x = torch.sigmoid(x)
        x = self.hidden2(x)
        x = torch.sigmoid(x)
        return x

3.3 模型實(shí)例化

model = NeuralNetwork().to(device)
print(model)

4. 訓(xùn)練與評(píng)估流程

4.1 訓(xùn)練函數(shù)

def train(dataloader, model, loss_fn, optimizer):
    model.train()
    batch_size_num = 1
    for X, y in dataloader:
        X, y = X.to(device), y.to(device)
        pred = model.forward(X)
        loss = loss_fn(pred, y)
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

        loss_value = loss.item()
        if batch_size_num % 100 == 0:
            print(f"loss: {loss_value:>7f} [number:{batch_size_num}]")
        batch_size_num += 1

訓(xùn)練步驟

  1. model.train():設(shè)置為訓(xùn)練模式(啟用Dropout)
  2. 前向傳播計(jì)算預(yù)測(cè)值
  3. 計(jì)算損失函數(shù)值
  4. optimizer.zero_grad():清空梯度
  5. loss.backward():反向傳播計(jì)算梯度
  6. optimizer.step():更新模型參數(shù)

4.2 測(cè)試函數(shù)

def test(dataloader, model, loss_fn):
    size = len(dataloader.dataset)
    num_batches = len(dataloader)
    model.eval()
    test_loss, correct = 0, 0
    with torch.no_grad():
        for X, y in dataloader:
            X, y = X.to(device), y.to(device)
            pred = model.forward(X)
            test_loss = loss_fn(pred, y)
            correct += (pred.argmax(1) == y).type(torch.float).sum().item()
            a = (pred.argmax(1) == y)
            b = (pred.argmax(1) == y).type(torch.float)
    test_loss /= num_batches
    correct /= size

    print(f"Test result:\n Accuracy:{(100 * correct):.2f}%, Avg loss: {test_loss}")

測(cè)試要點(diǎn)

  • model.eval():設(shè)置為評(píng)估模式(禁用Dropout)
  • torch.no_grad():禁用梯度計(jì)算,節(jié)省內(nèi)存
  • pred.argmax(1):獲取預(yù)測(cè)類別

5. 損失函數(shù)配置

loss_fn = nn.CrossEntropyLoss()

損失函數(shù)說明

  • 使用CrossEntropyLoss,適用于多分類問題
  • 結(jié)合了LogSoftmax和NLLLoss,直接輸出分類概率

6. 模型訓(xùn)練與評(píng)估

6.1 優(yōu)化器配置

# 原始優(yōu)化器
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)

6.2 單次訓(xùn)練與測(cè)試

train(train_dataloader, model, loss_fn, optimizer)
test(train_dataloader, model, loss_fn)

6.3 多輪訓(xùn)練(可選)

epochs = 10
for t in range(epochs):
    print(f"Epoch {t+1}\n----------------------")
    train(train_dataloader, model, loss_fn, optimizer)
print("Done!")
test(test_dataloader, model, loss_fn)

7. 提高準(zhǔn)確率的優(yōu)化方式

  1. 層數(shù)增加:從2層隱藏層增加到3層,增強(qiáng)模型表達(dá)能力
  2. 神經(jīng)元增加:第一層從128個(gè)神經(jīng)元增加到512個(gè)
  3. 激活函數(shù):用ReLU替代sigmoid,緩解梯度消失問題
  4. 正則化:添加Dropout層(0.2丟棄率),防止過擬合
  5. 改進(jìn)優(yōu)化器:降低學(xué)習(xí)率
        # 改進(jìn)架構(gòu)
        self.hidden1 = nn.Linear(28 * 28, 512)  # 增加神經(jīng)元
        self.dropout1 = nn.Dropout(0.2)  # 添加Dropout
        self.hidden2 = nn.Linear(512, 256)
        self.dropout2 = nn.Dropout(0.2)  # 添加Dropout
        self.hidden3 = nn.Linear(256, 128)  # 增加一層
        self.out = nn.Linear(128, 10)
        # 改進(jìn)的前向傳播
        x = self.flatten(x)
        x = self.hidden1(x)
        x = torch.relu(x)  # 使用ReLU替代sigmoid
        x = self.dropout1(x)  # 訓(xùn)練時(shí)隨機(jī)丟棄
        x = self.hidden2(x)
        x = torch.relu(x)  # 使用ReLU替代sigmoid
        x = self.dropout2(x)  # 訓(xùn)練時(shí)隨機(jī)丟棄
        x = self.hidden3(x)
        x = torch.relu(x)
        x = self.out(x)
# 改進(jìn)優(yōu)化器
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)  # 降低學(xué)習(xí)率

到此這篇關(guān)于PyTorch基于MNIST的手寫數(shù)字識(shí)別的文章就介紹到這了,更多相關(guān)PyTorch MNIST手寫數(shù)字識(shí)別內(nèi)容請(qǐng)搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家!

相關(guān)文章

  • Python使用graphviz畫流程圖過程解析

    Python使用graphviz畫流程圖過程解析

    這篇文章主要介紹了Python使用graphviz畫流程圖過程解析,文中通過示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友可以參考下
    2020-03-03
  • Python中的global與nonlocal關(guān)鍵字詳解

    Python中的global與nonlocal關(guān)鍵字詳解

    在Python編程中變量作用域是一個(gè)非常重要的概念,global和nonlocal關(guān)鍵字就能派上用場(chǎng)了,本文將詳細(xì)介紹這兩個(gè)關(guān)鍵字的用法、區(qū)別及適用場(chǎng)景,幫助大家徹底弄懂global與nonlocal關(guān)鍵字
    2025-07-07
  • Python中Matplotlib圖像添加標(biāo)簽的方法實(shí)現(xiàn)

    Python中Matplotlib圖像添加標(biāo)簽的方法實(shí)現(xiàn)

    本文主要介紹了Python中Matplotlib圖像添加標(biāo)簽的方法實(shí)現(xiàn),文中通過示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來一起學(xué)習(xí)學(xué)習(xí)吧
    2023-04-04
  • 在python中實(shí)現(xiàn)求輸出1-3+5-7+9-......101的和

    在python中實(shí)現(xiàn)求輸出1-3+5-7+9-......101的和

    這篇文章主要介紹了在python中實(shí)現(xiàn)求輸出1-3+5-7+9-......101的和,具有很好的參考價(jià)值,希望對(duì)大家有所幫助。一起跟隨小編過來看看吧
    2020-04-04
  • 用python畫一只帥氣的皮卡丘

    用python畫一只帥氣的皮卡丘

    大家好,本篇文章主要講的是用python畫一只帥氣的皮卡丘,感興趣的同學(xué)趕快來看一看吧,對(duì)你有幫助的話記得收藏一下
    2022-01-01
  • 詳解MySQL數(shù)據(jù)類型int(M)中M的含義

    詳解MySQL數(shù)據(jù)類型int(M)中M的含義

    int(M)拆分來說,int是代表整型數(shù)據(jù)那,么中間的M應(yīng)該是代表多少位了,后來查mysql手冊(cè)也得知了我的理解是正確的,下面這篇文章小編就來舉例詳細(xì)說明。 文中介紹的很詳細(xì),相信對(duì)大家的理解和學(xué)習(xí)很有幫助,有需要的朋友們下面就來學(xué)習(xí)學(xué)習(xí)吧。
    2016-11-11
  • Python attrs提高面向?qū)ο缶幊绦试敿?xì)

    Python attrs提高面向?qū)ο缶幊绦试敿?xì)

    Python是面向?qū)ο蟮恼Z(yǔ)言,一般情況下使用面向?qū)ο缶幊虝?huì)使得開發(fā)效率更高,軟件質(zhì)量更好,并且代碼更易于擴(kuò)展,可讀性和可維護(hù)性也更高,但是Python的類寫起來是真的累,這是可以在創(chuàng)建類的時(shí)候自動(dòng)添加上attrs模塊,下面文章我們就來介紹這個(gè)東西,需要的朋友可參考一下
    2021-09-09
  • Python加密方法小結(jié)【md5,base64,sha1】

    Python加密方法小結(jié)【md5,base64,sha1】

    這篇文章主要介紹了Python加密方法,結(jié)合實(shí)例形式總結(jié)分析了md5,base64,sha1的簡(jiǎn)單加密方法,需要的朋友可以參考下
    2017-07-07
  • Python超簡(jiǎn)單容易上手的畫圖工具庫(kù)(適合新手)

    Python超簡(jiǎn)單容易上手的畫圖工具庫(kù)(適合新手)

    這篇文章主要給大家介紹了關(guān)于Python超簡(jiǎn)單容易上手的畫圖工具庫(kù)的相關(guān)資料,文中通過圖文介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來一起學(xué)習(xí)學(xué)習(xí)吧
    2021-05-05
  • TensorFlow實(shí)現(xiàn)自定義Op方式

    TensorFlow實(shí)現(xiàn)自定義Op方式

    今天小編就為大家分享一篇TensorFlow實(shí)現(xiàn)自定義Op方式,具有很好的參考價(jià)值,希望對(duì)大家有所幫助。一起跟隨小編過來看看吧
    2020-02-02

最新評(píng)論

行唐县| 林芝县| 枝江市| 丹阳市| 连江县| 奉化市| 上虞市| 陆良县| 轮台县| 南华县| 海丰县| 大竹县| 瑞金市| 常州市| 渝中区| 民丰县| 灯塔市| 区。| 剑河县| 长兴县| 邳州市| 津南区| 广宁县| 塔河县| 佛冈县| 嘉峪关市| 吴堡县| 庄河市| 饶河县| 浙江省| 错那县| 武宁县| 新巴尔虎右旗| 图木舒克市| 蓬溪县| 富阳市| 洛浦县| 丁青县| 涟水县| 泸州市| 贵阳市|