pytorch的Backward過程用時太長問題及解決
pytorch Backward過程用時太長
問題描述
使用pytorch對網(wǎng)絡(luò)進行訓(xùn)練的時候遇到一個問題,forward階段很快(只需要幾毫秒),backward階段卻用時很長(需要十多秒)。
導(dǎo)致這個問題的原因很容易被大家忽視,而且網(wǎng)上基本上沒有直接的解決方案,經(jīng)過一天的折騰,總算把導(dǎo)致這個問題的原因搞清楚了。
解決方案
導(dǎo)致這個問題的原因在于訓(xùn)練數(shù)據(jù)的淺拷貝,由于backward過程中的梯度是和模型推理過程中的張量相關(guān)的,如果這些張量在被模型使用之前沒有被深拷貝,意味著backward過程的會重復(fù)從這些張量的原始內(nèi)存地址中取值,這個過程非常耗時。所以為了避免這個問題,需要養(yǎng)成一個好習(xí)慣,就是將張量數(shù)據(jù)輸入模型之前進行深拷貝
pytorch的深拷貝方式如下:
tensor_a = tensor_b.clone().detach()
Pytorch backward()簡單理解
backward()是反向傳播求梯度,具體實現(xiàn)過程如下
import torch x=torch.tensor([1,2,3],requires_grad=True,dtype=torch.double) y=x**2 z=y.mean() z.backward() print(x.grad)
結(jié)果
tensor([0.6667, 1.3333, 2.0000], dtype=torch.float64)
有幾個重要的點
1.必須要加上requires_grad=True才能求
2. 一般來說,需要標(biāo)量才能求梯度。
3.具體過程如下:

z是一個標(biāo)量(1*1矩陣)分別對x1,x2,x3求偏導(dǎo), 再代入x1,x2,x3的數(shù)值,就是如上程序輸出的結(jié)果
總結(jié)
以上為個人經(jīng)驗,希望能給大家一個參考,也希望大家多多支持腳本之家。
相關(guān)文章
Pycharm內(nèi)置終端及遠程SSH工具的使用教程圖文詳解
這篇文章主要介紹了Pycharm內(nèi)置終端及遠程SSH工具的使用教程,本文通過圖文并茂的形式給大家介紹的非常詳細,對大家的學(xué)習(xí)或工作具有一定的參考借鑒價值,需要的朋友可以參考下2020-03-03
Python實現(xiàn)刪除當(dāng)前目錄下除當(dāng)前腳本以外的文件和文件夾實例
這篇文章主要介紹了Python實現(xiàn)刪除當(dāng)前目錄下除當(dāng)前腳本以外的文件和文件夾的方法,涉及Python針對目錄及文件的刪除技巧,具有一定參考借鑒價值,需要的朋友可以參考下2015-07-07

