开发者
下载
FSDP中使用反向重计算时显存异常问题分析与优化

FSDP中使用反向重计算时显存异常问题分析与优化

模型训练

发表于: 2026/06/04

背景概述

在大规模深度学习模型训练中,显存资源是制约模型规模扩展的关键瓶颈。为应对这一挑战,Fully Sharded Data Parallel(FSDP)作为一种高效的分布式训练策略,通过将模型参数分片存储于多张GPU上,显著降低了单卡显存占用,支持超大模型的训练。然而,在结合反向重计算(Checkpointing)技术以进一步节省显存时,部分场景下仍可能出现显存占用异常升高的问题。本文聚焦于FSDP与反向重计算协同使用时的显存管理问题,分析其根本原因并提出有效解决方案。

问题现象  

在使用FSDPmindspeed-mm模型进行分片并启用反向重计算机制后,观察到反向传播阶段的显存占用约为前向传播的两倍,显存峰值显著上升,影响训练稳定性与资源利用率。通过设备内存快照工具分析,发现反向计算过程中存在两处显存未释放的异常堆积点。

根因分析  

1. 显存快照定位:通过PyTorch提供的设备内存快照功能,定位到反向计算阶段存在两个持续增长的显存占用点。  

  • 第一处为reduce操作,用于在FSDP中聚合各卡梯度,属于正常行为。  
  • 第二处为allgather操作,用于在重计算过程中收集分片参数以恢复前向计算状态,其申请的显存未被及时释放,构成异常。

2. 调用栈分析:深入分析调用栈发现,allgather操作在反向重计算中被重复触发,且在fully_shard被嵌套于重计算逻辑内部时,触发了pre_forwardpre_backward钩子函数,导致参数被多次unshard,但未在后续及时reshard,造成显存累积。

3. 核心问题:当fully_shard被包裹在checkpoint_wrapper(反向重计算)内部时,FSDP的默认行为会在每次反向传播前自动执行unshard操作以获取完整参数,但未在计算完成后主动触发reshard,导致分片参数未及时释放,显存持续累积。

解决方案  

1. 显式启用reshard_after_backward

  在调用fully_shard后,显式设置set_reshard_after_backward(True),确保在反向传播完成后立即对分片参数执行reshard操作,释放临时占用的显存。

2. 框架级优化

  已在最新版本中合入针对FSDP与重计算协同场景的优化补丁,修复了在checkpoint_wrapper内调用fully_shard时重复触发allgather及显存未释放的问题。该优化通过控制pre_forwardpre_backward的执行逻辑,避免不必要的参数展开与内存残留。(feat: torch fully_shard patch, optimize interaction between fully_shard and checkpoint_wrapper-MindSpeed-MM-AtomGit | GitCode

优化效果  

应用上述方案后,反向传播阶段的显存占用恢复至与前向传播相近水平,显存峰值下降约50%,训练稳定性显著提升,支持更大规模模型的高效训练。

技术建议与总结  

  • FSDP核心机制:FSDP全称Fully Sharded Data Parallel完全分片数据并行,通过分片模型参数降低每张卡的显存占用来拼接超大模型。每张卡持有1/DP卡数的参数,在需要参数计算时通过allgther收集这一层计算所需参数,每张卡参数需要的显存峰值为完整参数/层数。

通过合理配置与框架优化,FSDP与反向重计算可高效协同,实现大模型训练的显存与性能双重优化。