概念

反向传播(Back Propagation, BP)算法是使用梯度下降法相关的算法来优化一个神经网络时计算每一层梯度的方法,主要使用了多元函数的链式法则

已知多元函数 u=g(y1,y2,...,ym) u = g ( y 1 , y 2 , . . . , y m ) <script type="math/tex" id="MathJax-Element-1">u=g(y_1,y_2,...,y_m)</script> ,且 yi=fi(x) y i = f i ( x ) <script type="math/tex" id="MathJax-Element-2">y_i=f_i(x)</script>,所有函数都可微,则

ux=i=1muyiyix ∂ u ∂ x = ∑ i = 1 m ∂ u ∂ y i ∂ y i ∂ x
<script type="math/tex; mode=display" id="MathJax-Element-3">\frac{\partial u}{\partial x}=\sum_{i=1}^{m}\frac{\partial u}{\partial y_i}\frac{\partial y_i}{\partial x}</script>

公式推导

1、模型

不失一般性,我们考虑以下4层结构的神经网络(全连接):
这里写图片描述

2、符号说明

符号含义
nl n l <script type="math/tex" id="MathJax-Element-4">n_l</script>网络层数
yj y j <script type="math/tex" id="MathJax-Element-5">y_j</script>输出层第 j j <script type="math/tex" id="MathJax-Element-6">j</script>类标签
Sl<script type="math/tex" id="MathJax-Element-7">S_l</script> l l <script type="math/tex" id="MathJax-Element-8">l</script>层神经元个数(不包括偏置)
g(x)<script type="math/tex" id="MathJax-Element-9">g(x)</script>激活函数
w(l)ij w i j ( l ) <script type="math/tex" id="MathJax-Element-10">w_{ij}^{(l)}</script> l l <script type="math/tex" id="MathJax-Element-11">l</script>层第j<script type="math/tex" id="MathJax-Element-12">j</script>个单元与第 l+1 l + 1 <script type="math/tex" id="MathJax-Element-13">l+1</script>层第 i i <script type="math/tex" id="MathJax-Element-14">i</script>个单元之间的链接参数
bi(l)<script type="math/tex" id="MathJax-Element-15">b_i^{(l)}</script> l l <script type="math/tex" id="MathJax-Element-16">l</script>层的偏置与第l+1<script type="math/tex" id="MathJax-Element-17">l+1</script>层第 i i <script type="math/tex" id="MathJax-Element-18">i</script>个单元之间的链接参数
zi(l)<script type="math/tex" id="MathJax-Element-19">z_i^{(l)}</script> l l <script type="math/tex" id="MathJax-Element-20">l</script>层第i<script type="math/tex" id="MathJax-Element-21">i</script>个单元的输入(加权和,包括偏置)
a(l)i a i ( l ) <script type="math/tex" id="MathJax-Element-22">a_i^{(l)}</script> l l <script type="math/tex" id="MathJax-Element-23">l</script>层第i<script type="math/tex" id="MathJax-Element-24">i</script>个单元的输出(激活函数的值)
δ(l)i δ i ( l ) <script type="math/tex" id="MathJax-Element-25">\delta_i^{(l)}</script> l l <script type="math/tex" id="MathJax-Element-26">l</script>层第i<script type="math/tex" id="MathJax-Element-27">i</script>个单元的输入的偏导(或称为灵敏度、残差)
J(θ) J ( θ ) <script type="math/tex" id="MathJax-Element-28">J(\theta)</script>代价函数

3、符号定义

z(l)ia(l)iJ(θ)δ(l)i=b(l1)i+j=1Sl1w(l1)ija(l1)j=g(z(l)i)=12j=1Sl(yja(l)j)2=J(θ)z(l)i z i ( l ) = b i ( l − 1 ) + ∑ j = 1 S l − 1 w i j ( l − 1 ) a j ( l − 1 ) a i ( l ) = g ( z i ( l ) ) J ( θ ) = 1 2 ∑ j = 1 S l ( y j − a j ( l ) ) 2 δ i ( l ) = ∂ J ( θ ) ∂ z i ( l )
<script type="math/tex; mode=display" id="MathJax-Element-36">\begin{align*} z_i^{(l)}&=b_i^{(l-1)}+\sum_{j=1}^{S_{l-1}}w_{ij}^{(l-1)}a_j^{(l-1)} \\ a_i^{(l)}&=g(z_i^{(l)}) \\ J(\theta)&=\frac{1}{2}\sum_{j=1}^{S_l}(y_j-a_j^{(l)})^2 \\ \delta_i^{(l)}&=\frac{\partial J(\theta)}{\partial z_i^{(l)}} \end{align*}</script>

4、推导过程

δ(nl)iδ(l)iJ(θ)w(l)ijJ(θ)b(l)i=J(θ)z(nl)i=12z(nl)ij=1Snl(yja(nl)j)2=12z(nl)ij=1Snl(yjg(z(nl)j))2=12z(nl)i(yjg(z(nl)i))2=(yia(nl)i)g(z(nl)i)=J(θ)z(l)i=j=1Sl+1J(θ)z(l+1)jz(l+1)jz(l)i=j=1Sl+1δ(l+1)jz(l+1)jz(l)i=j=1Sl+1δ(l+1)jz(l)i(b(l)j+k=1Slw(l)jka(l)k)=j=1Sl+1δ(l+1)jz(l)i(b(l)j+k=1Slw(l)jkg(z(l)k))=j=1Sl+1δ(l+1)jz(l)i(w(l)jig(z(l)i))=j=1Sl+1δ(l+1)jw(l)jig(z(l)i)=g(z(l)i)j=1Sl+1δ(l+1)jw(l)ji=J(θ)z(l+1)iz(l+1)iw(l)ij=δ(l+1)iz(l+1)iw(l)ij=δ(l+1)iw(l)ij(b(l)i+k=1Slw(l)ika(l)k)=δ(l+1)ia(l)j=δ(l+1)ib(l)i(b(l)i+k=1Slw(l)ika(l)k)=δ(l+1)i δ i ( n l ) = ∂ J ( θ ) ∂ z i ( n l ) = 1 2 ∂ ∂ z i ( n l ) ∑ j = 1 S n l ( y j − a j ( n l ) ) 2 = 1 2 ∂ ∂ z i ( n l ) ∑ j = 1 S n l ( y j − g ( z j ( n l ) ) ) 2 = 1 2 ∂ ∂ z i ( n l ) ( y j − g ( z i ( n l ) ) ) 2 = − ( y i − a i ( n l ) ) g ′ ( z i ( n l ) ) δ i ( l ) = ∂ J ( θ ) ∂ z i ( l ) = ∑ j = 1 S l + 1 ∂ J ( θ ) ∂ z j ( l + 1 ) ∂ z j ( l + 1 ) ∂ z i ( l ) = ∑ j = 1 S l + 1 δ j ( l + 1 ) ∂ z j ( l + 1 ) ∂ z i ( l ) = ∑ j = 1 S l + 1 δ j ( l + 1 ) ∂ ∂ z i ( l ) ( b j ( l ) + ∑ k = 1 S l w j k ( l ) a k ( l ) ) = ∑ j = 1 S l + 1 δ j ( l + 1 ) ∂ ∂ z i ( l ) ( b j ( l ) + ∑ k = 1 S l w j k ( l ) g ( z k ( l ) ) ) = ∑ j = 1 S l + 1 δ j ( l + 1 ) ∂ ∂ z i ( l ) ( w j i ( l ) g ( z i ( l ) ) ) = ∑ j = 1 S l + 1 δ j ( l + 1 ) w j i ( l ) g ′ ( z i ( l ) ) = g ′ ( z i ( l ) ) ∑ j = 1 S l + 1 δ j ( l + 1 ) w j i ( l ) ∂ J ( θ ) ∂ w i j ( l ) = ∂ J ( θ ) ∂ z i ( l + 1 ) ∂ z i ( l + 1 ) ∂ w i j ( l ) = δ i ( l + 1 ) ∂ z i ( l + 1 ) ∂ w i j ( l ) = δ i ( l + 1 ) ∂ ∂ w i j ( l ) ( b i ( l ) + ∑ k = 1 S l w i k ( l ) a k ( l ) ) = δ i ( l + 1 ) a j ( l ) ∂ J ( θ ) ∂ b i ( l ) = δ i ( l + 1 ) ∂ ∂ b i ( l ) ( b i ( l ) + ∑ k = 1 S l w i k ( l ) a k ( l ) ) = δ i ( l + 1 )
<script type="math/tex; mode=display" id="MathJax-Element-30">\begin{align*} \delta_i^{(n_l)}&=\frac{\partial J(\theta)}{\partial z_i^{(n_l)}}\\ &=\frac{1}{2}\frac{\partial}{\partial z_i^{(n_l)}}\sum_{j=1}^{S_{n_l}}(y_j-a_j^{(n_l)})^2 \\ &=\frac{1}{2}\frac{\partial}{\partial z_i^{(n_l)}}\sum_{j=1}^{S_{n_l}}(y_j-g(z_j^{(n_l)}))^2 \\ &=\frac{1}{2}\frac{\partial}{\partial z_i^{(n_l)}}(y_j-g(z_i^{(n_l)}))^2 \\ &=-(y_i-a_i^{(n_l)})g'(z_i^{(n_l)})\\ \delta_i^{(l)}&=\frac{\partial J(\theta)}{\partial z_i^{(l)}}\\ &=\sum_{j=1}^{S_{l+1}}\frac{\partial J(\theta)}{\partial z_j^{(l+1)}}\frac{\partial z_j^{(l+1)}}{\partial z_i^{(l)}}\\ &=\sum_{j=1}^{S_{l+1}}\delta_j^{(l+1)}\frac{\partial z_j^{(l+1)}}{\partial z_i^{(l)}}\\ &=\sum_{j=1}^{S_{l+1}}\delta_j^{(l+1)}\frac{\partial}{\partial z_i^{(l)}}(b_j^{(l)}+\sum_{k=1}^{S_l}w_{jk}^{(l)}a_k^{(l)}) \\ &=\sum_{j=1}^{S_{l+1}}\delta_j^{(l+1)}\frac{\partial}{\partial z_i^{(l)}}(b_j^{(l)}+\sum_{k=1}^{S_l}w_{jk}^{(l)}g(z_k^{(l)})) \\ &=\sum_{j=1}^{S_{l+1}}\delta_j^{(l+1)}\frac{\partial}{\partial z_i^{(l)}}(w_{ji}^{(l)}g(z_i^{(l)})) \\ &=\sum_{j=1}^{S_{l+1}}\delta_j^{(l+1)}w_{ji}^{(l)}g'(z_i^{(l)}) \\ &=g'(z_i^{(l)})\sum_{j=1}^{S_{l+1}}\delta_j^{(l+1)}w_{ji}^{(l)} \\ \frac{\partial J(\theta)}{\partial w_{ij}^{(l)}}&=\frac{\partial J(\theta)}{\partial z_i^{(l+1)}}\frac{\partial z_i^{(l+1)}}{\partial w_{ij}^{(l)}}\\ &=\delta _i^{(l+1)}\frac{\partial z_i^{(l+1)}}{\partial w_{ij}^{(l)}}\\ &=\delta _i^{(l+1)}\frac{\partial}{\partial w_{ij}^{(l)}}(b_i^{(l)}+\sum_{k=1}^{S_l}w_{ik}^{(l)}a_k^{(l)}) \\ &=\delta _i^{(l+1)}a_j^{(l)}\\ \frac{\partial J(\theta)}{\partial b_i^{(l)}}&=\delta _i^{(l+1)}\frac{\partial}{\partial b_i^{(l)}}(b_i^{(l)}+\sum_{k=1}^{S_l}w_{ik}^{(l)}a_k^{(l)}) \\ &=\delta _i^{(l+1)} \end{align*}</script>

向量形式的公式

δ(l)J(θ)W(l)J(θ)b(l)=(W(l))Tδ(l+1)g(z(l))=δ(l+1)(a(l))T=δ(l+1) δ ( l ) = ( W ( l ) ) T δ ( l + 1 ) ∘ g ′ ( z ( l ) ) ∂ J ( θ ) ∂ W ( l ) = δ ( l + 1 ) ( a ( l ) ) T ∂ J ( θ ) ∂ b ( l ) = δ ( l + 1 )
<script type="math/tex; mode=display" id="MathJax-Element-31">\begin{align*} \boldsymbol{\delta}^{(l)}&=(\boldsymbol{W}^{(l)})^T\boldsymbol{\delta}^{(l+1)}\circ g'(\boldsymbol{z}^{(l)})\\ \frac{\partial J(\theta)}{\partial \boldsymbol{W}^{(l)}}&=\boldsymbol{\delta}^{(l+1)}(\boldsymbol{a}^{(l)})^T\\ \frac{\partial J(\theta)}{\partial \boldsymbol{b}^{(l)}}&=\boldsymbol{\delta}^{(l+1)} \end{align*}</script>
其中, <script type="math/tex" id="MathJax-Element-32">\circ</script>表示每个元素相乘,粗体的小写符号表示列向量,粗体的大写符号表示矩阵。

参考

([1] 中的公式推导有错误,本文已纠正)
[1] https://www.cnblogs.com/nowgood/p/backprop.html
[2] Bouvrie J. Notes on convolutional neural networks[J]. 2006.

Logo

有“AI”的1024 = 2048,欢迎大家加入2048 AI社区

更多推荐