机器学习知识总结—— 3. 计算图与反向传播(Computational Graph and Back Propagation)
计算图与正向运算(Computational Graph)
计算图的概念其实很好理解,比方说对于如下的运算过程
- 首先计算: Z 0 = x + B Z_0 = x+B Z0=x+B
- 其次计算: Z 1 = Z 0 × A Z_1 = Z_0 \times A Z1=Z0×A
- 最后计算: y = Z 1 + C y = Z_1 + C y=Z1+C
我们可以用计算图的形式表示这个过程:

这样,从图上看,涉及到的运算过程一共是三层(每一步的运算过程都可以作为计算图的一层运算节点),涉及到的变量只有1个,而常量有3个。
所以,计算图的本质就是—— 「运算过程的结构图」 。此外,在某些计算图中你或许会看到常量跟运算符绑定在一起,共同作为运算节点出现,就像下面这样:

不过,垂直的方式一般不太符合工程的习惯,我们通常用遵循 「从左往右」 的顺序表示运算过程。所以对于上面这个运算过程,你更多的会看到它是以下面这个形式表现的。

接下来,我们可以把这个网络做复杂点,加入更多的运算,于是得到下面这图

例如对于上面这图来说,它的运算过程有两个:
- 其中黑色箭头表示的,是 正向运算 过程;
- 其中橙红色箭头部分,则是反向传播过程。
正向传播比较好理解,对于 y = a x + b y = ax + b y=ax+b 执行 a x + b ax + b ax+b 得到 y y y 的过程就是 正向运算 过程。那么反向传播过程又是什么呢?
反向传播与导数(Back Propagation)
为了说明什么是反向传播,我们先假定存在这样一个线性模型
y ^ = x × ω \hat y = x \times \omega y^=x×ω
然后再给定一组数据:
| x | y y y |
|---|---|
| 1 | 1 |
| 2 | 2 |
| 3 | 3 |
| … | … |
我们现在比较想知道,如果 x = 9 x=9 x=9 时, y y y 应该等于多少。假定我们不知道模型里权重 ω \omega ω 的值,我们应该用什么方法求 ω \omega ω 的值呢?
补充知识:
这里涉及到梯度下降算法,如果你不清楚梯度下降过程,那么请参考我前面的文章,对于你理解这部分内容会有帮助:
梯度下降算法——1. 什么是梯度下降
梯度下降算法——2.梯度下降算法实现
为了说明情况,现在我们用计算图来表示:

这里的 Y ∗ Y* Y∗ 因为绘图软件的关系,其实是 y ^ \hat y y^。我们得到了新的值后,需要先与观测值进行比对,于是产生了新的一层:

为了更好便于我们思考,我们把上述问题简化为 「一个点」 的情况,并且选用方差作为损失函数(评价函数),那么就有如下公式:
L o s s = 1 n ∑ ( y ^ − y ) 2 Loss = \frac{1}{n} \sum (\hat y - y) ^ 2 Loss=n1∑(y^−y)2
如果假设此时得到的 y ^ x ≠ y x \hat y_x \neq y_x y^x=yx, 即每一项 y ^ \hat y y^ 与 y之间都存在误差,自然MSE方程的和不可能为0,为了令 L o s s → 0 Loss \rightarrow 0 Loss→0,就会在这个时候激活反向传播过程。

由于令 L o s s → 0 Loss \rightarrow 0 Loss→0 的这个过程,很像以前学过的关于极限的形式。换句话说,要想让MSE方程为0,就需要让MSE方程中涉及的权重尽可能的接近被观测的数据值 ,而推动这个计算的,最有效的数学工具就是导数,于是
∂ L ∂ ω = ∂ L ∂ y ∂ y ∂ ω \frac{\partial L}{\partial \omega} = \frac{\partial L}{\partial y} \frac{\partial y}{\partial \omega} ∂ω∂L=∂y∂L∂ω∂y
当然,这个链可以非常长,如果涉及到的参数足够多,它还会像波一样向每个参数传播出去,所以叫 「反向传播」。

那么对于一个简单的节点来说,它的参数值是怎么更新的呢?
从梯度下降开始聊起
使用链式法则更新参数
尽管我们已经在前面的章节里介绍了什么是 梯度下降 并且实现了一个简单的 梯度下降程序 。在这一节里,我们将进一步扩展内容,介绍梯度下降和反向传播是如何帮助我们快速锁定解的。
首先,我们引入最常见的线性方程,它也是神经元网络使用最频繁的基础函数。
y = ω x + b y = \omega x + b y=ωx+b
它被表述为权重 ω \omega ω 和实验参数 x x x 的乘积,和偏见 b b b 之和。现在我们通过上述公式得到了推测值 y ^ \hat y y^,要与观测值 y y y 进行比对,判断权重 ω \omega ω 好坏的标准,落到了名为「损失函数」的测试函数,它可以是均方差函数,也可以是交叉墒,或者其他什么函数。
常见均方差函数
L o s s = 1 n ∑ ( y − y ^ ) 2 Loss = \frac{1}{n} \sum (y - \hat y)^2 Loss=n1∑(y−y^)2
总之,这类函数的名字叫 L o s s ( y , y ^ ) Loss(y, \hat{y}) Loss(y,y^) 就对了。由于 y = ω x y = \omega x y=ωx,所以 Loss 函数可以被改写为:
L o s s ( ω ) = 1 n ∑ ( y − ω x ) Loss(\omega) = \frac{1}{n} \sum (y - \omega x) Loss(ω)=n1∑(y−ωx)
现在我们希望,新的权重 ω n e w \omega_{new} ωnew 比 之前的 ω o l d \omega_{old} ωold 有更好的表现力,使得
L o s s ( ω n e w ) < L o s s ( ω o l d ) Loss(\omega_{new}) < Loss(\omega_{old}) Loss(ωnew)<Loss(ωold)
换句话说,我们希望能找到一个十分接近理想值 ω o \omega_o ωo 的权重 ω o + Δ ω \omega_o + \Delta \omega ωo+Δω 使的斜率接近0,于是有
lim Δ ω → 0 L o s s ( ω o + Δ ω ) − L o s s ( ω o ) Δ ω → 0 \lim_{\Delta \omega \rightarrow 0} \frac{Loss(\omega_o + \Delta \omega) - Loss(\omega_o)}{\Delta \omega} \rightarrow 0 Δω→0limΔωLoss(ωo+Δω)−Loss(ωo)→0
于是我们得到了关于Loss的导数
知识补充
导数公式的定义为:
f ′ ( x 0 ) = lim Δ x → 0 f ( x o + Δ x ) − f ( x o ) Δ x f'(x_0) = \lim_{\Delta x \rightarrow 0} \frac{f(x_o + \Delta x) - f(x_o)}{\Delta x} f′(x0)=Δx→0limΔxf(xo+Δx)−f(xo)
于是我们可以通过「链式法则」得到下面这串东西
d L d ω = d L d y d y d ω \frac{dL}{d\omega} = \frac{dL}{dy} \frac{d y}{d \omega} dωdL=dydLdωdy
然后我们带入公式,可以得到
∂ L ∂ ω = ∑ 2 n ( ω x − y ) ⋅ ∂ ∂ ω ( ω x − y ) = ∑ 2 x n ( ω x − ω ∗ x ) \frac{\partial L}{\partial \omega} = \sum \frac{2}{n} (\omega x - y) \cdot \frac{\partial}{\partial \omega} (\omega x - y) = \sum \frac{2x}{n} (\omega x - \omega^* x) ∂ω∂L=∑n2(ωx−y)⋅∂ω∂(ωx−y)=∑n2x(ωx−ω∗x)
这里的 2 x 2 n \frac{2x^2}{n} n2x2 是一个常数,所以我们可以用 λ \lambda λ 表示学习率,然后这里还有个很关键的问题,就是确定梯度下降的方向。

对于上面这个红色的点来说,它可以朝两个方向更新自己的权重。我们需要让它下降到图中最低点,所以,我们应该令
ω ^ = ω − Δ ω \hat \omega = \omega - \Delta \omega ω^=ω−Δω
于是得到最终的更新函数
ω ^ = ω − λ ∂ L ∂ ω \hat \omega = \omega - \lambda \frac{\partial L}{\partial \omega} ω^=ω−λ∂ω∂L
通过这个函数,我们终于可以让权重朝着设计的方向收敛,并找出最优解了。

参数的更新流程
基本上来说,参数的更新过程就是上面这个流程。
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐

所有评论(0)