深度学习分布式训练:当一张卡不够用的时候
训练深度学习模型,早期很多人是在单张显卡上完成的。数据不算太大,模型也不算太深,跑一晚上往往就有结果。后来模型越来越大,数据越来越多,一张卡的显存和算力很快成为瓶颈。有人换更贵的卡,有人开始把任务拆到多张卡、多台机器上。这就是分布式训练要解决的问题:怎么让多份计算资源一起干活,又尽量少地互相拖后腿。
这篇文章用比较直白的方式,聊聊分布式训练在做什么、常见怎么拆、以及实际用起来会遇到什么。
为什么需要分布式
原因其实很具体。
显存不够是最常见的。大模型的参数、中间激活、优化器状态都要占空间。一张消费级或甚至数据中心级的卡,经常装不下完整的模型和足够大的 batch。把模型切开,或者把数据拆开并行算,就能继续往下做。
时间太长是另一个原因。即便显存勉强够,用很小的 batch 慢慢跑,可能要几周。把数据分到多张卡上同时算,等效增大了吞吐,墙钟时间会明显缩短。对需要反复试超参、赶进度的团队来说,这很实际。
还有规模本身的问题。有些任务数据量极大,单机读写和预处理都吃力;有些模型大到必须跨机器才能放下。分布式不再是“可选加速”,而成了“能不能训起来”的前提。
当然,分布式不是免费午餐。机器多了,通信、同步、故障都会变成新的成本。只有当单卡或单机的极限已经被摸到,再上分布式,通常更划算。
几种主要的拆分思路
分布式训练里,最常听到的是数据并行、模型并行,以及两者的组合。
数据并行相对好理解。每张卡都留一份完整的模型,但吃到的数据不同。各卡算完自己的梯度后,把梯度汇总、平均,再一起更新参数。这样等效于用更大的 batch 在训练,模型副本之间保持一致。实现相对成熟,框架支持也多,是大多数人的第一选择。
它的瓶颈往往在通信。每次更新前都要在卡之间同步梯度,网络慢或卡多的时候,通信时间可能占掉很大比例。batch 越大,计算相对通信的比重通常越高,所以数据并行很吃“算得够多、传得够快”。
模型并行则是把模型本身拆开。有的按层切,前面几层在一批卡上,后面几层在另一批卡上,数据像流水一样往下传,这常叫流水线并行。有的把同一层的大矩阵切开,分到不同卡上算,再拼回来,属于更细的张量并行。模型并行主要解决“模型太大,一张卡放不下”的问题,但拆得越细,卡之间的依赖和通信越复杂,实现和调优都更难。
实际的大模型训练很少只用一种。常见做法是混合:在节点内用张量并行把一层拆开,节点间用流水线并行把不同层分到不同阶段,再在更大范围用数据并行复制多份。这样既能放下模型,又能利用更多数据集中算力。具体怎么切,取决于模型结构、硬件拓扑和团队的工程能力。
还有一些相关思路,比如零冗余优化器状态(把优化器状态和梯度分片存),用来进一步省显存。它们往往和上面的并行策略配合使用,而不是单独存在。
通信与同步在说什么
多卡一起算,就一定要交换信息。数据并行里最关键的是梯度的全局同步,常用 AllReduce 一类集合通信:每张卡贡献自己的梯度,最终每张卡都得到完整的平均梯度。实现可以走各种网络和算法,目标都是降低延迟、提高带宽利用率。
如果机器之间网络很慢,或者跨了多个机架、多个节点,通信就会成为明显瓶颈。这时有人会用梯度压缩、延迟更新、局部SGD 等变通办法,用一点精度或收敛性的代价换通信减少。是否值得,要看任务对精度的敏感程度。
同步方式也有讲究。严格同步要求所有卡都算完并交换完再进入下一步,实现简单、结果可复现性较好,但快的卡要等慢的卡。异步更新让快的卡先走,整体吞吐可能更高,却容易引入梯度陈旧,收敛行为更难分析。生产里多数大规模训练仍偏向同步或受控的同步变体,为的是稳定和可预期。
实际落地时会碰到的问题
理论说起来整齐,真正跑起来会有一堆细节。
负载是否均衡很关键。如果某张卡分到的数据特别难算,或者某段流水线特别慢,其他卡就得空等。数据采样、padding、动态形状都可能造成不均衡,需要在数据管道和调度上花心思。
故障是另一件事。卡越多、机器越多,硬件出问题的概率越高。训练跑到一半若某张卡挂了,最好能从最近的检查点恢复,而不是从头再来。断点续训、弹性资源、任务重试,都是工程上必须考虑的。
软件栈也不省心。驱动、CUDA、通信库、框架版本要匹配;多机时还要配网络、共享存储、时钟同步。环境没理顺之前,大量时间会花在“为什么连不上”“为什么特别慢”上,而不是模型本身。
可复现性在分布式下更难。随机种子、数据顺序、浮点累加顺序都可能让每次结果有细微差别。对科研和需要严格对比的实验,要额外约定规范;对业务迭代,有时接受小范围波动也无妨,但要心里有数。
成本同样现实。多卡多机的费用、电费、人力,都要算进总账。有时优化单机效率、换更合适的模型结构,比盲目堆机器更划算。分布式是工具,不是目的。
什么时候该上,什么时候可以再等等
如果你的模型在单卡上已经能舒服地训练,batch 也够大,收敛时间和资源都可接受,就没必要急着上分布式。先把数据和单机训练打磨好,收益通常更直接。
当出现这些信号时,可以认真考虑:
- 模型或激活已经放不进单卡显存,缩小 batch 或改结构也很难继续;
- 单机训练时间长到无法接受,而你有多卡资源;
- 需要系统性地做大规模预训练或超大实验,单机吞吐明显不够。
上分布式之后,建议仍从小规模验证起:先两张卡、再单机多卡、再多机。每一步确认 loss 曲线、指标和单卡基线大致对齐,再放大。这样出了问题容易定位,是通信、是数据,还是实现细节。
框架选择上,主流深度学习框架都提供了分布式接口,也有专门面向大模型的工具链。对初学者,先熟悉一种数据并行的标准用法,再按需学模型并行和更复杂的切分,路径会平滑一些。文档和社区示例值得仔细跑通,而不是只看概念文章。
写在后面
分布式训练解决的是规模问题:模型太大、数据太多、时间太紧。它把单卡上已经验证过的训练过程,扩展到多卡多机,同时引入通信、同步和工程复杂度。概念上并不神秘,难在细节和稳定性。
对大多数做应用的人来说,不需要一上来就钻进各种并行策略的论文。先理解数据并行在做什么、通信大概开销在哪里、什么时候该加机器什么时候该优化单机,就足够应对很多实际决策。等真正遇到模型放不下或集群规模上来的时候,再深入模型并行和混合策略,也不迟。
硬件在进步,框架在封装,以前很难的事在逐渐变简单。但“把任务合理拆开、让多份资源高效协作”这件事本身不会消失。理解背后的取舍,比记住几个名词更有用。当你的训练任务开始撞上单卡天花板时,回头看看这些基本思路,会更容易判断下一步该往哪走。
补充一点关于“有效 batch size”的直觉。数据并行时,全局 batch 等于单卡 batch 乘以卡数。全局 batch 变大后,往往需要同步调整学习率,否则收敛行为会和单卡时差很多。有人用线性缩放规则,有人用更细致的预热与衰减策略。重要的是:不要假设“卡数乘上去,其他超参完全不动,效果自动一样”。做对比实验时,尽量记录全局 batch、学习率和其他关键超参,避免把并行策略的影响和超参影响混在一起。
另一类常见误解是“机器越多一定越快”。在通信密集或同步等待明显的情况下,加卡带来的加速比会逐渐下降,有时甚至出现负优化。这时需要看性能剖析:计算时间、通信时间、数据加载时间各占多少。若通信占比过高,可考虑增大单卡计算量、优化网络拓扑、或改用通信更少的算法。盲目扩容不如先找到瓶颈。
对于小团队或资源有限的情况,还可以关注单机多卡是否已经够用。很多中等规模的任务,一块机器上的 4 卡或 8 卡数据并行就能带来可观加速,却不必承受多机通信和部署的复杂度。等单机方案也吃紧了,再迈向多机,路径更稳。
最后,文档与复现实验值得单独强调。分布式脚本里多了进程组、设备映射、端口与超时等配置,环境稍有差异就可能跑不起来或结果对不上。把启动方式、依赖版本、关键环境变量和一次成功的运行日志留存好,能减少后续大量重复排障时间。对需要对外交付或长期维护的项目,这点尤其重要。
总的来说,分布式训练是工程和算法的结合体。算法侧要保证并行之后仍然能收敛到合理的模型;工程侧要保证在真实硬件和网络上跑得稳、跑得够快。两者都做好,大规模训练才可持续。若你正在从单卡走向多卡,不妨先把数据并行跑通、把基线对齐,再逐步引入更复杂的切分。步子稳一点,往往比一开始就上最炫的方案更省时间。
- 点赞
- 收藏
- 关注作者
评论(0)