最近在尝试用FSDP跑一个7B的LoRA微调,单卡A100 80G。按照文档设了sharding_strategy=ShardingStrategy.SHARD_GRAD_OP,但发现训练刚开始显存占用就飙到70多G,比不用FSDP的DDP还高。我理解FSDP应该把参数和梯度分片到各卡,但单卡显存反而更高了?另外,forward_prefetch和backward_prefetch的几种选项我都试过,效果差异不大。
PyTorch FSDP训练时显存不降反升,是配置问题还是预期行为?
全部回复
共 33 条这问题我踩过坑,单卡跑FSDP本来就会这样,分片省的是多卡间的显存,单卡上激活值和临时buffer反而可能更多。你试试把cpu_offload打开,或者调小bucket_cap_mb到25左右,显存能掉不少。另外LoRA的话,FSDP对冻结参数的处理有时会额外复制权重,检查下是不是把auto_wrap_policy设得太激进了。
FSDP的SHARD_GRAD_OP本身就只分片梯度,参数和优化器状态还是每卡一份的,7B模型光参数就要14G左右,加上LoRA的额外开销和激活值,70多G其实不算离谱。你对比DDP时可能没算上优化器状态,AdamW一开就是参数量的两倍显存,FSDP单卡反而要把整个模型的梯度都攒一轮再分片,峰值自然更高。建议你直接看下torch.profiler的内存快照,确认是参数、梯度还是激活占大头,另外试试把activation checkpointing打开,那个对显存影响比prefetch大得多。
这个现象我之前也踩过坑,SHARD_GRAD_OP其实只分片梯度,参数和优化器状态还是全量驻留的,7B模型光是参数+梯度+Adam状态就差不多要60G了,再加上LoRA的激活值,70多G不奇怪。你如果想让单卡显存真的降下来,得用FULL_SHARD,但代价是通信量翻倍,小batch下可能反而更慢。另外forward_prefetch对单卡场景基本没帮助,它主要是为了跨卡流水线并行设计的,你不如把cpu_offload打开试试,LoRA微调场景下能省不少。
看到你这个情况我第一反应是sharding_strategy设成SHARD_GRAD_OP其实只分片了梯度,参数和优化器状态还是全量复制,所以显存比DDP高完全正常,尤其是LoRA这种本身参数就少的情况,分片收益根本抵不过FSDP自身的通信缓冲和activation开销。你把策略改成FULL_SHARD试试,理论上参数和优化器状态都分片后应该能明显降下来,但代价是通信量会大不少,小batch下不一定划算。另外你用的是LoRA,其实可以考虑干脆不开FSDP,直接DDP配gradient checkpointing,7B在80G上完全能跑,还省去一堆调参的麻烦。
刚踩过类似的坑,试试关掉activation checkpointing或者调小batch size,这玩意分片后通信开销反而会放大峰值显存。
说实话你这情况我踩过一模一样的坑,最后发现大概率是预期行为。SHARD_GRAD_OP只分片了梯度和优化器状态,参数本身还是每卡都留一份完整副本,所以前向激活值加上参数副本,显存自然比DDP还高,尤其是7B这种规模,光参数就占14G多,LoRA虽然省了训练参数但基座权重没法省。另外你检查过activation_checkpointing开了没?FSDP的显存收益很大程度靠它把中间激活换到CPU,不然大batch下激活值才是吞显存的大头,比参数分片影响还猛。还有个小坑是forward_prefetch如果设置的时机不对,反而会提前把下一层的权重搬进来,导致峰值显存更难看,你可以试试把limit_all_gathers=True加上,强制延迟同步,我这边效果比调prefetch明显。最后建议你盯着torch.cuda.max_memory_allocated()看峰值,别只看nvidia-smi的实时占用,因为FSDP的all-gather和释放有滞后性,那个数字会误导人。如果还想压,就把sharding_strategy换成FULL_SHARD配合auto_wrap_policy,代价是通信开销大一点,但显存能掉到40G左右。
你这配置看起来没啥问题,但SHARD_GRAD_OP本来就是只分片梯度,参数和优化器状态还在每张卡上全量复制,7B模型光参数就占14G(半精度),加上LoRA的梯度和激活值,70多G其实不奇怪。想省显存得用FULL_SHARD,但代价是通信开销大,7B单卡跑本来就不太适合FSDP,分片收益要跨节点才明显。你DDP能跑的话,不如先查下是不是activation checkpointing没开,那个对显存影响比FSDP配置大得多。
大概率是FSDP把优化器状态也放在第一张卡上做初始化了,试试cpu_offload或者把batch size调小点看曲线。
说实话看到你这个现象我第一反应是检查一下你是不是把activation checkpointing和FSDP一起用了,因为光分片参数和梯度的话,7B模型在A100上激活值才是显存大头。SHARD_GRAD_OP这个策略下,前向传播时每层参数还是要全量物化到当前设备上的,如果你没有配合activate_checkpointing,激活值累积起来轻松超过DDP的占用。另外LoRA的话,base model的权重是冻结的,但FSDP的sharding单位是layer还是整个model会影响内存峰值,你可以试试把auto_wrap_policy设成按transformer block来分片,有时候能缓解瞬时峰值。还有个小坑是forward_prefetch在数据加载不均匀时反而会提前拉取下一层参数,导致当前层还没释放又占一份内存,你可以试着关掉它然后把limit_all_gathers=True打开,这个对峰值控制挺明显的。最后问一下,你确认过是纯参数+梯度+优化器状态的统计吗,因为PyTorch profiler里看起来的“显存”有时候会包含缓存分配器没释放的碎片,用torch.cuda.reset_peak_memory_stats()再测一次更准。我上次跑13B也遇到过类似情况,最后发现是mixed precision下master weight在每张卡上复制了一份,你要是开了fp16记得把cast_forward_inputs也检查下。
单卡跑FSDP本来就容易这样,分片没省多少反而多了all-gather开销,试试开CPU offload或者确认下是不是没启用LoRA的requires_grad。
单卡跑FSDP其实不太能体现分片优势,SHARD_GRAD_OP只切梯度,参数和优化器状态还是全量复制,显存自然下不来。7B模型光参数就14G,加上LoRA的激活值和临时buffer,70G挺正常的。想省显存可以试试FULL_SHARD配合CPU offload,或者干脆用QLoRA量化。另外prefetch那几个参数主要影响通信重叠,对峰值显存帮助有限,别指望靠它降。
显存涨这么多不太正常,检查下是不是没包住auto_wrap,或者LoRA层没被FSDP管到,参数其实还是全量复制了。
LoRA微调7B的话,SHARD_GRAD_OP其实没太大必要,因为LoRA的可训练参数本来就少,分片收益有限,反而FSDP的all-gather通信和临时buffer会额外吃显存。你可以试试FULL_SHARD配合把LoRA层单独排除在外,或者干脆用DDP加gradient checkpointing,显存可能更省。另外显存飙高也可能是optimizer state没被正确分片,检查下use_orig_params有没有开。