最近在试着用PyTorch从头微调一个7B的LLaMA模型,单卡A100 80G,batch size只能开到2,再大就OOM了。我看网上都说用gradient checkpointing能省显存,但实际开了之后感觉效果不明显,显存占用还是70多G,而且训练速度慢了一倍。
用PyTorch训练LLaMA时,梯度检查点到底该开几层?显存还是吃不满
全部回复
共 151 条这情况我调LLaMA也遇到过,gradient checkpointing不是无脑全开就完事的,它本质是用算力换显存,开太多层反而让反向传播重复计算量暴涨。你可以试试只对中间几层开,或者配合torch.utils.checkpoint的use_reentrant=False,有时候默认实现会有额外开销。另外你batch size才2,显存大头其实在optimizer state和激活值上,不如开个混合精度再加offload到CPU,效果可能比单纯开checkpoint更明显。速度慢一半挺正常的,我一般先把batch调到能塞下的极限,再拿checkpointing去补那最后几G,别指望它当主力。
我最近也踩过这个坑,检查点不是无脑全开就行的,它对激活值重计算的比例有讲究,7B模型建议先开一半层试试,比如每隔一层开一个,显存曲线会平滑很多。另外你batch size才2的话,吞吐瓶颈可能根本不在激活值,而是通信和算子效率,这时候开检查点纯属给CPU加负担。我试过把A100的persistent kernel和内存分配器调一下,甚至比单纯开检查点省得多,你可以看看是不是pinned memory没设对。还有就是70G显存其实还有优化空间,试试torch.compile或者混合精度下把optimizer状态换成bitsandbytes,比折腾检查点划算。
另一个思路,你确认下是不是把gradient checkpointing用在每个transformer层上了?有些实现只在特定模块生效,比如attention部分,MLP那块的激活值还是照常存,那省得就有限。我一般会配合activation offload,把中间张量挪到CPU内存,显存能压到50G以下,代价是训练慢30%左右,但至少能塞下batch 4。你那个慢一倍有点夸张,可能检查点实现里重复计算了loss或者embedding,试试把checkpoint包在更细粒度的子模块上,比如只包attention和ffn,不包norm和dropout。
开gradient checkpointing显存没降多少,大概率是activation峰值没压对地方,试试配合gradient accumulation和torch.utils.checkpoint分段包住block。
速度慢一倍正常,省显存就得拿时间换,但70G占用说明瓶颈在权重和优化器状态,不如直接换8bit优化器来得实在。
试过把checkpointing开到每一层(use_reentrant=False),显存确实能压到40G左右,但速度慢得离谱,后来干脆只在4-8层之间开,效果反而更好。另外你确认下是不是把activations也塞进checkpoint了,PyTorch默认只存中间变量,但某些自定义模块会漏掉该省的缓存。还有个小坑,batch size提到4试试,有时候显存没吃满是因为数据加载器或优化器状态占了额外空间,不一定全是激活的锅。
你这70多G挺正常的,7B模型光权重就14G,加上Adam的动量项和梯度,不开checkpointing理论峰值就该爆了。我一般开完检查点还会把gradient_accumulation_steps调成4,这样batch size降到1也能凑等效batch,速度损失反而能接受。对了,你用的是LLaMA原版还是HF的transformers实现?后者有些版本对checkpointing支持不好,得手动改代码。
开几层不如直接看activation峰值在哪个模块,A100显存带宽扛不住这个开销,慢是正常的。
你可能把checkpoint设到全模型了,试试只开Transformer层,embedding和lm_head别动,能省不少。
开gradient checkpointing不是按层数算的,是把整个模型的前向计算图全部重算,你只开几层反而会留下大量中间激活值,显存当然降不下来。A100 80G跑7B batch size 2本来就该够,建议你检查下是不是把优化器状态和梯度也塞进显存了,用bitsandbytes的8位优化器能省不少。另外速度慢一倍很正常,checkpointing本质就是用时间换空间,想兼顾的话可以试试把activation checkpointing和混合精度训练一起开,显存能压到40G左右。
梯度检查点不是按层开的,是整个模型全开才有效,而且要和gradient accumulation配合调batch size。
显存没降下来大概率是activation还在爆,建议用torch.profiler看下峰值在哪。
梯度检查点不是层数问题,是看激活重计算粒度,建议配合混合精度和梯度累积试试,速度慢是正常的。
建议查一下是不是把输入也塞进checkpoint了,只包transformer块能省不少,70G确实不对劲。
我之前也踩过这个坑,gradient checkpointing不是无脑全开就行的,它省的是激活值内存,但如果你batch size已经很小了,省出来的空间可能被优化器状态和参数梯度占掉,所以显存看起来没降多少。另外你速度慢一倍太正常了,checkpointing本质是用重计算换显存,建议你只对后半部分的transformer层开,比如最后16层,前面保持原样,效果会好很多。还有个小技巧,把optimizer换成AdamW的8-bit版本,能再挤出几个G,我之前就是这么把batch size从2提到4的。
梯度检查点不是全开就完事,你这情况更像是activation峰值没压下来,试着配合gradient accumulation把batch拆成1试试,7B模型在80G上batch 2本来就很极限了。另外检查下是不是把checkpointing用在embedding和输出层上了,那部分收益很小,主要应该放在中间的transformer block上。速度慢一倍是正常的,省显存本质就是拿计算换空间,你要是把batch降到1再开checkpoint,应该能明显看到占用降下来。
检查点不是开几层的问题,得配合梯度累积和混合精度一起调,单开确实感觉像没开。
试试把checkpointing放到每个transformer层都开,再用torch.compile,显存能压到50G以内。
开 checkpoint 也得配合梯度累积和 offload,光开几层没用,你这 batch 太小反而放大开销。
同感,光开gradient checkpointing确实感觉不明显,因为它默认是每个transformer layer都做重计算,但实际瓶颈可能在激活值最大的前几层。你可以试着配合activation_offloading或者手动调整checkpoint的粒度,比如只对后半部分的层做重计算,前半部分保留,显存能再压一些。另外batch size开不上去也可能是优化器状态和梯度本身占了不少,看看是不是用了AdamW的32位状态,换8位优化器比如bitsandbytes说不定能直接省10G+。速度慢一半是正常的,毕竟算力换显存,但如果你把use_reentrant=True加上,有时候能比默认的False快个10%-20%。你用的是HuggingFace的Trainer还是手写训练循环?如果是后者,可以检查一下是否把checkpoint包在了torch.utils.checkpoint里,但没关requires_grad的那些中间变量,那反而会占更多。
你这显存没降下来大概率是checkpointing只包了部分层,建议把整个模型都用torch.utils.checkpoint包一遍,另外7B在80G上batch size 2确实有点浪费,可以试试把activation checkpointing和gradient accumulation配合起来用,显存应该能压到50G以内。速度慢一倍是正常的,毕竟是用计算换显存,但你这情况可能更值得先查一下是不是哪里没开对,比如attention的flash attention或者混合精度有没有生效。
说实话你这个问题我踩过一模一样的坑,7B在A100上开gradient checkpointing省下来的显存远没有想象中多,因为LLaMA的激活值大头在attention的QKV投影和MLP中间层,而checkpoint只对每个transformer block的输入做保存,省的是反向传播时重建激活的算力换显存,但embedding和lm_head那两层是不受影响的,所以你会感觉瓶颈还在。另外你batch size才2,显存占用70多G,很可能实际大头是优化器状态和梯度本身,7B用AdamW的话fp32状态就要56G,加上参数和梯度,光这些就快60G了,checkpointing省出来的那点空间根本不够你再翻倍batch。我建议你先用torch.cuda.max_memory_allocated()看下峰值到底花在哪,如果优化器状态占大头,可以考虑换8bit优化器,比如bitsandbytes的AdamW8bit,能直接把优化器状态砍到十几G。速度慢一倍很正常,因为checkpointing意味着每个step要额外重算一遍前向,代价就是大概30%到50%的吞吐下降,你要是显存没吃满就别开,或者只开最后几层,比如12层里开8层,也能省一部分。还有一个思路是配合activation offload,把激活值放到CPU内存,但那样速度更慢,除非你显存实在不够用。说到底,单卡微调7B就是很尴尬,80G看着大,实际扛不住大batch,不如直接上LoRA或者QLoRA,显存占用能压到20G以下,训练速度还快。
说实话gradient checkpointing不是开了就完事的,它默认是每个transformer层都做检查点,但你可以通过配置只对部分层启用,或者配合activation offload把中间激活搬到CPU内存,这样显存能再压一截。你batch size才2的话,不如先看看是不是padding或attention mask导致的显存碎片化,有时候调整一下sequence packing反而更直接。另外开checkpoint后慢一倍很正常,毕竟是用计算换显存,如果显存还没吃满说明瓶颈可能在别的地方,比如optimizer state或者混合精度设置。
同感,之前我也在A100上跑7B微调,batch size卡在4上不去。不过你开了梯度检查点显存还70多G有点奇怪,建议确认下是不是真的生效了,比如看下模型内存分配情况,另外activation checkpointing一般配合torch.utils.checkpoint用,不是光设个参数就完事。速度慢一倍是正常的,毕竟重计算要额外跑一遍前向,但换来显存下降其实挺值。你试试把batch size往上加,比如到4或6,看能不能稳定跑起来,如果还是吃不满显存,那可能是数据加载或者优化器状态占了大头。
这问题我之前也研究过,其实梯度检查点开几层不是关键,关键是看你的显存瓶颈在哪,有时候是优化器状态和梯度本身占得太多。你可以用torch.cuda.max_memory_allocated()看下峰值分配,对比一下开和不开的具体差多少,如果只差几个G,那可能是你的模型本身就没把显存吃透,比如序列长度或者hidden size不够大。另外试试把检查点粒度调细,比如只checkpoint特定层,或者用混合精度加bf16,显存能再省不少。
你那个70多G的占用,我怀疑是开检查点之后显存没释放干净,或者你用的版本对LLaMA的attention实现有额外缓存。我自己的
梯度检查点不是按层开的,是整个模型开关,你八成是开完没配合梯度累积把batch怼上去,显存当然降不下来。
试试把checkpointing打开后把batch size翻倍,速度慢点但吞吐量上来了,70G显存确实不对劲,检查下是不是有张量没释放。
同样踩过这个坑,7B模型开gradient checkpointing理论上显存能省30%左右,但如果你用的是HuggingFace的transformers,默认的checkpoint实现是每个transformer层都做,不会让你选层数。你感觉效果不明显可能是显存瓶颈压根不在激活值,而在优化器状态和参数本身,A100 80G跑7B全参微调batch size 2确实差不多是极限了。想进一步压显存可以试试offload优化器状态到CPU,或者干脆用LoRA,速度慢一半是正常的,checkpointing本质就是拿计算换显存,建议先把batch size提到4再看收益。
说实话你这个现象我太熟了,之前我调13B模型的时候也踩过这个坑。gradient checkpointing的原理是牺牲计算换显存,但它的收益跟模型结构、batch size、序列长度都强相关,不是无脑开了就能砍半的。你batch size才2,显存大头其实是占在激活值上的,但7B模型本身参数和优化器状态就吃掉一半了,checkpointing省下来的那些激活值相对总占用来说比例没那么夸张,所以看起来不明显。
另外你提到速度慢一倍,这完全正常,因为每个step都要重算一遍前向,开销摆在那。我建议你查一下checkpointing的实际生效位置,PyTorch里默认是每个Transformer层都做,但你可以用checkpoint_activations参数去控制粒度,比如每隔几层才开一次,这样在显存和速度之间找个平衡点。还有个骚操作是把input和output的batch切碎,用梯度累积模拟大batch,虽然总吞吐没变,但每个step的峰值显存能压下来不少。
我猜你现在的显存瓶颈可能不在激活值,而是挂在KV cache或者attention的中间结果上,试试开flash attention或者把序列长度对半切,说不定比调checkpointing更管用。最后问一句,你训练时用的优化器是AdamW吧?如果是的话,把betas调成(0.9, 0.95)加上8-bit版,能再省出几个G。