最近开始尝试用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 条我之前也卡在这过,2k长度确实挺吃显存的,LoRA虽然省了训练参数但激活值一点没少。建议先试下Flash Attention,能把注意力那块的内存砍掉大半,而且跟PyTorch 2.1配合挺好,基本改两行就行。DeepSpeed ZeRO-3对这种单卡场景其实帮助不大,更多是多卡才划算,torch.compile倒是可以试试,但有时候跟某些算子不兼容反而更慢。另外检查下是不是把eval和train的batch都设成1了,还有输入pad到固定长度的话,用动态padding能省不少。
这配置爆显存太正常了,2k长度对8B来说就是分水岭,我试过同样设置卡在16k直接崩。你先把flash attention开了,基本能省30%左右,torch.compile在2.1版本对LLM支持还一般,别抱太大期望。DeepSpeed ZeRO-3配置麻烦但值得搞,不过我更建议先试试把输入截断到1.5k看看效果,很多时候领域数据没那么多关键信息。另外检查下是不是lora的target modules设太宽了,只冻结原模型改几层就够了。
你这个配置跑8B LoRA按理说24G够用,问题八成出在序列长度上,2k tokens的激活值比想象中吃显存,试试把max_seq_len先砍到1k看看还爆不爆。Flash Attention确实值得搞,它能直接把显存占用打下来一大截,而且跟PyTorch 2.x兼容性挺好,别怕配置麻烦,照着官方文档半小时能搞定。torch.compile在训练阶段收益不大,反而可能拖慢速度,建议先别开,等推理再考虑。另外检查下是不是优化器状态没走LoRA,把全参数都塞进AdamW了,那才是真正的显存杀手。
你这配置LoRA还爆显存大概率就是序列长度卡的,2k tokens对8B模型来说activation占大头,试试把max_length砍到1k看下峰值能降多少。Flash Attention值得搞一下,直接省一半activation内存,而且跟DeepSpeed不冲突,ZeRO-3反而可能拖慢速度,Stage 2就够用了。torch.compile建议先别开,跟某些自定义算子兼容性有坑,等稳定了再说。另外检查下是不是把validation也塞进显存了,eval的时候记得关梯度。
你这个长度确实吃显存,flash attention必开,再试试把lora的r降到8,3090跑8B全量微调本来就勉强。
你这配置跑8B LoRA还爆显存其实挺正常的,2k序列长度加上Llama 3.1的GQA结构,激活值比想象中吃得多。gradient checkpointing只省了中间激活,但optimizer状态和LoRA参数本身也占地方,尤其你如果用了AdamW,24G确实紧巴巴。DeepSpeed ZeRO-3对这种单卡场景帮助不大,ZeRO主要是多卡分片用的,单卡反而增加通信开销,不如直接试ZeRO-2或者干脆offload到CPU。Flash Attention值得装,能把注意力部分的显存从O(n²)降到O(n),你的2k长度提升会很明显,而且现在flash-attn库安装也不复杂,pip装个预编译wheel就行。torch.compile对显存没直接帮助,但能省点显存碎片,如果编译时间能忍可以开着,不过注意它和某些flash-attn版本可能冲突。我怀疑你还有个坑是没开gradient accumulation,虽然batch size=1,但LoRA的微调通常建议accumulate几步再更新,这样能稳定训练但显存不变。要不先试试把序列截断到1.5k,再配合flash-attn+bf16混合精度,大概率能跑起来,等熟悉了再上序列压缩或者QLoRA(4bit量化)那条路。
我之前也踩过这个坑,2k长度对8B模型来说确实挺吃显存的,LoRA虽然省了训练参数但激活值照样爆。建议你先试试在tokenizer里把max_length设成实际长度,别让padding塞满2k,能省不少。Flash Attention值得折腾一下,代码改动不大但显存能降30%往上,torch.compile我倒觉得收益一般,反而容易出兼容问题。另外你3090跑8B全量微调本来就很极限,实在不行就换QLoRA或者把序列截断到1k,效果损失没那么大。
跟你情况差不多,我一开始也是8B+LoRA在24G上爆显存,后来发现问题确实在序列长度上,2k tokens的激活值比想象中吃显存。建议你先试下Flash Attention,能省不少,而且现在改动很小,基本就是换个attention实现;DeepSpeed ZeRO-3配置确实麻烦,但对单卡场景帮助有限,不如直接上ZeRO-2或者offload。torch.compile我试过,对显存帮助不大,主要是提速,而且偶尔会跟某些算子冲突,建议等稳定了再开。另外检查下是不是把eval时的gradient也关了,还有优化器状态也占不少,可以试试adamw的8bit版本。
输入2k确实挺吃显存的,8B模型LoRA建议直接上8-bit优化器,能省不少。另外torch.compile对显存帮助不大,不如先把flash-attn配上试试。
8B全参微调24G本来就紧,试试4bit量化LoRA加flash attention,显存能降一半。
8B全参微调24G确实紧,你试下把max_seq_len砍到1024,或者换4bit QLoRA,立省一大截。
同款问题,序列长度影响很大,试试把max_len砍到1k,ZeRO-3配置其实没那么难,torch.compile留着后面再调。
3090跑8B LoRA按理说够,你查下是不是max_seq_len没设对,2k确实吃显存,开flash attention能省不少。
我跟你说,2k序列长度在8B模型上真的挺吃显存的,LoRA虽然省了训练参数,但激活值该占的内存一点没少,24G卡跑满序列加上大batch确实容易爆。我之前也卡在这,后来把序列长度砍到1k,显存立马降了4G,但如果你任务必须长上下文,那还是得上Flash Attention,它能省掉attention矩阵的显存占用,效果立竿见影。DeepSpeed ZeRO-3的话,单卡其实没啥必要,那是多卡才划算的,单卡上反而可能因为通信开销拖慢速度,你不如先试下ZeRO-2或者干脆把优化器状态切到CPU。torch.compile我试过,对显存优化帮助不大,主要是提速,而且有时候编译时间长得离谱,建议你等稳定跑通再考虑。还有个容易忽略的坑是输入padding,如果你用动态padding而不是把整个batch pad到max length,能省不少显存。另外检查下你的LoRA配置,是不是把bias也训练了,或者target_modules设得太大,这些都偷偷吃显存。你现在先试试把gradient checkpointing和Flash Attention一起开,再把batch size减到1,如果还爆就看下是不是加载模型时用的torch_dtype不是bf16。
我之前也卡在这步,最后发现是attention的seq len在作怪,2k确实不算短,flash attention能省不少,而且现在flash-attn库直接pip装就行,不用自己编译。DeepSpeed ZeRO-3对单卡其实帮助不大,主要省的是多卡通信内存,单卡建议先试ZeRO-2或者干脆关掉offload。torch.compile我试过,收益不明显还容易出兼容问题,不如先把输入长度砍到1k试试,或者用gradient accumulation模拟大batch,有时候反而更稳。
24G跑8B微调还爆显存,大概率是2k序列长度在作祟,Llama的attention是平方级涨的,你这输入长度直接吃掉了大半显存。Flash Attention值得装,能把KV cache和中间激活压一大截,torch.compile对显存帮助有限但能提速,建议先别碰。另外检查下是不是把LoRA的target modules设成全部线性层了,只微调query和value能省不少。我之前跑7B也遇过类似情况,后来把max_length临时砍到1k验证了下,显存立刻降了6G,你可以先试试这个方向。
这配置爆显存八成是2k序列长度+8B基座本身太吃激活值,先试Flash Attention,比ZeRO好上手多了。
torch.compile能省点但别抱太大希望,关键还是得把序列截短或换4bit QLoRA。
说实话你这个配置跑8B LoRA还爆显存,大概率不是batch size的锅,2k序列长度加上梯度检查点,理论上24G是够的。我怀疑你可能是把LoRA的target modules设得太宽了,或者没冻结原模型参数,试着检查一下是不是所有linear层都被注入了LoRA,只挑attention的QKV会省不少。另外DeepSpeed ZeRO-3在这个场景下其实有点杀鸡用牛刀,它主要是解决多卡显存不均的,单卡的话不如直接开ZeRO-2或者干脆用offload,把优化器状态扔到CPU上,能腾出好几个G。Flash Attention倒是值得装,它不只是省显存,还能加速,不过要确认你的CUDA版本跟flash-attn编译对得上,不然容易踩坑。torch.compile的话,我建议先别开,它跟梯度检查点有时候会有兼容性问题,而且首次编译还会额外吃显存,等你把基础跑通再优化性能不迟。还有个容易被忽略的点,你的输入长度2k,如果用的是动态padding,实际batch里最长的那条样本会决定显存峰值,可以考虑用static padding配合truncation到固定长度,或者用更激进的梯度累积步数来变相降低batch size。最后建议你装个nvidia-smi监控一下,看是激活值爆了还是梯度爆了,有时候是中间变量没释放,加个torch.cuda.empty_cache()在step之间会有改善。
显存爆掉八成是序列长度和注意力机制的锅,2k tokens在8B模型上确实吃紧,Flash Attention能直接砍掉一大块激活内存,建议优先上这个,比ZeRO-3好配多了。torch.compile对显存优化帮助不大,主要是提速,但跟Flash Attention一起用可能踩坑,先别开。另外checkpointing记得配合gradient accumulation用,不然反向传播算力浪费挺亏的。我之前在24G卡上跑7B,序列砍到1.5k就稳了,你可以试试动态padding或者截断。
这配置跑8B LoRA其实挺极限的,2k序列长度确实是大头,24G显存基本被激活值吃光了。建议先试试把序列长度砍到1k或者用梯度累积模拟更大batch,看能不能跑起来。Flash Attention值得搞,能省不少显存,但torch.compile对微调帮助不大,反而可能增加编译开销。另外检查下是不是把embedding和lm_head也加进LoRA了,这两个参数很占显存,默认只调attention和mlp就行。