---
title: FSDP中使用反向重计算时显存异常问题分析与优化-官方技术文章-Atlas Community
description: 在大规模深度学习模型训练中，显存资源是制约模型规模扩展的关键瓶颈。为应对这一挑战，Fully Sharded Data Parallel（FSDP）作为一种高效的分布式训练策略，通过将模型参数分片存储于多张GPU上，显著降低了单卡显存占用，支持超大模型的训练。然而，在结合反向重计算（Checkpointing）技术以进一步节省显存时，部分场景下仍可能出现显存占用异常升高的问题。本文聚焦于FSDP与
keywords: FSDP,中使用反向重,计算时显存异,常问题分析与,优化,官方技术文章,Atlas,Community
url: https://www.hiascend.com/developer/techArticles/20260604-2?envFlag=1
section: (其他)
---

# FSDP中使用反向重计算时显存异常问题分析与优化-官方技术文章-Atlas Community

URL: https://www.hiascend.com/developer/techArticles/20260604-2?envFlag=1
描述: 在大规模深度学习模型训练中，显存资源是制约模型规模扩展的关键瓶颈。为应对这一挑战，Fully Sharded Data Parallel（FSDP）作为一种高效的分布式训练策略，通过将模型参数分片存储于多张GPU上，显著降低了单卡显存占用，支持超大模型的训练。然而，在结合反向重计算（Checkpointing）技术以进一步节省显存时，部分场景下仍可能出现显存占用异常升高的问题。本文聚焦于FSDP与
关键词: FSDP,中使用反向重,计算时显存异,常问题分析与,优化,官方技术文章,Atlas,Community

官方技术文章 [了解详情](https://www.hiascend.com/zh/developer/techArticles)

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

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

模型训练

Released: 2026/06/04

## 背景概述

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

问题现象

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

根因分析

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

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

2. 调用栈分析：深入分析调用栈发现，`allgather`操作在反向重计算中被重复触发，且在`fully_shard`被嵌套于重计算逻辑内部时，触发了`pre_forward`与`pre_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_forward`与`pre_backward`的执行逻辑，避免不必要的参数展开与内存残留。（feat: torch fully_shard patch, optimize interaction between fully_shard and checkpoint_wrapper-MindSpeed-MM-AtomGit | GitCode [了解详情](https://gitcode.com/Ascend/MindSpeed-MM/pull/2290)）

优化效果

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

技术建议与总结

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

- 最佳实践：在使用`FSDP`与反向重计算结合的场景时，建议始终显式启用`set_reshard_after_backward(True)`，以确保显存及时释放。
- 调试工具推荐：可借助内存快照功能（如`torch.cuda.memory_snapshot()`）进行显存使用分析，定位异常占用点。具体使用方式可参考（https://www.hiascend.com/document/detail/zh/Pytorch/730/ptmoddevg/Frameworkfeatures/docs/zh/framework_feature_guide_pytorch/memory_snapshot.md [了解详情](https://www.hiascend.com/document/detail/zh/Pytorch/730/ptmoddevg/Frameworkfeatures/docs/zh/framework_feature_guide_pytorch/memory_snapshot.md)）。

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