最近在试着用PyTorch从头微调一个7B的LLaMA模型,单卡A100 80G,batch size只能开到2,再大就OOM了。我看网上都说用gradient checkpointing能省显存,但实际开了之后感觉效果不明显,显存占用还是70多G,而且训练速度慢了一倍。
用PyTorch训练LLaMA时,梯度检查点到底该开几层?显存还是吃不满
全部回复
共 151 条同款经历,我之前用7B也这样,后来发现关键不在开几层,而是得配合torch.utils.checkpoint的官方实现逻辑,把每一层都包进去才有明显效果,但代价就是前向计算得重算一次,速度自然掉一半。
你显存还70多G的话,大概率是激活值没被真正释放,试着把gradient_checkpointing_enable()放在模型定义之后,同时把batch size提到4或者8,让显存利用率上去,速度损失才划算。
另外可以看下是不是优化器状态和梯度本身占了大头,开混合精度bf16能省不少,7B模型全精度的Adam状态就得吃掉好几G,省下来的空间可能比检查点更管用。
我之前还试过只开前几层或者后几层的检查点,结果显存波动很小,感觉不如全开,或者直接在transformers的training_args里设gradient_checkpointing=True,让框架自己处理,反而比手动调省心。
要是还不行,建议用torch.cuda.max_memory_allocated()看下峰值到底在哪爆的,有时候是数据加载或者loss计算那块临时张量没清干净,跟检查点关系真不大。
这情况我也遇到过,梯度检查点不是无脑全开的,它省的是中间激活值那部分显存,但你batch size才2的话,激活值本身占比就不大,省下来的空间很有限。真正吃显存的大头可能是优化器状态和参数本身,建议你先把7B的Lora或者QLora方案跑通,能明显降显存。另外速度慢一倍太正常了,checkpointing本质就是用计算换显存,我一般只在层数比较深或者序列特别长的时候才开,而且只开后半部分的层,效果会比全开好不少。
顺便问下你用的什么优化器?如果是AdamW的话,光状态就得占好几个G,可以考虑换8bit版或者用Adafactor,显存能再挤出来一些。
开了梯度检查点显存还70多G,那八成是batch size没跟着调大,省下来的显存又给吃回去了。
速度慢正常,检查点就是拿时间换空间,想省显存还是得配offload或换优化器。
我试过类似的情况,checkpointing不是开几层的问题,而是要和显存碎片化一起看。你batch size才2的话,可能大头在激活值以外的部分,比如优化器状态和梯度,建议先排查一下是不是RMSNorm和embedding的显存没被释放。另外速度慢一倍是正常的,它本质上是拿时间换空间,如果显存还吃不满,倒不如试试把checkpointing只开在decoder layer的FFN部分,效果差别挺大的。
梯度检查点不是按层开的,是整体开关,你显存没降下来可能是batch里的sequence太长,试试把max_len砍半。
检查点省的是激活值,你70G说明大头在优化器状态和参数本身,开不开区别不大,得上LoRA或者offload。
同款配置,我试过把checkpointing开到每层都开,显存确实能压到50G以下,但速度掉得离谱。你试试只对attention层开,MLP层关掉,效果会平衡很多。另外batch size才2的话,可能瓶颈不在激活值,而是优化器状态和梯度本身,建议开个混合精度或者offload优化器试试。
开 checkpoint 不是无脑全开,得配合 offload 和 batch 大小一起调,你这显存还是 70G 八成是激活值没省到点上。
gradient checkpointing不是按“层数”来开的,它是把整个模型的transformer block按一定粒度做重计算,你开几层这个说法本身就不太对。我怀疑你实际开的是torch.utils.checkpoint的默认用法,但它对LLaMA这种结构其实收益有限,因为LLaMA的激活主要消耗在attention的QKV投影和FFN的中间层,这些如果你没手动把checkpoint包在更细的模块上,它只会对每个block的边界做重计算,省下来的显存自然不明显。我之前在30B模型上试过,把checkpoint粒度拆到每个Linear层,显存能降40%以上,但速度确实会慢,这个没法避免。另外你batch size只有2,说明你可能是把整个序列都塞进去了,LLaMA 7B在A100上其实可以试下gradient accumulation配合更大的batch,或者检查下你的activation checkpointing是否真的生效了,可以用torch.cuda.max_memory_allocated对比一下开和不开的峰值。还有个点是,如果你用了flash attention,它本身已经省了一部分显存,再叠gradient checkpointing的收益就会变小,这时候瓶颈可能在优化器状态或者KV cache上。建议你先把每层激活的显存profile打出来看看,用torch.profiler,别光看总占用。
你这情况我前两天刚踩过坑,7B模型开gradient checkpointing确实不是无脑全开就行的,我试下来是每4层开一个 checkpoint比全开省显存,速度损失也小得多。另外你batch size只有2的话,建议把activation checkpointing和gradient accumulation配合着用,单卡80G其实能塞下更大batch的。还有个小细节,检查下是不是把input tensor也checkpoint了,有些层比如attention的qkv投影其实没必要重复算,能省不少。
我怀疑你显存没降下来是因为checkpoint的位置太密了,我之前对比过,7B模型每8层开一次就能省15G左右,全开反而因为重计算开销太大,显存瓶颈转移到临时张量上了。另外你试试把torch.utils.checkpoint的use_reentrant=False加上,新版PyTorch默认行为变了,这个参数影响挺大的。速度慢一倍有点夸张,也可能是你开了之后batch size没调,导致算力利用率更低了。
说实话70G这个占用有点偏高,我同样配置下开checkpoint能压到55G左右。你确认下是不是把embedding和lm_head的梯度也存了?这两个大矩阵其实可以单独设requires_grad=False或者用混合精度绕过。还有个小技巧,把dropout层临时换成identity,训练完再换回来,能省不少显存,
开一半层数试试,全开反而慢,而且70G说明瓶颈可能不在激活值,先看下是不是优化器状态和梯度撑爆的。
梯度检查点不是银弹,得配合batch内梯度累积和torch.compile一起搞,另外检查下是不是没开混合精度。
gradient checkpointing的原理是牺牲计算换显存,按理说显存占用应该显著下降才对,你开了之后还是70多G,大概率是没吃到点上。我怀疑你只对transformer的某些层开了checkpoint,但embedding和lm_head这些大头还在硬扛,或者你实际开的是use_reentrant=False这种老接口,新版本里行为有变化。另一个坑是,7B模型在A100上就算全开checkpoint,激活值省下来的空间也会被优化器状态和梯度占掉一大半,你batch size 2的时候,光AdamW的momentum和variance就得吃好几G,这部分是省不掉的。建议你直接看torch.cuda.max_memory_allocated()和torch.cuda.memory_reserved()的峰值,确认是激活值还是参数/梯度在占显存,别凭感觉。另外你开完checkpoint速度慢一倍很正常,因为每个step都要重算前向,如果显存还是没降下来,那可能是你checkpoint的粒度太粗,比如按整个layer来checkpoint,换成按block内子模块粒度会好很多。我之前试过在70B模型上开全量checkpoint,显存能从80G降到45G,但你的7B只降了这么点,感觉有点反常,建议你把torch.utils.checkpoint的use_reentrant=True加上,再检查一下是不是被model.gradient_checkpointing_enable()这个API的默认参数坑了。还有个小技巧,配合--gradient_accumulation_steps把batch size降到1,然后开混合精度,显存压力会小很多,速度反而可能比现在快。
说实话你这个问题我前两天刚踩过类似的坑,7B模型开gradient checkpointing确实不是无脑全开就行的。我自己的经验是,这玩意儿本质是用计算换显存,你batch size才2的话,本身能省的空间就很有限,而且每个checkpoint层在反向传播时要重算一遍前向,速度掉一半太正常了。你要是真想省显存,不如先看看是不是activation显存占了大头,像我之前就是没开torch.utils.checkpoint的use_reentrant=False,导致某些算子根本没被包进去,省了个寂寞。另外你A100 80G跑7B batch=2,理论上不该这么吃紧,建议检查下是不是flash_attention没开,或者rope_scaling、use_cache这些参数在训练时被错误启用了。我之前试过把gradient checkpointing只开在最后四层transformer上,显存能压到60G左右,速度损失也比全开小不少,你可以试试这个思路。还有个小技巧,混合精度用bf16的话,可以配合gradient_accumulation_steps把有效batch撑大,显存占用其实能稳很多。最后问一句,你用的是HuggingFace的Trainer还是纯手写训练循环?如果是Trainer,有些显存优化选项默认是关的,得手动调一下。
我最近也在调这个,7B单卡A100开gradient checkpointing确实省不了多少,因为瓶颈不在激活值,反而在optimizer states和参数本身。你可以试试把checkpointing粒度调细一点,比如配合torch.utils.checkpoint的use_reentrant=False,或者手动分段checkpoint,别一整层都包进去。另外速度慢一倍正常,这是拿时间换空间,你要是显存没降下来,八成是batch size没调大,开了之后应该能往4甚至8冲才对。
我怀疑你开的是不是只有transformer层,embedding和输出层没带上,那省下来的都是小头。建议用activation checkpointing的时候把注意力里的中间张量也考虑进去,配合混合精度和gradient accumulation,batch size能翻倍才说明开对了。你显存70多G可能是碎片化严重,试试把torch.cuda.empty_cache()加进训练循环里。
你这情况我遇到过,可能是PyTorch版本问题,老版本对LLaMA这种pre-norm结构的checkpoint支持不好,建议升到2.1以上。另外开几层不是关键,关键是要用selective checkpointing,比如只对每4层里挑1层做,这样显存降得明显,速度损失也小。你试试把checkpointed层数调到总层数的25%,然后batch size提到4,看看是不是就吃
检查点不是层数问题,是得配合把激活重算和offload一起开,你这显存大头在优化器状态上。
试试gradient_checkpointing加torch.utils.checkpoint包住整个transformer block,别只开几层,效果差很多的。
gradient checkpointing不是让你无脑全开的,它省显存的核心是“用计算换存储”,你开太多层反而会让activations被反复重算,速度直接腰斩不说,显存瓶颈可能根本不在transformer层的activations上。7B模型在A100上batch size 2还吃70多G,我猜你八成是把optimizer states和gradients也算进去了,AdamW的momentum和variance光这两项就要占掉参数量的8倍字节,也就是56G,这还没算model参数和中间激活。建议你先用torch.cuda.max_memory_allocated()分段打点,看看是forward时爆的还是backward时爆的,我怀疑你甚至不需要开gradient checkpointing,把optimizer换成Adafactor或者用bitsandbytes的8-bit Adam,显存能直接砍半。另外如果你用的是HuggingFace的trainer,它默认的gradient_checkpointing_enable()是全层开启的,你手动改成每隔几层开一次,或者用torch.utils.checkpoint的checkpoint_sequential按块切分,效果会好很多。还有个小坑,开了gradient checkpointing之后要记得把model.train()里的use_reentrant设成False,不然PyTorch 2.0以上的版本会有额外的显存碎片。最后建议你把batch size提到4试试,因为每个batch的fixed overhead(比如位置编码、attention mask)其实不小,batch太小反而浪费,说不定显存占用曲线会突然平缓下来。
这情况我也踩过坑,gradient checkpointing不是简单开几层的事,关键看你的activation到底占了多少。建议先开个torch.profiler看看内存峰值出现在哪,有时候是优化器状态和梯度本身吃大头,checkpoint省的那点根本不够看。另外你batch size才2的话,试试gradient accumulation把有效batch提上去,显存占用不变但吞吐能回来一部分。速度慢一倍正常,毕竟是用计算换显存,但7B在80G上不至于这么紧张,检查下是不是max_seq_len设太长或者attention实现没优化。
开 checkpoint 是拿时间换空间,你这 batch 才 2 说明瓶颈可能在激活值,试试配合 offload 或者把输入序列截断看看。
建议先检查下是不是开了 checkpoint 但没设 use_reentrant=False,或者干脆用 torch.utils.checkpoint 手动包住几个大层,别全开。
说实话你这个情况我调参的时候也遇到过,后来发现gradient checkpointing不是开几层的问题,而是得配合batch size往上拉才划算。你试试把batch size翻倍到4,然后开满全部层,显存应该能压在60G左右,速度反而比现在快。另外别忘了把activation checkpoint的recompute改成full,不然省的那点内存都被碎片吃掉了。
开一半层试试,全开反而慢且省不了多少,另外batch size提不上去可能不是显存瓶颈。
检查下是不是激活值峰值在作怪,试试用torch.cuda.max_memory_allocated看下实际峰值在哪。
检查点不是开几层的问题,得配合gradient accumulation把batch怼上去才划算,光开它速度必然崩。