尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

ZeRO技术解析:从优化器状态到参数分片,如何高效训练大模型

ZeRO技术解析:从优化器状态到参数分片,如何高效训练大模型 1. 从单卡到多卡为什么我们需要参数分片如果你在训练一个大型语言模型比如现在动辄几百亿甚至上千亿参数的模型你可能会发现一个尴尬的现实一张顶级的消费级显卡比如RTX 4090其24GB的显存可能连模型本身都装不下更别提训练过程中需要的优化器状态、梯度和激活值了。这就像你想用一辆家用轿车去运输一个集装箱根本无从下手。这就是分布式训练特别是参数分片技术在今天变得至关重要的根本原因。分布式训练的核心思想很简单既然一张卡不够那就用多张卡。但怎么用这里面门道就多了。最朴素的想法是数据并行我把训练数据分成几份每张卡上放一份数据和一份完整的模型副本各自独立计算梯度然后大家把梯度汇总一下平均一下再更新各自卡上的模型。这个方法很直观对于模型参数量小于单卡显存的情况扩展性很好。但它的瓶颈也很明显每张卡上都必须存一份完整的模型参数、优化器状态和梯度。当模型大到单卡存不下完整副本时数据并行就失效了。于是人们开始思考如何把模型本身“拆开”。这就是模型并行。它把模型的不同层或者同一层的不同部分放到不同的卡上。一张卡只负责计算模型的一部分卡与卡之间需要频繁通信传递中间结果激活值。这种方法确实能训练非常大的模型但代价是引入了复杂的工程实现和严重的通信开销计算效率往往不高对网络带宽要求极高。那么有没有一种方法能结合数据并行的简单性和模型并行的内存效率呢这就是ZeROZero Redundancy Optimizer零冗余优化器要解决的问题。它本质上是一种数据并行的增强版或者更准确地说是一种优化器状态、梯度和参数的分片策略。它的核心洞见在于在标准数据并行中冗余的不仅仅是模型参数还有占内存大头的优化器状态和梯度。ZeRO通过精妙的分片消除了这些冗余让多张显卡的内存能像搭积木一样拼起来共同承载一个巨大的模型同时保持了数据并行编程的简洁性。2. ZeRO 的核心思想消除内存冗余的三阶段策略ZeRO 不是一种全新的并行范式而是对现有数据并行范式的深度优化。要理解它我们得先看看标准数据并行训练时每张显卡上到底存了些什么以及它们各自占多少内存。在一个典型的混合精度训练场景中例如使用 Adam 优化器和 FP16/BF16 精度每张卡上的内存占用主要来自四部分模型参数模型本身的权重假设是 FP16 格式。梯度反向传播后计算出的梯度通常也是 FP16 格式。优化器状态这是内存大户。以 Adam 优化器为例它为每个 FP16 参数需要维护一份 FP32 的参数副本用于高精度更新。一份 FP32 的一阶动量。一份 FP32 的二阶动量。 这意味着每个 FP16 参数在优化器状态中对应着3倍的 FP32 数据。由于 FP32 是 4 字节FP16 是 2 字节所以优化器状态的内存开销是原始 FP16 参数的3 * 4 / 2 6倍。激活值前向传播过程中产生的中间结果用于反向传播计算梯度。这部分内存与模型结构、批次大小序列长度强相关在 Transformer 类模型中尤其庞大。在标准数据并行下上述所有数据在每张卡上都有完整的副本这是巨大的内存浪费。ZeRO 通过分片来消除这种冗余它提供了三个渐进式的优化阶段你可以根据你的内存瓶颈和通信开销容忍度来选择启用。2.1 ZeRO 阶段一优化器状态分片这是 ZeRO 的入门级优化也是内存收益最显著的一步。在这个阶段只有优化器状态被分片。假设我们有N张 GPU。原本每张卡上都要保存完整的优化器状态大小是参数的6倍。现在ZeRO 将这完整的优化器状态平均分成N份每张卡只保存其中的一份。同时每张卡仍然保存完整的模型参数和梯度。它是如何工作的前向传播每张卡用自己完整的模型参数副本独立处理一份数据计算损失。反向传播每张卡独立计算相对于自己完整模型参数的梯度。此时每张卡上都有完整的梯度。梯度聚合通过all-reduce操作对所有卡上的完整梯度进行求和与平均使得每张卡都得到一份全局平均后的完整梯度。这一步和标准数据并行一样。参数更新这是关键。每张卡只根据自己负责的那一份优化器状态分片来更新对应分片的那一部分模型参数。例如卡0负责更新参数的第0到第k个它就用全局平均梯度中对应位置的部分以及自己保存的这部分参数的优化器状态一阶动量、二阶动量等来进行 Adam 更新。更新完成后卡0上就有了最新的第0到第k个参数。参数同步由于每张卡只更新了一部分参数为了进行下一次前向传播所有卡必须同步完整的模型参数。这通过一个all-gather操作实现每张卡把自己更新好的那部分参数广播给所有其他卡同时也从其他卡接收它们更新好的参数。最终所有卡都获得了完整的、更新后的模型参数。内存与通信权衡内存节省优化器状态的内存占用减少了N倍。这是巨大的提升因为优化器状态通常是内存瓶颈。通信开销相比标准数据并行只需要一次all-reduce用于梯度阶段一增加了一次all-gather操作用于同步全部参数。通信量从只同步梯度变成了同步梯度同步全部参数。但由于参数通常是 FP16而优化器状态是 FP32节省的内存往往远大于增加的通信成本。2.2 ZeRO 阶段二梯度分片在阶段一的基础上ZeRO 阶段二进一步将梯度也进行分片。此时每张卡上保存完整的模型参数FP16。一份优化器状态分片FP32。一份梯度分片FP16。工作流程变化前向传播不变每张卡用完整参数计算。反向传播每张卡计算完整梯度但计算完成后立即通过一个reduce-scatter操作进行梯度聚合与分片。reduce-scatter可以理解为all-reduce和scatter的结合所有卡将各自计算的完整梯度进行求和reduce然后将求和后的完整梯度按卡数切分每张卡只保留属于自己那一份scatter。这样在反向传播结束时每张卡上持有的不再是完整梯度而是全局平均梯度中属于自己负责的那一份分片。参数更新每张卡使用自己持有的那份梯度分片以及自己负责的优化器状态分片来更新对应部分的模型参数。参数同步和阶段一一样通过all-gather同步全部参数。内存与通信权衡内存节省梯度内存占用也减少了N倍。现在每张卡上只有一份参数是完整的模型参数优化器状态和梯度都是1/N。通信开销梯度聚合从all-reduce同步全部梯度变成了reduce-scatter同步并分发梯度。reduce-scatter的通信量和all-reduce基本属于同一量级但模式不同。参数同步的all-gather通信量不变。总体通信量与阶段一相差不大但内存节省更多。2.3 ZeRO 阶段三参数分片这是 ZeRO 的完全体也是内存效率最高的模式。在此阶段模型参数、梯度和优化器状态全部被分片。每张卡上只保存一份模型参数分片FP16。一份对应的优化器状态分片FP32。一份对应的梯度分片FP16。工作流程的彻底改变前向传播由于每张卡只有一部分参数要进行计算必须先把所需的参数收集齐。在计算某一层时通过all-gather操作所有卡贡献出自己持有的该层参数分片在每张卡上临时重建该层的完整参数。用重建的完整参数执行该层的前向计算。计算完成后立即释放掉这些临时聚合的完整参数以节省内存。只保留该层输出的激活值用于反向传播。反向传播过程是前向传播的镜像。计算某一层的梯度时同样需要先通过all-gather临时重建该层的完整参数。利用该层的激活值和上游传来的梯度计算该层参数的梯度。计算完成后通过reduce-scatter操作将计算出的完整梯度进行聚合和分片每张卡只保留属于自己参数分片的那部分梯度然后释放临时重建的参数。参数更新每张卡使用自己持有的梯度分片和优化器状态分片更新自己持有的那部分模型参数分片。由于参数本来就是分片持有的更新后无需额外的同步操作因为参数的所有权是固定的。内存与通信权衡内存节省达到极致。模型参数内存也减少了N倍。理论上N张卡的总内存可以容纳一个单卡内存N倍大的模型。这使得训练千亿级参数模型成为可能。通信开销急剧增加。在前向和反向传播的每一层都可能需要一次all-gather和一次reduce-scatter。通信频率和通信量都非常高。这极大地考验集群的网络带宽和延迟。计算效率由于频繁的通信计算效率MFU模型浮点运算利用率会受到影响。这本质上是用通信换内存。注意在实际实现中如 DeepSpeed 的 ZeRO-3为了缓解通信压力会采用“流水线”式的策略并非严格在每层计算前都同步而是尽量重叠通信和计算。例如在计算当前层时异步地收集下一层所需的参数。3. ZeRO 的工程实现与关键技巧理解了原理我们来看看如何在实际中使用 ZeRO以及有哪些“坑”需要避开。目前ZeRO 最成熟、应用最广的实现是微软 DeepSpeed 库。PyTorch 本身也通过FullyShardedDataParallel模块提供了类似 ZeRO-3 的功能。3.1 使用 DeepSpeed 配置 ZeRO一个典型的 DeepSpeed 配置文件ds_config.json中关于 ZeRO 的部分可能长这样{ train_batch_size: 32, train_micro_batch_size_per_gpu: 4, gradient_accumulation_steps: 2, zero_optimization: { stage: 3, // 启用 ZeRO 阶段三 offload_optimizer: { device: cpu, // 将优化器状态卸载到 CPU 内存 pin_memory: true }, offload_param: { device: cpu, // 将参数分片卸载到 CPU 内存 pin_memory: true }, overlap_comm: true, // 重叠通信和计算 contiguous_gradients: true, // 梯度连续内存提升通信效率 sub_group_size: 1e9, reduce_bucket_size: 5e8, // 通信桶大小影响通信效率 stage3_prefetch_bucket_size: 5e8, // 阶段三的参数预取桶大小 stage3_param_persistence_threshold: 1e6, // 多大参数以上的层会持久化在GPU上 stage3_max_live_parameters: 1e9, stage3_max_reuse_distance: 1e9, stage3_gather_16bit_weights_on_model_save: true // 保存模型时自动 gather 为16位精度 }, fp16: { enabled: true, loss_scale: 0, loss_scale_window: 1000, initial_scale_power: 16, hysteresis: 2, min_loss_scale: 1 }, gradient_clipping: 1.0 }关键配置解析stage: 核心选项选择 ZeRO 的阶段123。offload_optimizer和offload_param: 这是 ZeRO-Offload 和 ZeRO-Infinity 的核心。当 GPU 显存仍然不足时可以将优化器状态甚至参数分片卸载到 CPU 内存甚至 NVMe 硬盘上用存储空间换显存。pin_memory可以加速 CPU 到 GPU 的数据传输。overlap_comm: 是否让通信all-gather,reduce-scatter与计算重叠。这是提升阶段三效率的关键技术务必开启。reduce_bucket_size和stage3_prefetch_bucket_size: 通信“桶”的大小。梯度/参数不是一个个单独通信的而是打包成一个个“桶”进行集体通信。这个大小需要调优太小会导致通信次数过多太大会增加单次通信的延迟和内存开销。通常可以设置为5e8500M左右开始尝试。stage3_param_persistence_threshold: 一个非常实用的参数。在阶段三参数是分片且按需加载的。但对于一些特别大的参数例如嵌入层频繁地加载/卸载反而效率低。这个参数指定一个阈值大于此大小的参数层会持久化保存在所有 GPU 上不参与分片和卸载牺牲一些内存来换取性能。3.2 实际部署中的经验与避坑指南1. 阶段选择不是越高越好很多新手会盲目开启阶段三认为内存省得最多就是最好的。这不对。你需要做一个权衡如果你的模型刚好能被阶段二或阶段一装下优先用阶段二。因为阶段三的通信开销非常大在带宽有限的集群上其训练速度可能远慢于阶段二。先用nvidia-smi或torch.cuda.memory_allocated()监控你的显存使用确定瓶颈在哪里。只有当你用阶段二显存依然不足OOM时才考虑阶段三。并且开启阶段三后一定要配合overlap_comm和调整bucket_size来优化性能。2. 通信瓶颈与网络要求ZeRO-3 对网络要求极高。理想环境是使用 InfiniBand 或高速 RoCE 网络。在普通的万兆以太网环境下ZeRO-3 的效率可能很低。一个简单的判断方法是观察 GPU 利用率。如果开启 ZeRO-3 后GPU 利用率例如通过nvtop或gpustat查看长期低于 30%很可能就是通信成了瓶颈。此时要么优化网络要么退回到 ZeRO-2要么尝试增大micro_batch_size来让每次通信传输的数据更有价值但会增大激活值内存。3. 激活值内存—— ZeRO 管不了的“漏网之鱼”ZeRO 解决了参数、梯度、优化器状态的内存但激活值内存它没有直接解决。对于 Transformer 模型激活值内存可能和模型参数内存一样大甚至更大。当启用 ZeRO-3 后参数内存不再是问题激活值内存就可能成为新的 OOM 原因。解决方案使用激活检查点。也叫梯度检查点它以前向传播时重计算部分激活值为代价换取极大的激活值内存节省。在 PyTorch 中可以用torch.utils.checkpoint.checkpoint。在 DeepSpeed 中可以在配置文件中配置activation_checkpointing。经验通常会对模型中每个 Transformer 层应用激活检查点。这会带来约 30% 的计算开销但能节省 70%-80% 的激活值内存在内存紧张时非常划算。4. 保存与加载模型的陷阱在 ZeRO-3 下模型参数是分片存储在各张卡上的。直接调用model.state_dict()得到的是分片的状态。如果你要保存一个完整的、可用于推理的模型需要特殊的处理。DeepSpeed在配置中设置“stage3_gather_16bit_weights_on_model_save”: true然后使用deepspeed.save_checkpoint或model_engine.save_checkpoint。它会自动将分片参数聚合起来保存。手动操作如果你需要更灵活的控制可以在保存前使用deepspeed.zero.GatheredParameters上下文管理器来临时聚合参数。但要注意这会瞬间消耗大量内存可能触发 OOM。5. 调试与日志ZeRO 的复杂性使得调试变得困难。建议在初期开启 DeepSpeed 的详细日志。在环境变量中设置DS_LOG_LEVELinfo或debug。在配置文件中设置“wall_clock_breakdown”: true可以输出前向、反向、通信等各阶段的时间占比帮你定位性能瓶颈。使用torch.distributed的barrier和print时要注意只在特定 rank如 rank0上输出避免日志刷屏。4. ZeRO 的演进Offload 与 Infinity即使使用了 ZeRO-3对于万亿参数级别的模型仅靠 GPU 显存即使是多卡也可能不够。为此DeepSpeed 团队进一步提出了 ZeRO-Offload 和 ZeRO-Infinity。ZeRO-Offload它的核心思想是“卸载”。在 ZeRO-2 的基础上它将优化器状态和梯度卸载到 CPU 内存中进行更新。GPU 只负责前向和反向传播中的计算密集型部分。由于 CPU 内存通常远大于 GPU 显存这进一步扩大了可训练模型的规模。但代价是 CPU 和 GPU 之间的数据移动PCIe 带宽成为新的瓶颈。它适用于优化器状态和梯度是主要内存瓶颈且模型计算量不是极端巨大的场景。ZeRO-Infinity这是 ZeRO 的终极形态。它不仅可以卸载到 CPU 内存还可以进一步将参数、梯度、优化器状态卸载到NVMe 固态硬盘上。NVMe 的容量可以非常大数TB这使得训练万亿甚至十万亿参数的模型成为可能。Infinity 采用了复杂的异步 I/O、内存映射和预取策略来缓解硬盘 I/O 的延迟问题。当然其训练速度相比纯 GPU 内存训练会慢很多这是一种用时间和存储空间换取模型规模的能力。如何选择这形成了一个清晰的技术选型路径单卡能放下 - 单卡训练。单卡放不下但多卡总显存放得下完整副本 - 标准数据并行。多卡总显存放不下完整副本但放下分片后的参数没问题 -ZeRO-2优先。参数分片后显存仍紧张但 CPU 内存充足 -ZeRO-3或ZeRO-Offload视通信和 PCIe 带宽权衡。需要训练极端大模型CPU 内存也不够 -ZeRO-Infinity。5. 对比与总结ZeRO 在分布式训练生态中的位置最后我们把 ZeRO 放回整个分布式训练的大图景里看看它和别的技术是什么关系。ZeRO vs. 标准数据并行ZeRO 是数据并行的超集。它保持了数据并行“每个卡处理不同数据”的简单逻辑但通过分片消除了内存冗余。你可以把标准数据并行看作是 ZeRO Stage 0。ZeRO vs. 模型并行张量并行、流水线并行互补而非替代这是最重要的认知。ZeRO数据并行和模型并行是正交的可以结合使用。例如Meta 训练 Llama 时就同时使用了流水线并行、张量并行和 ZeRO 数据并行。分工不同模型并行张量/流水线解决的是“单个层或操作太大一张卡算不了”的问题。比如一个 100B 参数的层必须拆开到多张卡上计算。ZeRO解决的是“模型状态参数、梯度、优化器状态太多一张卡存不下”的问题。结合使用常见的模式是在节点内使用张量并行来处理大层在节点间使用流水线并行来切分模型层然后在所有副本间使用 ZeRO 数据并行来增加批量大小和提升训练稳定性。DeepSpeed 和 Megatron-LM 的集成就是这种混合并行的典范。个人实践体会从我自己的项目经验来看ZeRO 极大地降低了大规模模型训练的门槛。在大多数情况下你不需要从零开始设计复杂的模型并行代码只需要在 DeepSpeed 配置文件中将stage从 1 改为 2 或 3就能让你的训练脚本支撑起大得多的模型。这种“配置即所得”的能力非常强大。然而它也不是银弹。最大的教训就是一定要做好监控不仅要监控显存还要监控 GPU 利用率和网络带宽。我曾在一个网络一般的集群上强行使用 ZeRO-3结果训练速度只有 ZeRO-2 的 1/3白白浪费了计算资源。所以从 ZeRO-2 开始逐步推进用数据内存和吞吐量来驱动你的配置选择这才是最稳妥的做法。对于超大规模训练混合并行是必然之路而 ZeRO 是其中不可或缺的、负责解决“内存墙”问题的关键组件。
返回列表