最近在试着用PyTorch从头微调一个7B的LLaMA模型,单卡A100 80G,batch size只能开到2,再大就OOM了。我看网上都说用gradient checkpointing能省显存,但实际开了之后感觉效果不明显,显存占用还是70多G,而且训练速度慢了一倍。
用PyTorch训练LLaMA时,梯度检查点到底该开几层?显存还是吃不满
全部回复
共 151 条我最近也在折腾这个,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层都打开。另外速度慢一倍是正常的,因为它是用计算换显存,但你这显存没降下来就很奇怪,是不是序列长度太长或者把中间变量都保留了?
开了梯度检查点但显存降不下来,可能是你只对部分层做了checkpoint,建议把所有transformer层都包进去试试,效果会明显很多。速度慢一半是正常的,毕竟这是用计算换空间,但7B模型在A100上batch size只到2确实有点不对劲,我猜是不是activation offloading没开或者数据加载有瓶颈?另外可以检查下是不是混合精度没配置好,bf16能省不少。
说实话7B的模型在A100上开gradient checkpointing效果确实没想象中那么神,我试过好几轮,发现它主要省的是中间激活值那部分,但如果你batch size已经压到2了,激活占的比例本来就不大,省不出多少空间。反倒是你把micro batch size再拆小一点,比如试试gradient accumulation,可能显存能降得更明显,就是训练时间也会拉长。另外检查一下你是不是把所有transformer层都打了checkpoint,其实只对后几层开效果就差不多了,全开反而把计算图拖慢了。
我也是7B微调用A100,batch size 2开梯度检查点确实省得不多,主要因为模型本身没那么大,显存瓶颈反而在中间激活值上。你可以试试把检查点只开在最后几层transformer,前面保持不检查,这样能平衡速度和显存。另外检查下是不是开了全参数微调,用LoRA之类的方法配合检查点会省很多。
我试过类似的情况,梯度检查点不是层数越多越好,开太多反而让计算图重算开销变大,速度掉得厉害。建议你先从每2层检查一次开始试,或者直接用PyTorch自带的activation checkpointing,只对attention和FFN部分做检查点。另外显存吃不满可能是不在batch size上,得看看是不是数据加载或者模型并行那块有小瓶颈。
梯度检查点不是层数越多越好,开太多反而会让计算图反复重算,显存省不了多少速度还崩了。
你开的是整个模型还是只开部分层?我试过只开后一半层反而更省显存。
老实说我也踩过这个坑,gradient checkpointing不是无脑全开就好的。LLaMA的decoder层里,每层都有self-attention和FFN,如果你把checkpointing粒度设成整个layer,那其实每个子模块的中间激活还是会被保留一部分,省不了太多。我试过把checkpointing细化到每个子模块(比如torch.utils.checkpoint分别包住attention和mlp),显存能降到50多G,batch size能拉到4,不过速度确实慢了差不多一倍,这个没办法,时间换空间嘛。
另外有个细节容易被忽略:你开checkpointing的时候,是不是把input也设成requires_grad=False了?有时候前向传播里有些输入变量本来就不需要梯度,但PyTorch默认会保留它们的中间结果,显存就白占了。我一般会手动把embedding层的输出detach一下,或者用更细粒度的checkpoint_segments控制。
还有,你显存70多G其实说明模型本身和optimizer states已经吃掉大部分了,checkpointing主要省的是中间激活那部分。7B模型光参数fp16就14G,加上momentum和variance的Adam状态又要28G,再加上梯度本身,已经快50G了,剩下才是留给激活的。如果你开的是full batch gradient checkpointing,可能省出来的20G都被optimizer吃回去了,要不你试试把batch size调小但梯度累积步数拉长?我目前是batch size=1,gradient accumulation=8,配合单层checkpointing,A100上勉强能跑。
这情况我也遇到过,梯度检查点不是无脑全开就行,关键要看模型结构和显存瓶颈在哪。建议你试试只对transformer的后半部分层启用检查点,前面几层保留全激活,有时候能省不少显存还不怎么降速度。另外batch size小的话,检查点带来的额外重计算开销占比会变大,所以速度翻倍也正常。你检查下是不是把整个模型都包进去了?
这情况我遇到过,试试把checkpoint的层数调到4-8层,别全开,再配合activation offloading,显存能再降点。
调成half精度了吗?fp16或bf16能省一大截。另外梯度检查点建议直接开全部层,别纠结几层,这玩意儿就是拿时间换空间,速度慢一倍太正常了。你batch size上不去有没有试试梯度累积?2步累积也能等效4的batch,显存压力小很多。
你这batch size才2,显存大头其实在优化器状态和激活值上,光靠梯度检查点省不了多少。
开多了计算图重构太频繁,速度肯定掉,试试只开最后4-6层,显存能压到50G左右。
我最近也在调7B模型,梯度检查点对微调其实没那么神,尤其是你batch size已经很小的情况下,显存大头是optimizer states和激活值缓存,光靠checkpointing省不了多少。建议你试试activation offloading或者干脆用deepspeed zero stage 2,显存能直接降到40多G,速度影响也比checkpointing小。不过你70多G的占用确实有点奇怪,会不会是dataloader或者自定义模块里有什么显存泄露?
梯度检查点开满反而慢,试试只开后半层,显存和速度能平衡点。