基于Pytorch的神經(jīng)網(wǎng)絡(luò)之Regression的實(shí)現(xiàn)
1.引言
我們之前已經(jīng)介紹了神經(jīng)網(wǎng)絡(luò)的基本知識(shí),神經(jīng)網(wǎng)絡(luò)的主要作用就是預(yù)測(cè)與分類(lèi),現(xiàn)在讓我們來(lái)搭建第一個(gè)用于擬合回歸的神經(jīng)網(wǎng)絡(luò)吧。
2.神經(jīng)網(wǎng)絡(luò)搭建
2.1 準(zhǔn)備工作
要搭建擬合神經(jīng)網(wǎng)絡(luò)并繪圖我們需要使用python的幾個(gè)庫(kù)。
import torch import torch.nn.functional as F import matplotlib.pyplot as plt x = torch.unsqueeze(torch.linspace(-5, 5, 100), dim=1) y = x.pow(3) + 0.2 * torch.rand(x.size())
既然是擬合,我們當(dāng)然需要一些數(shù)據(jù)啦,我選取了在區(qū)間 內(nèi)的100個(gè)等間距點(diǎn),并將它們排列成三次函數(shù)的圖像。
2.2 搭建網(wǎng)絡(luò)
我們定義一個(gè)類(lèi),繼承了封裝在torch中的一個(gè)模塊,我們先分別確定輸入層、隱藏層、輸出層的神經(jīng)元數(shù)目,繼承父類(lèi)后再使用torch中的.nn.Linear()函數(shù)進(jìn)行輸入層到隱藏層的線(xiàn)性變換,隱藏層也進(jìn)行線(xiàn)性變換后傳入輸出層predict,接下來(lái)定義前向傳播的函數(shù)forward(),使用relu()作為激活函數(shù),最后輸出predict()結(jié)果即可。
class Net(torch.nn.Module):
def __init__(self, n_feature, n_hidden, n_output):
super(Net, self).__init__()
self.hidden = torch.nn.Linear(n_feature, n_hidden)
self.predict = torch.nn.Linear(n_hidden, n_output)
def forward(self, x):
x = F.relu(self.hidden(x))
return self.predict(x)
net = Net(1, 20, 1)
print(net)
optimizer = torch.optim.Adam(net.parameters(), lr=0.2)
loss_func = torch.nn.MSELoss()網(wǎng)絡(luò)的框架搭建完了,然后我們傳入三層對(duì)應(yīng)的神經(jīng)元數(shù)目再定義優(yōu)化器,這里我選取了Adam而隨機(jī)梯度下降(SGD),因?yàn)樗荢GD的優(yōu)化版本,效果在大部分情況下比SGD好,我們要傳入這個(gè)神經(jīng)網(wǎng)絡(luò)的參數(shù)(parameters),并定義學(xué)習(xí)率(learning rate),學(xué)習(xí)率通常選取小于1的數(shù),需要憑借經(jīng)驗(yàn)并不斷調(diào)試。最后我們選取均方差法(MSE)來(lái)計(jì)算損失(loss)。
2.3 訓(xùn)練網(wǎng)絡(luò)
接下來(lái)我們要對(duì)我們搭建好的神經(jīng)網(wǎng)絡(luò)進(jìn)行訓(xùn)練,我訓(xùn)練了2000輪(epoch),先更新結(jié)果prediction再計(jì)算損失,接著清零梯度,然后根據(jù)loss反向傳播(backward),最后進(jìn)行優(yōu)化,找出最優(yōu)的擬合曲線(xiàn)。
for t in range(2000):
prediction = net(x)
loss = loss_func(prediction, y)
optimizer.zero_grad()
loss.backward()
optimizer.step()3.效果
使用如下繪圖的代碼展示效果。
for t in range(2000):
prediction = net(x)
loss = loss_func(prediction, y)
optimizer.zero_grad()
loss.backward()
optimizer.step()
if t % 5 == 0:
plt.cla()
plt.scatter(x.data.numpy(), y.data.numpy(), s=10)
plt.plot(x.data.numpy(), prediction.data.numpy(), 'r-', lw=2)
plt.text(2, -100, 'Loss=%.4f' % loss.data.numpy(), fontdict={'size': 10, 'color': 'red'})
plt.pause(0.1)
plt.ioff()
plt.show()

最后的結(jié)果:

4. 完整代碼
import torch
import torch.nn.functional as F
import matplotlib.pyplot as plt
x = torch.unsqueeze(torch.linspace(-5, 5, 100), dim=1)
y = x.pow(3) + 0.2 * torch.rand(x.size())
class Net(torch.nn.Module):
def __init__(self, n_feature, n_hidden, n_output):
super(Net, self).__init__()
self.hidden = torch.nn.Linear(n_feature, n_hidden)
self.predict = torch.nn.Linear(n_hidden, n_output)
def forward(self, x):
x = F.relu(self.hidden(x))
return self.predict(x)
net = Net(1, 20, 1)
print(net)
optimizer = torch.optim.Adam(net.parameters(), lr=0.2)
loss_func = torch.nn.MSELoss()
plt.ion()
for t in range(2000):
prediction = net(x)
loss = loss_func(prediction, y)
optimizer.zero_grad()
loss.backward()
optimizer.step()
if t % 5 == 0:
plt.cla()
plt.scatter(x.data.numpy(), y.data.numpy(), s=10)
plt.plot(x.data.numpy(), prediction.data.numpy(), 'r-', lw=2)
plt.text(2, -100, 'Loss=%.4f' % loss.data.numpy(), fontdict={'size': 10, 'color': 'red'})
plt.pause(0.1)
plt.ioff()
plt.show()到此這篇關(guān)于基于Pytorch的神經(jīng)網(wǎng)絡(luò)之Regression的實(shí)現(xiàn)的文章就介紹到這了,更多相關(guān) Pytorch Regression內(nèi)容請(qǐng)搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家!
相關(guān)文章
Python通過(guò)cron或schedule實(shí)現(xiàn)爬蟲(chóng)的自動(dòng)定時(shí)運(yùn)行
自動(dòng)定時(shí)運(yùn)行爬蟲(chóng)是很多數(shù)據(jù)采集項(xiàng)目的基本需求,通過(guò) Python 實(shí)現(xiàn)定時(shí)任務(wù),可以保證數(shù)據(jù)采集的高效和持續(xù)性,本文將帶大家了解如何在 Python 中使用 cron 和 schedule 來(lái)實(shí)現(xiàn)爬蟲(chóng)的自動(dòng)定時(shí)運(yùn)行,需要的朋友可以參考下2024-12-12
Pandas提高數(shù)據(jù)分析效率的13個(gè)技巧匯總
這篇文章主要是為大家歸納整理了13個(gè)工作中常用到的pandas使用技巧,方便更高效地實(shí)現(xiàn)數(shù)據(jù)分析,感興趣的小伙伴可以跟隨小編一起學(xué)習(xí)一下2022-05-05
Numpy實(shí)現(xiàn)矩陣運(yùn)算及線(xiàn)性代數(shù)應(yīng)用
這篇文章主要介紹了Numpy實(shí)現(xiàn)矩陣運(yùn)算及線(xiàn)性代數(shù)應(yīng)用,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來(lái)一起學(xué)習(xí)學(xué)習(xí)吧2021-03-03
python實(shí)現(xiàn)批量監(jiān)控網(wǎng)站
本文給大家分享的是一個(gè)非常實(shí)用的,python實(shí)現(xiàn)多網(wǎng)站的可用性監(jiān)控的腳本,并附上核心點(diǎn)解釋?zhuān)邢嗤枨蟮男』锇榭梢詤⒖枷?/div> 2016-09-09
Pycharm連接遠(yuǎn)程服務(wù)器并實(shí)現(xiàn)遠(yuǎn)程調(diào)試的實(shí)現(xiàn)
這篇文章主要介紹了Pycharm連接遠(yuǎn)程服務(wù)器并實(shí)現(xiàn)遠(yuǎn)程調(diào)試的實(shí)現(xiàn),文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來(lái)一起學(xué)習(xí)學(xué)習(xí)吧2019-08-08
python實(shí)現(xiàn)web方式logview的方法
這篇文章主要介紹了python實(shí)現(xiàn)web方式logview的方法,涉及Python基于web模塊操作Linux命令的技巧,具有一定參考借鑒價(jià)值,需要的朋友可以參考下2015-08-08最新評(píng)論

