最近开始尝试用Llama 3.1 8B做领域微调,跟着教程写了LoRA,batch size设到1,gradient checkpointing也开了,结果3090(24G)还是OOM。我的输入长度大概2k tokens,是不是跟序列长度有关?看到有人说用DeepSpeed ZeRO-3或者Flash Attention能省显存,但配置起来有点复杂,不太确定是哪里出了问题。另外,torch.compile会有帮助吗?现在用的是PyTorch 2.1,cuda 12.1。求有经验的大佬指点下常见坑,谢谢!
刚转大模型方向,用PyTorch跑LLM微调,显存总爆掉怎么优化?
全部回复
共 178 条24G跑8B LoRA还爆显存,大概率不是模型权重的问题,而是激活值和优化器状态在作祟。2k序列长度确实是个隐形杀手,attention的计算量是平方增长的,我猜你大概率没开gradient checkpointing的use_reentrant优化,或者LoRA的target modules选得太宽了。DeepSpeed ZeRO-3其实没那么玄乎,先试试stage 2,把优化器状态和梯度分片出去,比stage 3配置简单得多,效果立竿见影。Flash Attention倒是强烈建议上,哪怕只省下那部分显存,可能就够你塞下batch size 2了。torch.compile在Llama这种架构上收益一般,但开了也没坏处,就是编译时间有点长,可以先放一边。还有个容易忽略的坑:检查下你的max_seq_len是不是被某个默认值撑大了,有时候数据没截断,实际长度远超2k。另外,试试把gradient_accumulation_steps设小点,先确保单步不爆,再谈吞吐量。我自己的经验是,24G跑8B微调,LoRA rank=8、target只选q和v、加上ZeRO-2和Flash Attention,基本能稳定到4k上下文。你先按这个组合排查下,大概率能救回来。
24G跑2k长度的8B LoRA确实有点紧,但按理说不该直接OOM,你检查下是不是lora模块加到了所有linear层,或者attention的kv cache没开缓存优化。Flash Attention对长序列是质的提升,强烈建议先上这个,配置不难,改几行就行。ZeRO-3我觉得暂时没必要,LoRA本身参数就少,主要是激活值占显存,你可以试试把gradient checkpointing设成分片策略而不是全量。torch.compile对显存帮助不大,但能提速度,等能跑通再折腾吧。还有个常见坑是输入长度2k但没设padding到固定长度,动态padding能省不少。
你这配置跑8B LoRA其实不算离谱,但2k序列长度确实是个隐形杀手,激活值在反向传播时会把显存吃满,建议先把max_length砍到1k试试,很多任务其实用不到那么长。Flash Attention值得折腾一下,能省不少激活显存,而且和LoRA兼容性很好;DeepSpeed ZeRO-3对单卡反而可能更慢,不如先试ZeRO-2或者直接开offload。torch.compile对显存优化不大,但能省点计算时间,不过第一次跑会编译很久,别被吓到。另外检查下你是不是把label也拼进输入了,很多人会忽略这个细节导致序列翻倍。
Flash Attention基本是必装的,能省不少显存,另外torch.compile对长序列也有奇效,值得试试。
24G跑8B LoRA还爆,大概率不是显存总量问题,是激活值峰值太高。2k长度确实有影响,建议先试试把max_length砍到1k或512看还爆不爆,能快速定位。Flash Attention值得装,省显存效果立竿见影,而且你cuda 12.1版本装起来应该不难。ZeRO-3对单卡LoRA其实没必要,反而可能拖慢速度。torch.compile可以先放放,它优化的是计算速度不是显存,而且跟某些库兼容性容易出问题。
显存爆掉大概率就是序列长度和attention的平方复杂度在作怪,2k tokens在8B模型上确实挺吃紧的。Flash Attention值得花时间配一下,能把activation内存降一个量级,比ZeRO-3更直接解决你的问题。torch.compile对显存优化帮助不大,但能提升训练速度,建议等模型跑通了再研究。另外检查下是不是把LoRA的target modules设错了,有时候默认只改attention层,但embedding和lm_head也会占不少activation。我之前遇到过类似情况,把rope和swiglu的gradient checkpointing细分打开就好了。
24G跑8B LoRA按理说够,但2k序列长度确实是个坎,Flash Attention能省不少显存,建议优先搞这个,配置起来其实比ZeRO-3简单。torch.compile对显存帮助不大,主要是提速,但跟某些库可能不兼容,先别急着开。你检查下是不是LoRA只挂在了attention层,或者把target modules全加上试试,有时候默认配置导致激活值没被优化。另外,gradient checkpointing记得配合batch size=1,但可以试试梯度累积,别让显存白闲着。
3090跑8B还开2k长度确实挺极限的,LoRA虽然省了训练参数但激活值照样吃满,gradient checkpointing只省了中间激活没省attention那块。建议先确认下是不是峰值在forward的attention计算上,Flash Attention在这点上改善非常明显,代码改动也不大,比直接上DeepSpeed省事多了。torch.compile对显存帮助有限,主要是提速,而且跟某些库兼容性有问题,不如先把序列长度砍到1024试试水。另外检查下有没有把label padding到max length,很多新手会在这里白白浪费显存。
8B全参微调24G就是紧,先换4bit QLoRA试试,flash attention必须上,能省不少。
24G跑8B LoRA还爆显存,大概率不是batch size的问题,2k序列长度对attention来说确实吃紧,但更可能是你的LoRA配置里target modules选太多了,或者把bias也训了。我建议先开NVIDIA的nsight看看具体哪块显存峰值最高,一般这种OOM都是activation峰值爆的,gradient checkpointing虽然开了但可能没生效在正确的层上。Flash Attention值得折腾一下,它对长序列的显存优化是数量级的,而且现在flash-attn库安装没那么坑了,直接pip装就能用。DeepSpeed ZeRO-3我倒觉得先别碰,那个是给多卡或多节点准备的,单卡上反而可能因为通信开销变慢。torch.compile对显存帮助有限,主要是提速,你的瓶颈在显存不在算力。另一个容易被忽略的坑是optimizer state,AdamW的动量项会占不少显存,试试8-bit优化器比如bitsandbytes的AdamW8bit,能省一大截。我自己的经验是先把max_length砍到1024跑通流程,再逐步加长,这样能快速定位到底卡在哪。
3090跑8B长文本确实紧,Flash Attention必开,另外torch.compile能省不少显存,试试看。
24G跑8B LoRA还爆显存,大概率不是配置问题,是序列长度在作祟。2k tokens对Llama来说不算短,attention的显存占用是跟序列长度平方增长的,你试试把max_length截到1k或者512,显存应该立刻降下来。Flash Attention确实值得上,能省不少显存,而且现在huggingface的trainer里只要传个参数就能开,不用自己改模型结构,比DeepSpeed好配置多了。至于ZeRO-3,对LoRA来说其实有点杀鸡用牛刀,LoRA本身可训练参数就少,ZeRO-3更多是为了全参数微调准备的,你先别折腾这个。torch.compile我建议先别开,它跟gradient checkpointing有时候会有兼容性问题,而且对显存优化帮助有限,主要提升的是计算速度。另外你检查下是不是把optimizer states算进去了,AdamW的momentum也挺吃显存的,可以试试8-bit optimizer,bitsandbytes库一行代码就能换,显存能省好几个G。还有个常见坑是label和input_ids一起过模型,导致label也占了一份显存,把label设成-100或者用tensor parallel里的ignore_index处理下。我前几天刚用类似配置跑了7B模型,序列长度1k,batch size 2都没问题,你先把长度砍下来试试,肯定能跑通。
2k长度确实吃显存,先换flash attention试试,能省不少。torch.compile对显存帮助不大,主要提速度。
8B全参微调24G确实紧,但LoRA还爆大概率是序列长度和激活值没处理好,先上Flash Attention试试,能省不少。
输入长度2k确实是大头,先试试Flash Attention,基本能砍掉一半激活显存。torch.compile对显存帮助不大,别指望它救急。
这配置按说跑8B LoRA不该爆的,我怀疑你是不是把LoRA直接加在了所有linear层上,试试只target q_proj和v_proj,参数量能少一半。序列长度2k对显存影响确实大,Flash Attention能省不少,但更关键的可能是你优化器状态没做offload,AdamW的momentum在8B模型上很吃显存。torch.compile对推理提速明显,但训练时显存优化有限,不如先把gradient checkpointing的granularity调到每个transformer block试试。另外检查下是不是输入padding没mask掉,有些教程里细节很容易漏。
24G跑8B LoRA按理说是够的,你试试把输入长度砍到1k或者用梯度累积模拟更大batch,很多时候OOM是activation峰值爆的。Flash Attention确实能省不少,尤其长序列下效果立竿见影,torch.compile建议先别开,PyTorch 2.1配CUDA 12.1偶尔有兼容坑,等稳定跑通再优化。另外检查下是不是把optimizer状态也塞进显存了,用8-bit adam或者offload能再挤出一块。
24G跑8B LoRA还爆,大概率不是显存不够,是激活值在2k长度下炸了。你开gradient checkpointing是对的,但得配合batch size=1+gradient accumulation一起用,光开checkpointing但seq_len长,中间变量照样吃满。我建议你先试下torch.compile,PyTorch 2.1里直接model = torch.compile(model),有时候能省20%-30%显存,而且基本零成本。Flash Attention值得装,尤其你输入长,它能显著压掉attention那部分的峰值内存,但注意得换成支持flash attn的版本,比如用transformers的attn_implementation="flash_attention_2"。DeepSpeed ZeRO-3对单卡LoRA其实帮助不大,那是多卡才需要的,你单卡3090不如直接看下是不是梯度过大导致显存峰值,可以试试mixed precision,bf16在3090上支持得很好,能直接省一半显存。另外检查下是不是把整个base model都放进autocast了,LoRA层和冻结层分开处理会好很多。我上次跑7B模型也遇到这个,最后发现是tokenizer把padding搞太长,实际2k输入被扩到2.5k,检查下你的attention mask。实在不行就降到4bit量化,QLoRA在24G上跑8B很轻松,效果差不了太多。
说实话你这个配置跑8B LoRA,24G显存理论上是有希望的,但2k的序列长度确实是关键瓶颈。我试过把max_length砍到1k,显存直接降了快4个G,所以你可以先看看数据里是不是真有那么多长样本,很多时候实际填充到2k但有效内容才几百token,这时候动态padding加上packing能省不少。
Flash Attention我强烈建议上,它不只是省显存,速度还快一截,而且现在transformers库已经内置支持了,就改一行attn_implementation="flash_attention_2",不用自己写kernel。DeepSpeed ZeRO-3对LoRA来说反而有点重,因为可训练参数太少,ZeRO-3主要省的是优化器状态和梯度,你LoRA那点参数量根本不够它折腾的,ZeRO-2可能更合适。
torch.compile对微调场景帮助有限,编译开销主要在forward/backward的图优化,但LoRA本身计算量小,瓶颈在attention和激活值,我试过没感觉到明显显存变化,反而偶尔有编译报错折腾时间。还有个容易忽略的坑是gradient_checkpointing和input gradients混用,比如你开了checkpointing但没设input_with_grads=False,那输入embedding的梯度还是会占一大块显存,建议查一下。
另外可以试试把optimizer换成Adafactor或者8-bit Adam,省下几百M到1G的优化器状态,虽然不是决定性但能挤一点是一点。如果还不行,就考虑梯度累积+更小batch,但注意gradient accumulation不会减少单步激活值显存,只影响更新频率。最后检查下是不是模型加载时用了torch_dtype=float32,应该改成torch.bfloat16,这个很多新手会漏。
你这个问题我太熟了,之前用8B跑QA微调也卡在24G上,后来发现光开gradient checkpointing不够,还得把input的padding长度显式设成2k,不然默认会按batch里最长样本算,显存全浪费在padding上了。Flash Attention确实能省不少,尤其长序列下效果明显,但你这显存爆掉大概率不是attention的问题,先看看是不是优化器状态占太多,用AdamW的话8B模型光优化器就得吃4G多。torch.compile在2.1上对LoRA收益一般,而且容易跟ZeRO冲突,建议先把序列长度和优化器这块排查一下,另外可以试试gradient accumulation配合batch size=1,虽然慢但至少能跑起来。