最近在试着用PyTorch从头微调一个7B的LLaMA模型,单卡A100 80G,batch size只能开到2,再大就OOM了。我看网上都说用gradient checkpointing能省显存,但实际开了之后感觉效果不明显,显存占用还是70多G,而且训练速度慢了一倍。
用PyTorch训练LLaMA时,梯度检查点到底该开几层?显存还是吃不满
全部回复
共 6 条我最近也在折腾这个,7B模型单卡A80想跑大batch确实挺吃力的。你开gradient checkpointing效果不明显,可能是默认只开了一部分层,建议试着把checkpointing的粒度调细,比如每个transformer block都做,或者配合activation offloading一起用。另外速度慢一倍挺正常的,毕竟是用时间换空间,我通常会在显存刚好够用的前提下尽量少开几层,平衡一下训练效率。
梯度检查点一般是隔几层开一次,全开反而慢,你试试每4层设一个checkpoint。
我最近也在搞类似的事,7B模型开梯度检查点其实不是层数越多越好,我试过每层都开反而显存回收效率不高,建议你试试只开后半部分的层或者每隔两层开一个。另外你batch size才2的话,显存吃不满可能是因为激活值本身占大头,检查点对这部分回收有限,不如把batch size试着提到4看看能不能跑动。速度慢一倍很正常,计算换显存嘛,可以调一下gradient accumulation步数来平衡一下训练效率。
实测LLaMA-7B开梯度检查点确实不会让显存降低太多,因为瓶颈主要在attention的KV cache和中间激活值上。你可以试试把checkpointing只用在transformer block的前几层,比如1-16层,后16层不开,这样显存能降不少。另外你batch size才2的话,检查点带来的额外计算量占比太大,速度肯定慢,建议先调低seq_len到1024或者512,把batch size拉上去再说。
说实话你这情况我太熟了,之前我调7B的时候也踩过这个坑。梯度检查点不是无脑全开就行的,它本质是用时间换空间,每多开一层计算图就得重算一次,速度降一半太正常了。我自己的经验是,LLaMA这种decoder-only结构,开中间8到12层左右性价比最高,既能把显存从70多G压到50多G,速度损失也控制在30%以内。另外你batch size才2就OOM,可能不光是激活值的问题,建议看看是不是优化器状态或者embedding层占了大头,比如用bitsandbytes的8-bit Adam能再省好几G。还有个骚操作是检查下你的hidden states有没有被不必要地保留,有时候DataLoader里多缓存几个中间变量就会把显存吃满。总之别指望单靠checkpointing一步到位,得和混合精度、activation offloading这些组合着来,我调完以后batch size能跑到6还不爆,代价就是代码改得跟千层饼一样。
你这batch size 2显存还70多G,感觉是梯度检查点没全开或者开的位置不对,llama默认的checkpointing只对某些层生效,你试试用model.gradient_checkpointing_enable()把所有transformer层都打开。另外速度慢一倍是正常的,因为它是用计算换显存,但你这显存没降下来就很奇怪,是不是序列长度太长或者把中间变量都保留了?