NLP预训练大模型是近年来AI领域最热门的方向之一。我们团队也在做预训练模型的研发,有一次在训练一个大模型的时候,出了一个严重的故障,导致训练了好几天的模型差点报废。本文是这次故障的完整复盘,包括故障的发生过程、排查过程、根本原因、以及我们总结的经验教训。如果你也在做深度学习或者大模型训练,希望这篇文章能给你一些参考和警示。
一、背景:我们的预训练项目
先说说我们的预训练项目吧。
2020年的时候,NLP预训练大模型非常火,BERT、GPT、T5等模型层出不穷。我们团队也决定做一个自己的预训练模型,用于公司内部的各种NLP任务。
我们的模型规模不算特别大,大概几亿参数,用了几十张GPU卡来训练。训练数据是我们收集的几十G的中文文本,包括新闻、百科、小说、论坛等各种来源。
训练流程是这样的:先对数据做清洗和预处理,然后用分布式训练框架在GPU集群上训练,训练过程中定期保存checkpoint,训练完成之后再做下游任务的微调验证。
这个项目我们已经做了一段时间,前面几次小规模的训练都很顺利。这次我们要训练一个更大的模型,用了更多的数据和更多的GPU卡,预计训练时间是一周左右。
训练开始之后,前几天都很顺利,loss在稳步下降,各项指标都正常。我们都觉得这次训练应该会很顺利,但是没想到,在训练到第五天的时候,出问题了。
二、故障发生:loss突然飙升
那是一个周四的晚上,我本来已经下班了,正在家吃饭。突然收到了监控系统的告警,说训练任务的loss突然飙升了。
我赶紧打开电脑,远程连接到训练集群,查看训练日志。发现从某个step开始,loss突然从正常的2.3左右飙升到了100多,而且还在继续上升。
这是一个非常严重的问题。在深度学习训练中,loss突然飙升通常意味着训练出了严重的问题,如果不及时处理,模型可能会完全废掉,之前几天的训练就白费了。
我赶紧做了以下几件事:
- 暂停了训练任务,防止情况进一步恶化
- 保存了当前的状态和日志,方便后续排查
- 通知了团队的其他成员,大家一起排查问题
那天晚上,我们团队几个人都在线上排查问题,一直搞到凌晨两三点,才找到问题的原因。现在回想起来,还是觉得惊心动魄。
三、排查过程:一步步定位问题
排查的过程非常曲折,我们走了很多弯路。下面说说我们是怎么一步步定位到问题的。
第一步:检查数据
loss突然飙升,首先想到的可能是数据出了问题。比如数据预处理出错了,或者某个batch的数据有异常。
我们检查了训练数据,发现数据都是正常的,没有异常值。而且数据预处理的流程和之前几次训练是一样的,之前都没问题,这次应该也不会有问题。
我们还检查了数据加载的代码,也没有发现问题。数据加载是正常的,每个batch的数据都是符合预期的。
所以数据问题的可能性被排除了。
第二步:检查模型结构
接下来我们怀疑是不是模型结构出了问题。比如某个层的参数初始化不对,或者某个层的计算有问题。
我们检查了模型的代码,和之前的版本对比,发现这次训练我们改了几个地方:
- 增加了模型的层数和隐层维度
- 调整了注意力机制的一些参数
- 用了一个新的优化器
这些改动都有可能导致训练不稳定。我们逐一排查,但是没有发现明显的代码bug。
我们还加载了之前保存的checkpoint,检查模型的参数,发现参数都是正常的,没有出现NaN或者Inf。在loss飙升之前,参数都是正常更新的。
所以模型结构的问题也暂时排除了。
第三步:检查训练超参数
然后我们检查了训练的超参数,比如学习率、batch size、梯度裁剪等。
我们发现,这次训练我们用了一个比较大的学习率,因为模型变大了,我们觉得应该用大一点的学习率才能收敛得快。而且我们用了学习率warmup,理论上应该不会有问题。
但是loss飙升的那个step,正好是学习率warmup结束、开始下降的那个step。这会不会是巧合?
我们仔细分析了一下,觉得学习率的变化可能是一个诱因,但是不应该导致loss突然飙升到100多。因为即使学习率大一点,有梯度裁剪的话,也不会出现这么严重的问题。
所以学习率可能是一个因素,但不是根本原因。
第四步:检查GPU和硬件
接下来我们怀疑是不是硬件出了问题。比如某张GPU卡出了故障,计算出错了。
我们检查了GPU的状态,发现所有GPU卡的温度、功耗、显存都是正常的,没有报错。而且分布式训练的框架也没有报通信错误。
我们还做了一个测试,用同样的代码和数据,在另一组GPU卡上跑了几个step,发现loss也是正常的,没有飙升。这说明硬件应该没问题。
所以硬件问题也排除了。
第五步:检查混合精度训练
这时候我们有点焦头烂额了,常见的原因都排除了,问题到底出在哪里呢?
这时候团队里一个同学提出,会不会是混合精度训练的问题?
我们这次训练用了混合精度训练(FP16),因为大模型训练很吃显存,用混合精度可以省显存,还能加快训练速度。之前的小规模训练也用了混合精度,都没问题。
但是这次模型变大了,batch size也变大了,会不会在某些情况下,混合精度的计算出现了溢出?
我们仔细检查了混合精度的代码,发现了一个问题:我们在计算loss的时候,没有做loss scaling。
混合精度训练中,因为FP16的表示范围比较小,梯度很容易下溢(变成0)。为了解决这个问题,通常需要做loss scaling,就是把loss乘以一个比较大的数,这样梯度也会放大,不会下溢。然后在更新参数的时候,再把梯度除回来。
我们之前的小规模训练,因为梯度比较大,即使不做loss scaling也不会下溢,所以一直没出问题。但是这次模型变大了,batch size也变大了,某些层的梯度变得很小,在FP16下就下溢了,变成了0。
梯度变成0之后,这些层的参数就不更新了。但是其他层还在正常更新,这就导致了模型的不平衡。随着训练的进行,这种不平衡越来越严重,最终导致了数值不稳定,loss突然飙升。
找到问题原因之后,我们都松了一口气。原来是混合精度训练的loss scaling没做好,导致了梯度下溢,最终引发了训练崩溃。
四、问题修复和恢复
找到原因之后,修复就简单了。
我们在训练代码中加上了loss scaling,用的是动态loss scaling,就是训练框架自动调整scaling factor,保证梯度既不会下溢也不会溢出。
然后我们从loss飙升之前的最后一个正常checkpoint恢复训练,用修复后的代码继续训练。恢复之后,loss又回到了正常的水平,稳步下降,没有再出现飙升的情况。
最终,这次训练顺利完成了,虽然中间出了一次故障,但是因为我们及时暂停了训练,并且从之前的checkpoint恢复了,所以只损失了几个小时的训练时间,之前几天的训练成果都保住了。
现在回想起来,如果当时没有及时发现问题,让训练继续跑下去,模型可能就完全废了,那几天的训练就白费了。几十张GPU卡跑几天的成本是很高的,想想都后怕。
五、根本原因分析
故障解决之后,我们做了一次深入的根本原因分析。
直接原因:混合精度训练没有做loss scaling,导致部分层的梯度在FP16下下溢变成0,模型参数更新不平衡,最终导致数值不稳定,loss飙升。
为什么之前没出问题:之前的小规模训练,模型小,batch size小,梯度比较大,即使不做loss scaling也不会下溢。所以这个bug一直潜伏着,没有被发现。
为什么这次出问题了:这次模型变大了,batch size也变大了,某些层的梯度变得很小,在FP16下就下溢了。而且因为用了更大的学习率,参数更新的幅度更大,不平衡的影响被放大了,最终导致了训练崩溃。
为什么没有及时发现:我们的监控系统只监控了loss,没有监控梯度的分布和参数更新的情况。如果我们监控了梯度,就能更早发现某些层的梯度变成0了,在loss飙升之前就发现问题。
测试不充分:我们在大规模训练之前,只做了小规模的测试,没有在大规模的配置下做充分的测试。如果我们在正式训练之前,用大规模的配置跑几个小时,可能就能发现这个问题。
六、我们采取的改进措施
这次故障给我们敲响了警钟。我们采取了一系列改进措施,防止类似的问题再次发生。
1. 完善混合精度训练的代码
我们把混合精度训练的代码做了完善,加上了动态loss scaling,并且做了充分的测试。现在不管模型大小、batch size大小,混合精度训练都是稳定的。
我们还封装了一个通用的混合精度训练工具,团队所有项目都用这个工具,避免每个人自己写的时候出错。
2. 加强监控
我们完善了训练监控系统,除了监控loss,还监控以下指标:
- 每层的梯度分布,包括梯度的均值、方差、最大值、最小值
- 每层的参数更新幅度
- 梯度中0的比例
- 混合精度的scaling factor
- 是否出现NaN或者Inf
这些指标一旦出现异常,就会触发告警,让我们能在问题恶化之前就发现和处理。
3. 充分的小规模测试
我们制定了规范,任何大规模训练之前,必须先做小规模的测试,验证代码的正确性和稳定性。小规模测试通过之后,再逐步扩大规模,不能一上来就用大规模配置训练。
而且每次修改了训练代码或者模型结构之后,都要重新做测试,不能想当然地觉得没问题。
4. 更频繁的checkpoint
我们把checkpoint的保存频率提高了,从原来的每几个小时保存一次,改成每小时保存一次。这样即使出了问题,也只会损失一个小时的训练时间,不会损失太多。
而且我们会保留最近几个checkpoint,不会只保留最新的一个,防止最新的checkpoint已经被污染了。
5. 训练前的检查清单
我们整理了一个训练前的检查清单,每次开始大规模训练之前,都要按照清单逐项检查:
- 数据是否正确
- 模型代码是否有bug
- 超参数是否合理
- 混合精度是否配置正确
- 分布式训练是否正常
- 监控是否配置好
- checkpoint保存是否正常
- 是否有回滚方案
检查清单大大减少了因为疏忽导致的故障。
七、大模型训练的其他坑
除了这次遇到的混合精度问题,我们在大模型训练中还遇到过其他一些坑,也分享给大家。
坑1:梯度爆炸和梯度消失
大模型训练中,梯度爆炸和梯度消失是常见问题。一定要做好梯度裁剪,合理设置学习率,用合适的参数初始化方法。
坑2:学习率调度不当
学习率太大容易导致训练不稳定,太小又收敛太慢。一定要用warmup,并且根据模型大小和batch size调整学习率。
坑3:数据质量问题
训练数据中的脏数据、重复数据、异常数据都会影响训练效果。一定要做好数据清洗,并且监控数据的质量。
坑4:分布式训练的通信问题
分布式训练中,多卡之间的通信很容易出问题。要做好通信的容错和重试,并且监控通信的性能和错误率。
坑5:显存不足
大模型训练很容易显存不足。可以用混合精度、梯度累积、模型并行、激活检查点等技术来节省显存。
坑6:训练时间太长
大模型训练动辄几天几周,中间很容易出问题。一定要做好checkpoint和容错,并且有完善的监控,出了问题能及时发现和恢复。
八、写在最后
这次NLP预训练故障,虽然过程惊心动魄,但是也让我们学到了很多。大模型训练是一个复杂的系统工程,任何一个小细节出问题,都可能导致严重的后果。
我们总结了几条经验:
- 不要想当然,任何改动都要充分测试
- 完善的监控是及时发现问题的关键
- 做好checkpoint和容错,出了问题能快速恢复
- 细节决定成败,混合精度、学习率这些小地方都不能忽视
- 团队协作很重要,遇到问题大家一起排查,效率更高
深度学习和大模型训练是一个快速发展的领域,新技术、新方法层出不穷。但是不管技术怎么发展,严谨的态度和完善的工程实践都是必不可少的。
希望我们的这次故障复盘能给正在做深度学习或者大模型训练的朋友一些参考和警示。如果你也遇到过类似的问题,欢迎在评论区交流讨论。
最后用一句话结束本文:"训练大模型就像走钢丝,一步走错,满盘皆输。"愿每一个AI工程师都能顺利训练出自己的模型,少踩坑,多出成果。
评论(0)
暂无评论,快来抢沙发~
评论功能仅对会员开放,请先登录
登录