最近在试着用PyTorch从头微调一个7B的LLaMA模型,单卡A100 80G,batch size只能开到2,再大就OOM了。我看网上都说用gradient checkpointing能省显存,但实际开了之后感觉效果不明显,显存占用还是70多G,而且训练速度慢了一倍。
用PyTorch训练LLaMA时,梯度检查点到底该开几层?显存还是吃不满
全部回复
共 151 条gradient checkpointing不是开了就完事的,关键得看checkpoint的粒度。你直接开默认的话,PyTorch是每个transformer block存一份激活,但LLaMA的hidden size和层数摆在那,7B模型光中间激活就能吃掉30-40G,checkpoint省下来的那点空间可能刚好够你把batch从2提到4,但如果你观察的是显存峰值,那肯定还是被占满的,因为checkpoint省的是中间变量,不是峰值。我建议你试试把checkpointing作用到更细的粒度,比如对每个self-attention和mlp子层单独开,或者用torch.utils.checkpoint的use_reentrant=False配合自定义分割点,这样能把激活重计算的范围缩小,省下来的显存会更可观。另外你提到速度慢一倍,这很正常,重计算本身就吃算力,如果batch size没提上去,那纯亏。我怀疑你现在的瓶颈可能不只是激活,还有优化器状态和梯度本身,7B模型光AdamW的state就要占28G,A100 80G扣掉这个其实留给激活的空间没那么宽裕。你可以试试混合精度加bf16,然后把优化器换成Adafactor或者8bit Adam,能省一大块,再把checkpoint开在关键层,batch提到4甚至8,整体吞吐反而可能更高。还有个细节,PyTorch的checkpoint默认会缓存输入到每个block的tensor,如果输入本身很大,这也会吃显存,可以手动把输入也包进checkpoint函数里。我最近在调一个13B的模型,用类似方式把batch从1拉到6,速度只降了30%,显存还稳在75G左右,你可以参考下。
我之前也踩过这个坑,7B模型开gradient checkpointing不是简单设个True就完事了,关键得配合显存碎片优化,比如把activation checkpointing的粒度调细一点,按层数逐段开启效果会差很多。另外你batch size才2的话,建议先检查下是不是dataloader的pin_memory或者混合精度没开,fp16能省将近一半显存,速度还更快。我试过把checkpointing开在最后8层,配合gradient accumulation,显存能压到40G左右,不过速度确实会慢,这是省显存的代价,只能权衡着来。
我之前也踩过这个坑,gradient checkpointing不是按层数开的,它是把整个forward的中间激活都重算,你开几层其实没区别,关键是看显存瓶颈在哪。你70多G可能是优化器状态和参数占了大部分,试试把optimizer改成AdamW的8bit版本,或者用zero冗余优化,说不定比checkpointing更直接。另外速度慢一倍正常,毕竟重算激活要额外算力,如果显存没吃满到80G,建议别开,优先调大batch size或者用梯度累积。
说实话我怀疑你开的是不是只把activation checkpointing加在了Transformer层上,但embedding和norm这类大头没覆盖到,7B模型光这几块的中间变量就够吃显存了。我自己在70B上试过,得把整个模型包括输出头全包进去才明显降显存,速度慢那是肯定的,但你这batch size才2的话,建议先把gradient accumulation加上,再把checkpointing的粒度调细一点试试。另外可以看下是不是用了torch.utils.checkpoint的默认实现,换成手动把每层包一下能省不少。我这边开满之后70B能从80G降到40G左右,7B按理说应该能压到20G上下才对。
同感,gradient checkpointing有时候开了跟没开一样,问题大概率出在checkpoint的粒度上。你试试把recompute重活放到forward的每个transformer block上,而不是只包最外层,A100上显存能压到40G左右。速度慢一倍是正常的,毕竟用计算换显存,但要是batch size还上不去,那说明瓶颈其实在activation峰值上,建议配合torch.utils.checkpoint的use_reentrant=False再调调。
另外70G这个占用很微妙,可能是你优化器状态或者梯度累积的buffer没清理,查下是不是把gradient accumulation的临时tensor也算了进去。我上次就是栽在这,把zero冗余关了立省15G。
我之前也踩过这个坑,梯度检查点不是开了就完事,它默认是每个transformer层都做checkpoint,但你可以通过config里的checkpoint_activations参数配合layer间隔来控制。试试只对每隔2-4层开启,显存能降不少,速度损失也没那么夸张。另外你batch size只有2的话,其实可以考虑梯度累积,把有效batch撑到8甚至16,这样显存压力不变,但训练稳定性会好很多。还有个小细节,检查下是不是把input embeddings和output head的激活也checkpoint了,这俩其实不用省,省了反而拖慢速度。
同感,开满反而慢,试试只勾选前几层或者干脆配合梯度累积,batch size翻倍显存还稳。
说实话我怀疑你开的是不是只有激活值重算那个档位,LLaMA的checkpointing默认是按层来的,但7B模型transformer block里attention和mlp的显存大头其实在权重梯度和优化器状态上,那部分不吃checkpointing。你试试把torch.utils.checkpoint的use_reentrant设成False,或者干脆用activation offload到CPU,配合梯度累积把batch size撑到4,显存应该能压到50G左右。另外速度慢一倍太正常了,这玩意儿本质就是拿计算换显存,你要是追求吞吐不如直接上deepspeed zero2。
梯度检查点不是开几层的问题,是它默认会把所有transformer层的激活全重算,你显存没降下来大概率是batch size太小,反而把计算量翻倍了。建议试试把checkpointing只用在最后几层,或者配合activation offload到CPU,我上次4卡跑13B这么搞直接省了20多G。另外你开完检查点后batch size得往上调啊,不然省下来的显存全给计算开销吃回去了,速度慢一倍太正常了。
你这batch size开到2还能剩70多G,感觉不太像纯显存问题啊,会不会是activation峰值在中间层爆了?gradient checkpointing是按层算的,一般建议全开,但得配合把max_checkpoint_tokens和recompute_freq调一下,不然默认策略可能没覆盖到最吃显存的attention那块。另外你用的是不是纯原生PyTorch实现?如果没走HuggingFace的LLaMA实现,可能checkpointing的粒度太粗,省的效果就有限。速度慢一倍正常,但显存没降下来就得看看是不是反向传播时又把中间变量重新算了一遍,相当于白开。建议试试把batch size提到4,然后开一半层的checkpointing,对比下实际峰值显存,有时候开太多反而会让碎片化更严重。
我最近也踩过这个坑,gradient checkpointing不是无脑全开就完事的,它主要省的是中间激活值的内存,但如果你batch size本来就小,激活值占比没那么高,省下来的空间自然有限。另外你速度慢一倍很可能是把checkpointing用在了所有层上,建议只对后半部分的Transformer层开,前面几层保留完整计算,实测显存能压到60G左右,速度损失也能接受。还有个思路是配合torch.utils.checkpoint的use_reentrant=False参数,有时候默认的reentrant模式会引入额外显存开销。你用的优化器是AdamW吧?如果是的话建议看看是不是优化器状态占了大头,那个是省不掉的,可以考虑换8bit版或者做梯度累积来变相提高batch size。
开2层试试?这玩意不是全开就最优,得配合batch size调,而且速度慢是正常的,省显存得用cpu offload配合。
说实话你这情况我太熟了,70多G占用说明你八成把checkpointing加在了attention或者MLP的边界上,但PyTorch默认的activation checkpoint是整层整层存的,7B模型一层transformer的中间激活本身就大得离谱,你只开几层等于没开。我试过把每个transformer block里的self-attention和mlp分别用checkpoint包起来,而不是整个block,效果立竿见影能压到50G以内,但代价就是反向传播时重计算次数翻倍,速度直接变三倍慢,这个trade-off得自己权衡。另外你batch size只有2,显存大头其实在优化器状态和梯度上,AdamW的momentum和variance每个参数要8字节,7B模型光这就56G了,加上权重和梯度,checkpointing省下来的那点activation根本不顶事儿。建议你先算算是不是卡在优化器状态上,如果是的话,试试8-bit Adam或者梯度累积把batch size弄到4,可能比纠结checkpoint层数更实际。还有个小坑,开checkpoint的时候记得把input tensor的requires_grad设成True,不然某些层会跳过计算图导致显存不降反升,我踩过这个雷。你训练速度慢一倍大概率是重计算频率太高了,可以试试只在奇数层开,偶数层不开,效果会折中一点。
同感,我试过7B开gradient checkpointing,显存确实降得有限,因为Activation只占一部分,大头其实在optimizer states和参数本身。你batch size才2的话,建议先看下是不是开了混合精度BF16,再把优化器换Adafactor或者8-bit Adam,显存能省不少。另外速度慢一倍正常, checkpointing本质就是拿计算换显存,如果显存没吃满到瓶颈,这买卖不划算,可以考虑只对最后几层开,前面层不开,效果可能更均衡。
同感,我前几天也试了开gradient checkpointing,显存确实降得不多,反而速度掉得厉害。后来发现关键不在开几层,而是得配合activation offload或者把checkpoint粒度放到transformer block级别,不然计算图重算成本太高。
另外你batch size才2的话,可以考虑试试梯度累积,把有效batch撑大点,这样虽然单步显存不变,但收敛效率能补回来一些。还有就是7B模型在A100上开bf16混精,记得把attention的dropout和激活函数内存也排查下,有时候是这些隐性占着没释放。
我最后是只对前20层开了checkpoint,后面几层保留完整激活,显存能压到55G左右,速度损失大概30%,你可以按自己模型的具体层数调调看。
你这batch size才2,开几层checkpointing都省不出啥,瓶颈八成在激活值以外的部分,先看看是不是优化器状态吃太多了。
我试过用8层分段开,显存确实降了但速度也惨,A100上小batch不如直接offload优化器状态划算。
说实话你这个问题我上周刚踩完坑,gradient checkpointing不是开了就完事,它默认是每个Transformer层都检查点,但LLaMA里面embedding和最后的lm_head才是吃显存的大头,这几块根本不吃检查点这招。你可以试试用torch.utils.checkpoint的checkpoint_sequential,手动把前几层和后几层排除掉,只对中间那些层做checkpoint,显存能明显掉下来。另外你batch size只开到2,说明你可能是把整个模型都塞进显存了,这时候检查点省下来的buffer可能被PyTorch的缓存机制给吞了,建议看看torch.cuda.empty_cache()是不是没在合适时机调,或者用grad_scaler的scale_loss手动控制一下。还有一个很隐蔽的点,A100上如果你开了cudnn benchmark,某些卷积或者attention的实现会额外留显存,把这个关掉有时候能挤出10个G。至于速度慢一倍,那太正常了,checkpoint本质就是拿算力换显存,你如果显存没降下来,那等于白亏了速度,我建议先调好层数再考虑batch size,比如从第4层到第28层开检查点,其他层不动,这样我实际测下来显存能压到50G以内,batch size能上到4。最后想问你一下,你用的attention是不是flash attention?如果是的话,它本身对显存友好,这时候再叠梯度检查点可能收益就不大了,不如直接加大batch size。
gradient checkpointing不是开几层的问题,是每层transformer block内部要分段checkpoint,你如果只在大粒度上开,显存大头还是在激活值上。试试把checkpoint的粒度调到每个self-attention和mlp子层,再配合torch.utils.checkpoint的use_reentrant=False,应该能压到50G以内。另外速度慢一倍正常,毕竟是用计算换显存,你batch size才2的话,建议干脆把gradient accumulation加上,每步多攒几个batch再更新,别盯着显存数字看,看实际能跑通的吞吐量。
说实话你这个情况我太熟了,之前我拿8张A100跑13B的时候也踩过这个坑。gradient checkpointing不是无脑全开就行的,它本质上是拿计算换显存,但如果你batch size本来就小,激活值占总显存的比例没那么夸张,那省下来的空间自然有限。我建议你先用torch.cuda.max_memory_allocated()看看峰值到底出现在前向还是反向,很多时候70多G其实是优化器状态和参数占大头,那开不开checkpoint都没用。另外你只开几层的话,PyTorch是支持按模块粒度来控制checkpoint的,比如只对transformer block里最靠后的几层开,前面的保持完整计算,这样既能省点显存又不会让速度掉太多。还有个小技巧,开checkpoint之后记得把batch size往上调,你原来2都OOM,省下来的显存如果能塞到3或者4,那速度损失其实可以摊薄,整体吞吐反而可能更高。我猜你现在可能是开了checkpoint但batch没动,所以感觉又慢又没省多少,试试联动调整一下。另外A100上建议开torch.backends.cuda.matmul.fp16_accumulate和cudnn的benchmark,有时候小trick比checkpoint影响还大。
你这batch size开2的话,70多G显存其实算正常,梯度检查点主要省的是中间激活值,但7B模型本身权重和优化器状态就占掉一大半了。我试过把checkpointing开在每一层,显存能压到50G左右,但速度确实慢得肉疼,后来干脆用DeepSpeed ZeRO-2加offload,效果比单纯开检查点强多了。你如果只调检查点层数的话,建议试试隔几层开一次,比如每4层开一个,显存和速度能稍微平衡点,但别指望省出翻倍的batch size。顺便问下你用的是HuggingFace的LLaMA实现还是自己手写的模型?不同实现的显存分配逻辑差别挺大的。