ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

07-FSDP分布式训练多卡跑大模型不再OOM

2026/8/3 9:40:17 拓冰建站 浏览量
07-FSDP分布式训练多卡跑大模型不再OOM

FSDP分布式训练:多卡跑大模型不再OOM

单卡显存不够,第一反应就是"加卡"。但多卡不是插上去就能用——数据怎么分、梯度怎么同步、显存怎么管,这三个问题搞不定,8张卡也跑不起来。

PyTorch FSDP(Fully Sharded Data Parallel)是目前最推荐的多卡训练方案。这篇把FSDP的原理、配置、踩坑全讲清楚。

为什么不推荐DDP

DDP(DistributedDataParallel)是PyTorch最早的多卡方案,原理简单:每张卡保存完整的模型副本,数据分片,各自前向反向后同步梯度。

问题在于:每张卡都要装下完整模型。7B模型fp16要14GB,加上梯度和优化器状态,单卡要56GB。A100 80GB才能跑,40GB都不够。

DDP每卡显存 = 模型权重 + 梯度 + 优化器状态 + 激活值 ≈ 14GB + 14GB + 28GB + ~10GB ≈ 66GB

FSDP的思路:模型参数也分片。每张卡只存1/N的参数,需要的时候从其他卡收集(all-gather),用完就扔掉。

FSDP每卡显存 = 模型权重/N + 梯度/N + 优化器状态/N + 激活值/N + 临时通信缓冲 ≈ 14GB/4 + 14GB/4 + 28GB/4 + ~10GB/4 + ~2GB ≈ 16.5GB (4卡FSDP)

4张A100 40GB就能跑7B模型全参数训练。

FSDP的执行流程

FSDP在每个Module级别做分片。前向传播时:

1. 当前层需要计算 → all-gather收集所有卡上该层的参数分片 → 拼出完整参数 2. 用完整参数做前向计算 3. 计算完 → 丢掉非本卡的参数分片,释放显存 4. 反向传播时同样:按需gather,用完释放

关键点:只有正在计算的层才占完整显存。其他层的参数都是分片状态,只占1/N。

这和梯度检查点(gradient checkpointing)的"用时间换空间"不同——FSDP不增加计算量,只是增加了通信开销。

FSDP实战代码

启动分布式训练

importosimporttorchimporttorch.distributedasdistfromtorch.distributed.fsdpimportFullyShardedDataParallelasFSDPfromtorch.distributed.fsdpimportMixedPrecision,ShardingStrategyfromtorch.distributed.fsdp.wrapimporttransformer_auto_wrap_policydefsetup_distributed():"""初始化分布式环境"""dist.init_process_group(backend="nccl")local_rank=int(os.environ["LOCAL_RANK"])torch.cuda.set_device(local_rank)returnlocal_rank,dist.get_rank(),dist.get_world_size()defcleanup_distributed():dist.destroy_process_group()

torchrun启动(不是python直接跑):

# 4卡训练torchrun--nproc_per_node=4train.py# 2机8卡(每机4卡)torchrun--nproc_per_node=4--nnodes=2--node_rank=0--master_addr=192.168.1.1--master_port=29500train.py

torchrun会自动设置RANKWORLD_SIZELOCAL_RANK等环境变量。

配置FSDP

defcreate_fsdp_model(model:nn.Module,rank:int)->FSDP:"""将模型包装为FSDP"""# 混合精度配置mp_policy=MixedPrecision(param_dtype=torch.bfloat16,