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

PyTorch搭建一維線性回歸模型(二)

 更新時間:2021年04月09日 14:27:04   作者:Liam Coder  
這篇文章主要為大家詳細介紹了PyTorch搭建一維線性回歸模型,文中示例代碼介紹的非常詳細,具有一定的參考價值,感興趣的小伙伴們可以參考一下

PyTorch基礎入門二:PyTorch搭建一維線性回歸模型

1)一維線性回歸模型的理論基礎

給定數據集,線性回歸希望能夠優(yōu)化出一個好的函數,使得能夠和盡可能接近。

如何才能學習到參數呢?很簡單,只需要確定如何衡量之間的差別,我們一般通過損失函數(Loss Funciton)來衡量:。取平方是因為距離有正有負,我們于是將它們變?yōu)槿钦摹_@就是著名的均方誤差。我們要做的事情就是希望能夠找到,使得:

均方差誤差非常直觀,也有著很好的幾何意義,對應了常用的歐式距離?,F在要求解這個連續(xù)函數的最小值,我們很自然想到的方法就是求它的偏導數,讓它的偏導數等于0來估計它的參數,即:

求解以上兩式,我們就可以得到最優(yōu)解。

2)代碼實現

首先,我們需要“制造”出一些數據集:

import torch
import matplotlib.pyplot as plt
 
 
x = torch.unsqueeze(torch.linspace(-1, 1, 100), dim=1)
y = 3*x + 10 + torch.rand(x.size())
# 上面這行代碼是制造出接近y=3x+10的數據集,后面加上torch.rand()函數制造噪音
 
# 畫圖
plt.scatter(x.data.numpy(), y.data.numpy())
plt.show()

我們想要擬合的一維回歸模型是。上面制造的數據集也是比較接近這個模型的,但是為了達到學習效果,人為地加上了torch.rand()值增加一些干擾。

上面人為制造出來的數據集的分布如下:

有了數據,我們就要開始定義我們的模型,這里定義的是一個輸入層和輸出層都只有一維的模型,并且使用了“先判斷后使用”的基本結構來合理使用GPU加速。

class LinearRegression(nn.Module):
  def __init__(self):
    super(LinearRegression, self).__init__()
    self.linear = nn.Linear(1, 1) # 輸入和輸出的維度都是1
  def forward(self, x):
    out = self.linear(x)
    return out
 
if torch.cuda.is_available():
  model = LinearRegression().cuda()
else:
  model = LinearRegression()

然后我們定義出損失函數和優(yōu)化函數,這里使用均方誤差作為損失函數,使用梯度下降進行優(yōu)化:

criterion = nn.MSELoss()
optimizer = torch.optim.SGD(model.parameters(), lr=1e-2)

接下來,開始進行模型的訓練。

num_epochs = 1000
for epoch in range(num_epochs):
  if torch.cuda.is_available():
    inputs = Variable(x).cuda()
    target = Variable(y).cuda()
  else:
    inputs = Variable(x)
    target = Variable(y)
 
  # 向前傳播
  out = model(inputs)
  loss = criterion(out, target)
 
  # 向后傳播
  optimizer.zero_grad() # 注意每次迭代都需要清零
  loss.backward()
  optimizer.step()
 
  if (epoch+1) %20 == 0:
    print('Epoch[{}/{}], loss:{:.6f}'.format(epoch+1, num_epochs, loss.data[0]))

首先定義了迭代的次數,這里為1000次,先向前傳播計算出損失函數,然后向后傳播計算梯度,這里需要注意的是,每次計算梯度前都要記得將梯度歸零,不然梯度會累加到一起造成結果不收斂。為了便于看到結果,每隔一段時間輸出當前的迭代輪數和損失函數。

接下來,我們通過model.eval()函數將模型變?yōu)闇y試模式,然后將數據放入模型中進行預測。最后,通過畫圖工具matplotlib看一下我們擬合的結果,代碼如下:

model.eval()
if torch.cuda.is_available():
  predict = model(Variable(x).cuda())
  predict = predict.data.cpu().numpy()
else:
  predict = model(Variable(x))
  predict = predict.data.numpy()
plt.plot(x.numpy(), y.numpy(), 'ro', label='Original Data')
plt.plot(x.numpy(), predict, label='Fitting Line')
plt.show()

其擬合結果如下圖:

附上完整代碼:

# !/usr/bin/python
# coding: utf8
# @Time  : 2018-07-28 18:40
# @Author : Liam
# @Email  : luyu.real@qq.com
# @Software: PyCharm
#            .::::.
#           .::::::::.
#           :::::::::::
#         ..:::::::::::'
#        '::::::::::::'
#         .::::::::::
#      '::::::::::::::..
#         ..::::::::::::.
#        ``::::::::::::::::
#        ::::``:::::::::'    .:::.
#        ::::'  ':::::'    .::::::::.
#       .::::'   ::::   .:::::::'::::.
#      .:::'    ::::: .:::::::::' ':::::.
#      .::'    :::::.:::::::::'   ':::::.
#     .::'     ::::::::::::::'     ``::::.
#   ...:::      ::::::::::::'       ``::.
#   ```` ':.     ':::::::::'         ::::..
#            '.:::::'          ':'````..
#           美女保佑 永無BUG
 
import torch
from torch.autograd import Variable
import numpy as np
import random
import matplotlib.pyplot as plt
from torch import nn
 
 
x = torch.unsqueeze(torch.linspace(-1, 1, 100), dim=1)
y = 3*x + 10 + torch.rand(x.size())
# 上面這行代碼是制造出接近y=3x+10的數據集,后面加上torch.rand()函數制造噪音
 
# 畫圖
# plt.scatter(x.data.numpy(), y.data.numpy())
# plt.show()
class LinearRegression(nn.Module):
  def __init__(self):
    super(LinearRegression, self).__init__()
    self.linear = nn.Linear(1, 1) # 輸入和輸出的維度都是1
  def forward(self, x):
    out = self.linear(x)
    return out
 
if torch.cuda.is_available():
  model = LinearRegression().cuda()
else:
  model = LinearRegression()
 
criterion = nn.MSELoss()
optimizer = torch.optim.SGD(model.parameters(), lr=1e-2)
 
num_epochs = 1000
for epoch in range(num_epochs):
  if torch.cuda.is_available():
    inputs = Variable(x).cuda()
    target = Variable(y).cuda()
  else:
    inputs = Variable(x)
    target = Variable(y)
 
  # 向前傳播
  out = model(inputs)
  loss = criterion(out, target)
 
  # 向后傳播
  optimizer.zero_grad() # 注意每次迭代都需要清零
  loss.backward()
  optimizer.step()
 
  if (epoch+1) %20 == 0:
    print('Epoch[{}/{}], loss:{:.6f}'.format(epoch+1, num_epochs, loss.data[0]))
model.eval()
if torch.cuda.is_available():
  predict = model(Variable(x).cuda())
  predict = predict.data.cpu().numpy()
else:
  predict = model(Variable(x))
  predict = predict.data.numpy()
plt.plot(x.numpy(), y.numpy(), 'ro', label='Original Data')
plt.plot(x.numpy(), predict, label='Fitting Line')
plt.show()

以上就是本文的全部內容,希望對大家的學習有所幫助,也希望大家多多支持腳本之家。

相關文章

  • DataFrame 將某列數據轉為數組的方法

    DataFrame 將某列數據轉為數組的方法

    下面小編就為大家分享一篇DataFrame 將某列數據轉為數組的方法,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2018-04-04
  • Python實現注冊登錄系統(tǒng)

    Python實現注冊登錄系統(tǒng)

    這篇文章主要為大家詳細介紹了適合初學者學習的Python3銀行賬戶登錄系統(tǒng),具有一定的參考價值,感興趣的小伙伴們可以參考一下
    2017-08-08
  • 基于asyncio 異步協程框架實現收集B站直播彈幕

    基于asyncio 異步協程框架實現收集B站直播彈幕

    本文給大家分享的是基于asyncio 異步協程框架實現收集B站直播彈幕收集系統(tǒng)的簡單設計,并附上源碼,有需要的小伙伴可以參考下
    2016-09-09
  • Python中檢查字符串是否僅包含字母的方法詳解

    Python中檢查字符串是否僅包含字母的方法詳解

    這篇文章主要為大家詳細介紹了Python中的多種方法來檢查字符串是否只由字母組成,以及它們的應用場景和優(yōu)劣,感興趣的小伙伴可以跟隨小編一起學習一下
    2023-11-11
  • python輸入、數據類型轉換及運算符方式

    python輸入、數據類型轉換及運算符方式

    這篇文章主要介紹了python輸入、數據類型轉換及運算符方式,具有很好的參考價值,希望對大家有所幫助。如有錯誤或未考慮完全的地方,望不吝賜教
    2022-07-07
  • Python實現Const詳解

    Python實現Const詳解

    這篇文章主要介紹了Python實現Const的方法的相關資料,需要的朋友可以參考下
    2015-01-01
  • Python斷言assert的用法代碼解析

    Python斷言assert的用法代碼解析

    這篇文章主要介紹了Python斷言assert的用法代碼解析,分享了相關代碼示例,小編覺得還是挺不錯的,具有一定借鑒價值,需要的朋友可以參考下
    2018-02-02
  • 淺談Python程序的錯誤:變量未定義

    淺談Python程序的錯誤:變量未定義

    這篇文章主要介紹了淺談Python程序的錯誤:變量未定義,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2020-06-06
  • 一文詳細NumPy中np.zeros的使用

    一文詳細NumPy中np.zeros的使用

    np.zeros是NumPy庫中一個非常實用的函數,用于快速創(chuàng)建指定形狀和大小的全零數組,本文主要介紹了NumPy中np.zeros的使用,感興趣的可以了解一下
    2024-03-03
  • 解決Windows下PowerShell無法進入Python虛擬環(huán)境問題

    解決Windows下PowerShell無法進入Python虛擬環(huán)境問題

    這篇文章主要介紹了解決Windows下PowerShell無法進入Python虛擬環(huán)境問題,具有很好的參考價值,希望對大家有所幫助,如有錯誤或未考慮完全的地方,望不吝賜教
    2024-02-02

最新評論

来凤县| 额济纳旗| 应用必备| 平江县| 临夏市| 汝南县| 区。| 怀仁县| 宽甸| 如东县| 延吉市| 吉水县| 广元市| 项城市| 新沂市| 禹州市| 广东省| 左贡县| 辽宁省| 崇文区| 额敏县| 江孜县| 齐齐哈尔市| 文昌市| 政和县| 剑川县| 云南省| 密山市| 佳木斯市| 恭城| 邹平县| 牟定县| 澄城县| 阿拉尔市| 额敏县| 阿尔山市| 巩义市| 当涂县| 静安区| 广宗县| 普陀区|