3.10 训练稳定性与计算性能
1.数值尺度要照看好:偏大偏小都出问题
说起来,前面几章我们一直说,神经网络里一层一层做的事,其实就是拿输入乘上参数、再过一个激活函数。这么一来,整张网络里到处都是数:输入是数,参数是数,每一层算出来的中间结果也是数,而每一个数都有自己的合理范围。这些数要是没照看好,麻烦往往很快就到,而且多半是大麻烦。
我先讲一种最常见的出问题的情况。某一层的输出要是偏大,下一层的激活函数很可能直接走进饱和区。什么叫饱和区呢,拿Sigmoid来举例,当输入很大或者很小的时候,它的输出几乎贴着 或者 不再动了,对输入的变化几乎没了反应。我们把Sigmoid的输出记作 ,其中 表示这一层算出来的输入, 就是Sigmoid这个函数本身。在这片区域里, 几乎不再变化,它对输入 的导数 也就跟着接近零( 就是 对 求导得到的结果)。导数一旦接近零,梯度信号就传不动了,这一层的参数几乎收不到学习信号。这情形有点像一场讲座听到后半段,人已经疲了,台上讲什么都没什么反应。反过来也一样,某一层的输出要是偏小,信号又微弱得快要消失,传不了几层就没影了。
所以这里就有两件特别关键的事,一件叫初始化,一件叫归一化。初始化说的是训练开始之前,要先给参数挑一组合适的初始值,可不能一上来就把数值弄得离谱,否则后面整张网都得跟着出问题。归一化说的是在网络中间的若干位置,主动把中间结果重新调整到一个合适的范围里,让每一层接收到的东西都不会太离谱。这两件事的目标是一致的,就是让信息在整个网络里始终在一个健康的尺度上流动,既不至于某一层突然冲得过高,也不至于某一层一路萎缩下去。
2.梯度在深层网络里能不能顺利传回去
数值尺度讲完了,我们再说说梯度。还记得梯度下降吧,我们靠梯度告诉每一层的参数该往哪个方向调。可深度学习之所以叫深度,就在于它层数多,梯度要从最后的损失那一头,一路传回到最早的浅层,中间得过几十层甚至上百层。
这件事靠链式法则完成。链式法则是微积分里的老朋友了,说穿了就是一层套一层地乘导数。每一层算梯度时,都要把后面传过来的梯度,再乘上本层的一个倍数,一层接一层往前传。这个倍数通常是激活函数的导数,也可能是权重矩阵的某种乘积。设网络一共有 层, 表示网络的总层数,第 层的这个倍数记作 , 是层的编号,那么最浅层最终收到的梯度,大约就是这些倍数全部乘在一起的结果:
这里 是连乘符号,意思是从第 层到第 层,每一层的 全部相乘。
这种连乘恰恰是最棘手的地方。要是每一层的 都稍稍小于 ,比方说都是 ,那连乘50层,,梯度直接缩到原来的千分之五,浅层参数几乎收不到什么学习信号,这叫梯度消失。要是每一层的 都稍稍大于 ,比方说都是 ,那 ,层数再多一些梯度就会膨胀到溢出,这叫梯度爆炸。打个比方,梯度消失有点像一场长途接力,每一棒都稍微掉一点速,传到最后几棒已经跑不动了。梯度爆炸则像每一棒都使猛了劲,到后面节奏彻底失控。两种情况都够难受,模型都学不进去。
前面几章我们提到的残差连接(ResNet里那种跳着接的做法)、合理的参数初始化、归一化层,做的都是同一件事,就是想办法让这条乘法链条尽量平稳,保证梯度能从损失一路顺利传回到最早的那一层,既不会在中途传没了,也不会传着传着膨胀过头。
3.计算和存储:训练比预测费得多
数值和梯度讲完了,我们再看看硬件这边到底在忙什么。深度学习的计算主力,其实就那么几种密集运算,包括矩阵乘法、卷积和注意力运算。GPU(显卡)特别擅长同时处理海量的、彼此互不干扰的同一种运算,所以这些密集运算交给GPU来跑最为合适。CPU则更擅长处理复杂的控制逻辑和各种数据准备工作,两者分工配合,才能把训练这一大摊子事跑顺。我之前翻过一本讲计算机体系结构的书,里头把GPU比作一支庞大的合唱团,每个人只唱同一个简单的音,但成千上万人一起开口,气势就出来了,这个比喻我觉得格外贴切。
训练和预测对存储的要求差得很远。做一次预测,只要把输入从前往后算到输出就行,中间结果用完就能丢。可训练还要做反向传播,这就要求把前向过程中每一层的中间结果都留着,等反向算梯度时再拿出来用。所以训练用的显存,通常比单次预测大得多,这也是为什么训练大模型那么吃显卡、那么吃显存。
这里还要再提一个特别关键的设置,叫批量大小(batch size)。批量大小决定了一次塞进去多少个样本一起算。批量越大,GPU的并行能力越吃得满,吞吐量(单位时间能处理多少样本)上去了,训练的整体速度也跟着快起来。可代价是占用的显存也水涨船高,因为每个样本都得各自存一份中间结果。所以衡量一个训练设置到底好不好,不能只盯一个数,要同时看吞吐量、延迟(一次前向加反向要多久)、显存占用和最终结果精度这四样,一起权衡。只看一个数很容易踩坑,可能速度是上去了,显存却当场撑爆。
4.精度和效率:用更少的比特也能跑
最后我们聊一个比较新的话题,精度。
数值可以用不同的精度来存,精度的高低体现在每个数用多少个比特来表示。常见的有fp32(每个数用32个比特)、fp16(每个数用16个比特),还有更新一些的bf16之类。低精度(每个数用更少的比特来表示)能省存储、算得也更快,这是它最讨喜的地方。代价是它能表示的数值范围变小了,数值误差也得想办法控制,否则训练过程容易抖来抖去稳不住,损失曲线一跳一跳的。
现在大家最常用的做法是混合精度(mixed precision),让不同的运算各用合适的精度,需要精确的地方留着高精度,能省的地方就压成低精度。这么一搭配,既提速又省显存,同时还要设法保证训练过程里梯度和损失依然稳得住。说到底,训练这件事就是精度、速度和显存这三者之间的反复取舍,要根据自己的显卡和具体任务,找到最划算的那个平衡点,并没有一套万能配方能通吃所有场景。
记得有一次,小张兴冲冲地拿来一份号称提速三倍的混合精度配置,结果一跑损失直接变成了NaN,排查半天才发现是有一段对数值范围敏感的运算没留高精度,低精度下直接溢出了。这件事说明,性能调优远不止照抄一份配置那么简单,得真正理解每一个数在网络里是怎么流动的,才能既跑得快,又跑得稳。
练习
Q1. 深层网络里梯度消失和梯度爆炸是怎么产生的?
梯度反传靠链式法则一层套一层地乘导数,每层乘一个倍数 ,最浅层收到的梯度大约是 。要是每层 都稍稍小于 1(比如都是 0.9),连乘几十层梯度就被压扁到几乎没了,浅层参数收不到学习信号,这是梯度消失;要是每层 都稍稍大于 1(比如都是 1.1),连乘起来梯度就膨胀溢出,这是梯度爆炸。残差连接、合理初始化、归一化层就是为了让这条乘法链条尽量平稳。
Q2. 手算一下:每层的倍数都是 ,连乘 50 层,梯度缩到原来的多少?每层都是 呢?
,梯度缩到原来的千分之五,浅层几乎收不到信号,典型的梯度消失。,梯度膨胀到上百倍,层数再多些就会溢出,这就是梯度爆炸。可见连乘几万倍地放大或缩小,对深层网络是致命的。
Q3. 为什么训练用的显存比单次预测大得多?小张那套号称提速三倍的混合精度配置为什么跑出 NaN?
做预测只要前向算到输出,中间结果用完就丢;可训练还要反向传播算梯度,得把前向每一层的中间结果都留着等反向用时取出来,所以训练显存大得多。小张那份混合精度配置里,有一段对数值范围敏感的运算没留高精度,低精度下数值范围小,直接溢出,损失就变成了 NaN。混合精度要的是需要精确的地方留高精度、能省的地方压低精度,照抄配置而不理解数值流动,很容易踩坑。