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

Pytorch實現(xiàn)LSTM和GRU示例

 更新時間:2020年01月14日 10:29:08   作者:winycg  
今天小編就為大家分享一篇Pytorch實現(xiàn)LSTM和GRU示例,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧

為了解決傳統(tǒng)RNN無法長時依賴問題,RNN的兩個變體LSTM和GRU被引入。

LSTM

Long Short Term Memory,稱為長短期記憶網絡,意思就是長的短時記憶,其解決的仍然是短時記憶問題,這種短時記憶比較長,能一定程度上解決長時依賴。

上圖為LSTM的抽象結構,LSTM由3個門來控制,分別是輸入門、遺忘門和輸出門。輸入門控制網絡的輸入,遺忘門控制著記憶單元,輸出門控制著網絡的輸出。最為重要的就是遺忘門,可以決定哪些記憶被保留,由于遺忘門的作用,使得LSTM具有長時記憶的功能。對于給定的任務,遺忘門能夠自主學習保留多少之前的記憶,網絡能夠自主學習。

具體看LSTM單元的內部結構:

在每篇文章中,作者都會使用和標準LSTM稍微不同的版本,針對特定的任務,特定的網絡結構往往表現(xiàn)更好。

GRU

上述的過程的線性變換沒有使用偏置。隱藏狀態(tài)參數不再是標準RNN的4倍,而是3倍,也就是GRU的參數要比LSTM的參數量要少,但是性能差不多。

Pytorch

在Pytorch中使用nn.LSTM()可調用,參數和RNN的參數相同。具體介紹LSTM的輸入和輸出:

輸入: input, (h_0, c_0)

input:輸入數據with維度(seq_len,batch,input_size)

h_0:維度為(num_layers*num_directions,batch,hidden_size),在batch中的

初始的隱藏狀態(tài).

c_0:初始的單元狀態(tài),維度與h_0相同

輸出:output, (h_n, c_n)

output:維度為(seq_len, batch, num_directions * hidden_size)。

h_n:最后時刻的輸出隱藏狀態(tài),維度為 (num_layers * num_directions, batch, hidden_size)

c_n:最后時刻的輸出單元狀態(tài),維度與h_n相同。

LSTM的變量:

以MNIST分類為例實現(xiàn)LSTM分類

MNIST圖片大小為28×28,可以將每張圖片看做是長為28的序列,序列中每個元素的特征維度為28。將最后輸出的隱藏狀態(tài) 作為抽象的隱藏特征輸入到全連接層進行分類。最后輸出的

導入頭文件:

import torch
import torch.nn as nn
import torch.optim as optim
import torchvision
from torchvision import transforms
class Rnn(nn.Module):
  def __init__(self, in_dim, hidden_dim, n_layer, n_classes):
    super(Rnn, self).__init__()
    self.n_layer = n_layer
    self.hidden_dim = hidden_dim
    self.lstm = nn.LSTM(in_dim, hidden_dim, n_layer, batch_first=True)
    self.classifier = nn.Linear(hidden_dim, n_classes)

  def forward(self, x):
    out, (h_n, c_n) = self.lstm(x)
    # 此時可以從out中獲得最終輸出的狀態(tài)h
    # x = out[:, -1, :]
    x = h_n[-1, :, :]
    x = self.classifier(x)
    return x

訓練和測試代碼:

transform = transforms.Compose([
  transforms.ToTensor(),
  transforms.Normalize([0.5], [0.5]),
])

trainset = torchvision.datasets.MNIST(root='./data', train=True, download=True, transform=transform)
trainloader = torch.utils.data.DataLoader(trainset, batch_size=128, shuffle=True)

testset = torchvision.datasets.MNIST(root='./data', train=False, download=True, transform=transform)
testloader = torch.utils.data.DataLoader(testset, batch_size=100, shuffle=False)

net = Rnn(28, 10, 2, 10)

net = net.to('cpu')
criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(net.parameters(), lr=0.1, momentum=0.9)

# Training
def train(epoch):
  print('\nEpoch: %d' % epoch)
  net.train()
  train_loss = 0
  correct = 0
  total = 0
  for batch_idx, (inputs, targets) in enumerate(trainloader):
    inputs, targets = inputs.to('cpu'), targets.to('cpu')
    optimizer.zero_grad()
    outputs = net(torch.squeeze(inputs, 1))
    loss = criterion(outputs, targets)
    loss.backward()
    optimizer.step()

    train_loss += loss.item()
    _, predicted = outputs.max(1)
    total += targets.size(0)
    correct += predicted.eq(targets).sum().item()

    print(batch_idx, len(trainloader), 'Loss: %.3f | Acc: %.3f%% (%d/%d)'
      % (train_loss/(batch_idx+1), 100.*correct/total, correct, total))

def test(epoch):
  global best_acc
  net.eval()
  test_loss = 0
  correct = 0
  total = 0
  with torch.no_grad():
    for batch_idx, (inputs, targets) in enumerate(testloader):
      inputs, targets = inputs.to('cpu'), targets.to('cpu')
      outputs = net(torch.squeeze(inputs, 1))
      loss = criterion(outputs, targets)

      test_loss += loss.item()
      _, predicted = outputs.max(1)
      total += targets.size(0)
      correct += predicted.eq(targets).sum().item()

      print(batch_idx, len(testloader), 'Loss: %.3f | Acc: %.3f%% (%d/%d)'
        % (test_loss/(batch_idx+1), 100.*correct/total, correct, total))




for epoch in range(200):
  train(epoch)
  test(epoch)

以上這篇Pytorch實現(xiàn)LSTM和GRU示例就是小編分享給大家的全部內容了,希望能給大家一個參考,也希望大家多多支持腳本之家。

相關文章

  • python語法教程之def()函數定義及用法

    python語法教程之def()函數定義及用法

    函數是組織好的,可重復使用的,用來實現(xiàn)單一,或相關聯(lián)功能的代碼段,下面這篇文章主要給大家介紹了關于python語法教程之def()函數定義及用法的相關資料,文中通過實例代碼介紹的非常詳細,需要的朋友可以參考下
    2023-01-01
  • python+pygame實現(xiàn)坦克大戰(zhàn)

    python+pygame實現(xiàn)坦克大戰(zhàn)

    這篇文章主要為大家詳細介紹了python+pygame實現(xiàn)坦克大戰(zhàn),具有一定的參考價值,感興趣的小伙伴們可以參考一下
    2019-09-09
  • 用python 批量操作redis數據庫

    用python 批量操作redis數據庫

    這篇文章主要介紹了如何用python 批量操作redis數據庫,幫助大家更好的理解和學習使用python,感興趣的朋友可以了解下
    2021-03-03
  • Django路由層URLconf作用及原理解析

    Django路由層URLconf作用及原理解析

    這篇文章主要介紹了Django路由層URLconf作用及原理解析,文中通過示例代碼介紹的非常詳細,對大家的學習或者工作具有一定的參考學習價值,需要的朋友可以參考下
    2020-09-09
  • Python爬蟲庫BeautifulSoup的介紹與簡單使用實例

    Python爬蟲庫BeautifulSoup的介紹與簡單使用實例

    BeautifulSoup是一個可以從HTML或XML文件中提取數據的Python庫,本文為大家介紹下Python爬蟲庫BeautifulSoup的介紹與簡單使用實例其中包括了,BeautifulSoup解析HTML,BeautifulSoup獲取內容,BeautifulSoup節(jié)點操作,BeautifulSoup獲取CSS屬性等實例
    2020-01-01
  • Python通過rembg實現(xiàn)圖片背景去除功能

    Python通過rembg實現(xiàn)圖片背景去除功能

    在圖像處理領域,背景移除是一個常見且重要的任務,Python中的rembg庫就是一個強大的工具,它基于深度學習技術,能夠準確、快速地移除圖像背景,本文將結合多個實際案例,詳細介紹rembg庫的安裝、基本用法、高級功能以及在實際項目中的應用,需要的朋友可以參考下
    2024-09-09
  • Python運行報錯UnicodeDecodeError的解決方法

    Python運行報錯UnicodeDecodeError的解決方法

    本文給大家分享的是在Python項目中經常遇到的關于編碼問題的一個小bug的解決方法以及分析方法,有相同遭遇的小伙伴可以來參考下
    2016-06-06
  • Python3.7安裝pyaudio教程解析

    Python3.7安裝pyaudio教程解析

    這篇文章主要介紹了Python3.7安裝pyaudio教程解析,文中通過示例代碼介紹的非常詳細,對大家的學習或者工作具有一定的參考學習價值,需要的朋友可以參考下
    2020-07-07
  • Python Numpy運行報錯:IndexError: too many indices for array的分析及解決

    Python Numpy運行報錯:IndexError: too many in

    在使用Numpy進行數組操作時,經常會遇到各種錯誤,其中,IndexError: too many indices for array是一種常見的錯誤,它通常發(fā)生在嘗試使用一個過多維度的索引來訪問一個較低維度的數組時,本文介紹了Python Numpy報錯的解決辦法,需要的朋友可以參考下
    2024-07-07
  • Python爬蟲之requests基礎用法詳解

    Python爬蟲之requests基礎用法詳解

    這篇文章主要介紹了Python爬蟲之requests基礎用法詳解,雖然Python的標準庫中urllib模塊已經包含了平常我們使用的大多數功能,但是它的API使用起來讓人感覺不太友好,而requests庫使用更簡潔方便,需要的朋友可以參考下
    2023-10-10

最新評論

舞钢市| 文成县| 吉木乃县| 鹤岗市| 年辖:市辖区| 惠州市| 华容县| 海南省| 黄浦区| 临泉县| 理塘县| 凌源市| 营山县| 伊金霍洛旗| 柏乡县| 阜城县| 玉林市| 星子县| 无棣县| 遵义市| 峨眉山市| 怀仁县| 象山县| 武强县| 永吉县| 商水县| 新余市| 宁远县| 周口市| 奉贤区| 那坡县| 丰镇市| 十堰市| 报价| 张家港市| 喀什市| 青田县| 温泉县| 丹凤县| 泸州市| 宝应县|