PyTorch中model.zero_grad()和optimizer.zero_grad()用法
廢話不多說(shuō),直接上代碼吧~
model.zero_grad()
optimizer.zero_grad()
首先,這兩種方式都是把模型中參數(shù)的梯度設(shè)為0
當(dāng)optimizer = optim.Optimizer(net.parameters())時(shí),二者等效,其中Optimizer可以是Adam、SGD等優(yōu)化器
def zero_grad(self): """Sets gradients of all model parameters to zero.""" for p in self.parameters(): if p.grad is not None: p.grad.data.zero_()
補(bǔ)充知識(shí):Pytorch中的optimizer.zero_grad和loss和net.backward和optimizer.step的理解
引言
一般訓(xùn)練神經(jīng)網(wǎng)絡(luò),總是逃不開(kāi)optimizer.zero_grad之后是loss(后面有的時(shí)候還會(huì)寫(xiě)forward,看你網(wǎng)絡(luò)怎么寫(xiě)了)之后是是net.backward之后是optimizer.step的這個(gè)過(guò)程。
real_a, real_b = batch[0].to(device), batch[1].to(device) fake_b = net_g(real_a) optimizer_d.zero_grad() # 判別器對(duì)虛假數(shù)據(jù)進(jìn)行訓(xùn)練 fake_ab = torch.cat((real_a, fake_b), 1) pred_fake = net_d.forward(fake_ab.detach()) loss_d_fake = criterionGAN(pred_fake, False) # 判別器對(duì)真實(shí)數(shù)據(jù)進(jìn)行訓(xùn)練 real_ab = torch.cat((real_a, real_b), 1) pred_real = net_d.forward(real_ab) loss_d_real = criterionGAN(pred_real, True) # 判別器損失 loss_d = (loss_d_fake + loss_d_real) * 0.5 loss_d.backward() optimizer_d.step()
上面這是一段cGAN的判別器訓(xùn)練過(guò)程。標(biāo)題中所涉及到的這些方法,其實(shí)整個(gè)神經(jīng)網(wǎng)絡(luò)的參數(shù)更新過(guò)程(特別是反向傳播),具體是怎么操作的,我們一起來(lái)探討一下。
參數(shù)更新和反向傳播

上圖為一個(gè)簡(jiǎn)單的梯度下降示意圖。比如以SGD為例,是算一個(gè)batch計(jì)算一次梯度,然后進(jìn)行一次梯度更新。這里梯度值就是對(duì)應(yīng)偏導(dǎo)數(shù)的計(jì)算結(jié)果。顯然,我們進(jìn)行下一次batch梯度計(jì)算的時(shí)候,前一個(gè)batch的梯度計(jì)算結(jié)果,沒(méi)有保留的必要了。所以在下一次梯度更新的時(shí)候,先使用optimizer.zero_grad把梯度信息設(shè)置為0。
我們使用loss來(lái)定義損失函數(shù),是要確定優(yōu)化的目標(biāo)是什么,然后以目標(biāo)為頭,才可以進(jìn)行鏈?zhǔn)椒▌t和反向傳播。
調(diào)用loss.backward方法時(shí)候,Pytorch的autograd就會(huì)自動(dòng)沿著計(jì)算圖反向傳播,計(jì)算每一個(gè)葉子節(jié)點(diǎn)的梯度(如果某一個(gè)變量是由用戶創(chuàng)建的,則它為葉子節(jié)點(diǎn))。使用該方法,可以計(jì)算鏈?zhǔn)椒▌t求導(dǎo)之后計(jì)算的結(jié)果值。
optimizer.step用來(lái)更新參數(shù),就是圖片中下半部分的w和b的參數(shù)更新操作。
以上這篇PyTorch中model.zero_grad()和optimizer.zero_grad()用法就是小編分享給大家的全部?jī)?nèi)容了,希望能給大家一個(gè)參考,也希望大家多多支持腳本之家。
相關(guān)文章
用Python寫(xiě)飛機(jī)大戰(zhàn)游戲之pygame入門(mén)(4):獲取鼠標(biāo)的位置及運(yùn)動(dòng)
這篇文章主要介紹了用Python寫(xiě)飛機(jī)大戰(zhàn)游戲之pygame入門(mén)(4):獲取鼠標(biāo)的位置及運(yùn)動(dòng),需要的朋友可以參考下2015-11-11
python密碼學(xué)文件解密實(shí)現(xiàn)教程
這篇文章主要為大家介紹了python密碼學(xué)文件解密實(shí)現(xiàn)教程,有需要的朋友可以借鑒參考下,希望能夠有所幫助,祝大家多多進(jìn)步,早日升職加薪2022-05-05
django注冊(cè)用郵箱發(fā)送驗(yàn)證碼的實(shí)現(xiàn)
這篇文章主要介紹了django注冊(cè)用郵箱發(fā)送驗(yàn)證碼的實(shí)現(xiàn),文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來(lái)一起學(xué)習(xí)學(xué)習(xí)吧2021-04-04
ubuntu?20.04系統(tǒng)下如何切換gcc/g++/python的版本
這篇文章主要給大家介紹了關(guān)于ubuntu?20.04系統(tǒng)下如何切換gcc/g++/python版本的相關(guān)資料,文中通過(guò)代碼介紹的非常詳細(xì),對(duì)大家學(xué)習(xí)或者使用ubuntu具有一定的參考借鑒價(jià)值,需要的朋友可以參考下2023-12-12
python常用數(shù)據(jù)結(jié)構(gòu)元組詳解
這篇文章主要介紹了python常用數(shù)據(jù)結(jié)構(gòu)元組詳解,文章圍繞主題展開(kāi)詳細(xì)的內(nèi)容介紹,具有一定的參考價(jià)值,需要的小伙伴可以參考一下2022-08-08
Python 實(shí)例進(jìn)階之預(yù)測(cè)房?jī)r(jià)走勢(shì)
買(mǎi)房應(yīng)該是大多數(shù)都會(huì)要面臨的一個(gè)選擇,當(dāng)前經(jīng)濟(jì)和政策背景下,未來(lái)房?jī)r(jià)會(huì)漲還是跌?這是很多人都關(guān)心的一個(gè)話題。今天分享的這篇文章,以波士頓的房地產(chǎn)市場(chǎng)為例,根據(jù)低收入人群比例、老師學(xué)生數(shù)量等特征,利用 Python 進(jìn)行了預(yù)測(cè),給大家做一個(gè)參考2021-11-11
最好的Python DateTime 庫(kù)之 Pendulum 長(zhǎng)篇解析
datetime 模塊是 Python 中最重要的內(nèi)置模塊之一,它為實(shí)際編程問(wèn)題提供許多開(kāi)箱即用的解決方案,非常靈活和強(qiáng)大。例如,timedelta 是我最喜歡的工具之一2021-11-11

