大模型千卡训练-经验指北
作者:你的真实姓名
知乎:https://www.zhihu.com/question/650979052/answer/3501160453
最近在知乎上看到一个回答,把千卡训练的难度吹得天花乱坠。但真用过了千卡的人都会知道,其实就那么几个关键点。所以,今天专门写一篇,把这事儿掰扯清楚。

全文分三部分:先聊聊千卡训练到底难在哪,什么情况下才应该上千卡;然后重点讲,真上了千卡怎么让它跑起来,并且跑出近乎线性的性能提升;最后,再展开说说那些至今仍悬而未决的问题——至少对开源社区来说是这样。
为什么千卡训练是困难的?
千卡训练和八卡训练的本质区别,就是显卡多了100多倍。
这直接带来了两个问题:
- 通信时间暴增
- 故障概率飙升
这两点其实很好理解。
先说通信。PyTorch支持的通信后端有NCCL、Gloo、MPI三个——请务必用NCCL。它的AllReduce操作会根据硬件走Ring或者Tree。Ring的时间复杂度是O(n),Tree是O(log n)。但不管哪种,128个节点理论上就比单节点慢至少7倍,实践中跨节点通讯的延迟更是远超单机。
再说故障。一个节点出问题的概率如果是p,那128个节点至少一个出问题的概率就是1 - (1-p)^128。简单算一下:如果单个操作在单机上出错的概率是1%,那么在128节点上出错的概率就高达72.37%。
而且规模一大,很多小问题都会被放大到难以忍受。比如数据增强多花了0.1秒,一亿条数据下来就是278个小时——当然这只是随口举的例子,实际有各种机制兜底,不会有这么大影响。
所以,钱多烧手绝不是用千卡的理由。闲得蛋疼可能算一个——但你得多蛋疼才能想出来这么折磨自己的主意?
千卡训练真正解决的是大模型&大数据问题。如果训练时间没超过8192 GPU日,那你绝对不需要一千张卡。
看到这里,绝大多数人已经可以关掉这篇文章了。除非你的模型和数据都以B(十亿)为单位。当然,如果你正在厕所里手机没电,想看点东西解闷儿——尽管我怀疑会不会有人真把它打出来——那可以继续往下看。
如何使用一千张卡训练?
如何提高计算效率?
这事儿得具体问题具体分析。因为通信、计算速度受硬件影响更大,而每个集群的硬件拓扑都不一样。同样是A100集群,我全是DGX节点,每张A100都是SXM接口,配一块专属IB网卡;你一个小破普惠服务器插8张PCI-E A100,IB卡一个节点只给一张。那咱们遇到的问题,就完全不是一个问题。
所以,要讨论怎么提高效率、减少耗时,首先要搞清楚训练耗在哪儿。一个训练步的耗时来自哪里?需要牢记:没有profile的优化,都是瞎忙活。
你可能会说:forward、backward、sync。很好,这说明你了解PyTorch的基本流程。但现实要复杂得多:
- dataset读取数据,构建输出
- dataloader collate数据,进行预处理
- 模型forward计算输出
- loss compute
- 模型backward计算梯度
- 模型sync梯度
- 优化器step更新权重
- 打印log
当然还可以无限细分下去,但这些基本够用了。需要注意的是,除了4-7是真正的耗时,其他都需要通过异步操作覆盖掉。这也是优化的目标。
PyTorch的dataloader、CUDA和分布式里都支持异步执行。前者可以通过设置num_workers和prefetch_count为0来关闭,后两者可以通过cuda.synchronize和dist.barrier手动同步。profile时,需要先测整个step的时长,然后在每次测量前执行手动同步,算出每个部分的耗时。如果前者的总耗时等于后者4-7的耗时之和,那通常不需要优化——但这种情况在千卡操作里几乎不可能发生。
第6步通信往往是个大头。所以,还需要进一步优化通信。
以下内容是对《PyTorch Distributed: Experiences on Accelerating Data Parallel Training》论文的概括,有兴趣的建议通读并背诵全文。
计算-通信重叠
在PyTorch中,梯度通信和反向传播是交叠进行的。每算完一层的梯度,都会立即触发当前层的同步。实现方式是:每个进程完成自己第k层的梯度计算后,触发一个钩子给计数器+1,当计数器达到进程数时,开火通信。很多同学在算梯度时遇到过RuntimeError: Expected to ha ve finished reduction in the prior iteration before starting a new one.,就是因为有些模块没参与loss计算,导致梯度同步卡住。注意,当find_unused_parameters=True时,PyTorch用nn.Module.__init__里定义的子模块的反向顺序作为梯度桶的构建顺序。所以,确保模块定义和调用的顺序一致,对高效训练很重要。
梯度合桶
理论上说,同步越及时,重合度越高,性能越好。但实际上每次通信都有开销。所以,同步不是越多越快越好。PyTorch引入了梯度合桶机制,把多个Tensor装在一个桶里再通信,减少通信次数来降低总耗时。合桶的Buffer Size等参数需要针对硬件和模型调优,才能取得最佳效果。PyTorch的默认参数是从0.x时代传下来的,通常都得调。
梯度累加
当你做完所有操作后,可能会发现:同步时间还是比单节点慢好几倍。这其实是正常情况。实际上,超过256卡的训练想把通信盖住,几乎不可能。你说“老师我看FB论文说他们256卡就是线性提升啊”——那这里就不得不提一个策略:梯度累加。梯度累加会执行k次forward+backward,然后再step。好处很多:一是大模型batch size通常不能开太大,梯度累加可以提升等效batch size;二是累加期间的backward不需要通信梯度,能加快训练速度。
少即是快
Python本身很慢。你说JIT trace+torch.compile有提升,我当然不反对。但对最高效率来说,只有“必须要存在的代码”和“不存在的代码”两种。
HuggingFace的Transformers就是一个反面例子:两个子模块就能写完的TransformerLayer,他们硬是能写出一堆。而且他们还信奉Single Model File Policy……你说你这完全不考虑继承,封装这么多层是要搞鸡毛?正例反而是PyTorch——笑死,我竟然会夸脸书代码写得好。具体来说,就是nn.functional里的各种实现,你会发现它们的第一行往往是handle_torch_func。熟悉Python装饰器的人会问:为什么不用装饰器统一一下?因为装饰器会引入额外的函数调用,而额外的函数调用就是额外的开销。
所以,要确保最高效率,写一个简单的训练代码和模型代码非常重要。毕竟,1%的效率提升,省下的可能是几百个GPU日。
如何平稳训练
这部分只讨论你能控制的问题。
捕捉不致命的异常
故障率高的问题其实很好解决。训练中大部分异常都是非致命的,抓住它们就行。这里推荐我之前写的一个装饰器https://danling.org/utils/decorators/#danling.utils.decorators.catch,它的作用是catch异常,然后调回调函数(默认就是把错误打印到log里)。你只需要用它装饰非fatal的操作就行了。
实际应用中,最常见的问题是什么?存ckpt写满了磁盘——别笑,从商汤到深势再到上海AI Lab,这个问题哪都有。也不知道为啥肯买那么多显卡,但就是不肯多插点儿硬盘。catch住所有保存操作,如果你有闲心,可以在回调里删一下之前的ckpt;没闲心的话……大不了重训一次(逃)。第二常见的问题:存log写满了硬盘。所以所有logging操作也要catch。这就是为什么我习惯用tmux开很长的缓存窗口——总能抢回一些log。
说点正经的:任何联网操作都需要catch,常见的有从ceph读取数据,还有写log到远程(逃)。其他就没什么了。我见过有人尝试恢复OOM,但效果似乎不太好,至少我自己没用过。简单来说,唯一不应捕捉的错误是:集群炸了。
那有的大兄弟就问了:集群没爆炸,但两张卡突然掉了怎么办?这个第三部分再讨论。
过程也很重要
有用过丹灵(http://danling.org)的同学可能比较熟悉。丹灵其他地方都很轻量,唯独实验管理写得很复杂。现代丹灵会创建一个三级实验目录:project/experiment-run/timestamp。其中project是用户给出的,experiment和run分别通过代码版本和配置计算出来,timestamp是运行开始时间。也就是说,如果代码和配置完全一样,丹灵就会认为这是同一个运行。在设置中打开auto_resum,就会自动找最新的检查点(这就是为什么最后一级用时间戳)来加载。其实微软的amlt更好用,它甚至会创建一个代码的diff文件夹,帮你回忆当初改了些啥。
收敛,收敛,收敛
训着训着模型发散,几乎是每个训大模型的人都会碰到的问题。输出和loss只要有nan,果断丢掉。梯度先clip by value再clip by norm,都是常规操作。哦对了,还有初始化……关于大模型收敛性的论文有一堆,此处不再赘述。
比更大,还更大,再更大
弹性训练
实际上,当训练超过2048 GPU日时,整个过程里单个GPU甚至单个节点下线,是再正常不过的事。
PyTorch在1.10就引入了torchelastic弹性训练机制——用过的都在骂娘。等下,让我先骂一遍,呸。好了,咱们继续。
印象中,在微软最后一轮面试被问到:如何设计一个弹性分布式系统?
我的回答很教科书:每k分钟,系统做一次AllReduce统计存活进程数,然后选举出一个主进程。主进程计算好每个进程的rank和local rank,进行broadcast。所有进程每次forward开始时向主进程发送一个心跳包汇报状态。主进程根据心跳包确定这个step参与同步的机器有多少。
但很可惜,2024年了,还是没人去写。
大小梯度同步
我一直认为,梯度同步不应该以GPU/进程为单位。而应该分两种:大同步(节点间同步)和小同步(节点内同步)。小同步可以更高频地进行,大同步则可以更慢地执行。这样不仅能提高实际梯度同步频率、降低总耗时,还能天然地结合小batch和大batch训练的优势——节点内小batch关注个体,节点间大batch关注整体。
-
- 关于宇宙的好的网名有哪些
- 角色扮演 | 1
- 网名