十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

MindSpore大模型训练显存优化与断点续训实践指南

MindSpore大模型训练显存优化与断点续训实践指南 MindSpore 大模型训练跑到一半显存爆掉或者断点续训恢复后发现 loss 对不上这两件事我猜你至少遇到过一件。今天这篇文章想聊聊我在这套框架里做高效显存管理和增量式断点续训的经验。我不会只丢一堆配置项而是把显存到底花在哪、每种优化手段的代价、检查点到底该存哪些东西以及恢复之后怎么验证一步步拆开讲。适合已经在用 MindSpore 训练大模型、想继续提升稳定性和资源利用率的同学参考。1. 大模型训练的显存账先搞清楚钱花在哪1.1 一张卡上的显存被谁吃了大模型训练里的显存消费方主要有四块模型权重、梯度、优化器状态、中间激活值。不少人只算权重比如 7B 的 BF16 模型权重差不多 14GB觉得在 80GB 的卡上跑绰绰有余结果 batch size 开到 2 就 OOM。原因很简单梯度通常和权重同量级也要 14GB如果用的是常规 Adam 优化器还要额外维护一个 FP32 的权重副本、一阶动量和二阶动量算下来每参数要 16 字节左右7B 模型光优化器状态就可能超过 100GB。再加上 Transformer 中间激活值尤其是序列长、层数多的情况下激活值可能比参数本身还占显存。所以单卡做 7B 全量微调基本是不现实的。显存优化的本质就是在这四块之间做取舍。我把它们整理成一张账方便后面讲优化手段时对照显存去向7B 模型典型量级主要降低手段模型权重约 14GBBF16低精度、模型并行、offload梯度约 14GBBF16梯度累积、ZeRO 分布式切分优化器状态约 84GBFP32 的 master weight m v优化器状态切分、CPU offload中间激活取决于 batch、seq_len、层数重计算、减小 batch、梯度累积这里要特别提醒很多人一开始只盯着 batch size以为把 batch 调小就能解决一切。实际上如果用的是全参数微调Adam 状态才是大头。遇到 OOM 先看训练脚本里的优化器配置再看激活值规模别盲目减 batch。1.2 MindSpore 静态图的内存池与显存复用MindSpore 有两种运行模式PYNATIVE_MODE 动态图和 GRAPH_MODE 静态图。小规模调试用 PyNative 确实方便可以像普通 Python 一样逐行打印、断点调试但大模型训练强烈建议切到 Graph。原因在于静态图模式下MindSpore 会把前向、反向整体编译成一张完整计算图。编译器能分析每个张量的生命周期在内存池里做显存复用同样一块显存前向结束后马上可以给反向用避免动态图每步都重新分配、释放造成的内存碎片。我实测过同一个模型PyNative 下 batch size 只能开到 8Graph 下能开到 12差距非常明显。开启方式很简单import mindspore as ms ms.set_context(modems.GRAPH_MODE, device_targetAscend, save_graphsFalse)这里save_graphsFalse很重要否则每跑一步都会把中间计算图 dump 到磁盘等发现的时候磁盘已经满了。另外静态图的编译阶段会有一点额外耗时但对大模型训练的长期收益来说这点编译成本完全值得。1.3 并行方案的内存视图单卡显存放不下时并行是绕不开的。数据并行会把模型和优化器状态完整复制到每张卡上通过通信合并梯度模型并行和张量并行则是把网络切到不同卡单卡显存压力直线下降流水线并行按层分段类似工厂流水线每张卡只负责其中几层。在 MindSpore 里最常见的是数据并行通过并行上下文可以设置from mindspore import context from mindspore.context import ParallelMode context.set_auto_parallel_context( parallel_modeParallelMode.DATA_PARALLEL, gradients_meanTrue )如果想更精细控制可以用半自动并行给关键算子手动指定shard策略。这里不展开太多但要记住显存管理不是某一招单独起作用而是单卡优化手段和集群并行策略的组合。后面讲到的混合精度、梯度累积、重计算都是单卡内先做的优化然后再考虑怎么切到多卡。2. MindSpore 高效显存管理的四种落地手段2.1 混合精度性价比最高的第一步混合精度的核心思想是前向和反向计算用 FP16 或 BF16但优化器保留 FP32 的权重副本。全 FP32 训练 7B光权重就是 28GB换到 BF16 之后权重直接减半。MindSpore 的Model接口可以直接指定混合精度等级from mindspore import Model from mindspore.amp import DynamicLossScaleManager model Model( net, loss_fnloss, optimizeroptimizer, amp_levelO2, loss_scale_managerDynamicLossScaleManager() )amp_level的 O2 一般表示大部分算子走半精度框架会自动插入 loss scaling。在昇腾上我更喜欢用 BF16因为它的动态范围和 FP32 基本一致不容易出现 FP16 那种小梯度下溢的问题。使用 FP16 时loss scaling 是必须开的否则 loss 会在几百步之后突然变成 NaN而且很难排查。实操里有个细节如果模型里有一些自定义算子amp_level可能不生效需要手动把算子加入黑名单或者白名单。验证方法很粗暴开混合精度后用固定种子跑一小段把每一步的 loss 和梯度范数打出来和纯 FP32 对比如果差异不超过一个很小的阈值基本就安全了。2.2 梯度累积batch 大不起来时的折中显存瓶颈很大一块在中间激活而激活大小和 batch size、序列长度成正比。如果单卡 batch 1 都放不下梯度累积是最直接的思路把一个大的 batch 拆成几个 micro batch分别前反向梯度累加后再统一更新权重。需要强调梯度累积不是改变优化器的更新次数而是凑够等效 batch size。例如原计划 batch size 32现在显存只够 batch 8那就拆成 4 个 micro batch每 4 步累加一次梯度后更新。MindSpore 里如果版本支持nn.MicroBatchInterleaved可以直接包装网络import mindspore.nn as nn net nn.MicroBatchInterleaved(backbone, micro_size4)如果不支持或者想手动控制就自己写循环每个 micro batch 前向、反向把梯度累加到一个变量里达到累积步数后再调用优化器更新。梯度累积的代价是训练时间增加因为反向传播仍然要执行只是每次更小。我一般会把 micro batch 调小到原来的 1/4配合重计算使用7B 模型在 80G 卡上能塞下等效 batch 8 左右。2.3 激活重计算用少量计算换回大量显存Transformer 前向过程中会保留每层激活值供反向使用这在长序列场景下特别吃显存。重计算Activation Checkpointing的思路是前向时不保存中间激活反向传播用到哪一层再重新算那一层的激活值。显存峰值能降 30% 到 50%代价是计算量增加训练时间多 20% 到 30%。这个时间换空间的做法不是把所有层都包上就一定好。实际使用中我会选最深的几个 Transformer Block 开启重计算浅层和输出层保持原样。因为重计算不是免费午餐开得越多额外计算越多。在长序列场景下收益最明显短序列反而没必要训练时间变长显存省得有限。MindSpore 对重计算的支持在不同版本里入口不太一样有些版本可以直接对Cell调用重计算相关方法有些需要通过配置把某些节点标记为 recompute。拿到一个新版本先用小模型验证 API再跑到全量训练否则很容易在几百行代码之后发现某个算子不支持重计算那种挫败感我体会过太多次了。2.4 CPU Offload 和优化器状态切片如果混合精度、梯度累积、重计算都上了显存还是不够就要考虑把部分状态搬到 CPU 内存。优化器状态是最适合 offload 的因为它只参与更新阶段不参与前反向的高频计算。把 Adam 的 m/v 放到 CPU显存立刻空出一大块但每次更新都要通过 PCIe 或总线传输训练速度会明显变慢。这种方案适合在一台机器上跑稍大模型做微调查参而不是追求极致吞吐的场景。如果有多卡更推荐 ZeRO 这类优化器状态切片机制把优化器状态按 rank 切分每卡只维护 1/N 份更新时通过集体通信拿到完整状态。显存占用下降通信量增加但通常比 CPU offload 高效。在 MindSpore 中数据并行和半自动并行能承担一部分状态切分逻辑具体 offload 开关不同硬件版本有差异。我建议大规模长期训练优先用分布式状态切分而不是简单 offload单卡实验、快速验证时再用 CPU offload。2.5 显存优化的检查顺序综合来说我落地显存优化的顺序是这样先开混合精度这是收益最大、改动最小的一步。再根据显存余量调整 micro batch用梯度累积凑等效 batch。开激活重计算重点处理 Transformer Block。还是不够再考虑并行切分或 CPU offload。每一步都要同时监控显存峰值和训练吞吐不能只看显存降了多少。之前有个项目把重计算开满显存是降下来了但训练速度慢了近一半最终整体吞吐反而变差。后来改成只重计算一半层显存刚好卡住吞吐也回来了。3. 增量式断点续训从“能保存”到“续得上”3.1 为什么说“重载了权重”不等于“断点续训”很多人的断点续训是这样的训练中断后用load_checkpoint把模型参数加载回来然后接着model.train。这样确实能从保存的时刻继续训练但训练行为往往已经变了。原因是权重恢复了优化器状态不一定恢复。如果用 Adam它的一阶动量、二阶动量直接决定下一步怎么走这些不恢复相当于优化器重新初始化学习率调度器如果从第 0 步重新开始warmup 会再来一遍学习率曲线完全错位数据流如果没有回到中断时的样本位置前面训练过的数据可能再来一遍后面没到过的数据却被跳过。所以增量式断点续训的核心是完整恢复“训练现场”而不是只把网络权重塞回去。3.2 增量式断点需要保存的状态清单我根据自己的项目经验整理了一份检查点状态清单状态是否必须说明模型权重必须网络参数少了它就没法恢复优化器状态必须包括一阶、二阶动量还有优化器内部的 step学习率调度器状态强烈建议当前 epoch、累计 step、warmup 阶段、当前 lrloss scale 状态必须混合精度时FP16 训练时动态 loss scale 需要恢复数据游标必须当前 epoch、当前 batch 偏移保证不重复不遗漏随机源状态尽量Python random、NumPy、框架的 seed、dataset shuffle 状态这里最容易被忽略的是数据游标。很多工程团队恢复权重后训练集从头开始跑虽然模型最终也能收敛但实际训练步数和设计不符学习率又没有对应调整后期会莫名其妙过拟合。3.
返回列表