FSDP2 使用指南
作者:昇腾实战派知识地图:【昇腾实战派】综合指导知识地图-CSDN博客
背景概述
随着大模型训练的普及,分布式训练技术成为提升训练效率、降低显存占用的关键手段。PyTorch 提供的 Fully Sharded Data Parallel(FSDP)是一种模型分片并行策略,能够将模型参数、梯度和优化器状态分布到多个设备上,从而支持训练超出单设备显存容量的大模型。
FSDP2 是 FSDP 的下一代实现,基于 DTensor 进行参数管理,相比 FSDP1 在显存利用和计算效率上均有显著提升,当前昇腾MindSpeed训练框架已适配FSDP2。
本文将从初始化、优化器部署、训练循环等方面详细介绍 FSDP2 的使用方法,并提供可直接运行的完整脚本。
官方文档参考:https://docs.pytorch.org/tutorials/intermediate/FSDP_tutorial.html#how-fsdp2-works
快速开始: 将可执行的分布式脚本保存为 train.py 后,运行以下命令即可启动训练:
torchrun --nproc_per_node 2 train.py
其中 2 为可用的设备数量(如 NPU、GPU),可根据实际情况调整。如需直接获取可测试脚本,请跳转至本文第3节。
1. 初始化
1.1 模型导入与参数初始化
在开始分布式训练前,需要准备好模型定义文件。本文以官方提供的 Transformer 实现为例,请先从 GitHub 下载或克隆 model.py 文件。为避免模块名与变量名冲突,建议将文件重命名为 model_Tr.py。此外,还需导入 ModelArgs 类来配置模型参数。
导入与初始化的方法如下:
import torch.distributed as dist
from model_Tr import Transformer, ModelArgs
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
# 初始化通信域
backend = "hccl" if torch_npu.npu.is_available() else "gloo"
dist.init_process_group(backend=backend, rank=rank, world_size=world_size)
# 绑定设备
device = torch.device(f"npu:{local_rank}")
# 参数初始化
args = ModelArgs()
vocab_size = args.vocab_size
seq_len = args.max_seq_len
batch_size = 4
epochs = 5
# 模型初始化
model = Transformer(args)
1.2 FSDP 分片初始化
与 DDP 不同,FSDP2 不仅需要在根模型上调用 fully_shard,还需要对子模块(如每个 Transformer 层)分别应用 fully_shard。这样,在计算某一层时,其余层的参数处于分片状态,从而大幅降低显存占用。
fully_shard(model) 会将未独立分片的参数归组,并为它们安排高效的 all-gather / reduce-scatter 操作。调用 fully_shard 后,分片后的模型会被自动放置到训练设备上。
from torch.distributed.fsdp import fully_shard, FSDPModule
# 按照官方文档顺序应用 FSDP2
for layer in model.layers:
fully_shard(layer)
fully_shard(model)
1.3 完整初始化代码
结合前文的参数初始化与模型导入,完整的初始化代码如下:
from model_Tr import Transformer, ModelArgs
import torch
import torch_npu
import torch.distributed as dist
from torch.distributed.fsdp import fully_shard, FSDPModule
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
# 初始化通信域
backend = "hccl" if torch_npu.npu.is_available() else "gloo"
dist.init_process_group(backend=backend, rank=rank, world_size=world_size)
# 绑定设备
device = torch.device(f"npu:{local_rank}")
# 参数初始化
args = ModelArgs()
vocab_size = args.vocab_size
seq_len = args.max_seq_len
batch_size = 4
epochs = 5
# 模型实例化
model = Transformer(args)
# 应用 FSDP2
for layer in model.layers:
fully_shard(layer)
fully_shard(model)
# 验证分片是否成功
assert isinstance(model, Transformer)
assert isinstance(model, FSDPModule)
print(model)
预期输出示例:
FSDPTransformer(
(tok_embeddings): Embedding(...)
...
(layers): 3 x FSDPTransformerBlock(...)
(output): Linear(...)
)
2. 优化器部署与训练测试
2.1 优化器部署
应用 fully_shard 后,模型参数会变为 DTensor,其分片策略为 Shard(0)(沿第 0 维切分,即按设备数等分)。创建优化器前,可通过以下循环验证参数类型与分片方式:
from torch.distributed.tensor import DTensor, Shard
for param in model.parameters():
assert isinstance(param, DTensor)
assert param.placements == (Shard(0),)
# 可通过 param.to_local() 查看分片后的参数
optim = torch.optim.Adam(model.parameters(), lr=1e-2)
注意: 优化器必须在 fully_shard 之后构造。此时模型和优化器的状态字典都将以 DTensor 表示,torch.optim.Adam 和 torch.nn.utils.clip_grad_norm_ 等 API 可直接用于 DTensor 参数,无需额外适配,使得单设备与分布式训练的代码保持一致。
2.2 前向/反向传播与预取(Prefetching)
fully_shard 会注册前向/反向钩子,在计算前自动执行 All-Gather 以获取完整参数,并在计算后立即重新分片(reshard)。All-Gather 是一种集合通信操作:它将各个 rank 上按 dim=0 切分的本地参数聚合为完整张量并广播至所有进程,实现“按需物化”。计算结束后,通过 reshard 丢弃冗余副本,仅保留本地分片。
这一机制将显存复杂度由 O(model_size) 降为 O(model_size / N),本质是用通信带宽换取显存容量。同时,All-Gather 可在独立 CUDA 流中异步执行,借助预取机制与计算重叠,通信延迟能够被有效掩盖。
隐式预取(默认)
CPU 线程会在计算第 i 层之前自动发起该层的 All-Gather,并放入非默认流中,使其与第 i 层的计算并行。反向传播时,All-Gather 按逆序发起。隐式预取无需额外配置,训练循环与单卡训练完全一致:
for _ in range(epochs):
x = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
loss = model(x).sum()
loss.backward()
optim.step()
optim.zero_grad()
建议先从隐式预取开始,以了解开箱即用的性能表现。
显式预取
通过 set_modules_to_forward_prefetch 和 set_modules_to_backward_prefetch 可手动控制预取顺序。在下例中,计算第 i 层时会提前发起第 i+1、i+2 层的 All-Gather。显式预取在以下场景中效果显著:
- CPU 密集型负载:隐式预取下,CPU 可能来不及在核函数执行前发起下一层通信,显式预取可提前调度。
- 预取多层:一次预取多个层可进一步提升重叠效果,但会占用更多显存。
- 提前发起首次 All-Gather:调用
model.unshard()可让通信更早开始,避免在model(x)调用时才暴露延迟。
num_to_forward_prefetch = 2
for i, layer in enumerate(model.layers):
if i >= len(model.layers) - num_to_forward_prefetch:
break
layers_to_prefetch = [
model.layers[i + j] for j in range(1, num_to_forward_prefetch + 1)
]
layer.set_modules_to_forward_prefetch(layers_to_prefetch)
num_to_backward_prefetch = 2
for i, layer in enumerate(model.layers):
if i < num_to_backward_prefetch:
continue
layers_to_prefetch = [
model.layers[i - j] for j in range(1, num_to_backward_prefetch + 1)
]
layer.set_modules_to_backward_prefetch(layers_to_prefetch)
for _ in range(epochs):
# 提前触发首次 all-gather,与 model(x) 之前的计算重叠
model.unshard()
x = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
loss = model(x).sum()
loss.backward()
optim.step()
optim.zero_grad()
3. 脚本总结
请确保已将
model.py下载并重命名为model_Tr.py,置于当前工作目录下。
完整脚本(隐式预取)
适合快速验证分布式训练,无需手动配置预取,训练循环与单卡一致。
from model_Tr import Transformer, ModelArgs
import torch
import torch_npu
from torch.distributed.fsdp import fully_shard, FSDPModule
from torch.distributed.tensor import DTensor, Shard
def main():
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
# 初始化通信域
backend = "hccl" if torch_npu.npu.is_available() else "gloo"
dist.init_process_group(backend=backend, rank=rank, world_size=world_size)
# 绑定设备
device = torch.device(f"npu:{local_rank}")
#参数初始化
args = ModelArgs()
vocab_size = args.vocab_size
seq_len = args.max_seq_len
batch_size = 4
epochs = 5
local_rank = int(os.environ["LOCAL_RANK"])
device = torch.device(f"npu:{local_rank}")
# 创建模型
model = Transformer(args)
# 按照官方文档顺序应用 FSDP2
for layer in model.layers:
fully_shard(layer)
fully_shard(model)
assert isinstance(model, Transformer)
assert isinstance(model, FSDPModule)
print(model)
#用于检测模型是否分片成功
#——————————————输出样例————————————————
# FSDPTransformer(
# (tok_embeddings): Embedding(...)
# ...
# (layers): 3 x FSDPTransformerBlock(...)
# (output): Linear(...)
# )
# 按照官方文档顺序应用 FSDP2
#—————————————————————————————————————
for param in model.parameters():
assert isinstance(param, DTensor)
assert param.placements == (Shard(0),)
# 用于检测参数是否分片成功
optim = torch.optim.Adam(model.parameters(), lr=1e-2)
for _ in range(epochs):
x = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
loss = model(x).sum()
loss.backward()
optim.step()
optim.zero_grad()
print(f"loss: {loss:.4f}")
if __name__ == "__main__":
main()
完整脚本(显式预取)
通过多级预取和提前 unshard 进一步重叠通信与计算,以额外显存换取更高的吞吐。
from model_Tr import Transformer, ModelArgs
import torch
import torch_npu
from torch.distributed.fsdp import fully_shard, FSDPModule
from torch.distributed.tensor import DTensor, Shard
def main():
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
# 初始化通信域
backend = "hccl" if torch_npu.npu.is_available() else "gloo"
dist.init_process_group(backend=backend, rank=rank, world_size=world_size)
# 绑定设备
device = torch.device(f"npu:{local_rank}")
#参数初始化
args = ModelArgs()
vocab_size = args.vocab_size
seq_len = args.max_seq_len
batch_size = 4
epochs = 5
local_rank = int(os.environ["LOCAL_RANK"])
device = torch.device(f"npu:{local_rank}")
# 创建模型
model = Transformer(args)
for layer in model.layers:
fully_shard(layer)
fully_shard(model)
assert isinstance(model, Transformer)
assert isinstance(model, FSDPModule)
print(model)
#——————————————输出样例————————————————
# FSDPTransformer(
# (tok_embeddings): Embedding(...)
# ...
# (layers): 3 x FSDPTransformerBlock(...)
# (output): Linear(...)
# )
# 按照官方文档顺序应用 FSDP2
#—————————————————————————————————————
for param in model.parameters():
assert isinstance(param, DTensor)
assert param.placements == (Shard(0),)
# inspect sharded parameters with param.to_local()
optim = torch.optim.Adam(model.parameters(), lr=1e-2)
num_to_forward_prefetch = 2
for i, layer in enumerate(model.layers):
if i >= len(model.layers) - num_to_forward_prefetch:
break
layers_to_prefetch = [
model.layers[i + j] for j in range(1, num_to_forward_prefetch + 1)
]
layer.set_modules_to_forward_prefetch(layers_to_prefetch)
num_to_backward_prefetch = 2
for i, layer in enumerate(model.layers):
if i < num_to_backward_prefetch:
continue
layers_to_prefetch = [
model.layers[i - j] for j in range(1, num_to_backward_prefetch + 1)
]
layer.set_modules_to_backward_prefetch(layers_to_prefetch)
for _ in range(epochs):
# trigger 1st all-gather earlier
# this overlaps all-gather with any computation before model(x)
model.unshard()
x = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
loss = model(x).sum()
loss.backward()
optim.step()
optim.zero_grad()
if __name__ == "__main__":
main()
将上述任一脚本保存为 train.py 后,运行以下命令即可开始训练:
torchrun --nproc_per_node 2 train.py更多推荐



所有评论(0)