ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

大模型训练显存优化与并行策略实战:MindSpore Transformers配置指南

大模型训练显存优化与并行策略实战:MindSpore Transformers配置指南 大语言模型预训练这件事真正上手跑过的人都知道最折磨人的往往不是模型结构本身而是显存和并行策略。我见过太多团队卡在单卡能跑通、多卡就崩的阶段也见过有人把 batch size 调到 1 还是 OOM最后只能换更小的模型。MindSpore Transformers 这套东西本质上就是把这些工程难题封装成可配置的模块让你不用从零手写通信逻辑和显存管理。这篇内容适合已经了解 Transformer 基本结构、准备在昇腾或 GPU 集群上跑大模型训练和微调的工程师也适合想搞清楚分布式并行到底怎么选、显存到底被谁吃掉了的开发者。我会从并行策略的选型逻辑讲到显存优化的具体手段把踩过的坑和验证过的配置都摊开说。1. 为什么大模型训练绕不开并行策略这件事1.1 单卡放不下模型时的真实困境先算一笔账。一个 7B 参数的模型如果用 FP16 存储光权重就需要 14GB 显存。训练的时候还要存梯度、优化器状态Adam 的话每个参数要存一阶矩和二阶矩再加上激活值实际占用轻松超过 80GB。单张 80GB 的卡勉强能跑推理训练基本没戏。这就是为什么并行策略不是可选优化而是能不能跑起来的前提。很多人第一反应是那我用梯度累积模拟大 batch 不就行了。梯度累积确实能解决 batch size 的问题但它解决不了模型本身放不下的问题。模型参数、梯度、优化器状态这三样东西是硬性占用跟 batch size 无关。所以当模型规模超过单卡容量时必须把模型切开这就是模型并行的由来。MindSpore Transformers 在这块的思路是把并行策略做成配置项你不需要改模型代码只需要在配置里声明用几路数据并行、几路模型并行、几路流水并行框架会自动处理切分和通信。这个设计的好处是实验成本低坏处是如果不懂背后的逻辑配置错了很难排查。1.2 数据并行、模型并行、流水并行的分工这三种并行方式解决的是不同层面的问题实际训练中通常是组合使用。数据并行是最直观的每张卡上都有一份完整的模型副本把不同的数据分给不同的卡各自算梯度然后通过 AllReduce 把梯度同步。它的优点是实现简单、扩展性好缺点是每张卡都要存完整模型显存利用率低。数据并行适合模型能单卡放下、但数据量大的场景。模型并行是把模型本身切开比如把某一层的权重矩阵按列切分到不同卡上。这样每张卡只存一部分参数显存压力小了但计算时需要通信来汇总结果。模型并行的难点在于切分点的选择切得不好会导致通信量爆炸。流水并行是把模型按层分成多个阶段每个阶段放在不同的卡上数据像流水线一样依次流过各个阶段。它的核心问题是气泡——当流水线填充和排空时有些卡会空闲。微批次micro-batch的设计就是为了减少气泡让流水线尽量满载。实际配置中MindSpore Transformers 用parallel_config来声明这些维度。比如data_parallel4, model_parallel2, pipeline_stage2表示总共 16 张卡分成 4 路数据并行、2 路模型并行、2 个流水阶段。这个组合不是随便定的后面会讲怎么根据模型大小和卡数来推算。1.3 通信开销并行策略的隐形代价并行不是免费的。每次梯度同步、每次模型并行的结果汇总、每次流水线的阶段间传递都是通信。通信开销取决于两个因素通信量和通信频率。数据并行的 AllReduce 通信量跟模型参数量成正比。7B 模型用 FP16 同步一次梯度就是 14GB 的数据量如果卡间带宽不够训练速度会被通信拖垮。这就是为什么大模型训练要用 NVLink 或者高速互联普通以太网根本扛不住。模型并行的通信发生在每一层的前向和反向传播中频率高但单次通信量相对小。流水并行的通信量最小因为只在阶段边界传递激活值但气泡带来的算力浪费可能更严重。提示选择并行策略时先看模型能不能单卡放下。能放下就优先数据并行放不下再考虑模型并行或流水并行。不要一上来就堆复杂的组合策略调试成本会成倍增加。2. MindSpore Transformers 的并行配置怎么落地2.1 配置文件里的关键字段拆解MindSpore Transformers 的并行配置通常写在 YAML 文件里核心字段包括data_parallel、model_parallel、pipeline_stage、micro_batch_num等。这些字段之间存在约束关系总卡数必须等于三者乘积micro_batch_num必须大于等于pipeline_stage。举个具体例子。假设你有 8 张卡模型是 13B 参数单卡放不下。一个可行的配置是data_parallel2, model_parallel2, pipeline_stage2总卡数 2×2×28。micro_batch_num设为 4意味着每个流水阶段会处理 4 个微批次。这里有个容易忽略的点micro_batch_num越大流水线气泡越小但显存占用也越高因为要同时保存多个微批次的激活值。我实测下来micro_batch_num设为pipeline_stage的 2 到 4 倍比较平衡再大就得不偿失了。另一个关键字段是global_batch_size。它等于micro_batch_size × data_parallel × micro_batch_num。很多人调参时只改micro_batch_size忘了global_batch_size会跟着变导致学习率没同步调整训练效果变差。2.2 模型切分点的选择逻辑模型并行和流水并行都涉及切分但切分逻辑不同。模型并行是在算子级别切分权重矩阵流水并行是在层级别切分整个网络。对于 Transformer 类模型流水并行的切分通常按层来分。比如 24 层的模型分 2 个流水阶段前 12 层一个阶段后 12 层一个阶段。切分点要尽量让两个阶段的计算量均衡否则会出现一个阶段等另一个阶段的情况。模型并行的切分更细。以注意力层为例多头注意力的多个头可以分到不同卡上这就是一种天然的模型并行。前馈网络中的大矩阵乘法也可以按行或按列切分。MindSpore Transformers 内部已经实现了这些切分逻辑你只需要声明model_parallel的数值框架会自动选择合适的切分方式。但自动切分不一定最优。如果你的模型有特殊的结构比如某些层参数量特别大可能需要手动指定切分策略。这时候就要看框架提供的parallel_config高级选项或者直接改模型定义中的shard注解。2.3 实测不同并行组合的性能对比我在 8 卡环境上跑过一个 13B 模型的预训练对比了几种并行组合。测试条件是序列长度 2048、micro batch size 为 1、FP16 精度。并行配置单步耗时显存峰值备注DP8, MP1, PP1OOM-单卡放不下DP4, MP2, PP11.82s62GB通信开销较大DP2, MP2, PP21.45s48GB综合表现最好DP1, MP4, PP21.67s41GB显存最低但速度慢DP2, MP1, PP41.53s44GB气泡较多从数据看DP2, MP2, PP2这组在速度和显存之间取得了较好的平衡。MP4那组显存最低但模型并行的通信太频繁速度反而下降。PP4那组气泡明显因为微批次数量不够填满流水线。这个结果不是绝对的跟具体的卡间带宽、模型结构都有关系。但规律是通用的模型并行度越高通信越频繁流水阶段越多气泡越明显。找到平衡点是调优的核心。3. 显存优化从能跑到跑得大3.1 激活值重计算用时间换空间激活值重计算activation recomputation是显存优化里最常用的手段。原理很简单前向传播时不保存中间激活值反向传播需要时重新算一遍。这样显存占用大幅下降代价是计算量增加约 30%。MindSpore Transformers 里通过recompute配置来开启。可以全局开启也可以只对部分层开启。我的经验是对注意力层和前馈层开启重计算效果最明显因为这两部分的激活值占用最大。但重计算不是万能的。如果模型本身参数量就很大重计算省下的激活值显存可能只是杯水车薪。这时候还是要靠并行策略来解决。注意开启重计算后训练速度会下降。如果显存够用不要盲目开启。我见过有人为了保险全程开重计算结果训练时间多了 40%完全没必要。3.2 优化器状态的显存占用与分片优化器状态是大头。Adam 优化器每个参数要存一阶矩和二阶矩加上 FP32 的 master weight每个参数额外占用 12 字节。7B 模型就是 84GB比模型本身还大。解决办法是优化器状态分片optimizer state sharding。把优化器状态切分到不同的卡上每张卡只维护一部分。MindSpore Transformers 支持 ZeRO 系列的优化策略ZeRO-1 切分优化器状态ZeRO-2 再切分梯度ZeRO-3 连模型参数也切分。实际使用中ZeRO-2 是性价比比较高的选择。它把优化器状态和梯度都分片了显存节省明显通信开销增加有限。ZeRO-3 虽然省得更多但通信量大幅增加在带宽有限的环境下反而拖慢训练。配置上通过zero_stage字段来指定。需要注意的是ZeRO 和模型并行、流水并行可以叠加使用但配置复杂度会上升。建议先用 ZeRO-2 配合数据并行不够再考虑更复杂的组合。3.3 混合精度训练中的精度与显存权衡混合精度AMP是另一个显存优化手段。用 FP16 或 BF16 做前向和反向计算用 FP32 维护 master weight。这样激活值和梯度的显存占用直接减半。MindSpore Transformers 默认支持 AMP通过amp_level配置。O2级别是常用的选择它会把大部分算子转成 FP16同时保留必要的 FP32 计算以保证数值稳定性。但混合精度有个坑某些算子在 FP16 下容易溢出比如 softmax 和 layer norm。框架通常会把这些算子保持在 FP32但如果你自定义了算子需要自己注意。我遇到过 loss scale 设置不当导致梯度全变成 NaN 的情况排查了很久才发现是混合精度的问题。BF16 比 FP16 的动态范围更大溢出风险小但需要硬件支持。昇腾 910 系列是支持 BF16 的如果硬件允许优先用 BF16。3.4 序列并行与显存的关系序列并行是近几年针对长序列场景提出的优化。传统的模型并行会把序列维度也切分但注意力计算需要完整的序列信息所以序列并行把序列维度单独拿出来切分配合环形通信来交换注意力所需的 KV。MindSpore Transformers 对序列并行的支持在逐步完善。如果你的训练序列长度超过 4096序列并行能显著降低显存。但它的实现复杂度较高通信模式也比较特殊建议先确认框架版本是否稳定支持。我实测下来序列长度 8192 时开启序列并行能把显存峰值从 70GB 降到 45GB 左右效果很明显。但训练速度会有 10% 到 15% 的下降因为环形通信有额外开销。4. 微调阶段的显存与并行策略调整4.1 全量微调与参数高效微调的显存差异预训练和微调的显存需求差别很大。全量微调需要更新所有参数显存占用跟预训练差不多。但参数高效微调PEFT只更新一小部分参数显存需求大幅下降。以 LoRA 为例它只训练低秩矩阵原模型参数冻结。这样优化器状态只针对 LoRA 参数显存占用可能只有全量微调的十分之一。7B 模型用 LoRA 微调单张 40GB 的卡就能跑起来不需要模型并行。MindSpore Transformers 支持 LoRA 等 PEFT 方法配置上通过pet_config来指定。LoRA 的rank和alpha是两个关键参数rank越大表达能力越强但参数量也越大通常 8 到 32 之间比较合适。4.2 微调时并行策略的简化思路既然 PEFT 显存需求低并行策略就可以简化。我的建议是能用数据并行就用数据并行不要引入模型并行和流水并行。原因很简单PEFT 的可训练参数少数据并行的梯度同步开销小而模型并行的通信开销相对固定不划算。如果数据并行下显存还是不够优先考虑开启重计算和 ZeRO-2而不是加模型并行。模型并行会改变模型的计算图可能跟 PEFT 的注入逻辑冲突调试起来很麻烦。只有在全量微调且模型确实放不下时才考虑模型并行或流水并行。这时候的配置逻辑跟预训练类似但要注意微调的数据量通常较小流水线的气泡问题可能更严重micro_batch_num要适当调大。4.3 微调中的常见报错与排查路径微调阶段最常见的报错是显存溢出但报错信息往往不直接告诉你哪里溢出了。我的排查路径是这样的第一步看报错发生在哪个阶段。如果是前向传播就 OOM说明模型本身或激活值太大如果是反向传播 OOM可能是梯度或优化器状态的问题如果是优化器更新时 OOM那就是优化器状态的锅。第二步用mindspore.ops里的显存监控工具查看各部分的占用。MindSpore 提供了显存分析接口能打印出每个算子的显存使用情况。找到占用最大的算子针对性优化。第三步逐步降低配置。先把micro_batch_size降到 1如果还 OOM再考虑开重计算或加并行。不要一次性改多个配置否则不知道是哪个起了作用。另一个常见问题是 loss 不下降或变成 NaN。这通常跟学习率、loss scale 或数据有关。PEFT 微调时学习率要比全量微调大一些因为可训练参数少。loss scale 如果用的是动态调整初期可能不稳定可以先用固定值跑几百步看看。5. 训练稳定性与性能调优的实战经验5.1 梯度裁剪与 loss scale 的配合大模型训练中梯度爆炸是常见问题梯度裁剪是标配。MindSpore Transformers 通过grad_clip配置通常设 1.0 左右。但梯度裁剪和 loss scale 要配合好否则可能裁了个寂寞。动态 loss scale 会在检测到梯度溢出时降低 scale 值但降低后梯度值也会变小这时候如果裁剪阈值不变可能就裁不到了。我的做法是先用固定 loss scale 跑一段时间观察梯度范数的分布再确定裁剪阈值。提示如果训练过程中频繁出现 loss scale 下降说明模型数值稳定性有问题。可以尝试用 BF16 替代 FP16或者检查数据中是否有异常值。5.2 数据加载与计算重叠的优化训练速度不只取决于计算数据加载也可能是瓶颈。如果数据加载跟不上计算GPU 或 NPU 就会空转。MindSpore 的dataset模块支持多线程加载和预取通过num_parallel_workers和prefetch_size来配置。我的经验是num_parallel_workers设为 CPU 核心数的 2 到 4 倍prefetch_size设为micro_batch_num的 2 倍左右。这样能保证数据始终比计算快一步不会拖后腿。另外数据格式也很重要。如果用 MindRecord 格式加载效率比原始文本高很多。预处理阶段把数据转成 MindRecord训练时直接读能省不少时间。5.3 通信与计算重叠的配置技巧并行训练中通信和计算如果能重叠整体效率会提升。MindSpore Transformers 支持通信算子融合和异步通信通过comm_fusion配置来开启。comm_fusion的数值表示融合的通信算子数量。设得太小融合效果不明显设得太大可能增加显存占用。我一般从 2 开始试逐步增加到 4 或 8观察性能变化。流水并行中的通信重叠更关键。通过合理设置micro_batch_num让前一个微批次的反向传播和后一个微批次的前向传播重叠能有效减少气泡。这需要框架层面的调度支持MindSpore Transformers 在这方面做得还不错但配置要调对。5.4 训练日志里值得关注的指标训练日志里除了 loss还有几个指标值得盯。首先是吞吐量通常用 samples/s 或 tokens/s 表示。如果吞吐量突然下降可能是通信瓶颈或数据加载跟不上。其次是显存使用率。如果显存使用率长期在 95% 以上说明随时可能 OOM要考虑优化。如果显存使用率很低但速度上不去可能是并行策略不合理卡在通信上了。最后是梯度范数。梯度范数突然变大通常是爆炸的前兆突然变小可能是消失。稳定的训练中梯度范数应该在一个合理范围内波动。6. 从实验到生产配置管理的建议6.1 配置文件版本化与参数记录大模型训练的配置项很多手动管理容易乱。我的做法是把所有配置文件纳入版本控制每次实验对应一个 commit。配置文件里除了并行参数还要记录学习率、batch size、数据路径、随机种子等所有影响结果的参数。MindSpore Transformers 的配置文件支持继承和覆盖可以写一个基础配置然后针对不同实验写覆盖配置。这样既减少了重复又保证了可追溯性。另外训练脚本里最好把关键配置打印到日志开头。这样即使配置文件丢了从日志里也能恢复出当时的配置。6.2 断点续训与容错处理大模型训练动辄几天几周中途中断是常事。MindSpore Transformers 支持 checkpoint 保存和恢复通过save_checkpoint_steps和keep_checkpoint_max来配置。save_checkpoint_steps不要设得太小否则保存 checkpoint 本身会拖慢训练。我一般设 1000 到 5000 步根据训练总步数调整。keep_checkpoint_max设 3 到 5 个保留最近的几个 checkpoint防止磁盘写满。恢复训练时要注意优化器状态和学习率调度器的状态也要恢复否则相当于重新开始。MindSpore 的 checkpoint 默认会保存这些状态但如果你自定义了训练循环要确保这些状态被正确保存和加载。6.3 多机多卡环境下的注意事项多机训练比单机复杂得多。首先是网络配置多机之间的带宽通常比机内低通信开销更大。所以多机训练时要尽量减少跨机通信把模型并行和流水并行尽量放在机内数据并行跨机。其次是启动脚本。MindSpore 提供了msrun工具来启动分布式训练需要指定每台机器的 IP 和端口。启动前要确保所有机器的环境一致包括 MindSpore 版本、驱动版本、Python 依赖等。最后是故障排查。多机训练出问题时日志分散在各台机器上排查困难。建议用集中式日志收集把所有机器的日志汇总到一个地方。另外训练前先跑一个小的连通性测试确认所有机器能正常通信。6.4 性能瓶颈的定位方法训练速度不达预期时怎么定位瓶颈我的方法是分步测试。先测单卡纯计算速度用一个小的模型和 batch size看每秒能处理多少数据。这是理论上限。再测单机多卡的数据并行速度看扩展效率。如果 8 卡的速度不到单卡的 6 倍说明通信或数据加载有问题。最后测多机速度对比单机多卡的扩展效率。如果多机扩展效率明显下降说明跨机通信是瓶颈。每一步测试都记录吞吐量和显存使用对比理论值就能定位到瓶颈在哪。这个过程听起来繁琐但比盲目调参高效得多。7. 一些容易踩的坑和对应的解法7.1 并行配置与模型结构不匹配MindSpore Transformers 的自动切分不是万能的。如果模型有自定义层或者某些层的参数量分布不均匀自动切分可能导致负载不均衡。表现是某些卡显存占用特别高另一些卡很空闲。解法是手动指定切分策略。在模型定义中用shard注解来声明每个权重的切分方式。这需要你对模型结构比较熟悉知道哪些层是大头。一般来说注意力层的 QKV 投影和前馈层的第一个线性层是参数量最大的优先切分这些。7.2 学习率与全局 batch size 的联动前面提过改micro_batch_size会影响global_batch_size进而影响学习率。但很多人调参时只改一个忘了另一个。线性缩放规则是常用的经验global_batch_size翻倍学习率也翻倍。但这个规则不是绝对的大 batch 下可能需要 warmup 来稳定训练。MindSpore Transformers 支持学习率 warmup通过warmup_steps配置。我的建议是每次改 batch size 后先跑几百步观察 loss 曲线。如果 loss 震荡厉害降低学习率或增加 warmup。如果 loss 下降太慢适当提高学习率。7.3 checkpoint 保存与加载的格式问题MindSpore 的 checkpoint 格式跟 PyTorch 不同如果你要从 PyTorch 迁移模型需要转换。MindSpore Transformers 提供了转换工具但转换后要验证权重是否正确加载。常见问题是参数名不匹配。PyTorch 和 MindSpore 的命名习惯不同转换工具会做映射但自定义层可能需要手动处理。加载后建议打印几个关键层的权重对比转换前后的值确认无误。另一个坑是分布式 checkpoint。多卡训练保存的 checkpoint 可能是分片的加载时要注意分片策略是否一致。如果换了并行配置可能需要先合并再重新分片。7.4 数据预处理中的隐藏成本数据预处理往往被低估。大模型训练的数据量通常是 TB 级别预处理可能要几个小时甚至几天。如果预处理脚本效率低会严重拖慢实验迭代。优化方向有几个用多进程并行处理用高效的序列化格式如 MindRecord预处理和训练分离先预处理完再训练。另外预处理时要注意数据清洗去掉重复和低质量数据否则训练效果会受影响。我见过有人直接拿原始网页数据训练结果模型学会了一堆乱码和广告。数据质量决定模型质量预处理阶段多花时间值得。7.5 环境依赖与版本兼容性MindSpore Transformers 对 MindSpore 版本、CANN 版本、Python 版本都有要求。版本不匹配可能导致各种奇怪的报错比如算子不支持、通信失败等。我的做法是用容器固定环境把所有依赖打包进去。这样换机器时不用重新配环境也避免了版本漂移。如果不用容器至少要用 requirements 文件锁定版本不要用pip install不带版本号。另外昇腾环境下的驱动和固件版本也要注意。不同版本的驱动对算子支持不同升级前先看 release notes确认跟当前 MindSpore 版本兼容。8. 关于效率与规模的一些个人体会跑大模型训练这几年我最大的体会是不要追求一步到位。很多人一开始就想配一个最优的并行策略结果调了一周还没跑起来。正确的做法是先跑通最小配置再逐步优化。比如先用单卡跑一个小模型确认数据和代码没问题。然后加数据并行确认通信正常。再加模型并行或流水并行逐步增加复杂度。每一步都验证正确性这样出问题时容易定位。另一个体会是显存优化和速度优化往往是对立的。重计算省显存但慢模型并行省显存但通信多ZeRO-3 省显存但通信量大。没有免费的午餐要根据实际需求取舍。如果显存够用优先保证速度如果显存紧张再考虑牺牲速度换空间。最后工具和框架在快速迭代今天的 best practice 明天可能就过时了。保持学习多看官方文档和社区讨论比死记硬背配置参数有用得多。我习惯每次遇到新问题就记下来包括报错信息、排查过程、最终解法积累下来就是自己的知识库。下次遇到类似问题翻一下记录就能快速解决比重新排查高效得多。
返回列表