最近在试着用PyTorch从头微调一个7B的LLaMA模型,单卡A100 80G,batch size只能开到2,再大就OOM了。我看网上都说用gradient checkpointing能省显存,但实际开了之后感觉效果不明显,显存占用还是70多G,而且训练速度慢了一倍。
用PyTorch训练LLaMA时,梯度检查点到底该开几层?显存还是吃不满
全部回复
共 151 条说实话你这情况我太熟了,之前我用7B模型在80G卡上跑也是这个感觉。gradient checkpointing不是说你开了就一定能压到很低,它省的是中间激活值的峰值,但如果你batch size本来就小,省下来的空间可能又被优化器状态和梯度给占了。我怀疑你显存大头根本不在激活值上,而是在Adam的动量项和方差项上,7B模型光这两项就要吃几十个G,你算算看是不是这个理。
另外你提到速度慢一倍,这太正常了,checkpointing本质是拿计算换显存,每层都得重算一遍前向。我自己的经验是别开全量,挑后半段网络开就行,比如从第20层开始开,前面层保持原样,这样显存能明显降下来,速度损失也小很多。你可以试试用torch.utils.checkpoint的checkpoint_sequential,配合layer数量动态调整,比手动包每一层灵活多了。
还有个小坑,你检查一下是不是把gradient_checkpointing_enable()放在model的device_map设置之前了,顺序不对的话有时候不生效,我踩过这个坑。另外batch size开到2的话,梯度累积步数调大点,比如4到8步,这样等效batch size变大,收敛也能稳一点,虽然单步慢但总时间未必差。最后建议你开个显存监控脚本,看看峰值到底出现在前向还是反向,定位清楚再调层数。
说实话你这个显存占用不太对,7B模型开gradient checkpointing之后activation应该能压到很低,70多G更像是没生效或者只开了一半。建议检查下是不是模型里某些子模块没包进checkpoint函数,比如embedding和norm层有时候会被漏掉。另外batch size 2的话,可以考虑配合offload或者混合精度一起用,单纯靠checkpointing省出来的空间有限,速度翻倍是正常的,毕竟用计算换显存。
说实话你这个问题我之前也踩过坑,gradient checkpointing不是说你开了就完事,它默认只对前几层生效,而且7B模型本身激活值大头在attention和FFN的中间层,你得用torch.utils.checkpoint把整个decoder block都包起来才行。另外显存吃不满70G可能是你batch size太小,反向传播时梯度累积占用的buffer没被释放,试试把gradient_accumulation_steps调大一点,或者用activation offloading到CPU。速度慢一倍是正常的,checkpoint本质就是拿计算换显存,但你如果显存还够用,不如把checkpoint层数减半,然后把batch size提到4,整体吞吐反而会更高。
说实话你这显存占用不正常,70多G基本等于没省下来。我怀疑你只开了模型层的checkpoint,但忘了把activation checkpointing也应用到embedding和norm这些模块上,这几个地方在7B模型里占的显存其实挺可观的。
另外你batch size才2的话,梯度检查点省下的显存会被额外保存的中间变量抵消一部分,建议你试试把gradient_checkpointing_enable加上use_reentrant=False,然后配合torch.cuda.empty_cache()在每步之后清一下碎片缓存。我自己的经验是7B模型开全部层能压到40G左右,你慢慢调吧。
说实话你这个现象我太熟了,刚接触gradient checkpointing的时候我也以为开了就能把显存打下来,结果一看nvidia-smi还是70多G,心态直接崩了。后来仔细看了下代码才发现,很多人光在模型外面包了一层checkpoint,但内部那些需要重新计算的激活值其实没被真正切分,尤其是LLaMA这种带了RMSNorm和旋转位置编码的结构,很多中间张量还是被完整保留在计算图里。
另一个坑是checkpoint的粒度,默认是按整个transformer层来的,但如果你显存瓶颈其实出在embedding或者最后的lm_head上,那再怎么开checkpoint都白搭。我建议你先用torch.cuda.memory_stats或者profiler看一眼是哪个模块占了peak memory,然后再决定是只对中间那几层开,还是把attention和mlp分开处理。
另外你提到batch size只能开2,这个其实不太正常,A100 80G跑7B理论上batch size到4或者8是没问题的,除非你序列长度特别长,或者用了很大的padding。你可以试试把序列长度动态裁剪、开flash attention,再把优化器改成AdamW的8bit版本,这几个组合拳下来显存能再省不少。
还有个小经验,gradient checkpointing对显存的节省是跟batch size强相关的,你batch size越小,它能省下来的比例就越低,反而计算开销是固定的,所以速度慢一倍完全正常。我之前试过把checkpoint和gradient accumulation配合起来用,视觉上显存是稳住了,但吞吐量反而更低了。
最后想问你一下,你用的是原生HuggingFace的LLaMA实现还是自己改过的模型类?如果你用的是transformers库,那个use_cache=True的默认配置在训练时其实会悄悄保留一堆缓存张量,很多人忘了把它关掉,白白多占好几个G。可以检查下这个细节,说不定显存还能再压一截。
检查点不是全开就完事,得配合重计算策略和batch梯度累积,试试只开中间几层再调大batch。
我试过类似的情况,关键不是开几层,而是看你的模型结构和优化器状态。7B模型开满梯度检查点理论上显存能压到40G左右,你如果只开部分层,那省下来的显存确实有限。
另外你batch size才2,说明瓶颈可能不在前向激活值,而是优化器状态和梯度本身。建议你查一下是不是用了AdamW的32位主权重,把优化器换成bitsandbytes的8位版本,或者开启梯度累积把有效batch size拉大,显存占用会更可控。
速度慢一倍是正常的,gradient checkpointing本质是拿计算换显存,你如果显存没降下来,那就是没开到位,或者根本没生效。可以打印一下模型显存分配看看具体哪块占大头。
我自己的经验是,如果显存卡在70G,先检查是不是忘了关gradient cache或者设了torch.no_grad,再考虑层数问题。你用transformers的from_pretrained加载的话,直接在config里设use_cache=False也很关键,这个经常被忽略。
你这batch size才2显存就70多G,检查点开几层都白搭啊,先看看是不是激活值缓存和优化器状态占了大头。
试试配合梯度累积把batch堆上去,再开全量检查点,速度慢点但显存能压到40G左右。
7B单卡A100开gradient checkpointing只省几个G确实正常,因为这东西的收益跟模型结构、activation大小强相关,尤其LLaMA的GQA设计本身就压了KV cache,省不出太多空间。你batch size开2显存还吃70多G,大概率是优化器状态和中间激活占了大头,建议用torch.cuda.memory_stats看下峰值分配在哪,别光盯着显存占用数值。另外checkpointing是拿算力换显存,你慢一倍是预期内的,但开了之后batch size应该能往4-6走才对,如果还只能开2,八成是没开对地方——比如只包了attention没包MLP块,或者没配合activation offload。还有个思路是试试flexible checkpoint,按层数比例动态切,但7B这种规模其实手写一个每4层checkpoint一次的配置就够了。以及确认下是不是用的HuggingFace的实现,有些版本对gradient_checkpointing的enable/disenable处理有bug,会重复计算。如果显存还是压不下来,干脆换AdamW的8-bit版,或者把序列长度裁短一点,比纠结checkpointing层数直接多了。
说实话我一开始也被gradient checkpointing给坑过,这玩意儿不是无脑全开就完事的。你开几层取决于你的计算图和激活值分布,7B模型transformer层中间那些大激活才是显存大头,但每层节省的量其实不太一样,我建议你用torch.cuda.max_memory_allocated去逐层打印一下,看看瓶颈具体卡在哪。
另外你batch size只有2,这本身就有点尴尬,gradient checkpointing是拿计算换显存,但如果你显存没压到临界点,它反而会拖慢速度,因为每层要重算前向。我自己的经验是,先把checkpointing开在最后几层(比如后20层),然后观察显存峰值变化,再逐步往前扩展,别一口气全开。
还有个容易忽略的点,你用的是不是最新的PyTorch版本?2.1以后对llama的checkpointing支持优化了不少,老版本可能用的是朴素实现,重算开销大。另外你可以试试配合activation offloading,或者把优化器换成Adafactor,8B模型用AdamW本身就要吃不少显存。
最后问一下,你用的是HuggingFace的Trainer还是自己写的训练循环?如果是Trainer,记得要把gradient_checkpointing_kwargs里的use_reentrant设成False,新版本的默认行为不一样,这个坑我踩过,显存能差好几个G。
说实话你这情况挺典型的,gradient checkpointing省的是中间激活值,但7B模型本身权重和优化器状态就占掉一大半了,80G卡跑2的batch也正常。你可以用torch.cuda.max_memory_allocated对比一下开前开后的真实峰值,如果只降了几个G说明瓶颈根本不在激活值上。另外试试把checkpointing只用在最后几层transformer block上,前面层保持全精度,这样能平衡速度和显存。速度慢一倍这个无解,本质是拿计算换空间,你要是训练步数不多不如直接开混合精度加梯度累积,可能反而更稳。
显存大头在优化器状态和激活值,你试试配合梯度累积把batch怼上去,速度损失能找补回来。
梯度检查点不是按层开的,它是把整个transformer block当成一个单元来重计算,7B模型全开的话省显存效果应该很明显才对,你显存还是70多G大概率是batch size或者序列长度本身就把显存占满了。另外你注意看下是不是开了checkpointing之后激活值确实降了,但优化器状态和参数占的那部分没变,A100 80G跑7B本来就很极限,建议先把序列长度砍半试试。速度慢一倍是正常的,毕竟算力换显存,但如果你发现显存没降下来,那可能是checkpointing没生效,检查下是不是只包了部分层或者模型forward里手动缓存了中间变量。
同感,梯度检查点这玩意儿开了确实慢,速度砍半很正常,但它省显存的逻辑是“用算力换空间”,如果你batch size只开2的话,省下来的显存可能根本不够再塞一个sample,所以体感不明显。我之前试过把gradient_checkpointing的粒度调细,比如只对前几层开,后几层保持原样,这样显存能降个10G左右,速度损失也小一些。另外你检查一下是不是把checkpoint用在了embedding和lm_head上,这两个地方其实没必要开,开了反而纯亏速度。还有个思路,试试torch.compile配合checkpointing,有时候能抵消一部分速度损失,不过7B模型编译时间挺久的。
开梯度检查点确实不是无脑开满就行,它省的是激活值重算的显存,但如果你batch size本来就小,省出来的空间可能被优化器状态和参数梯度占掉了。建议你查一下是不是把checkpoint包在了整个transformer层上,一般按层粒度开,比如每隔两层开一次,效果和速度能平衡些。另外70多G说明瓶颈可能不在激活值,试试看把optimizer换成Adafactor或者用offload,说不定比调checkpoint更直接。速度慢一倍正常,毕竟重算要时间,你这batch size才2,计算密度太低,慢是必然的。
我试过类似的情况,gradient checkpointing不是简单开了就完事,它默认是每个transformer层都做检查点,但实际可以配合use_reentrant=False或者手动指定层数,比如隔几层开一次,能平衡显存和速度。你显存没降下来,可能是激活值峰值还在,试试把batch size提到4或者8,配合gradient accumulation,让显存真正被榨干。另外A100上开bf16混合精度,比纯FP16省显存还稳,你用的是哪种精度?
这题我踩过坑,gradient checkpointing不是按层数开的,是按transformer block粒度设的,默认应该是每个block都开。你显存没降下来可能是activation checkpointing没生效,得看下模型代码里是不是用了torch.utils.checkpoint,或者检查下是不是梯度累积把batch撑大了。另外速度慢一倍正常,checkpointing本质是拿算力换显存,如果显存没降下来那肯定有地方没对。
我试过在llama上直接调use_cache=False,这个对显存影响也挺大,但微调时候其实可以关掉。你70多G是不是把优化器状态和梯度都算进去了?开混合精度bf16试试,A100对bf16支持很好,光这一项就能省不少。
说实话梯度检查点这玩意儿不是开几层的问题,是得配合显存分析和重计算策略一起看。你开完检查点但激活值还是全量存着,那肯定省不下来,建议先用torch.cuda.memory_snapshot看看具体哪块占着。另外你batch size才2,显存吃满是正常的,7B模型光权重和优化器状态就快30G了,剩下基本都耗在激活上。真嫌慢的话试试把checkpoint粒度放到每个transformer block内部,而不是整个模块,虽然代码麻烦点但能平衡速度和显存。
说实话你这情况我太熟了,之前用3090调7B的时候也踩过这个坑。gradient checkpointing不是无脑全开就行的,它本质是拿算力换显存,你开太多层反而会让反向传播重复计算一堆中间激活,速度掉一半太正常了。我自己的经验是先只开Transformer层里那些比较深的block,比如后一半层,前面浅层其实激活占不了多少空间。另外你batch size只有2的话,可以试试把gradient accumulation加上,显存瓶颈不一定全在激活值上,优化器状态和梯度可能才是大头,尤其你用AdamW的话每个参数要存两份动量。还有个容易忽略的点是LLaMA的embedding和最后的lm_head是共享权重的,这俩的反向传播也会攒不少tensor,可以手动对这部分做checkpointing。你要是方便的话,用torch.cuda.memory的堆栈分析看看到底哪块占的显存最多,别只看总占用,我之前发现有次是attention的key/value缓存没释放干净。最后想问你用的是HuggingFace的Trainer还是纯手写训练循环?后者的话可以自己控制checkpoint的粒度,有时候能省出好几个G。
说实话70多G这个占用不太正常,你确认一下是不是把整个模型参数和优化器状态都算进显存统计了?7B模型用AdamW光优化器就要28G左右,加上梯度、激活值,如果checkpointing没真正生效,那肯定还是吃满。
我试过类似配置,gradient checkpointing一般能把激活显存压到原来的1/3甚至更低,但前提是得配合activation checkpointing的recompute策略,而且batch size要相应调大才能体现优势。你速度慢一倍很正常,那是用计算换显存。
建议你先看看torch.utils.checkpoint是否真的包住了所有transformer层,同时把gradient_checkpointing_enable()放到模型加载之后调用。另外显存吃不满不一定是坏事,可能是你数据加载或mixed precision没设置对,检查下autocast和scaler。