假设我们有一个固定样本集 (x(1),y(1)),,(x(m),y(m))<script type="math/tex" id="MathJax-Element-1">{ (x^{(1)}, y^{(1)}), \ldots, (x^{(m)}, y^{(m)}) }</script>。具体来讲,对于单个样例(x,y)<script type="math/tex" id="MathJax-Element-2">(x,y)</script>,其代价函数为:

J(W,b;x,y)=12hW,b(x)y2.
<script type="math/tex; mode=display" id="MathJax-Element-3">\begin{align} J(W,b; x,y) = \frac{1}{2} \left\| h_{W,b}(x) - y \right\|^2. \end{align}</script>
代价函数是其均方误差。对于整个样例我们可以定义整体代价函数为:

J(W,b)=[1mi=1mJ(W,b;x(i),y(i))]+λ2l=1nl1i=1slj=1sl+1(W(l)ji)2=[1mi=1m(12hW,b(x(i))y(i)2)]+λ2l=1nl1i=1slj=1sl+1(W(l)ji)2
<script type="math/tex; mode=display" id="MathJax-Element-4">\begin{align} J(W,b) &= \left[ \frac{1}{m} \sum_{i=1}^m J(W,b;x^{(i)},y^{(i)}) \right] + \frac{\lambda}{2} \sum_{l=1}^{n_l-1} \; \sum_{i=1}^{s_l} \; \sum_{j=1}^{s_{l+1}} \left( W^{(l)}_{ji} \right)^2 \\ &= \left[ \frac{1}{m} \sum_{i=1}^m \left( \frac{1}{2} \left\| h_{W,b}(x^{(i)}) - y^{(i)} \right\|^2 \right) \right] + \frac{\lambda}{2} \sum_{l=1}^{n_l-1} \; \sum_{i=1}^{s_l} \; \sum_{j=1}^{s_{l+1}} \left( W^{(l)}_{ji} \right)^2 \end{align}</script>
以上公式中的第一项 J(W,b)<script type="math/tex" id="MathJax-Element-5">J(W,b)</script> 是一个均方差项。第二项是一个正则化项(也叫权重衰减项),其目的是减小权重的幅度,防止过拟合。

[注:通常权重衰减的计算并不使用偏置项b(l)i<script type="math/tex" id="MathJax-Element-6">b^{(l)}_i</script>,比如我们在 J(W,b)<script type="math/tex" id="MathJax-Element-7">J(W, b)</script>的定义中就没有使用。一般来说,将偏置项包含在权重衰减项中只会对最终的神经网络产生很小的影响。这个权重衰减实际上是贝叶斯正则化方法的变种。在贝叶斯正则化方法中,我们将高斯先验概率引入到参数中计算MAP(极大后验)估计(而不是极大似然估计)。]

权重衰减参数λ<script type="math/tex" id="MathJax-Element-8">\lambda</script>用于控制公式中两项的相对重要性。在此重申一下这两个复杂函数的含义:J(W,b;x,y)<script type="math/tex" id="MathJax-Element-9">J(W,b;x,y)</script>是针对单样例的方差代价函数;J(W,b)<script type="math/tex" id="MathJax-Element-10">J(W,b)</script> 是整体样本代价函数,它包含权重衰减项。
以上的代价函数经常被用于分类和回归问题。在分类问题中,我们用 y=0<script type="math/tex" id="MathJax-Element-11">y = 0</script>或1<script type="math/tex" id="MathJax-Element-12">1</script>,来代表两种类型的标签。

我们的目标是针对参数W<script type="math/tex" id="MathJax-Element-13">W</script>和b<script type="math/tex" id="MathJax-Element-14">b</script>来求其函数J(W,b)<script type="math/tex" id="MathJax-Element-15">J(W,b)</script>的最小值。为了求解神经网络,我们需要将每一个参数W(l)ij<script type="math/tex" id="MathJax-Element-16">W^{(l)}_{ij}</script> 和 b(l)i<script type="math/tex" id="MathJax-Element-17">b^{(l)}_i</script> 初始化为一个很小的、接近零的随机值(比如说,使用正态分布Normal(0,ϵ2)<script type="math/tex" id="MathJax-Element-18">{Normal}(0,\epsilon^2)</script>生成的随机值,ϵ<script type="math/tex" id="MathJax-Element-19">\epsilon</script> 设置为 0.01),之后对目标函数使用诸如批量梯度下降法的最优化算法。因为 J(W,b)<script type="math/tex" id="MathJax-Element-20">J(W, b)</script>是一个非凸函数,梯度下降法很可能会收敛到局部最优解;但是在实际应用中,梯度下降法通常能得到令人满意的结果。最后,需要再次强调的是,要将参数进行随机初始化,而不是全部置为0。如果所有参数都用相同的值作为初始值,那么所有隐藏层单元最终会得到与输入值有关的、相同的函数(也就是说,对于所有 i<script type="math/tex" id="MathJax-Element-21">i</script>,W(1)ij<script type="math/tex" id="MathJax-Element-22">W^{(1)}_{ij}</script>都会取相同的值,那么对于任何输入x<script type="math/tex" id="MathJax-Element-23">x</script> 都会有:a(2)1=a(2)2=a(2)3=<script type="math/tex" id="MathJax-Element-24">a^{(2)}_1 = a^{(2)}_2 = a^{(2)}_3 = \ldots</script>)。随机初始化的目的是使对称失效。

梯度下降法中每一次迭代都按照如下公式对参数W<script type="math/tex" id="MathJax-Element-25">W</script> 和b<script type="math/tex" id="MathJax-Element-26">b</script>进行更新:

W(l)ijb(l)i=W(l)ijαW(l)ijJ(W,b)=b(l)iαb(l)iJ(W,b)
<script type="math/tex; mode=display" id="MathJax-Element-123">\begin{align} W_{ij}^{(l)} &= W_{ij}^{(l)} - \alpha \frac{\partial}{\partial W_{ij}^{(l)}} J(W,b) \\ b_{i}^{(l)} &= b_{i}^{(l)} - \alpha \frac{\partial}{\partial b_{i}^{(l)}} J(W,b) \end{align}</script>
其中α<script type="math/tex" id="MathJax-Element-124">\alpha</script>是学习速率。其中关键步骤是计算偏导数。我们现在来讲一下反向传播算法,它是计算偏导数的一种有效方法。

我们首先来讲一下如何使用反向传播算法来计算 W(l)ijJ(W,b;x,y)<script type="math/tex" id="MathJax-Element-125">\frac{\partial}{\partial W_{ij}^{(l)}} J(W,b; x, y)</script> 和b(l)iJ(W,b;x,y)<script type="math/tex" id="MathJax-Element-126">\frac{\partial}{\partial b_{i}^{(l)}} J(W,b; x, y)</script>,这两项是单个样例(x,y)<script type="math/tex" id="MathJax-Element-127">(x,y)</script>的代价函数 J(W,b;x,y)<script type="math/tex" id="MathJax-Element-128">J(W,b;x,y)</script>的偏导数。一旦我们求出该偏导数,就可以推导出整体代价函数J(W,b)<script type="math/tex" id="MathJax-Element-129">J(W,b)</script>的偏导数:

W(l)ijJ(W,b)b(l)iJ(W,b)=1mi=1mW(l)ijJ(W,b;x(i),y(i))+λW(l)ij=1mi=1mb(l)iJ(W,b;x(i),y(i))
<script type="math/tex; mode=display" id="MathJax-Element-130">\begin{align} \frac{\partial}{\partial W_{ij}^{(l)}} J(W,b) &= \left[ \frac{1}{m} \sum_{i=1}^m \frac{\partial}{\partial W_{ij}^{(l)}} J(W,b; x^{(i)}, y^{(i)}) \right] + \lambda W_{ij}^{(l)} \\ \frac{\partial}{\partial b_{i}^{(l)}} J(W,b) &= \frac{1}{m}\sum_{i=1}^m \frac{\partial}{\partial b_{i}^{(l)}} J(W,b; x^{(i)}, y^{(i)}) \end{align}</script>
以上两行公式稍有不同,第一行比第二行多出一项,是因为权重衰减是作用于W<script type="math/tex" id="MathJax-Element-131">W</script> 而不是b<script type="math/tex" id="MathJax-Element-132">b</script>。

反向传播算法的思路如下:给定一个样例(x,y)<script type="math/tex" id="MathJax-Element-133">(x,y)</script>,我们首先进行“前向传导”运算,计算出网络中所有的激活值,包括hW,b(x)<script type="math/tex" id="MathJax-Element-134">h_{W,b}(x)</script> 的输出值。之后,针对第l<script type="math/tex" id="MathJax-Element-135">l</script> 层的每一个节点 i<script type="math/tex" id="MathJax-Element-136">i</script>,我们计算出其“残差”δ(l)i<script type="math/tex" id="MathJax-Element-137">\delta^{(l)}_i</script>,该残差表明了该节点对最终输出值的残差产生了多少影响。对于最终的输出节点,我们可以直接算出网络产生的激活值与实际值之间的差距,我们将这个差距定义为 δ(nl)i<script type="math/tex" id="MathJax-Element-138">\delta^{(n_l)}_i </script>(第nl<script type="math/tex" id="MathJax-Element-139">n_l </script>层表示输出层)。对于隐藏单元我们如何处理呢?我们将基于节点(第 l+1<script type="math/tex" id="MathJax-Element-140">l+1</script> 层节点)残差的加权平均值计算δ(l)i<script type="math/tex" id="MathJax-Element-141">\delta^{(l)}_i</script>,这些节点以a(l)i<script type="math/tex" id="MathJax-Element-142">a^{(l)}_i</script>作为输入。下面将给出反向传导算法的细节:

进行前馈传导计算,利用前向传导公式,得到L2,L3,<script type="math/tex" id="MathJax-Element-143">L_2, L_3, \ldots </script> 直到输出层Lnl<script type="math/tex" id="MathJax-Element-144">L_{n_l}</script>的激活值。
对于第nl<script type="math/tex" id="MathJax-Element-145">n_l</script> 层(输出层)的每个输出单元 i<script type="math/tex" id="MathJax-Element-146">i</script>,我们根据以下公式计算残差:

δ(nl)i=z(nl)i12yhW,b(x)2=(yia(nl)i)f(z(nl)i)
<script type="math/tex; mode=display" id="MathJax-Element-50">\begin{align} \delta^{(n_l)}_i = \frac{\partial}{\partial z^{(n_l)}_i} \;\; \frac{1}{2} \left\|y - h_{W,b}(x)\right\|^2 = - (y_i - a^{(n_l)}_i) \cdot f'(z^{(n_l)}_i) \end{align}</script>

δ(nl)i=znliJ(W,b;x,y)=znli12yhW,b(x)2=znli12j=1Snl(yja(nl)j)2=znli12j=1Snl(yjf(z(nl)j))2=(yif(z(nl)i))f(z(nl)i)=(yia(nl)i)f(z(nl)i)
<script type="math/tex; mode=display" id="MathJax-Element-51">\begin{align} \delta^{(n_l)}_i &= \frac{\partial}{\partial z^{n_l}_i}J(W,b;x,y) = \frac{\partial}{\partial z^{n_l}_i}\frac{1}{2} \left\|y - h_{W,b}(x)\right\|^2 \\ &= \frac{\partial}{\partial z^{n_l}_i}\frac{1}{2} \sum_{j=1}^{S_{n_l}} (y_j-a_j^{(n_l)})^2 = \frac{\partial}{\partial z^{n_l}_i}\frac{1}{2} \sum_{j=1}^{S_{n_l}} (y_j-f(z_j^{(n_l)}))^2 \\ &= - (y_i - f(z_i^{(n_l)})) \cdot f'(z^{(n_l)}_i) = - (y_i - a^{(n_l)}_i) \cdot f'(z^{(n_l)}_i) \end{align}</script>
l=nl1,nl2,nl3,,2<script type="math/tex" id="MathJax-Element-52">l = n_l-1, n_l-2, n_l-3, \ldots, 2</script>的各个层,第 l<script type="math/tex" id="MathJax-Element-53">l</script>层的第 i<script type="math/tex" id="MathJax-Element-54">i</script> 个节点的残差计算方法如下:
δ(l)i=j=1sl+1W(l)jiδ(l+1)jf(z(l)i)
<script type="math/tex; mode=display" id="MathJax-Element-55"> \delta^{(l)}_i = \left( \sum_{j=1}^{s_{l+1}} W^{(l)}_{ji} \delta^{(l+1)}_j \right) f'(z^{(l)}_i)</script>

δ(nl1)i=znl1iJ(W,b;x,y)=znl1i12yhW,b(x)2=znl1i12j=1Snl(yja(nl)j)2=12j=1Snlznl1i(yja(nl)j)2=12j=1Snlznl1i(yjf(z(nl)j))2=j=1Snl(yjf(z(nl)j))z(nl1)if(z(nl)j)=j=1Snl(yjf(z(nl)j))f(z(nl)j)z(nl)jz(nl1)i=j=1Snlδ(nl)jz(nl)jznl1i=j=1Snlδ(nl)jznl1ik=1Snl1f(znl1k)Wnl1jk=j=1Snlδ(nl)jWnl1jif(znl1i)=j=1SnlWnl1jiδ(nl)jf(znl1i)
<script type="math/tex; mode=display" id="MathJax-Element-56"> \begin{align} \delta^{(n_l-1)}_i &=\frac{\partial}{\partial z^{n_l-1}_i}J(W,b;x,y) = \frac{\partial}{\partial z^{n_l-1}_i}\frac{1}{2} \left\|y - h_{W,b}(x)\right\|^2 = \frac{\partial}{\partial z^{n_l-1}_i}\frac{1}{2} \sum_{j=1}^{S_{n_l}}(y_j-a_j^{(n_l)})^2 \\ &= \frac{1}{2} \sum_{j=1}^{S_{n_l}}\frac{\partial}{\partial z^{n_l-1}_i}(y_j-a_j^{(n_l)})^2 = \frac{1}{2} \sum_{j=1}^{S_{n_l}}\frac{\partial}{\partial z^{n_l-1}_i}(y_j-f(z_j^{(n_l)}))^2 \\ &= \sum_{j=1}^{S_{n_l}}-(y_j-f(z_j^{(n_l)})) \cdot \frac{\partial}{\partial z_i^{(n_l-1)}}f(z_j^{(n_l)}) = \sum_{j=1}^{S_{n_l}}-(y_j-f(z_j^{(n_l)})) \cdot f'(z_j^{(n_l)}) \cdot \frac{\partial z_j^{(n_l)}}{\partial z_i^{(n_l-1)}} \\ &= \sum_{j=1}^{S_{n_l}} \delta_j^{(n_l)} \cdot \frac{\partial z_j^{(n_l)}}{\partial z_i^{n_l-1}} = \sum_{j=1}^{S_{n_l}} \left(\delta_j^{(n_l)} \cdot \frac{\partial}{\partial z_i^{n_l-1}}\sum_{k=1}^{S_{n_l-1}}f(z_k^{n_l-1}) \cdot W_{jk}^{n_l-1}\right) \\ &= \sum_{j=1}^{S_{n_l}} \delta_j^{(n_l)} \cdot W_{ji}^{n_l-1} \cdot f'(z_i^{n_l-1}) = \left(\sum_{j=1}^{S_{n_l}}W_{ji}^{n_l-1}\delta_j^{(n_l)}\right)f'(z_i^{n_l-1}) \end{align} </script>
将上式中的nl1<script type="math/tex" id="MathJax-Element-57">n_l-1</script>与nl<script type="math/tex" id="MathJax-Element-58">n_l</script>的关系替换为l<script type="math/tex" id="MathJax-Element-59">l</script>与l+1<script type="math/tex" id="MathJax-Element-60">l+1</script>的关系,就可以得到:
δ(l)i=j=1sl+1W(l)jiδ(l+1)jf(z(l)i)
<script type="math/tex; mode=display" id="MathJax-Element-61"> \delta^{(l)}_i = \left( \sum_{j=1}^{s_{l+1}} W^{(l)}_{ji} \delta^{(l+1)}_j \right) f'(z^{(l)}_i)</script>
以上逐次从后向前求导的过程即为“反向传导”的本意所在。
下面计算我们需要的偏导数,计算方法如下:

W(l)ijJ(W,b;x,y)b(l)iJ(W,b;x,y)=a(l)jδ(l+1)i=δ(l+1)i.
<script type="math/tex; mode=display" id="MathJax-Element-62">\begin{align} \frac{\partial}{\partial W_{ij}^{(l)}} J(W,b; x, y) &= a^{(l)}_j \delta_i^{(l+1)} \\ \frac{\partial}{\partial b_{i}^{(l)}} J(W,b; x, y) &= \delta_i^{(l+1)}. \end{align}</script>

最后,我们用矩阵-向量表示法重写以上算法。我们使用“<script type="math/tex" id="MathJax-Element-63">\bullet</script>” 表示向量乘积运算符(在Matlab或Octave里用“.*”表示,也称作阿达马乘积)。若 a=bc<script type="math/tex" id="MathJax-Element-64">a = b \bullet c</script>,则 ai=bici<script type="math/tex" id="MathJax-Element-65">a_i = b_ic_i</script>。在上一个教程中我们扩展了 f()<script type="math/tex" id="MathJax-Element-66">f(\cdot)</script> 的定义,使其包含向量运算,这里我们也对偏导数 f()<script type="math/tex" id="MathJax-Element-67">f'(\cdot) </script>也做了同样的处理(于是又有f([z1,z2,z3])=[f(z1),f(z2),f(z3)]<script type="math/tex" id="MathJax-Element-68">f'([z_1, z_2, z_3]) = [f'(z_1), f'(z_2), f'(z_3)]</script>)。

那么,反向传播算法可表示为以下几个步骤:
进行前馈传导计算,利用前向传导公式,得到L2,L3,<script type="math/tex" id="MathJax-Element-69">L_2, L_3, \ldots</script>直到输出层 Lnl<script type="math/tex" id="MathJax-Element-70">L_{n_l}</script>的激活值。
对输出层(第nl<script type="math/tex" id="MathJax-Element-71">n_l</script>层),计算:

δ(nl)=(ya(nl))f(z(nl))
<script type="math/tex; mode=display" id="MathJax-Element-72">\begin{align} \delta^{(n_l)} = - (y - a^{(n_l)}) \bullet f'(z^{(n_l)}) \end{align}</script>

对于 l=nl1,nl2,nl3,,2<script type="math/tex" id="MathJax-Element-73">l = n_l-1, n_l-2, n_l-3, \ldots, 2 </script>的各层,计算:

δ(l)=((W(l))Tδ(l+1))f(z(l))
<script type="math/tex; mode=display" id="MathJax-Element-74">\begin{align} \delta^{(l)} = \left((W^{(l)})^T \delta^{(l+1)}\right) \bullet f'(z^{(l)}) \end{align}</script>
计算最终需要的偏导数值:
W(l)J(W,b;x,y)b(l)J(W,b;x,y)=δ(l+1)(a(l))T,=δ(l+1).
<script type="math/tex; mode=display" id="MathJax-Element-75">\begin{align} \nabla_{W^{(l)}} J(W,b;x,y) &= \delta^{(l+1)} (a^{(l)})^T, \\ \nabla_{b^{(l)}} J(W,b;x,y) &= \delta^{(l+1)}. \end{align}</script>

实现中应注意:在以上的第2步和第3步中,我们需要为每一个 i<script type="math/tex" id="MathJax-Element-76">i</script> 值计算其f(z(l)i)<script type="math/tex" id="MathJax-Element-77">f'(z^{(l)}_i)</script>。假设f(z)<script type="math/tex" id="MathJax-Element-78">f(z)</script>是sigmoid函数,并且我们已经在前向传导运算中得到了a(l)i<script type="math/tex" id="MathJax-Element-79">a^{(l)}_i</script>。那么,使用我们早先推导出的 f(z)<script type="math/tex" id="MathJax-Element-80">f'(z)</script>表达式,就可以计算得到f(z(l)i)=a(l)i(1a(l)i)<script type="math/tex" id="MathJax-Element-81">f'(z^{(l)}_i) = a^{(l)}_i (1- a^{(l)}_i)</script>。

最后,我们将对梯度下降算法做个全面总结。在下面的伪代码中,ΔW(l)<script type="math/tex" id="MathJax-Element-82">\Delta W^{(l)} </script>是一个与矩阵W(l)<script type="math/tex" id="MathJax-Element-83">W^{(l)}</script>维度相同的矩阵,Δb(l)<script type="math/tex" id="MathJax-Element-84">\Delta b^{(l)}</script> 是一个与b(l)<script type="math/tex" id="MathJax-Element-85">b^{(l)}</script> 维度相同的向量。注意这里ΔW(l)<script type="math/tex" id="MathJax-Element-86">\Delta W^{(l)}</script>是一个矩阵,而不是Δ<script type="math/tex" id="MathJax-Element-87">\Delta</script> 与W(l)<script type="math/tex" id="MathJax-Element-88">W^{(l)}</script>相乘”。下面,我们实现批量梯度下降法中的一次迭代:

对于所有l<script type="math/tex" id="MathJax-Element-89">l</script>,令ΔW(l):=0<script type="math/tex" id="MathJax-Element-90">\Delta W^{(l)} := 0</script> , Δb(l):=0<script type="math/tex" id="MathJax-Element-91">\Delta b^{(l)} := 0</script> (设置为全零矩阵或全零向量)
对于 i=1<script type="math/tex" id="MathJax-Element-92">i = 1 </script>到 m<script type="math/tex" id="MathJax-Element-93">m</script>,
使用反向传播算法计算W(l)J(W,b;x,y)<script type="math/tex" id="MathJax-Element-94">\nabla_{W^{(l)}} J(W,b;x,y)</script> 和 b(l)J(W,b;x,y)<script type="math/tex" id="MathJax-Element-95">\nabla_{b^{(l)}} J(W,b;x,y)</script>。
计算 ΔW(l):=ΔW(l)+W(l)J(W,b;x,y)<script type="math/tex" id="MathJax-Element-96">\Delta W^{(l)} := \Delta W^{(l)} + \nabla_{W^{(l)}} J(W,b;x,y)</script>。
计算 Δb(l):=Δb(l)+b(l)J(W,b;x,y)<script type="math/tex" id="MathJax-Element-97">\Delta b^{(l)} := \Delta b^{(l)} + \nabla_{b^{(l)}} J(W,b;x,y)</script>。
更新权重参数:

W(l)b(l)=W(l)α[(1mΔW(l))+λW(l)]=b(l)α[1mΔb(l)]
<script type="math/tex; mode=display" id="MathJax-Element-98">\begin{align} W^{(l)} &= W^{(l)} - \alpha \left[ \left(\frac{1}{m} \Delta W^{(l)} \right) + \lambda W^{(l)}\right] \\ b^{(l)} &= b^{(l)} - \alpha \left[\frac{1}{m} \Delta b^{(l)}\right] \end{align}</script>
现在,我们可以重复梯度下降法的迭代步骤来减小代价函数 J(W,b)<script type="math/tex" id="MathJax-Element-99">J(W,b) </script>的值,进而求解我们的神经网络。

注:本文参考Ufldl教程

Logo

魔乐社区(Modelers.cn) 是一个中立、公益的人工智能社区,提供人工智能工具、模型、数据的托管、展示与应用协同服务,为人工智能开发及爱好者搭建开放的学习交流平台。社区通过理事会方式运作,由全产业链共同建设、共同运营、共同享有,推动国产AI生态繁荣发展。

更多推荐