3.2 深度学习计算

1.把计算画成一张图:计算图

说起来,前面几章我们一直在说算梯度、调参数,可模型里头那么多数,到底是怎么一步步算出来、又怎么一步步把梯度传回去的呢?这一章我们不妨把这套机制拆开来看看。干这件事的主角是计算图(computation graph),它是深度学习框架(比如PyTorch)在背后默默维护的一张图。

那么计算图究竟是什么?说穿了,它就是把一次复杂的计算,拆成一个个最小的零件,再用线把它们连起来。图里头有两种东西,一种叫节点,一个节点要么是一个变量(其实就是一个数),要么是一次运算(比如相加、相乘)。另一种叫边,边表示数据从哪个节点流到哪个节点。你要是写过代码,完全可以把它类比成一张变量依赖图,每个变量都是从哪几个变量算来的,一目了然。

举个例子。假设有输入 xx 和权重 ww,先做一次乘法得到 u=wxu=wx,再加一个偏置 bb 得到 v=u+bv=u+b。这条小链子里,uu 是中间结果,vv 是最后的结果。计算图把这几步全记下来,谁依赖谁、先算谁后算谁,清清楚楚。

说起来,计算图真正值钱的地方就在这里。它把每个结果是怎么来的全都留了档,等模型最后算出损失,系统就能顺着这张图原路折返,一个一个地揪出每个参数对损失到底有多大影响。要是没有这张图当向导,反向求导根本就无从下手。

再往深里说一点,计算图本身还顺手解决了一个特别烦人的问题,那就是重复计算。你想,模型里头一个中间结果经常要被后面好几条路用到,要是没有这张图统一管着,这个中间结果可能得反反复复算上好几遍,纯属浪费时间。我记得之前翻过一本讲编译原理的书,里头讲到一种叫公共子表达式消除的优化,说的其实就是同一回事,把算过的结果存起来,下次要用直接取。计算图做的就是这件事,它把每个节点只算一次,结果存起来,谁要用就直接来取,整张图的执行顺序也按依赖关系排得明明白白(这种按依赖排好的执行顺序,在图论里叫拓扑序,你大概知道有这么个词就行)。这也就是为什么PyTorch这种框架算梯度能算得那么快,基础打得好,上层才稳得住。

2.链式法则:拆解复合函数的瑞士军刀

要算每个参数对损失的影响,靠的是微积分里头一个叫链式法则(chain rule)的工具。这套工具专门对付那种一层套一层的复合函数,而神经网络恰恰就是一层套一层堆出来的,所以它俩天生一对。

LL 是损失,vv 是计算链上的某个中间变量,ww 是某个参数。如果损失是通过 vv 才影响到 ww 的(也就是说 LL 依赖 vv,而 vv 又依赖 ww),那么链式法则告诉我们:

Lw=Lvvw\frac{\partial L}{\partial w}=\frac{\partial L}{\partial v}\frac{\partial v}{\partial w}

这里头的 \partial 是偏导数符号,Lv\frac{\partial L}{\partial v} 念作损失 LL 对中间变量 vv 的偏导数,vw\frac{\partial v}{\partial w} 念作中间变量 vv 对参数 ww 的偏导数。

翻译成大白话就是,你想知道 ww 抖一下、LL 会跟着抖多少,可以拆成两步来看。第一步,先看 ww 抖一下、vv 会跟着抖多少,第二步,再看 vv 抖一下、LL 会跟着抖多少。把这两步的抖动幅度乘起来,就是你要的答案。

再举个例子,带点数字的。假设 v=3wv=3w(那么 vw=3\frac{\partial v}{\partial w}=3,意思是 ww 涨1,vv 就涨3),又假设 L=2vL=2v(那么 Lv=2\frac{\partial L}{\partial v}=2vv 涨1,LL 涨2)。那么 wwLL 的影响就是 3×2=63\times2=6,也就是说 ww 涨1、最终 LL 涨6。链式法则就是这么个一层层往下乘的套路。

这里不妨多说一句,链式法则还有个容易被新手忽略的细节。要是一个节点有好几个下游,也就是说它同时影响了后面好几个中间变量,那么它对损失的梯度,就得把每条路上的贡献都累加起来。比方说 ww 既进了 v1v_1 又进了 v2v_2,而 LL 又同时依赖 v1v_1v2v_2,那就有 Lw=Lv1v1w+Lv2v2w\frac{\partial L}{\partial w}=\frac{\partial L}{\partial v_1}\frac{\partial v_1}{\partial w}+\frac{\partial L}{\partial v_2}\frac{\partial v_2}{\partial w},注意这里是加号,每条路独立算一份再相加。神经网络里头一个权重往往同时拽着好多个输出,这个加法的细节后面反向传播会反复用到,先在心里留个印象,免得以后看到梯度凭空多出一项,还以为是框架出了bug。

真实的网络里,计算步骤多到数不清,链式法则就把每一步的局部导数一路乘下去。反向传播(backpropagation)干的就是这件事,只不过它聪明地倒过来,从最后的损失出发,一层一层往回算,每路过一个参数,就把属于它的那段梯度顺手算好填上,绝不重复劳动。

3.前向和反向:一次更新的两个半场

一次完整的参数更新,可以分成前后两个半场。

前半场叫前向计算,也叫前向传播。它从输入开始,顺着网络一路往前算,经过一层又一层,最后算出预测值,再拿预测值跟真实值比一比,得到损失。这一半场回答的问题是,模型现在到底预测得怎么样。

后半场叫反向计算,也就是反向传播。它从损失出发,沿着计算图往回传梯度,把梯度一个不落地送到每个参数手里。这一半场回答的问题是,每个参数该往哪边调、调多少。

为什么反向传播非要倒着来呢?这里头有个特别实在的考量。你看,前向算的时候,每个节点都已经把自己的输出值算出来了,这些值反向的时候可以原样拿来用。要是反过来从前往后求梯度,每算一个参数的梯度,就得把后面那一长串的局部导数从头乘到尾,同一个中间结果会被反复重算,算起来又慢又费劲。倒过来就不一样了,从损失出发,每往后退一步,就把当前节点对损失的累计梯度记下来,下一个节点直接拿这个累计梯度跟自己的局部导数一乘就完事,每个中间梯度只算一次,每个节点也只访问一次。这套拆解在算法里头还有个学名叫反向模式自动微分(reverse-mode autodiff),名字唬人,道理其实就是这么个倒着累乘的活儿。

梯度到手之后,就该优化器出场了,它按照梯度(再配上学习率)把参数实际改一改。于是前向、反向、优化这三件事,正好对应模型产生结果、模型发现误差、模型调整自己,一个完整的训练步就这么跑完了。这套流程你以后写代码会反复用到,先在脑子里过一遍,免得到时候看着框架的API摸不着头脑。

4.完整走一遍:带数字的反传小例子

光说套路有点虚,我们不妨拿一组具体的数字从头到尾走一遍,你就彻底通透了。

假设有一个超级迷你的网络,就一条线,没有乱七八糟的层。输入 x=2x=2,权重 w=3w=3,偏置 b=1b=1,这条网络干的事儿是先算 u=wxu=wx,再算 v=u+bv=u+bvv 就是模型吐出来的预测值。假设现在手头这个样本的真实值 y=10y=10,损失函数我们挑最好懂的那种,叫平方损失,写成 L=12(vy)2L=\frac{1}{2}(v-y)^2(前面那个 12\frac{1}{2} 纯粹是为了求导的时候把系数 22 抵消掉,让式子好看,没有别的深意)。

前向这一趟先跑起来。第一步算 u=wx=3×2=6u=wx=3\times2=6,第二步算 v=u+b=6+1=7v=u+b=6+1=7,最后算损失 L=12(710)2=12×9=4.5L=\frac{1}{2}(7-10)^2=\frac{1}{2}\times9=4.5。前半场结束,模型这次预测出来是 77,离真实值 1010 差了 33,损失是 4.54.5。这个数字摆在这里不算大也不算小,正好拿来当例子。

后半场开始,从损失往回退。第一站,先算损失对预测值的偏导数 Lv\frac{\partial L}{\partial v}。把 L=12(vy)2L=\frac{1}{2}(v-y)^2vv 求一下导,得到 Lv=vy=710=3\frac{\partial L}{\partial v}=v-y=7-10=-3。这个 3-3 的意思是,vv 涨1,LL 会跌3(因为是负号),把这个数先记下来。

第二站,算损失对 bb 的偏导数。因为 v=u+bv=u+b,所以 vb=1\frac{\partial v}{\partial b}=1,于是 Lb=Lv×vb=(3)×1=3\frac{\partial L}{\partial b}=\frac{\partial L}{\partial v}\times\frac{\partial v}{\partial b}=(-3)\times1=-3。也就是说,偏置 bb 涨1,LL 跌3。

第三站,算损失对 ww 的偏导数。这里要多过一节,因为 ww 先通过 uu 影响到 vv,再由 vv 影响到 LL。先算 vu=1\frac{\partial v}{\partial u}=1(因为 v=u+bv=u+b),再算 uw=x=2\frac{\partial u}{\partial w}=x=2(因为 u=wxu=wx,对 ww 求导就是把 xx 提出来当系数),所以 Lw=Lv×vu×uw=(3)×1×2=6\frac{\partial L}{\partial w}=\frac{\partial L}{\partial v}\times\frac{\partial v}{\partial u}\times\frac{\partial u}{\partial w}=(-3)\times1\times2=-6。结论就是,权重 ww 涨1,LL 会跌6。

发现没有,这两个梯度都是负数,说明参数往大了走,损失是往下掉的。这正好就是梯度下降想干的事儿,于是优化器按照 θθηθL\theta\leftarrow\theta-\eta\nabla_\theta L 的套路来更新参数。这里头 θ\theta 表示要更新的参数(可以代入 ww 或者 bb),η\eta 是学习率,θL\nabla_\theta L 是损失对参数 θ\theta 的梯度。设学习率 η=0.1\eta=0.1,那 ww 就更新成 ww0.1×(6)=3+0.6=3.6w\leftarrow w-0.1\times(-6)=3+0.6=3.6bb 更新成 bb0.1×(3)=1+0.3=1.3b\leftarrow b-0.1\times(-3)=1+0.3=1.3。你看,两个参数都被推着往大了走了一小步,正好对应损失下降的方向。再顺手验算一下,用新参数重跑一次前向,v=3.6×2+1.3=8.5v=3.6\times2+1.3=8.5,离真实值 1010 只差 1.51.5 了,损失也从 4.54.5 降到了 12×1.52=1.125\frac{1}{2}\times1.5^2=1.125,确确实实是在往好了走。这就是一次完整的参数更新,从头到尾就是这么点活儿。

5.回到PyTorch:自动求导是怎么替你代劳的

把这些原理捋清楚之后,再回头看PyTorch里的自动求导(autograd),你就会觉得特别亲切。你在代码里写的前向那套计算,PyTorch在背后一声不吭地全记成了一张计算图。等你算完损失,调一下反向传播那一行,框架就沿着这张图,把所有梯度都算好,整整齐齐填进每个参数的梯度字段里,根本不用你手写一行求导。你要操心的,就剩搞清楚什么时候该跑前向、什么时候该调反向、拿到梯度之后用什么策略更新参数这点事儿了。

说点实操层面的细节,免得你第一次写代码就卡壳。你想让PyTorch帮某个张量求梯度,得先把它标成需要梯度的,也就是代码里常见的requires_grad_(True)那一行。你只要在前向计算里把这个张量用进去,PyTorch就会偷偷把整条链路记进计算图里。等你算完损失,写一句loss.backward(),框架就开始沿着计算图反向跑,跑完之后每个参数的梯度就老老实实躺在它的.grad字段里,你直接读就行。还有个坑得提前打个预防针,PyTorch默认会把梯度累加到.grad里,而不会自动覆盖。所以每跑完一次参数更新,你得手动把梯度清零,对应的就是optimizer.zero_grad()那一行。这一步要是忘了写,是新手最容易踩的雷之一,梯度一累加,训练就直接乱套了,调到怀疑人生都不知道哪里出了问题。

还有一点值得提,PyTorch用的是动态计算图,意思是图是跟着你前向代码一步步现搭起来的,每跑一次前向,图就重新搭一遍,跑完反向还能再重搭。这跟早期一些框架用的静态图不一样,静态图得先把整张图定义死再喂数据。动态图的好处是你写代码就跟写普通Python一样随意,循环和判断随便加,调试起来也直观,print一下哪一步都看得见,对新手确实友好。这有点像一位江南的老厨师试菜,随手尝一口、随时调咸淡,而不是把一整桌菜按固定方子做完才端上桌,灵活劲儿就在这里。

说句题外话,这套靠计算图自动算梯度的本事,是现代深度学习框架能火起来的根基之一。早年练神经网络,反向传播的公式得自己一行行推导、一行行实现,写错一个符号能让你调上三天三夜的bug,调到怀疑人生。现在框架把这苦力活全包了,我们才得以把脑子省下来,去琢磨更上层的模型结构和训练技巧,说穿了就是站在了巨人的肩膀上。

6.训练循环的几个行话:iteration、epoch和batch size

最后这一节,我们不妨把训练循环里头几个外文词一起捋清楚,这几个词你以后看任何教程、读任何论文都会撞上,现在搞明白了以后就再也不会混。

第一个是batch size,叫批量大小。前面那章我们说过,训练的时候不会一个样本一个样本地喂,而是一次抓一小撮样本一起算平均损失再更新参数。这一小撮里头有多少个样本,就是batch size的值。打个比方,小明复习考试的时候总不会一道一道慢慢磨,肯定是凑一沓卷子一起写,那一沓卷子有多少张,差不多就是batch size的意思。batch size取多大很讲究,取太小吃不饱、训练慢,取太大显存又装不下,常见的取值一般在32到256之间,具体得看你显卡多大、模型多大,属于经验之谈。

第二个是iteration,叫迭代。它的定义特别干脆,拿一个batch跑一次前向、一次反向、再更新一次参数,这就叫一次iteration。换句话说,一次iteration就等于做了一次完整的参数更新。你训练的时候日志里看到的step 100那种计数,基本就是iteration在数。

第三个是epoch,叫回合(也有译作轮次的)。一个epoch的意思是把整个训练集从头到尾完整地过一遍。如果你的训练集有10000张图,batch size是100,那一个epoch里头就得跑 10000÷100=10010000\div100=100 次 iteration,才能把所有样本都轮一遍。所以一个epoch里有多少次iteration,约等于训练集样本总数除以batch size。

把这三个词串起来,总更新次数就好算了,总更新次数 \approx epoch数 ×\times 每个epoch的iteration数。比方说你训练50个epoch,每个epoch跑100次iteration,那一共就是 50×100=500050\times100=5000 次参数更新。这数你心里有谱之后,看日志就不会再被一堆step和epoch搞得云里雾里了。真实训练里头,这几个值都是要提前手动设好的超参数,调得好不好直接决定模型最后的成绩,建议大家多练多试,熟能生巧。

今天就先聊到这里,下一章见。

练习

Q1. 计算图是什么?它为什么能让反向求导变得可行又高效?

计算图把一次复杂计算拆成一个个最小节点(变量或运算),用边表示数据流向,谁依赖谁一目了然。它留了档,等算出损失就能顺着图原路折返揪出每个参数对损失的影响;同时它把每个中间结果只算一次并存起来,避免重复计算(类似公共子表达式消除),按拓扑序执行,所以算梯度又可行又快。

Q2. 迷你网络 u=wxu=wxv=u+bv=u+bL=12(vy)2L=\frac12(v-y)^2,已知 x=2,w=3,b=1,y=10x=2,w=3,b=1,y=10。手算前向得到 LL,再算 Lw\frac{\partial L}{\partial w}Lb\frac{\partial L}{\partial b}

前向:u=3×2=6u=3\times2=6v=6+1=7v=6+1=7L=12(710)2=4.5L=\frac12(7-10)^2=4.5。反向:Lv=vy=710=3\frac{\partial L}{\partial v}=v-y=7-10=-3Lb=(3)×1=3\frac{\partial L}{\partial b}=(-3)\times1=-3Lw=(3)×1×x=(3)×2=6\frac{\partial L}{\partial w}=(-3)\times1\times x=(-3)\times2=-6(多过一节 v=u+bv=u+bu=wxu=wx)。若学习率 0.1,则 w3+0.6=3.6w\leftarrow3+0.6=3.6b1+0.3=1.3b\leftarrow1+0.3=1.3

Q3. PyTorch 自动求导里为什么每次更新前要写 optimizer.zero_grad()?动态图和静态图有什么区别?

PyTorch 默认把梯度累加到 .grad 而不自动覆盖,所以每次更新前必须清零,否则梯度越加越多、训练直接乱套,这是新手最容易踩的雷。动态图是跟着前向代码一步步现搭、跑一次前向搭一遍,写代码像普通 Python 一样随意、调试直观;静态图得先把整张图定义死再喂数据,灵活度差但便于底层优化。

相关标签
深度学习自动求导反向传播PyTorch