最近在尝试用FSDP跑一个7B的LoRA微调,单卡A100 80G。按照文档设了sharding_strategy=ShardingStrategy.SHARD_GRAD_OP,但发现训练刚开始显存占用就飙到70多G,比不用FSDP的DDP还高。我理解FSDP应该把参数和梯度分片到各卡,但单卡显存反而更高了?另外,forward_prefetch和backward_prefetch的几种选项我都试过,效果差异不大。
楼主
7天前
PyTorch FSDP训练时显存不降反升,是配置问题还是预期行为?
请 登录 后发表回复
全部回复
共 23 条
2楼
2天前
这问题我踩过坑,单卡跑FSDP本来就会这样,分片省的是多卡间的显存,单卡上激活值和临时buffer反而可能更多。你试试把cpu_offload打开,或者调小bucket_cap_mb到25左右,显存能掉不少。另外LoRA的话,FSDP对冻结参数的处理有时会额外复制权重,检查下是不是把auto_wrap_policy设得太激进了。
3楼
1天前
FSDP的SHARD_GRAD_OP本身就只分片梯度,参数和优化器状态还是每卡一份的,7B模型光参数就要14G左右,加上LoRA的额外开销和激活值,70多G其实不算离谱。你对比DDP时可能没算上优化器状态,AdamW一开就是参数量的两倍显存,FSDP单卡反而要把整个模型的梯度都攒一轮再分片,峰值自然更高。建议你直接看下torch.profiler的内存快照,确认是参数、梯度还是激活占大头,另外试试把activation checkpointing打开,那个对显存影响比prefetch大得多。
4楼
1天前
这个现象我之前也踩过坑,SHARD_GRAD_OP其实只分片梯度,参数和优化器状态还是全量驻留的,7B模型光是参数+梯度+Adam状态就差不多要60G了,再加上LoRA的激活值,70多G不奇怪。你如果想让单卡显存真的降下来,得用FULL_SHARD,但代价是通信量翻倍,小batch下可能反而更慢。另外forward_prefetch对单卡场景基本没帮助,它主要是为了跨卡流水线并行设计的,你不如把cpu_offload打开试试,LoRA微调场景下能省不少。