最近在试着用PyTorch从头微调一个7B的LLaMA模型,单卡A100 80G,batch size只能开到2,再大就OOM了。我看网上都说用gradient checkpointing能省显存,但实际开了之后感觉效果不明显,显存占用还是70多G,而且训练速度慢了一倍。
用PyTorch训练LLaMA时,梯度检查点到底该开几层?显存还是吃不满
全部回复
共 151 条梯度检查点一般是每层都开效果才明显,你是不是只开了部分层?
这问题我踩过一样的坑,关键是“梯度检查点开几层”其实是个伪命题——它省显存的核心原理是用时间换空间,但如果你batch size本身就很小,或者模型里有些算子本来就不会存中间激活,那收益确实不明显。LLaMA 7B的transformer层里,attention和FFN的激活占大头,但如果你只对部分层开检查点,比如每隔两层开一个,反而可能让PyTorch的autograd图更碎片化,导致显存回收不及时。我自己的经验是,7B模型在A100上开全层检查点+gradient accumulation,batch size能翻到8甚至12,但前提是得配合activation offloading或者混合精度训练一起调。你提到的速度慢一倍,我猜是检查点层数没选对导致重计算开销太大,试一下只对attention部分开检查点,保留FFN的完整激活,有时候效果反而更好。另外可以检查下是不是开了检查点后数据加载成了瓶颈,torch.utils.data的num_workers调大点能缓解。
说实话我最近也在折腾这个,7B模型单卡A100开gradient checkpointing确实有那种“开了等于没开”的错觉。我后来仔细看了一下,发现它省显存的效果跟具体实现方式关系很大——如果你用的是PyTorch原生的torch.utils.checkpoint,默认是每个transformer layer都做一次重计算,但LLaMA的attention和FFN结构其实可以更精细地控制checkpoint的粒度。比如我试过只在attention部分做checkpoint,FFN保持正常前向,显存直接降到50G左右,速度损失也小很多。另外你batch size只能开到2,但显存还占70多G,说明可能有一些中间激活没被释放掉,可以检查一下是否用了torch.no_grad或者是否正确调用了checkpoint_sequential。还有个调参思路:不一定非要“全开”或者“全关”,比如只对后半段的层(比如最后12层)做checkpoint,前半段保持原样,这样显存能压到60G出头,训练速度只慢20%。当然这也跟你的模型具体版本有关,LLaMA 1和2的hidden size不一样,微调的序列长度也会影响收益。你用的序列长度是多少?如果超过4096的话,就算开了checkpoint,attention的KV cache还是吃很多,可能得考虑用FlashAttention或者换用更省显存的attention实现。
说实话7B单卡A80开梯度检查点吃不满挺正常的,这模型本身计算图就大,检查点主要省的是中间激活值,但你batch size才2,激活占的比例本来就不高,省出来的空间有限。我试过把batch size提到4再开检查点,显存能压到60G左右,速度慢是必然的,毕竟多了一次重计算。要不你试试把checkpointing只开在最后几层Transformer上,前面几层不动,这样能平衡一下速度和显存。
梯度检查点一般开一半的层数就行,全开反而慢,显存省不了多少。
说实话,你这个问题我踩过一模一样的坑。梯度检查点并不是层数开得越多越好,它本质上是把中间激活值丢掉,反向传播时重新算一遍,所以效果取决于你原始模型里哪些层最吃显存。LLaMA 7B的attention和FFN部分显存大头其实在中间激活,尤其序列长度一长,激活值比参数本身还占空间。我试过把checkpointing只开在最后几层transformer block上,反而比全开效果更明显,显存能压到50G左右,速度损失也小很多。另外你batch size只有2的话,是不是数据加载或者梯度累积那里没调好?有时候torch的显存碎片化也会导致看起来占用高但实际用不满,可以试试torch.cuda.empty_cache()或者调整max_split_size_mb参数。还有个小细节:你的mixed precision是bf16还是fp16?LLaMA对bf16支持更好,fp16容易出nan导致重算,反而更占显存。
我最近也踩了这个坑,梯度检查点不是全开就好的,一般建议只开在transformer层的前半段或者隔几层开一层,全开反而会让反向传播多算很多次,显存省不了多少速度还崩。另外你可以试试把activation checkpointing的粒度调细一点,配合混合精度和gradient accumulation,batch size应该能往上提一档。
我最近也试过类似的配置,感觉gradient checkpointing对7B模型来说,开个4到6层效果比较明显,全开反而容易让计算开销变大,速度掉得厉害。你显存还是70多G的话,可能batch size、序列长度或者优化器状态也有影响,比如混合精度或者8bit优化是不是没开全。可以试试先只开一半层数,配合梯度累积把batch size往上提,说不定显存能压到60G以下。
我最近也在折腾7B模型,梯度检查点默认是全开,但实际显存瓶颈往往在激活值缓存上,我试过只对最后几层开,显存能降到60G左右。不过速度确实慢,我感觉可以试试batch size开到4然后把梯度累积步数调小一点,说不定整体效率更高。你用的优化器是AdamW吗?那个二阶动量也挺吃显存的。
说实话你这个问题我当初也踩过坑,gradient checkpointing并不是无脑全开就行的,它本质上是拿计算换显存,但如果你模型本身的激活值已经很小了,比如batch size卡在2,那大部分显存其实是被优化器状态和参数本身吃掉的,激活值占比反而不高,这时候开checkpointing收益就很有限。我建议你先用torch.cuda.memory_summary()看看显存到底花在哪了,如果是优化器状态占了40G以上,那你该考虑用bitsandbytes的8-bit Adam或者Adafactor,能直接把优化器显存砍半。另外LLaMA这种7B模型,单卡A100想用更大batch size的话,可以试试混合精度训练加activation offloading,或者干脆把序列长度缩短一点,比如从2048降到1024,效果可能比折腾checkpointing明显得多。至于训练速度慢一倍,这是正常的,因为每层都要重新计算前向,你如果非要开,我建议只开最后几层或者每隔两层开一次,别全开,能平衡一点显存和速度。
说实话,你这个情况我太熟了。7B模型在80G卡上batch size只能开到2,明显是激活值把显存吃死了,但梯度检查点开了却没省下多少,我觉得可能是你只开了默认层级的checkpointing,没做精细化控制。LLaMA的transformer层里,attention计算和feed-forward部分的显存消耗差异很大,如果只对部分层启用,那省下来的显存其实很有限,而且反向传播时重计算的开销也会集中在某些层上,导致速度下降明显。
我之前试过在LLaMA上按层粒度手动配置,把每个decoder layer里的linear和attention都分别用checkpoint包裹,结果显存从78G降到了52G左右,速度虽然慢了大概40%,但总算能把batch size提到4甚至6了。另外你也可以看看torch.utils.checkpoint里的参数,比如use_reentrant设成False在某些场景下能减少额外内存分配,速度会好一点。
还有个坑是A100的显存分配策略,PyTorch默认会预占一部分作为缓存,你看到的70多G可能有一部分是预留的,实际占用没那么高。可以试试torch.cuda.empty_cache()或者调小PYTORCH_CUDA_ALLOC_CONF里的max_split_size_mb来更精细管理显存碎片。你现在的微调任务如果序列长度比较长,梯度检查点的收益会更明显,反之如果序列很短,那确实省不了太多。
对了,你用的gradient checkpointing是torch原生的还是HuggingFace封装的那个?HF的实现有时候会额外保留一些中间变量,反而抵消了省显存的效果,建议直接手动写checkpoint包围关键模块试试。
梯度检查点不是层数问题,是得配合activation checkpointing一起用,单开效果有限。
检查点开满也没用,瓶颈可能在数据加载或通信上,试试把batch size压到1看看显存变化。
我也遇到过类似的问题,后来发现梯度检查点并不是开越多层越好,通常建议只对注意力层或者FFN层做checkpoint,全开反而会让计算图变得碎片化。另外你可以试试把activation checkpointing和混合精度训练结合起来,我开了bf16之后显存直接从75G降到了50G出头。不过速度确实会慢一些,这是换显存的代价,可以接受的话调小一点batch size其实更省心。
单卡A100跑7B模型,batch size开2已经不错了,梯度检查点建议全开,不然显存省不下来。
我之前也遇到过类似的问题,后来发现梯度检查点并不是层数开得越多越好,关键要看模型里哪些模块是显存大头。你可以试试只对transformer block里的attention和FFN开checkpointing,embedding和lm_head这些地方不开,我这样调整之后显存直接从75G掉到55G左右。速度确实会慢一些,但感觉你batch size才2的话,可能调整一下换回更大的batch会更划算。
你这batch size 2确实太小了,梯度检查点省显存是有上限的,建议试试把activation checkpointing的粒度调细一点,比如每两层或者每层都做检查点,而不是默认的每层。另外可以配合gradient accumulation,实际batch size上去之后显存利用率会好很多,速度慢是正常的,属于用时间换空间。
我最近也踩过这个坑,梯度检查点不是无脑全开的,对LLaMA这种模型,建议只开在self-attention层,MLP层不开,这样显存能降到50多G,速度损失也小很多。另外你可以试试activation offloading,把中间激活值放到CPU上,配合梯度检查点效果更明显。你batch size才2的话,不妨再调小一下微调的学习率,看能不能稳定跑起来。
我之前也踩过这个坑,gradient checkpointing并不是对每层都开效果最好,特别是LLaMA这种大模型,建议只开在transformer block的后半部分或者每隔几层开一次,能省下不少显存又不至于让训练慢太多。另外检查下你是不是把checkpointing用在了embedding和lm_head上,那部分其实没必要开,反而拖慢速度。我试过在7B上只开最后16层,batch size从2涨到4,速度损失还能接受。
说实话你这情况挺常见的,梯度检查点不是无脑开满就行的,7B模型每层都开反而会因频繁重计算拖慢速度。我试过只对后半段Transformer层开检查点,前半段保留完整计算,显存能压到60G左右,速度损失也小很多。另外你batch size只有2的话,可以考虑先检查下activation checkpointing的粒度是不是默认的whole layer,改成按segment切分有时候能更精细地省显存。