欢迎光临

PyTorch FSDP(Fully Sharded Data Parallel)源码级深度解析:从 ZeRO-3 实现原理到生产环境调优实战

一、引言:当模型参数放不进单卡显存时

随着大语言模型规模的不断膨胀,从 BERT-base 的 1.1 亿参数到 Llama 3 的 4050 亿参数,单张 GPU 的显存早已无法承载完整的模型训练。即便是拥有 80GB HBM3 显存的 H100,在面对超过 30B 参数的模型时也显得捉襟见肘。分布式训练不再是可选方案,而是必修课。

在分布式训练策略中,Data Parallelism(数据并行)是最直观的思路:每个 GPU 持有完整的模型副本,各自处理不同的数据批次。但当模型规模超过单卡显存时,数据并行就失效了。ZeRO(Zero Redundancy Optimizer)系列方法由 Microsoft 在 2019 年提出,通过消除冗余的显存占用来突破单卡容量天花板。PyTorch 的 FSDP(Fully Sharded Data Parallel)正是 ZeRO-3 策略的工程化实现,它已经成为 PyTorch 生态中训练大模型的事实标准。

本文将从源码层面深入剖析 FSDP 的核心机制,覆盖分片策略、通信调度、内存管理与生产环境调优等关键维度。阅读本文之前,建议对 PyTorch 的 DistributedDataParallel(DDP)有基本了解,因为 FSDP 可以理解为”在 DDP 之上叠加了模型状态的显存分片”。

PyTorch 分布式训练架构图

二、从 DDP 到 FSDP:显存冗余的本质

2.1 DDP 的显存分布与问题

在标准的 DDP 训练中,每张 GPU 上存储了完整的模型参数(Parameters)、梯度(Gradients)和优化器状态(Optimizer States,如 Adam 中的 momentum 和 variance)。对于一个参数量为 Ψ 的模型:

  • 模型参数:占用 4Ψ 字节(FP32)或 2Ψ 字节(BF16/FP16)
  • 梯度:与参数大小相同,占用 4Ψ 或 2Ψ 字节
  • 优化器状态(以 Adam 为例):两个状态各 4Ψ 字节,共 8Ψ 字节(FP32 下)
  • 中间激活值(Activations):取决于 batch size 和序列长度,通常与参数显存相当甚至更大

以 70B 参数模型为例,仅参数 + 梯度 + Adam 状态在 FP32 下就需要 4×70×4 = 1120 GB 显存——远超过任何单张 GPU 的容量。即便使用混合精度训练(BF16 参数 + FP32 优化器状态),也需要约 2Ψ + 2Ψ + 8Ψ = 12Ψ = 840 GB。这显然行不通。

2.2 ZeRO 的三阶段优化

ZeRO 的核心洞察是:数据并行中,每个 GPU 都存储了完整的优化器状态和参数,但实际上它们都是同步的——我们只需要在每个 GPU 上保留自己负责的那一份,其他 GPU 的数据可以通过通信来获取。ZeRO 将优化分为三个阶段:

阶段 分片内容 每 GPU 显存(70B 模型, BF16+FP32 Adam) 通信开销
ZeRO-1 优化器状态(Optimizer States) ~120 GB 与 DDP 相同
ZeRO-2 优化器状态 + 梯度(Gradients) ~60 GB 与 DDP 相同
ZeRO-3 优化器状态 + 梯度 + 参数(Parameters) ~30 GB(64 GPU 时) 增加 ~50%

FSDP 实现了 ZeRO-3 的全部功能。在 N 张 GPU 上,每张 GPU 只持有约 1/N 的完整模型状态,从而实现了接近线性的显存缩放。

三、FSDP 核心原理:前向与反向传播中的动态参数收集

3.1 分片的粒度:Flattening 与 FlatParameter

FSDP 并没有对每个独立参数进行分片——那样会产生大量的细粒度通信操作(All-Gather),通信效率极低。相反,FSDP 将整个模型的参数进行展平(Flatten),形成一个或几个连续的 FlatParameter 张量,然后在这些大块上执行分片。


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
# 简化源码示意:FlatParameter 的构造
# 来源:torch/distributed/fsdp/_flat_param.py

class FlatParameter(nn.Parameter):
    """将多个 nn.Parameter 展平为一个连续张量。"""
    def __init__(self, params: List[nn.Parameter]):
        # 获取所有参数的数据指针
        flat_data = torch.cat([p.data.view(-1) for p in params])
        super().__init__(flat_data, requires_grad=True)
        self._param_infos = []  # 记录每个子参数在 flat 中的偏移
        offset = 0
        for p in params:
            numel = p.numel()
            self._param_infos.append({
                "start": offset,
                "end": offset + numel,
                "shape": p.shape,
                "dtype": p.dtype,
            })
            offset += numel

这样做的好处有两个:一是减少了通信元的数量(All-Gather 一个大张量比对数千个小张量分别 All-Gather 高效得多);二是方便将 FlatParameter 均匀切分到各个 rank。

3.2 前向传播:All-Gather + 计算 + Discard

在前向传播过程中,FSDP 将整个计算图划分为多个”单元”,每个单元对应模型中的一个子模块(通常是 Transformer Block)。对每个单元:

  1. All-Gather:从所有 rank 上收集完整的参数,形成完整的 FlatParameter
  2. 前向计算:用收集到的完整参数执行该子模块的前向传播
  3. Discard(丢弃):前向计算完成后,立即释放那些不属于本 rank 的参数片段,只保留本 rank 拥有的原始分片

这意味着在前向传播过程中,同一时刻只有一个”单元”的参数是完整的。这大幅降低了峰值显存占用。


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
# 前向传播中的 All-Gather 流程(简化伪代码)
def forward(self, input):
    if self._is_root or self._has_params:
        # Step 1: All-Gather 收集完整参数
        with torch.cuda.stream(self._streams["all_gather"]):
            self._all_gather_params()
       
        # Step 2: 同步流,等收集完成
        torch.cuda.current_stream().wait_stream(self._streams["all_gather"])
   
    # Step 3: 执行前向计算
    output = self._module(input)
   
    if not self._is_root:
        # Step 4: 释放非本 rank 的参数
        self._free_unsharded_params()
   
    return output

3.3 反向传播:重新收集参数 + 计算梯度 + Reduce-Scatter

反向传播的逻辑更加精巧。由于前向传播结束后参数已被丢弃,反向传播时需要重新通过 All-Gather 收集当前单元的参数来计算梯度。这也是 ZeRO-3 “以通信换显存”的核心代价:前向和反向各需要一次 All-Gather。

  1. 重新 All-Gather:再次从所有 rank 收集完整的参数
  2. 计算梯度:使用收集到的参数计算梯度
  3. Reduce-Scatter:对所有 rank 上的梯度执行 Reduce-Scatter 操作——先 All-Reduce 求和,再将结果按分片分散到各 rank,使得每个 rank 只保留自己分片的梯度
  4. 释放完整参数:与参数类似,释放非本 rank 的参数

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
# 反向传播中的 Reduce-Scatter(简化伪代码)
# 来源:torch/distributed/fsdp/_flat_param.py::_reduce_scatter

def _reduce_scatter(self):
    fsdp_group = self.process_group
    world_size = fsdp_group.size()
    shard_size = self.full_numel // world_size
   
    # 对梯度做 Reduce-Scatter
    # 每个 rank 得到 1/world_size 的梯度碎片
    output = torch.empty(shard_size, dtype=self.dtype, device=self.device)
   
    # 使用 NCCL 的 reduce_scatter
    dist.reduce_scatter_tensor(
        output=output,
        input=self._gradient,
        group=fsdp_group,
        op=dist.ReduceOp.AVG  # 或 SUM,取决于配置
    )
   
    # 释放完整梯度
    self._gradient = None
    self._sharded_gradient = output

四、分片策略(Sharding Strategy):如何选择最优粒度

FSDP 提供了三种分片策略,通过

1
ShardingStrategy

枚举来配置:

策略 枚举值 描述 适用场景
FULL_SHARD 1 模型参数、梯度、优化器状态全部分片(ZeRO-3) 超大模型训练,显存极度紧张
SHARD_GRAD_OP 2 仅分片梯度和优化器状态,参数副本全部保留(ZeRO-2) 显存够用但想降低优化器状态开销
NO_SHARD 3 不分片,相当于标准 DDP 模型参数可放入单卡,仅需数据并行
HYBRID_SHARD 4 节点内 FULL_SHARD,节点间 NO_SHARD(即跨节点复制完整模型) 多节点训练,节点内用 NVLink,节点间用 IB/RoCE

HYBRID_SHARD 是一个特别实用的策略。在多节点训练场景中,同一节点内的 GPU 通过 NVLink 互联(带宽可达 600 GB/s),但跨节点只有 InfiniBand(通常 200-400 Gbit/s ≈ 25-50 GB/s)。HYBRID_SHARD 在节点内部做 FULL_SHARD(利用高带宽进行频繁通信),而在节点之间做 NO_SHARD(只在节点间同步梯度,通信量与 DDP 相同)。这很好地在通信效率和显存效率之间取得了平衡。

五、CPU Offload:将显存压力转移到内存

当显存仍然不够时,FSDP 支持将优化器状态甚至参数卸载(Offload)到 CPU 内存中。这是突破 GPU 显存物理极限的最后一招。


1
2
3
4
5
6
7
8
9
10
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import CPUOffload

model = FSDP(
    model,
    cpu_offload=CPUOffload(offload_params=True),
    # 参数卸载到 CPU
    # forward/backward 时 All-Gather 从 CPU 取参数
    # 计算完成后参数又释放回 CPU
)

CPU Offload 的工作原理:

  • 参数卸载:分片后的参数存储在 CPU 固定的(pinned)内存上,通过
    1
    cudaMemcpyAsync

    异步拷贝到 GPU 进行计算

  • 优化器状态卸载:优化器中的 momentum 和 variance 存储在 CPU 内存中,优化器更新步骤在 CPU 上执行
  • 梯度卸载:Reduce-Scatter 后的梯度碎片也存放在 CPU 上

代价是什么?PCIe 带宽。一张 GPU 到 CPU 的 PCIe 4.0 x16 带宽约为 32 GB/s,而 NVLink 带宽是 600 GB/s——相差近 20 倍。启用 CPU Offload 后,All-Gather 的延迟会急剧增加,训练吞吐可能下降 50%-70%。

因此,CPU Offload 应作为”最后的逃生出口”:只有当模型尺寸实在无法靠 GPU 显存容纳时,才考虑启用。一般来说,优先增加 GPU 数量或使用梯度检查点(Gradient Checkpointing)来降低激活值显存,而不是直接启用 CPU Offload。

六、通信调度与重叠(Communication Overlap)

每次 All-Gather 和 Reduce-Scatter 都是同步操作,如果顺序执行会严重拖慢训练速度。FSDP 通过以下机制来减少通信对计算的影响:

6.1 预取(Prefetch)机制

FSDP 在前向传播时会预取下一个单元的完整参数,与当前单元的计算重叠执行。这通过

1
_prefetch_forward_params

1
_prefetch_backward_params

实现:


1
2
3
4
5
6
7
8
9
# 前向预取逻辑(简化)
def _prefetch_forward_params(self, next_unit):
    if next_unit is not None and next_unit._has_params:
        # 使用单独的计算流执行 All-Gather
        torch.cuda.current_stream().record_event(self._prefetch_event)
        next_unit._all_gather_params(
            stream=self._prefetch_stream,
            wait_event=self._prefetch_event
        )

通过这种方式,当 FSDP 在处理第 N 个 Transformer Block 的前向计算时,第 N+1 个 Block 的参数 All-Gather 已经在后台完成了。理想情况下,通信开销可以被完全隐藏。

6.2 梯度同步的时机:NO_WAIT / WAIT / POST_ORDER

在反向传播中,Reduce-Scatter 的调度时机对性能影响很大。FSDP 的

1
BackwardPrefetch

参数控制:

  • BACKWARD_PRE(默认):在当前单元的梯度计算完成后立即启动 Reduce-Scatter,同时预取下一个单元的参数。这是推荐配置,通信和计算重叠效果最好。
  • BACKWARD_POST:在当前单元的计算和梯度处理完全结束后才启动通信,重叠效果较差,但显存占用稍低。

1
2
3
4
5
6
from torch.distributed.fsdp import BackwardPrefetch

FSDP(
    model,
    backward_prefetch=BackwardPrefetch.BACKWARD_PRE,
)

6.3 Limit All-Gathers(限制并行通信数)

虽然预取能有效隐藏通信延迟,但过多的并行 All-Gather 操作也会占用显存(因为被预取到 GPU 上的完整参数需要空间暂存)。FSDP 通过

1
limit_all_gathers=True

来限制同时进行的 All-Gather 数量,防止显存溢出:


1
2
3
4
5
6
FSDP(
    model,
    limit_all_gathers=True,  # 默认 False
    # 启用后,最多同时执行 2 个 All-Gather
    # 防止预取过多导致 OOM
)

七、生产环境调优实战指南

7.1 Auto Wrap Policy:自动决定 FSDP 单元边界

FSDP 需要知道在哪里进行参数 All-Gather 的边界。手动指定每个子模块是否 wrap 是繁琐且容易出错的。PyTorch 提供了自动封装策略(Auto Wrap Policy):


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
from torch.distributed.fsdp.wrap import (
    size_based_auto_wrap_policy,
    transformer_auto_wrap_policy,
    lambda_auto_wrap_policy,
)

# 策略一:基于参数大小自动 wrap
# 当一个模块的参数超过 min_num_params 时自动封装
auto_wrap_policy = partial(
    size_based_auto_wrap_policy,
    min_num_params=1e8,  # 1 亿参数为阈值
)

# 策略二:基于模块类型
# 只对 Transformer Block 类型进行封装
auto_wrap_policy = partial(
    transformer_auto_wrap_policy,
    transformer_layer_cls={
        LlamaDecoderLayer,
        GPT2Block,
        T5Block,
    },
)

FSDP(
    model,
    auto_wrap_policy=auto_wrap_policy,
)

最佳实践:对于 Transformer 架构,推荐使用

1
transformer_auto_wrap_policy

按层(Layer)划分 FSDP 单元。每个 Transformer Block 作为一个 FSDP 单元,这样通信粒度和计算粒度匹配良好。

7.2 Mixed Precision 配置:FP32 Master Weights 与 BF16 计算

FSDP 内建了混合精度支持,无需外挂 AMP(Automatic Mixed Precision):


1
2
3
4
5
6
7
8
9
10
11
12
13
14
from torch.distributed.fsdp import MixedPrecision

mp_config = MixedPrecision(
    param_dtype=torch.bfloat16,       # 前向传播时 All-Gather 到 GPU 的参数 dtype
    reduce_dtype=torch.bfloat16,      # Reduce-Scatter 时梯度的 dtype
    buffer_dtype=torch.bfloat16,      # buffer 的 dtype
)

FSDP(
    model,
    mixed_precision=mp_config,
    # 优化器仍使用 FP32 精度的主权重
    # FSDP 内部自动维护 FP32 的 shadow weights 在 optimizer states 中
)

需要注意的是:即使前向传播使用 BF16 参数计算,FSDP 内部仍然为优化器保存了 FP32 版本的参数分片。这意味着优化器步骤仍然在 FP32 精度下更新,保证了训练精度不下降——这正是 ZeRO 的理念:用通信保精度。

7.3 执行单元测试:诊断通信瓶颈

下面是一个完整的 FSDP 训练脚本框架,包含关键的 profiler 工具:


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
import torch
import torch.distributed as dist
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
from torch.profiler import profile, record_function, ProfilerActivity
from functools import partial

def train_fsdp():
    dist.init_process_group(backend="nccl")
    local_rank = int(os.environ["LOCAL_RANK"])
    torch.cuda.set_device(local_rank)
   
    model = LlamaForCausalLM.from_pretrained("meta-llama/Llama-2-7b")
   
    auto_wrap_policy = partial(
        transformer_auto_wrap_policy,
        transformer_layer_cls={LlamaDecoderLayer},
    )
   
    fsdp_model = FSDP(
        model,
        auto_wrap_policy=auto_wrap_policy,
        sharding_strategy=ShardingStrategy.FULL_SHARD,
        mixed_precision=MixedPrecision(
            param_dtype=torch.bfloat16,
            reduce_dtype=torch.bfloat16,
        ),
        device_id=local_rank,
    )
   
    optim = torch.optim.AdamW(fsdp_model.parameters(), lr=3e-5)
   
    # Profiling:观察 All-Gather / Reduce-Scatter 耗时
    with profile(
        activities=[ProfilerActivity.CUDA, ProfilerActivity.CPU],
        profile_memory=True,
        record_shapes=True,
    ) as prof:
        for step, batch in enumerate(dataloader):
            with record_function("forward"):
                loss = fsdp_model(**batch).loss
            with record_function("backward"):
                loss.backward()
            with record_function("optimizer"):
                optim.step()
                optim.zero_grad()
           
            if step >= 10:  # warm up后采样
                break
   
    # 打印通信相关 op 的耗时统计
    print(prof.key_averages().table(
        sort_by="cuda_time_total",
        row_limit=20
    ))
   
    # 观察 AllGather 和 ReduceScatter 是否被计算隐藏
    # 理想状况下它们在 time_line 上应与 compute 重叠

7.4 常见陷阱与解决策略

问题 现象 根因与解决
OOM(显存溢出) 启动后不久即报 CUDA OOM 启用

1
limit_all_gathers=True

;减小 batch size;启用 Gradient Checkpointing;确认 forward_prefetch 未过度预取

训练速度慢 吞吐远低于预期 MFU 检查通信是否与计算重叠(Profiler);确认 NCCL 环境变量已优化(NCCL_IB_HCA, NCCL_SOCKET_IFNAME);尝试 HYBRID_SHARD 减少跨节点通信
Loss 不收敛 Loss 震荡或发散 检查 MixedPrecision 配置是否丢失精度;确认 reduce_dtype 与 param_dtype 匹配;验证学习率是否需要缩放
All-Gather 死锁 训练卡住,GPU 利用率 0% 常见于嵌套 FSDP + activation checkpointing 场景。尝试设置

1
forward_prefetch=False

1
backward_prefetch=BackwardPrefetch.BACKWARD_POST
显存碎片 总显存有剩余但分配大张量仍 OOM 启用

1
torch.cuda.memory._set_allocator_settings('expandable_segments:True')

(PyTorch 2.0+)

八、FSDP 与 DeepSpeed ZeRO-3 的对比

DeepSpeed ZeRO-3 是 FSDP 的主要竞争对手。两者在算法层面等价,但工程实现上有显著差异:

维度 PyTorch FSDP DeepSpeed ZeRO-3
集成方式 原生 PyTorch 组件,pip install 即用 需要安装 deepspeed 库,模型需封装为 DeepSpeedEngine
分片粒度 FlatParameter(整体展平后分片) 参数组(Param Group)级别分片
通信重叠 通过预取流,支持前向/反向预取 更激进的梯度 Bucketing 策略
CPU Offload 支持(CPUOffload) 支持(ZeRO-Infinity/NVMe Offload)
训练优化器 需用户自行创建,FSDP 自动 params() 接口 DeepSpeed 自动管理优化器状态
Hugging Face 兼容性 优秀(被 Trainer 原生支持) 良好(需 DeepSpeed 配置 JSON 文件)
梯度检查点 由 PyTorch 原生的 activation checkpointing 提供 内建 deepspeed.checkpointing 模块
显存效率 相近(实现细节差异引入 ±5%) 相近
易用性 ⭐⭐⭐⭐(API 简洁,文档完善) ⭐⭐⭐(配置略复杂)

选择建议

  • 如果你正在使用 Hugging Face Transformers Trainer,FSDP 是最自然的选择——几行配置即可启用
  • 如果你需要更极致的优化(如 ZeRO-Infinity + NVMe Offload),或者正在进行非常规的分布式训练拓扑,DeepSpeed 提供了更多底层控制
  • 对于大多数场景,FSDP 已经足够成熟,而且减少了对外部库的依赖

九、未来展望:FSDP 的演进方向

PyTorch 社区没有停下 FSDP 的迭代脚步。在即将发布的 PyTorch 2.x 系列中,以下几个方向值得关注:

  • FSDPv2(torch.distributed.fsdp_v2):更高效的通信调度,原生支持异步 Reduce-Scatter 和更智能的预取策略,减少调度器层面的开销
  • HSDP(Hybrid Sharded Data Parallel)正式化:HYBRID_SHARD 策略将在 v2 中获得更完善的支持,包括自动检测节点拓扑
  • DTensor 集成:将 FSDP 的底层分片逻辑与 DTensor(Distributed Tensor)对齐,实现统一的分布式张量抽象
  • 自动调优:自动探测最优的分片策略(FULL_SHARD vs HYBRID_SHARD)、batch size 和通信预取参数

十、总结

PyTorch FSDP 是当前训练大语言模型的基础性工具。它将 ZeRO-3 的理论优势转化为可用的工程实现,让拥有数十甚至数百张 GPU 的研究团队能够训练出千亿参数级别的模型。

理解 FSDP 的核心原理——FlatParameter 分片、前向/反向的 All-Gather 与 Reduce-Scatter 动态调度、通信与计算的重叠——是从事 AI Infra 工作的必备知识。在排查分布式训练的性能问题时,熟练使用 PyTorch Profiler 来观察通信操作的时间线、判断通信是否被计算掩盖,是 AI Infra 工程师的核心技能。

最后,回到本文开篇的问题:当模型参数放不进单卡显存时?答案是:分而治之(Shard & Conquer)。FSDP 正是这一思想的工程典范——它让大规模分布式训练不再是少数几家巨头公司的专利,而是每一个 AI 工程师都能触达的技术能力。

希望本文能帮助各位读者更深入地理解 FSDP 的内在机制,在实际工作中更高效地训练你的大模型。

【本站文章皆为原创,未经允许不得转载】:汤不热吧 » PyTorch FSDP(Fully Sharded Data Parallel)源码级深度解析:从 ZeRO-3 实现原理到生产环境调优实战
分享到: 更多 (0)