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

基于TensorFlow中自定義梯度的2種方式

 更新時間:2020年02月04日 11:07:28   作者:FesianXu  
今天小編就為大家分享一篇基于TensorFlow中自定義梯度的2種方式,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧

前言

在深度學習中,有時候我們需要對某些節(jié)點的梯度進行一些定制,特別是該節(jié)點操作不可導(比如階梯除法如 ),如果實在需要對這個節(jié)點進行操作,而且希望其可以反向傳播,那么就需要對其進行自定義反向傳播時的梯度。在有些場景,如[2]中介紹到的梯度反轉(zhuǎn)(gradient inverse)中,就必須在某層節(jié)點對反向傳播的梯度進行反轉(zhuǎn),也就是需要更改正常的梯度傳播過程,如下圖的 所示。

在tensorflow中有若干可以實現(xiàn)定制梯度的方法,這里介紹兩種。

1. 重寫梯度法

重寫梯度法指的是通過tensorflow自帶的機制,將某個節(jié)點的梯度重寫(override),這種方法的適用性最廣。我們這里舉個例子[3].

符號函數(shù)的前向傳播采用的是階躍函數(shù)y=sign(x) y = \rm{sign}(x)y=sign(x),如下圖所示,我們知道階躍函數(shù)不是連續(xù)可導的,因此我們在反向傳播時,將其替代為一個可以連續(xù)求導的函數(shù)y=Htanh(x) y = \rm{Htanh(x)}y=Htanh(x),于是梯度就是大于1和小于-1時為0,在-1和1之間時是1。

使用重寫梯度的方法如下,主要是涉及到tf.RegisterGradient()和tf.get_default_graph().gradient_override_map(),前者注冊新的梯度,后者重寫圖中具有名字name='Sign'的操作節(jié)點的梯度,用在新注冊的QuantizeGrad替代。

#使用修飾器,建立梯度反向傳播函數(shù)。其中op.input包含輸入值、輸出值,grad包含上層傳來的梯度
@tf.RegisterGradient("QuantizeGrad")
def sign_grad(op, grad):
 input = op.inputs[0] # 取出當前的輸入
 cond = (input>=-1)&(input<=1) # 大于1或者小于-1的值的位置
 zeros = tf.zeros_like(grad) # 定義出0矩陣用于掩膜
 return tf.where(cond, grad, zeros) 
 # 將大于1或者小于-1的上一層的梯度置為0
 
#使用with上下文管理器覆蓋原始的sign梯度函數(shù)
def binary(input):
 x = input
 with tf.get_default_graph().gradient_override_map({"Sign":'QuantizeGrad'}):
 #重寫梯度
  x = tf.sign(x)
 return x
 
#使用
x = binary(x)

其中的def sign_grad(op, grad):是注冊新的梯度的套路,其中的op是當前操作的輸入值/張量等,而grad指的是從反向而言的上一層的梯度。

通常來說,在tensorflow中自定義梯度,函數(shù)tf.identity()是很重要的,其API手冊如下:

tf.identity(
 input,
 name=None
)

其會返回一個形狀和內(nèi)容都和輸入完全一樣的輸出,但是你可以自定義其反向傳播時的梯度,因此在梯度反轉(zhuǎn)等操作中特別有用。

這里再舉個反向梯度[2]的例子,也就是梯度為 而不是

import tensorflow as tf
x1 = tf.Variable(1)
x2 = tf.Variable(3)
x3 = tf.Variable(6)
@tf.RegisterGradient('CustomGrad')
def CustomGrad(op, grad):
#  tf.Print(grad)
 return -grad
 
g = tf.get_default_graph()
oo = x1+x2
with g.gradient_override_map({"Identity": "CustomGrad"}):
 output = tf.identity(oo)
grad_1 = tf.gradients(output, oo)
with tf.Session() as sess:
 sess.run(tf.global_variables_initializer())
 print(sess.run(grad_1))

因為-grad,所以這里的梯度輸出是[-1]而不是[1]。有一個我們需要注意的是,在自定義函數(shù)def CustomGrad()中,返回的值得是一個張量,而不能返回一個參數(shù),比如return 0,這樣會報錯,如:

AttributeError: 'int' object has no attribute 'name'

顯然,這是因為tensorflow的內(nèi)部操作需要取返回值的名字而int類型沒有名字。

PS:def CustomGrad()這個函數(shù)簽名是隨便你取的。

2. stop_gradient法

對于自定義梯度,還有一種比較簡潔的操作,就是利用tf.stop_gradient()函數(shù),我們看下例子[1]:

t = g(x)
y = t + tf.stop_gradient(f(x) - t)

這里,我們本來的前向傳遞函數(shù)是f(x),但是想要在反向時傳遞的函數(shù)是g(x),因為在前向過程中,tf.stop_gradient()不起作用,因此+t和-t抵消掉了,只剩下f(x)前向傳遞;而在反向過程中,因為tf.stop_gradient()的作用,使得f(x)-t的梯度變?yōu)榱?,從而只剩下g(x)在反向傳遞。

我們看下完整的例子:

import tensorflow as tf

x1 = tf.Variable(1)
x2 = tf.Variable(3)
x3 = tf.Variable(6)

f = x1+x2*x3
t = -f

y1 = t + tf.stop_gradient(f-t)
y2 = f

grad_1 = tf.gradients(y1, x1)
grad_2 = tf.gradients(y2, x1)
with tf.Session(config=config) as sess:
 sess.run(tf.global_variables_initializer())

 print(sess.run(grad_1))
 print(sess.run(grad_2))

第一個輸出為[-1],第二個輸出為[1],顯然也實現(xiàn)了梯度的反轉(zhuǎn)。

以上這篇基于TensorFlow中自定義梯度的2種方式就是小編分享給大家的全部內(nèi)容了,希望能給大家一個參考,也希望大家多多支持腳本之家。

相關(guān)文章

  • Python中pygame游戲模塊的用法詳解

    Python中pygame游戲模塊的用法詳解

    Pygame是一組用來開發(fā)游戲軟件的 Python 程序模塊,Pygame 在 SDL(Simple DirectMedia Layer) 的基礎(chǔ)上開發(fā)而成,它提供了諸多操作模塊,本文給大家介紹了Python中pygame游戲模塊的用法,需要的朋友可以參考下
    2024-01-01
  • ubuntu 18.04 安裝opencv3.4.5的教程(圖解)

    ubuntu 18.04 安裝opencv3.4.5的教程(圖解)

    這篇文章主要介紹了ubuntu 18.04 安裝opencv3.4.5的教程,本文圖文并茂給大家介紹的非常詳細,具有一定的參考借鑒價值,需要的朋友可以參考下
    2019-11-11
  • Python中高效的json對比庫deepdiff詳解

    Python中高效的json對比庫deepdiff詳解

    deepdiff模塊常用來校驗兩個對象是否一致,包含3個常用類,DeepDiff,DeepSearch和DeepHash,其中DeepDiff最常用,可以對字典,可迭代對象,字符串等進行對比,使用遞歸地查找所有差異,今天我們就學習一下快速實現(xiàn)代碼和文件對比的庫–deepdiff
    2022-07-07
  • 使用python復制PDF中的頁面的操作代碼

    使用python復制PDF中的頁面的操作代碼

    操作PDF文檔時,復制其中的指定頁面可以幫助我們從PDF文件中提取特定信息,如文本、圖表或數(shù)據(jù)等,以便在其他文檔中使用,本文將介紹如何使用Python 在同一文檔中復制PDF頁面,或者復制頁面到另一PDF文檔中,需要的朋友可以參考下
    2024-09-09
  • 用Python寫一段用戶登錄的程序代碼

    用Python寫一段用戶登錄的程序代碼

    下面小編就為大家分享一篇用Python寫一段用戶登錄的程序代碼,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2018-04-04
  • Python 詞典(Dict) 加載與保存示例

    Python 詞典(Dict) 加載與保存示例

    今天小編就為大家分享一篇Python 詞典(Dict) 加載與保存示例,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2019-12-12
  • python解析xml簡單示例

    python解析xml簡單示例

    這篇文章主要介紹了python解析xml,結(jié)合簡單實例形式分析了Python針對城市信息xml文件的讀取、解析相關(guān)操作技巧,需要的朋友可以參考下
    2019-06-06
  • python中ThreadPoolExecutor線程池和ProcessPoolExecutor進程池

    python中ThreadPoolExecutor線程池和ProcessPoolExecutor進程池

    這篇文章主要介紹了python中ThreadPoolExecutor線程池和ProcessPoolExecutor進程池,文章圍繞主題相關(guān)資料展開詳細的內(nèi)容介紹,具有一定的參考價值,感興趣的小伙伴可以參考一下
    2022-06-06
  • python列表:開始、結(jié)束、步長值實例

    python列表:開始、結(jié)束、步長值實例

    這篇文章主要介紹了python列表:開始、結(jié)束、步長值實例,具有很好的參考價值,希望對大家有所幫助。如有錯誤或未考慮完全的地方,望不吝賜教
    2022-05-05
  • 淺談Python協(xié)程

    淺談Python協(xié)程

    這篇文章主要介紹了Python協(xié)程的的相關(guān)資料,文中講解非常細致,代碼幫助大家更好的理解和學習,感興趣的朋友可以了解下
    2020-06-06

最新評論

东平县| 龙里县| 琼中| 自贡市| 射阳县| 石台县| 繁昌县| 建平县| 河东区| 洛扎县| 高州市| 通化县| 会宁县| 清涧县| 元谋县| 蒙山县| 玉龙| 定州市| 凤庆县| 德州市| 华安县| 芒康县| 临漳县| 大余县| 大荔县| 手游| 林州市| 南澳县| 延长县| 福贡县| 芒康县| 炉霍县| 吴川市| 蒲江县| 资溪县| 岳西县| 搜索| 加查县| 毕节市| 萍乡市| 左权县|