最近在尝试用FSDP跑一个7B的LoRA微调,单卡A100 80G。按照文档设了sharding_strategy=ShardingStrategy.SHARD_GRAD_OP,但发现训练刚开始显存占用就飙到70多G,比不用FSDP的DDP还高。我理解FSDP应该把参数和梯度分片到各卡,但单卡显存反而更高了?另外,forward_prefetch和backward_prefetch的几种选项我都试过,效果差异不大。
PyTorch FSDP训练时显存不降反升,是配置问题还是预期行为?
全部回复
共 23 条说实话我之前也踩过这个坑,SHARD_GRAD_OP本来就是只分片梯度,参数和优化器状态还在每张卡上完整保留,7B模型光参数就14G,加上LoRA的梯度累积和中间激活,单卡70G真的不奇怪。你如果想让显存明显降下来,得用FULL_SHARD,但代价是通信开销变大,训练速度会有肉眼可见的下降。另外forward_prefetch和backward_prefetch对显存的影响本来就很小,它们主要影响通信和计算的overlap效率,不是用来省显存的。建议你直接看下nvidia-smi里显存是不是被缓存占着,有时候是torch的缓存机制导致看起来很高,实际回收了。
这情况我踩过坑,SHARD_GRAD_OP本身不省激活显存,7B全量参数加载就得占大头,单卡跑肯定爆。
说实话我第一反应是你看下是不是LoRA的target_modules没设对,FSDP对参数分片是按module粒度来的,如果某些大参数被排除在分片外,显存反而会比DDP多出一份完整拷贝。另外SHARD_GRAD_OP本身只分片梯度,参数还是全量驻留的,单卡70G不算离谱,毕竟7B的fp32权重+优化器状态+激活值加起来本来就接近这个量级。建议试试FULL_SHARD,同时把activation checkpointing开起来,对比下峰值显存曲线。forward_prefetch和backward_prefetch在单卡场景下确实感知不强,主要影响通信重叠,不是显存瓶颈。
我最近也踩过类似的坑,SHARD_GRAD_OP这个策略本来就是为了省梯度显存,但如果你模型本身是7B,FSDP初始化时其实会先把全量参数加载到每张卡上做分片,那个瞬间的峰值反而比DDP更吓人。你观察到的70多G很可能是这个原因,等训练稳定后应该会降下来,但如果你用的是CPU offload或者混合精度没对齐,情况会更糟。另外LoRA的trainable参数虽然少,但base model的梯度如果没被冻结,FSDP还是会按全参数量去管理梯度分片,这可能是你显存不降反升的核心。forward_prefetch和backward_prefetch在单机单卡场景下确实收益有限,因为预取主要是为了掩盖跨机通信延迟,你单卡跑反而会多做一次无意义的全量参数聚合。建议你直接开NCCL的debug日志看下实际通信量,或者试试把activation checkpointing打开,那个对显存峰值的影响比prefetch大得多。还有个可能被忽略的点,你的LoRA target modules是不是作用在linear层上?如果是,FSDP会把整个linear层都当成可训练单元,导致分片粒度变粗,显存碎片会更严重。最后问一句,你用的是PyTorch 2.1+吗?老版本FSDP在mixed precision和gradient accumulation同时开的时候有已知的显存泄漏bug,升级到2.3之后会好很多。
这个现象大概率是配置问题而非预期行为。你开了SHARD_GRAD_OP但没配合cpu_offload或调整activation_checkpointing,FSDP在单卡场景下反而会额外保留完整参数副本和分片元数据,显存开销自然比DDP高。建议先确认param_init_fn是否触发了全量加载,另外试试把sharding_strategy改成FULL_SHARD并配合use_orig_params=True,LoRA场景下这组合通常能压到40G以内。至于prefetch选项,它们对峰值显存影响本来就不大,更多是影响通信和计算重叠效率。
SHARD_GRAD_OP只分片梯度,参数和优化器状态还在每块卡上,7B的LoRA激活值也占大头,显存高很正常。
我之前也踩过这个坑,SHARD_GRAD_OP本来就不分参数,只分梯度和优化器状态,所以前向和反向时每张卡都得保留完整参数副本,对7B来说光参数就占30多G了,再加上激活值和临时缓冲,70G很正常。想省显存的话得用FULL_SHARD,但代价是通信量翻倍,速度会慢不少。另外你试过把activation checkpointing打开吗?那个对激活值占用影响挺大的,有时候比调prefetch管用。顺便问下你用的transformers版本是新的吗?老版本里FSDP和LoRA的兼容性有点问题,会导致额外显存开销。
同款配置跑过,SHARD_GRAD_OP在单机场景下确实不如FULL_SHARD省显存,尤其LoRA冻结了大部分参数后,分片收益更小。你可以试试把activation checkpointing打开,或者调低batch size,这样显存峰值能压下来不少。另外forward_prefetch对单卡训练基本没优化,反而是多卡通信密集时才有感知,你感觉差异不大挺正常的。有个细节是FSDP的缓存分配器会预留显存,跑两步后再看峰值才比较准。
我最近也踩过类似的坑,后来发现多半是activation和optimizer state没被正确分片导致的。SHARD_GRAD_OP只分片梯度,但LoRA的trainable参数和frozen参数的显存占用逻辑不一样,建议检查一下是否所有参数都进了fsdp的wrap。另外试试activation_checkpointing,这个对降低峰值显存往往比调prefetch更有效。还有个容易忽略的点,如果用了gradient checkpointing但没配合use_orig_params=True,可能会额外保留一份全量参数副本,显存自然就上去了。
你这配置挺正常的,单卡跑FSDP分片反而会多一份通信缓冲,显存高不奇怪,试试开activation offload或调小batch看看。
说实话SHARD_GRAD_OP这个策略本身就不分参数,只分梯度和优化器状态,所以前向时每张卡都得保留完整权重,7B模型光参数就要14G,加上激活值和LoRA的临时张量,70多G真不奇怪。你要是想省显存得上FULL_SHARD,不过代价是通信开销会明显涨一波。另外forward_prefetch在这种单卡场景下基本没意义,因为不存在跨卡预取,backward_prefetch倒是能稍微压一下峰值,但幅度有限。我怀疑你对比DDP的时候是不是没把激活检查点算进去?这俩方案在内存管理上的差异其实挺大的。
7B LoRA用FSDP确实容易出现这种“反向优化”,因为LoRA本身可训练参数很少,FSDP分片节省的那点显存远不够抵消它额外存的原始权重和通信缓冲。你试试开CPU_OFFLOAD或者把auto_wrap_policy换成按层包装,有时候能压下来。另外确认下你设的sharding_strategy是不是真生效了,可以打印一下model.dp_state看看分片情况,我之前遇到过配置被后续的accelerate覆盖的情况。
这个现象我调FSDP时也踩过坑,SHARD_GRAD_OP只分片梯度,参数和优化器状态还是每卡全量复制,7B的LoRA虽然只训练部分层,但中间激活值才是大头,A100 80G跑满很正常。你试试把activation checkpointing打开,配合分片参数(FULL_SHARD),显存能直接砍半。另外forward_prefetch对单卡场景基本没用,它主要是跨卡通信优化,你多卡试试才看得出差别。
FSDP的SHARD_GRAD_OP本来就不分片参数,只分片梯度,所以峰值显存比DDP高是正常的,因为还要额外存分片状态和通信缓冲区。你试试FULL_SHARD,那个才是连参数一起分片的,LoRA场景下通常能压到40G左右。另外检查下是不是把use_orig_params设成True了,这个会影响参数分片后是否保留原图,有时候会多出不少显存开销。forward_prefetch这俩选项在单机多卡上确实感知不强,瓶颈多半在all-gather的通信开销上,不如调低bucket_cap_mb到25试试。
这问题我之前也踩过,SHARD_GRAD_OP这个策略本身就只分片梯度,参数和优化器状态还是全量驻留的,7B模型光fp16参数就14G,加上LoRA的额外激活和临时buffer,70G真不意外。你DDP显存低可能因为没开gradient checkpointing,或者batch size恰好卡在某个阈值下,FSDP的通信开销和分片元数据反而把峰值顶高了。我试过SHARD_OP策略,参数和梯度一起分片,单卡能压到40G左右,但通信频率明显上升,吞吐掉得厉害,小规模实验不值当。forward_prefetch和backward_prefetch在单机单卡场景下确实看不出差距,它们主要解决多机跨节点通信延迟,你这环境大概率不是瓶颈。建议先看下是不是activation checkpointing没开,7B的激活值在长序列下比参数还吃显存。另外检查下mixed_precision配置,如果FSDP用了bf16但LoRA层还是fp32,那部分参数会重复存储,显存就爆了。我最后是直接把sharding_strategy改成FULL_SHARD,配合gradient_checkpointing,单卡A100能稳在55G,虽然慢点但至少不OOM。你可以试试把bucket_cap_mb调小到10,减少通信峰值,有时候默认值太大反而拖累显存规划。
SHARD_GRAD_OP本来就不分参数,7B全量参数加激活值70G很正常,想省显存得开FULL_SHARD。
单卡跑FSDP分片反而多一套通信开销,你试试把activation checkpointing开开,显存立马能降不少。
这现象我之前也踩过坑,FSDP的SHARD_GRAD_OP只分片梯度,参数和优化器状态还是每卡全量复制,7B的LoRA虽然只训adaptor,但基座模型参数本身就得占满显存,再加上activation和碎片化内存,70多G真不奇怪。你可以试试把sharding_strategy换成FULL_SHARD,或者检查一下是否把auto_wrap_policy设成了按transformer层包装,否则FSDP只在最外层分片,效果会大打折扣。另外forward_prefetch对单卡场景基本没帮助,它主要是为了跨卡通信重叠,你如果卡数少,感知不明显很正常。
说实话这现象我踩过一模一样的坑,问题多半不在sharding策略上,而是FSDP的state_dict和优化器状态在初始化阶段会全量加载到当前rank上。你试试把limit_all_gathers=True打开,再把cpu_offload打开看看,显存能立刻掉下来。另外LoRA的target_modules如果覆盖了太多层,FSDP的forward预抓取反而会多放一份激活,跟prefetch关系不大。
遇到过类似情况,SHARD_GRAD_OP本来就不分参数,只分梯度和优化器状态,所以7B的权重还是全量驻留在每张卡上,显存自然比DDP高。另外LoRA本身可训练参数少,FSDP分片收益更小,但激活值、临时buffer这些反而可能因为通信调度多占一些显存。建议先确认下是不是开了activation checkpointing,没开的话激活显存在大batch下很夸张。另外可以试下把sharding_strategy换成FULL_SHARD,虽然通信开销大点,但单卡显存能明显降下来。forward_prefetch那些选项主要影响通信和计算重叠,对峰值显存影响确实不大。
这现象我上周刚踩过坑,SHARD_GRAD_OP本来就是只分片梯度,参数和优化器状态还在每张卡上,7B全量参数加LoRA的梯度累积峰值算下来确实比DDP还夸张。你试试把sharding_strategy换成FULL_SHARD,再配合activation_checkpointing,显存能降一大截。另外forward_prefetch这种开关对显存影响本来就不大,真正吃显存的是通信峰值和临时buffer,可以查下P2P的buffer大小设置。