pytorch-神經(jīng)網(wǎng)絡(luò)擬合曲線(xiàn)實(shí)例
代碼已經(jīng)調(diào)通,跑出來(lái)的效果如下:

# coding=gbk
import torch
import matplotlib.pyplot as plt
from torch.autograd import Variable
import torch.nn.functional as F
'''
Pytorch是一個(gè)擁有強(qiáng)力GPU加速的張量和動(dòng)態(tài)構(gòu)建網(wǎng)絡(luò)的庫(kù),其主要構(gòu)建是張量,所以可以把PyTorch當(dāng)做Numpy
來(lái)用,Pytorch的很多操作好比Numpy都是類(lèi)似的,但是其能夠在GPU上運(yùn)行,所以有著比Numpy快很多倍的速度。
訓(xùn)練完了,發(fā)現(xiàn)隱層越大,擬合的速度越是快,擬合的效果越是好
'''
def train():
print('------ 構(gòu)建數(shù)據(jù)集 ------')
# torch.linspace是為了生成連續(xù)間斷的數(shù)據(jù),第一個(gè)參數(shù)表示起點(diǎn),第二個(gè)參數(shù)表示終點(diǎn),第三個(gè)參數(shù)表示將這個(gè)區(qū)間分成平均幾份,即生成幾個(gè)數(shù)據(jù)
x = torch.unsqueeze(torch.linspace(-1, 1, 100), dim=1)
#torch.rand返回的是[0,1]之間的均勻分布 這里是使用一個(gè)計(jì)算式子來(lái)構(gòu)造出一個(gè)關(guān)聯(lián)結(jié)果,當(dāng)然后期要學(xué)的也就是這個(gè)式子
y = x.pow(2) + 0.2 * torch.rand(x.size())
# Variable是將tensor封裝了下,用于自動(dòng)求導(dǎo)使用
x, y = Variable(x), Variable(y)
#繪圖展示
plt.scatter(x.data.numpy(), y.data.numpy())
#plt.show()
print('------ 搭建網(wǎng)絡(luò) ------')
#使用固定的方式繼承并重寫(xiě) init和forword兩個(gè)類(lèi)
class Net(torch.nn.Module):
def __init__(self,n_feature,n_hidden,n_output):
#初始網(wǎng)絡(luò)的內(nèi)部結(jié)構(gòu)
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):
#一次正向行走過(guò)程
x=F.relu(self.hidden(x))
x=self.predict(x)
return x
net=Net(n_feature=1,n_hidden=1000,n_output=1)
print('網(wǎng)絡(luò)結(jié)構(gòu)為:',net)
print('------ 啟動(dòng)訓(xùn)練 ------')
loss_func=F.mse_loss
optimizer=torch.optim.SGD(net.parameters(),lr=0.001)
#使用數(shù)據(jù) 進(jìn)行正向訓(xùn)練,并對(duì)Variable變量進(jìn)行反向梯度傳播 啟動(dòng)100次訓(xùn)練
for t in range(10000):
#使用全量數(shù)據(jù) 進(jìn)行正向行走
prediction=net(x)
loss=loss_func(prediction,y)
optimizer.zero_grad() #清除上一梯度
loss.backward() #反向傳播計(jì)算梯度
optimizer.step() #應(yīng)用梯度
#間隔一段,對(duì)訓(xùn)練過(guò)程進(jìn)行可視化展示
if t%5==0:
plt.cla()
plt.scatter(x.data.numpy(),y.data.numpy()) #繪制真是曲線(xiàn)
plt.plot(x.data.numpy(),prediction.data.numpy(),'r-',lw=5)
plt.text(0.5,0,'Loss='+str(loss.data[0]),fontdict={'size':20,'color':'red'})
plt.pause(0.1)
plt.ioff()
plt.show()
print('------ 預(yù)測(cè)和可視化 ------')
if __name__=='__main__':
train()
以上這篇pytorch-神經(jīng)網(wǎng)絡(luò)擬合曲線(xiàn)實(shí)例就是小編分享給大家的全部?jī)?nèi)容了,希望能給大家一個(gè)參考,也希望大家多多支持腳本之家。
相關(guān)文章
Python中的反射知識(shí)點(diǎn)總結(jié)
在本篇文章里小編給大家整理了一篇關(guān)于Python中的反射知識(shí)點(diǎn)總結(jié)內(nèi)容,有需要的朋友們可以跟著學(xué)習(xí)參考下。2021-11-11
圖文詳解牛頓迭代算法原理及Python實(shí)現(xiàn)
牛頓迭代法又稱(chēng)為牛頓-拉夫遜(拉弗森)方法,它是牛頓在17世紀(jì)提出的一種在實(shí)數(shù)域和復(fù)數(shù)域上近似求解方程的方法。本文將利用圖文詳解牛頓迭代算法原理及實(shí)現(xiàn),需要的可以參考一下2022-08-08
SpringBoot中的@MessageMapping注解詳解
這篇文章主要介紹了SpringBoot中的@MessageMapping注解詳解,Spring Boot 提供了對(duì) WebSocket 的支持,其中 @MessageMapping 注解是一個(gè)常用的注解,它可以將一個(gè) Java 方法標(biāo)記為 WebSocket 的消息處理器,需要的朋友可以參考下2023-08-08
使用Python中的線(xiàn)程進(jìn)行網(wǎng)絡(luò)編程的入門(mén)教程
這篇文章主要介紹了使用Python中的線(xiàn)程進(jìn)行網(wǎng)絡(luò)編程的入門(mén)教程,本文來(lái)自于IBM官方網(wǎng)站技術(shù)文檔,需要的朋友可以參考下2015-04-04
Python中關(guān)于面向?qū)ο蟾拍畹脑敿?xì)講解
要了解面向?qū)ο笪覀兛隙ㄐ枰戎缹?duì)象到底是什么玩意兒。關(guān)于對(duì)象的理解很簡(jiǎn)單,在我們的身邊,每一種事物的存在都是一種對(duì)象??偨Y(jié)為一句話(huà)也就是:對(duì)象就是事物存在的實(shí)體2021-10-10
Python 幾行代碼即可實(shí)現(xiàn)人臉識(shí)別
Python中實(shí)現(xiàn)人臉識(shí)別功能有多種方法,依賴(lài)于python膠水語(yǔ)言的特性,我們通過(guò)調(diào)用包可以快速準(zhǔn)確的達(dá)成這一目的,本文給大家分享使用Python實(shí)現(xiàn)簡(jiǎn)單的人臉識(shí)別功能的操作步驟,感興趣的朋友一起看看吧2022-02-02
python使用ctypes調(diào)用第三方庫(kù)時(shí)出現(xiàn)undefined?symbol錯(cuò)誤詳解
python中時(shí)間的庫(kù)有time和datetime,pandas也有提供相應(yīng)的時(shí)間處理函數(shù),下面這篇文章主要給大家介紹了關(guān)于python使用ctypes調(diào)用第三方庫(kù)時(shí)出現(xiàn)undefined?symbol錯(cuò)誤的相關(guān)資料,需要的朋友可以參考下2023-02-02

