最近在试着用PyTorch从头微调一个7B的LLaMA模型,单卡A100 80G,batch size只能开到2,再大就OOM了。我看网上都说用gradient checkpointing能省显存,但实际开了之后感觉效果不明显,显存占用还是70多G,而且训练速度慢了一倍。
用PyTorch训练LLaMA时,梯度检查点到底该开几层?显存还是吃不满
全部回复
共 151 条说实话gradient checkpointing本来就是拿时间换空间,7B模型开满理论上能省不少,但你这batch size才2,说明瓶颈可能不在激活值,反倒是在优化器状态和梯度上。试试用8-bit adam或者offload到CPU,把优化器状态瘦身一下,显存应该能降下来一大截。另外检查下是不是把checkpointing用在了embedding和norm层上,那些层其实没必要开,只对transformer block开就行,能省一半的计算开销。速度慢一倍是正常的,毕竟每个block要重算两次前向,但你要是显存没降下来那肯定哪儿没对,建议看一眼显存都被谁占了。
你这batch size太小了,开多少层checkpoint都白搭,先把梯度累积加上再试。
说实话70多G这个占用有点不对劲,7B模型就算不开梯度检查点,A100 80G跑batch size 2也不该这么吃紧。你确认一下是不是把optimizer state和gradient也算进显存统计了,或者检查下是不是用了全参数微调而不是LoRA,那玩意儿省显存效果比gradient checkpointing明显多了。
另外gradient checkpointing不是按层开的,是整网生效的,你只开几层反而会让某些层的激活值保留导致内存碎片化,建议直接设use_reentrant=True或者用torch.utils.checkpoint包住整个transformer block试试。速度慢一倍是正常的,它本质上是拿算力换显存,但如果你显存瓶颈没解决,说明卡点在别的地方,比如attention的KV cache或者输入序列太长。
我上次微调13B模型时,把gradient_checkpointing_enable()和model.gradient_checkpointing_enable()都调了,再把batch_size提到4,显存才真正降下来。你试试把torch.cuda.empty_cache()加在每step之后,有时候PyTorch的缓存不会自动释放,看着占用高但实际新tensor还能分配。
说实话我一开始也踩过这个坑,梯度检查点不是无脑全开就行的。它本质上是拿计算换显存,每个block都检查点的话,反向传播时重算的activation太多了,速度掉得厉害。我自己的经验是,7B模型在A100上,开一半层数,比如每隔一层开一个检查点,显存能省个10-15G,速度损失能控制在30%以内,比全开划算多了。
另外你batch size开不到2这个问题,可能不光是activation的锅。我记得LLaMA的embedding和最后的lm_head显存占用也很夸张,尤其是vocab size大的时候。你可以试试把这两个也包进gradient checkpointing里,或者干脆用tied embedding,能省不少。
还有个小技巧,如果你用的是PyTorch 2.0以上,可以看看torch.utils.checkpoint的use_reentrant参数,新版默认是False,但有时候手动设成True反而更省显存,虽然速度会更慢一点。你显存还吃不满的话,我猜可能是你把gradient checkpointing和mixed precision一起用了?这两个配合不好的话,会有一些隐藏的显存碎片。
最后想确认下,你用的flash attention吗?如果没开,光attention的中间矩阵就能吃掉一大块。开了flash attention之后,再配合选择性gradient checkpointing,我跑13B都能单卡勉强塞进去,你可以试试看。
我最近也踩过这个坑,7B模型开gradient checkpointing确实不是无脑全开就完事。你batch size才2的话,显存大头其实在优化器状态和中间激活值上,建议试试同时开offload optimizer到CPU,或者把checkpointing只放在前几层transformer block,后几层保持原样,效果会明显些。另外你观察下是不是PyTorch版本的问题,2.1之后对checkpointing的显存回收策略有优化,升级一下可能比调层数更直接。速度慢一倍是正常的,这玩意本质是拿计算换显存,如果显存还有冗余空间,不如把batch size提到4然后关掉checkpointing,说不定总吞吐反而更高。
gradient checkpointing不是这么用的,你开几层其实不是关键,关键是得配合activation checkpointing的粒度来调。7B模型在A100上batch size 2已经挺极限了,70多G的占用说明你checkpoint的层数可能压根没生效,或者你开的是full checkpoint而不是selective。我试过在LLaMA上把checkpoint开在attention和MLP的边界,而不是每个transformer block都开,显存能压到50G左右,但速度损失确实明显,这是正常的,毕竟是用计算换显存。你如果batch size必须上到4以上,建议把gradient accumulation加上,哪怕batch size 2配4步accumulation,效果也比硬开checkpoint强。另外你确认一下是不是开了torch.utils.checkpoint之后,前向传播里的activations没有被正确释放,有时候需要手动把checkpoint的keep_rng_state设为False,不然默认会保留RNG状态,显存反而更高。我怀疑你另一个问题是optimizer state占了很大头,7B模型光AdamW的fp32状态就快28G了,你试试用bitsandbytes的8-bit optimizer,配合checkpoint,显存能明显降下来。速度慢一倍这个无解,除非你换FlashAttention或者用torch.compile,但微调场景下compile的收益也不稳定。
说实话你这个情况我太熟了,之前用32G卡跑13B模型的时候也踩过这个坑。gradient checkpointing不是开了就完事,它默认是每个transformer层都检查点,但实际瓶颈往往在激活值最大的那几个层,比如embedding和最后的lm head附近,所以你得先profile一下看看显存到底被谁吃掉了。像你7B模型A100只开batch 2,我猜可能是序列长度太长,激活值集中在attention那块,这时候单纯开checkpointing确实省不了多少,反而因为重计算把训练时间翻倍了。你可以试试只对后半部分的层开checkpointing,或者配合梯度累积把有效batch提上去,显存占用不变但吞吐能拉回来一点。另外别忘了开torch.compile和flash attention,这两个对显存和速度的优化比checkpointing直接多了。还有个小技巧,把optimizer的momentum放到CPU上offload,能省出好几个G,虽然会慢一点但比OOM强。总之先别急着全开,用nsight或者pytorch的profiler看看每层的峰值显存,再决定哪些层值得重计算。
速度慢一倍正常,检查点按层开收益递减,试试只开后半段+offload优化器,能压到50G以内。
显存大头在激活值和优化器状态,光靠梯度检查点省不了多少,试试开offload或者换8bit优化器。
gradient checkpointing省的是激活值,优化器状态和参数还是大头,7B模型光这两项就占不少。试试把优化器换成8-bit Adam,显存能降一大截。
gradient checkpointing省的是激活值,不是模型参数和优化器状态。7B模型fp16参数加Adam的动量方差,光这些就占了大头,激活那块省下来的空间相对有限,所以显存看着没降多少很正常。你这70多G里,优化器状态估计就吃掉快40G了。想再压可以试试8-bit Adam或者deepspeed的offload,比单纯调checkpointing管用。