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

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

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

1. 深度學習與PyTorch簡介

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

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

2.1 環(huán)境檢查

首先檢查PyTorch及相關庫的版本,確保環(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ù)加載與預處理

MNIST數(shù)據(jù)集包含60,000個訓練樣本和10,000個測試樣本,每個樣本為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ù)存儲路徑
  • train:是否為訓練集
  • download:是否自動下載
  • transform:數(shù)據(jù)預處理轉換,ToTensor()將PIL圖像轉換為張量并歸一化到[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實現(xiàn)數(shù)據(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)絡模型設計

3.1 設備選擇

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

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)絡架構

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

class NeuralNetwork(nn.Module):
    def __init__(self):
        super().__init__()
        self.a = 10
        self.flatten = nn.Flatten()
        原始架構
        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 模型實例化

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

4. 訓練與評估流程

4.1 訓練函數(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

訓練步驟

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

4.2 測試函數(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}")

測試要點

  • model.eval():設置為評估模式(禁用Dropout)
  • torch.no_grad():禁用梯度計算,節(jié)省內存
  • pred.argmax(1):獲取預測類別

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

loss_fn = nn.CrossEntropyLoss()

損失函數(shù)說明

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

6. 模型訓練與評估

6.1 優(yōu)化器配置

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

6.2 單次訓練與測試

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

6.3 多輪訓練(可選)

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. 提高準確率的優(yōu)化方式

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

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

相關文章

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

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

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

    Python中的global與nonlocal關鍵字詳解

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

    Python中Matplotlib圖像添加標簽的方法實現(xiàn)

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

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

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

    用python畫一只帥氣的皮卡丘

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

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

    int(M)拆分來說,int是代表整型數(shù)據(jù)那,么中間的M應該是代表多少位了,后來查mysql手冊也得知了我的理解是正確的,下面這篇文章小編就來舉例詳細說明。 文中介紹的很詳細,相信對大家的理解和學習很有幫助,有需要的朋友們下面就來學習學習吧。
    2016-11-11
  • Python attrs提高面向對象編程效率詳細

    Python attrs提高面向對象編程效率詳細

    Python是面向對象的語言,一般情況下使用面向對象編程會使得開發(fā)效率更高,軟件質量更好,并且代碼更易于擴展,可讀性和可維護性也更高,但是Python的類寫起來是真的累,這是可以在創(chuàng)建類的時候自動添加上attrs模塊,下面文章我們就來介紹這個東西,需要的朋友可參考一下
    2021-09-09
  • Python加密方法小結【md5,base64,sha1】

    Python加密方法小結【md5,base64,sha1】

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

    Python超簡單容易上手的畫圖工具庫(適合新手)

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

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

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

最新評論

综艺| 保亭| 建湖县| 米泉市| 横峰县| 始兴县| 谷城县| 黄梅县| 庆云县| 大城县| 比如县| 阜阳市| 墨江| 武宣县| 镇江市| 阿拉尔市| 凯里市| 砚山县| 吉安市| 泾源县| 修水县| 姚安县| 无极县| 定南县| 上林县| 鄂温| 合江县| 武强县| 晋江市| 大邑县| 南郑县| 巍山| 洮南市| 霍邱县| 禹城市| 安仁县| 礼泉县| 西安市| 西乌| 丰原市| 新源县|