PyTorch FSDP分布式训练是什么
PyTorch FSDP(Fully Sharded Data Parallel)是PyTorch 2.0引入的全分片数据并行方案,将模型参数、梯度和优化器状态按Rank维度切分,每个GPU只保留全局参数的1/N。相比DDP(DistributedDataParallel)只在反向传播时同步梯度、每个副本仍持有完整模型参数,FSDP把显存占用从O(模型大小)降到O(模型大小/N),是单机多卡和跨节点训练70B以上大模型的主流选择。
FSDP的前身是FairScale库的FSDP实现,PyTorch 2.0将其合并进torch.distributed.fsdp并持续迭代。对大模型开发团队而言,掌握FSDP意味着不需要购买8张A100即可在4张卡上完成同样参数的训练,AI模型部署和深度学习框架选型时的成本结构完全不同。
FSDP分片策略配置:全分片与分层分片对比
FSDP支持两种分片策略,通过sharding_strategy参数控制:
from torch.distributed.fsdp import ShardingStrategy
# 全分片:参数、梯度、优化器状态全部切分(默认)
SHARD_GRAD_OP = ShardingStrategy.SHARD_GRAD_OP
# 分层分片:每两个GPU一组共享完整参数,组间全分片
HYBRID_SHARD = ShardingStrategy.HYBRID_SHARD
SHARD_GRAD_OP将参数、梯度、优化器状态全部切分,显存占用最低但前向/反向时需All-Gather临时聚合参数;HYBRID_SHARD先在节点内做复制分片、节点间做全分片,适合单机8卡(节点间走NVLink、节点间走RDMA)的物理机架设场景,通信开销比纯全分片低30%-50%。
FSDP与DeepSpeed ZeRO对比:选型依据
DeepSpeed ZeRO-3同样把参数/梯度/优化器分片,二者在训练效果上等价,差异在工程生态:
- FSDP由PyTorch官方维护,与torch.compile、DTensor、torch.distributed.checkpoint原生兼容,API改动风险低
- DeepSpeed在ZeRO-Offload(CPU/NVMe卸载)和混合精度(bf16+fp8)的成熟度更高,offload参数配置更细
- FSDP从PyTorch 2.1起支持参数卸载(CPU offload),但CPU卸载与HF Trainer的DeepSpeed集成相比仍有差距
选型建议:团队技术栈以PyTorch为主、模型规模在7B-70B,选FSDP;需要极低显存运行(如单卡跑30B)或依赖DeepSpeed已有调优参数的,保留ZeRO。
FSDP训练循环完整实现
以一个7B模型、8卡A100(80G)的分布式训练配置为例,实现FSDP封装与训练循环:
import torch
import torch.distributed as dist
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp.fully_sharded_data_parallel import CPUOffload
from torch.distributed.fsdp import ShardingStrategy, MixedPrecision
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
from torch.utils.data.distributed import DistributedSampler
def build_fsdp_model(model, transformer_cls):
auto_wrap = transformer_auto_wrap_policy(
transformer_layer_cls={transformer_cls}
)
return FSDP(
model,
auto_wrap_policy=auto_wrap,
sharding_strategy=ShardingStrategy.SHARD_GRAD_OP,
cpu_offload=CPUOffload(offload_params=False),
mixed_precision=MixedPrecision(
param_dtype=torch.bfloat16,
reduce_dtype=torch.bfloat16,
buffer_dtype=torch.bfloat16,
),
device_id=torch.cuda.current_device(),
)
def train_loop(model, dataloader, optimizer, epochs, rank):
for epoch in range(epochs):
dataloader.sampler.set_epoch(epoch) # 数据切分,保证每卡数据不同
for step, (inputs, labels) in enumerate(dataloader):
inputs = inputs.to("cuda")
labels = labels.to("cuda")
outputs = model(inputs, labels=labels)
loss = outputs["loss"]
loss.backward()
optimizer.step()
optimizer.zero_grad()
if rank == 0 and step % 50 == 0:
print(f"epoch {epoch} step {step} loss {loss.item():.4f}")
if __name__ == "__main__":
dist.init_process_group(backend="nccl")
torch.set_num_threads(4)
rank = dist.get_rank()
model = create_7b_model() # 业务侧定义,须含transformer层类
fsdp_model = prepare_fsdp_example(model, LlamaDecoderLayer)
optimizer = torch.optim.AdamW(fsdp_model.parameters(), lr=3e-5)
loader = build_dataloader(distributed=True) # 用DistributedSampler
train_loop(fsdp_model, loader, optimizer, epochs=2, rank=rank)
dist.destroy_process_group()
关键点:transformer_auto_wrap_policy必须指定模型内部的Transformer层类,FSDP据此递归切分;混合精度建议直接bfloat16(A100及以上无损失),避免fp16损失缩放和溢出的调参。
FSDP显存优化配置与常见报错排查
显存打满时按顺序检查:
- sharding_strategy是否为SHARD_GRAD_OP,默认即可,不要误改成NO_SHARD(等同DDP)
- 开启CPUOffload,但注意offload后step速度可能下降15%-25%,显存换耗时
- 激活函数checkpointing(activation checkpointing)与FSDP兼容,配auto_wrap_policy一起用可再省30%显存:
model.gradient_checkpointing_enable() - torch.compile(fsdp_model)可加速,但FSDP+compile在PyTorch 2.2前曾存在参数all_gather重算问题,升级到最新版本
常见报错:Expected all tensors to be on the same device多数因dataloader数据未迁移到GPU;AssertionError: Invalid flat parameter是模型含未注册参数,检查自定义module是否调用super().__init__()。Loss为NaN先排查学习率与bf16范围,70B模型常将lr降到1e-4以下。
FSDP性能基准与调优结论
在8张A100(80G)上训练Llama-2-13B,FSDP SHARD_GRAD_OP相对ZeRO-3在吞吐上持平(约95%-102%),而配置HF Trainer的FSDP + bf16 + activation checkpointing,可达每秒约1300 tokens/GPU。把CUDA缓存分配器显存预留(PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True)可在长序列场景下再减5%-10%峰值显存。结论是:新项目默认用FSDP,官方维护与DDP/TP(张量并行)的融合越来越顺滑,PyTorch 2.4之后与DTensor(分布式张量)的组合已经是PyTorch生态里的大模型并行标准路线。
原创文章,作者:小编,如若转载,请注明出处:https://www.yunthe.com/pytorchfsdp-fen-bu-shi-xun-lian-shi-zhan-quan-fen-pian-shu/