最近在试着用PyTorch从头微调一个7B的LLaMA模型,单卡A100 80G,batch size只能开到2,再大就OOM了。我看网上都说用gradient checkpointing能省显存,但实际开了之后感觉效果不明显,显存占用还是70多G,而且训练速度慢了一倍。
用PyTorch训练LLaMA时,梯度检查点到底该开几层?显存还是吃不满
全部回复
共 151 条开gc后显存还70多G不太对劲,你checkpoint的粒度是不是设成整个block了,试试按层拆开。
同款经历,我上次拿7B做QLoRA的时候也遇到这问题。梯度检查点不是万能的,它主要省的是中间激活值的显存,但你batch size只有2的话,可能瓶颈根本不在激活值上,而是优化器状态和参数本身占了很大头。你算算,7B模型光是fp16权重就14G,Adam优化器状态还要翻倍,这还没算梯度呢。所以开了检查点省下来的那点空间,根本不够你把batch size从2提到3的。另外你提到速度慢一倍,这太正常了,检查点本质是拿计算换显存,每个block都得重算一遍前向,而且我怀疑你是不是把检查点粒度设得太细了,比如每个transformer层都开了,其实隔几层开一个就能有不错的收益。还有个思路,你可以试试torch.utils.checkpoint的use_reentrant=False参数,有时候默认的reentrant模式在LLaMA这种结构上反而会额外占内存。再就是查一下是不是有显存碎片,A100上跑大模型经常这样,可以设PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128试试,说不定能挤出来几个G。反正我这边的经验是,梯度检查点更适合batch size已经比较大了,比如8或者16的时候去压显存,你这种本来就卡在2的情况,不如优先考虑offload优化器状态到CPU,或者干脆用8bit的AdamW。
之前也踩过这个坑,检查点不是开几层的问题,而是得配合activation checkpointing的粒度设置,PyTorch里默认是每个transformer block都存,你试下把checkpoint函数包在更细的算子层面,比如attention和mlp分开,能多省不少。另外你batch size才2,显存大头其实在优化器状态和中间激活值,检查点省的那点不够看,不如把input batch sequence长度砍半,梯度累积补回来,速度反而可能更快。还有个小技巧,开检查点后把model的memory_format换成channels_last,有的卡上能再挤点空间,不过A100上提升不一定明显。
同款问题,我之前调7B也是这个感觉,开了gc显存从75G降到68G,但速度直接腰斩。后来发现关键在配合batch size一起调,你试试把batch size翻倍到4,同时开gc,显存占用其实差不多,但吞吐量能拉回来一些。
另外检查下是不是所有层都开了,PyTorch里默认是每层都checkpoint,但有些层(比如embedding)没必要省,可以手动排除掉。还有个小技巧,把input checkpointing也打开,有时候比单纯开gc效果更明显。
不过说实话,70多G占用对7B来说确实有点反常,你是不是还把optimizer state算进去了?AdamW光这个就要吃掉好几G,试试8-bit优化器,能再挤出来十几G。
我最近也踩过这个坑,gradient checkpointing不是无脑全开就行的,它主要是省激活值显存,但7B模型本身权重和优化器状态就占大头,你这70多G估计大头不在激活值上。我建议先看看是不是batch size太小导致显存碎片化,或者把optimizer换成AdamW + 8bit,再配合梯度累积,可能比单纯开checkpointing更管用。另外你checkpointing是不是把每个transformer层都开了?试试只开一半层,比如每隔两层开一次,显存和速度能平衡很多。
gradient checkpointing这玩意儿对7B模型来说确实不是万能药,它省的是激活值内存,但你的batch size才2,说明大头可能根本不在激活值上,而是优化器状态和参数本身占了绝大部分。A100 80G跑7B微调,光AdamW的fp32状态就得吃掉差不多28G,加上模型权重和梯度,70多G很合理,checkpointing能抠出来的那点激活值空间真不够看。
我倒是建议你先用torch.cuda.memory_summary()看看内存到底分配在哪块了,如果发现是优化器状态占大头,那直接换8-bit Adam或者Adafactor,效果比开checkpointing立竿见影得多。另外你开checkpointing速度慢一半很正常,因为每个step都要重算前向,这本质是拿时间换空间,但你现在空间没省下来多少,纯亏。
如果你非要开,可以试试只对transformer层的后半部分开,比如最后16层,前面的层激活值复用率低,省不了多少,反而拖慢速度。还有一种思路是配合activation memory offload到CPU,但那样数据搬运也会变慢,得看你的瓶颈到底在哪。说实话,batch size 2跑7B微调已经算正常操作了,很多人用LoRA不就是为了绕开这个尴尬吗,你要是想硬刚全参数微调,不如考虑下梯度累积,把有效batch size提上去,显存占用不变但训练稳定性会好很多。
你这情况我太熟了,之前用32G卡跑13B的时候也踩过这坑。gradient checkpointing不是无脑全开就行的,它本质是用算力换显存,每开一层就多一次前向重计算,开太多反而把显存省下来的空间又用在了激活值缓存上,训练速度还直线往下掉。我自己的经验是,7B模型的话,从第4层或者第6层开始开比较合适,后面那些层的激活值占大头,开前面的层收益很小。另外你batch size只有2,说明瓶颈可能根本不在激活值,而在优化器状态和梯度本身,试试把优化器换成Adafactor或者8bit版,能省下好几个G。还有个小技巧,用torch.utils.checkpoint的时候别直接包整个block,把attention和mlp拆开分别checkpoint,这样能更精细地控制内存分配。你显存70多G是不是还开了什么别的?比如gradient accumulation没设对,或者max_length拉太长,有时候序列长度对显存的影响比batch size还夸张。最后说句实在的,A100 80G跑7B微调本来就该够用,如果checkpointing没效果,多半是别的地方漏了内存,建议跑个nsys看看内存峰值到底在哪一步爆的。
gradient checkpointing这个事儿吧,关键不在于开几层,而在于你把它跟什么搭配着用。你只开了checkpoint但没配合显存优化器或者微调技巧的话,效果确实会被吞掉大半——我猜你用的是AdamW吧?这玩意儿本身就要存两份动量,7B模型光优化器状态就得占掉差不多20G,你batch size 2加上激活值,70多G挺正常的。
我自己的经验是,先别纠结checkpoint层数,试试把优化器换成Adafactor或者8bit版,直接能省出十几G。另外你把checkpoint的粒度从整个transformer block改成attention和FFN分开试试,PyTorch 2.0以后支持这种细粒度控制,有时候省显存效果比无脑开全层更明显。
至于速度慢一倍,这个无解,本质就是拿时间换空间。但你可以试试把checkpoint的开销藏在数据加载的异步操作里,或者用torch.compile做一下算子融合,我这边实测能挽回大概20%的速度损失。
还有个坑你查一下,就是你的gradient checkpointing是不是真的在每层都生效了?有时候模型代码里写了但没hook住,或者被torch.no_grad()包住的部分跳过计算图,那等于白开。建议你打印下每层的activation memory分布,看看瓶颈到底在哪个模块。
最后说句实在的,7B单卡A100想跑大batch,光靠checkpoint不够,建议配合梯度累积把batch size顶到8以上,显存吃满但吞吐量反而上去了,训练稳定性也会好很多。
我之前也踩过这个坑,gradient checkpointing不是无脑全开就行的,它本质是用重计算换显存,但7B模型全开的话,每个transformer层都存两份激活值,反而可能让显存碎片化更严重。建议你试试只对后半部分层开,或者配合offload到CPU,我这边是开12层+offload,显存能压到50G左右。另外你batch size才2,是不是梯度累积步数设太大了?这也会让显存里的中间变量堆积,建议把累积步数调小点,直接加大有效batch。速度慢一倍是正常的,毕竟每个step都在做两遍前向,但如果你显存还没吃满,说明检查点层数选多了,可以拿nvidia-smi监控一下每层峰值,再慢慢调。
梯度检查点不是按层开的,是全模型开关,你开了之后batch还是2说明瓶颈在激活值之外,查查优化器状态和中间变量吧。
这情况我太熟了,gradient checkpointing不是无脑全开就行的,它省的是激活值内存,但7B模型光权重和优化器状态就快40G了,你batch size才2,溢出瓶颈根本不在激活层。我建议先看下你static_graph和model_parallel有没有开,另外试试把checkpoint只用在每个transformer block的前半段,别全层都开,速度能回来不少。还有就是A100上开torch.compile配合checkpoint有时候效果意外得好,你可以先跑个profiler看看哪块内存最吃紧。
开gradient checkpointing确实不是无脑省钱,它省的是中间激活值的峰值,但如果你batch size已经小到2,激活占比没那么大,省下来的就被计算开销吃回去了。建议你把checkpointing粒度从每层改成每隔几层开一次,比如2或4,再配合混合精度和optimizer state offload,显存能明显降下来。另外你确认一下是不是把整个模型都包进去了,有时候只开transformer block就够,embedding和lm head不用管。还有个小细节,开checkpointing后尽量把batch size往上调,不然速度损失全浪费了,我试过同样显存把batch翻倍,总吞吐反而差不多。
我之前也踩过这个坑,gradient checkpointing不是按层数开的,它是把整个forward的激活值按checkpoint点重新计算,你开几层其实意义不大,关键看有没有配合混合精度和offload。另外你这70多G显存可能被优化器状态和梯度占了大头,光开检查点肯定压不下来,建议先看看是不是bf16没开,或者试试把batch size拆成梯度累积。速度慢一倍是正常的,它就是用时间换空间,但你这显存没降下来就很怪,可能是checkpoint的粒度没设置对,试试把每个transformer block都包进去。
速度慢一倍正常,但显存没降说明你checkpoint粒度设太大了,试试把整个block包进去。
70多G还叫吃不满?开检查点省的是激活值,你batch size才2,省出来的空间本来就有限,不如直接冲batch 4试试。
检查点分层不是关键,你先确认下是不是把activation checkpointing设在每个transformer block上了,或者试试配合梯度累积把batch撑到4,速度损失换显存才划算。
说实话你这个情况我太有同感了,之前我调13B模型的时候也卡在类似问题上,折腾了一周才搞明白。梯度检查点这玩意儿本质上是拿算力换显存,但它的省显存效果跟层数选择关系特别大,不是无脑全开就行的。我之前看过一些benchmark,LLaMA这种结构里,开在注意力层和前馈层的交界处效果最好,你要是均匀分布反而会让显存曲线变得很怪。另外你提到batch size只能开2,我怀疑是不是激活值峰值出现在某个特定层,你可以用torch.cuda.max_memory_allocated打点看看具体是哪一段爆的,有时候光看总占用会误导。至于速度慢一倍,这个正常,检查点本身要重算前向,但我建议你试试把检查点跟混合精度、梯度累积配合起来,比如把有效batch size靠累积撑到8,这样吞吐量其实能回来不少。还有个细节,PyTorch 2.0以上版本对checkpointing有优化,你确认下是不是用了torch.utils.checkpoint而不是手写的那种。最后想问下,你用的是HuggingFace的transformers还是纯手写训练循环?后者的话灵活度更高,能自己控制哪些模块开检查点,实测能省到60G以下。
检查点不是灵丹妙药,你batch size太小,省下的显存不够塞新样本,速度还翻倍,不如直接开梯度累积。
说实话我一开始也踩过这个坑,gradient checkpointing不是让你无脑全开的,它本质上是拿计算换显存,你开太多层反而会让显存碎片化更严重,而且backward的时候重复计算特别耗时。我之前试过在7B模型上只开最后几层transformer block,效果比全开要好不少,显存能压到55G左右,速度损失也没那么夸张。另外你batch size才2的话,建议先查一下是不是activation的显存峰值出现在forward的中间层,而不是整个网络均匀分布的,可以试着用torch.cuda.max_memory_allocated去定位一下具体是哪块吃掉的显存。还有个细节,如果你用了flash attention,它本身对显存优化就很大,这时候再叠加梯度检查点可能收益就不明显了,我怀疑你可能是这个情况。还有就是你训练的时候有没有开混合精度?bf16和fp16的显存差距能有30%以上,这个比梯度检查点管用多了。最后想问问你用的是HuggingFace的Trainer还是自己写的训练循环?如果是Trainer的话,gradient_checkpointing_kwargs里可以传use_reentrant=False,这个选项在很多版本下能减少额外的显存开销,很多人不知道。
同款问题,我试过把checkpointing开到每一层,显存确实能压到50G左右,但训练慢到怀疑人生。后来发现关键是配合activation offload和混合精度,光开checkpointing不调别的收益很有限。另外你这batch size卡在2的话,可以试试把优化器状态换成分片式,比如用zero-stage2,有时比单纯开检查点更管用。还有个小坑,注意别把embedding和norm层也放进checkpoint范围,那些不占大头但拖速度。
我前几天也折腾过这个,7B用A100的话,gradient checkpointing不是全开就完事,得配合显存碎片优化看实际峰值。你把batch size提到4试试,开一半层数,再配上torch的max_memory缓存限制,效果可能比全开更明显。另外速度慢一倍正常,毕竟是用计算换显存,但70G占用确实偏高,检查下是不是activation offload没生效。