最近开始尝试用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按理说是够的,你这个情况大概率是序列长度和attention计算的问题。2k tokens对于原生Llama 3.1来说其实不算长,但如果你用了padding或者没有做动态批处理,实际显存占用会远高于理论值。我建议先检查下tokenizer有没有把输入统一pad到某个固定长度,改成动态padding能省不少。Flash Attention确实是个大杀器,尤其是长序列场景,它直接优化了attention计算的内存占用,而且现在的版本配置起来没那么复杂,装好包之后在模型forward里改一下就行,不用动太多代码。DeepSpeed ZeRO-3对8B模型来说有点重,而且跟LoRA的适配性有时候会出问题,我反而建议先试试ZeRO-2,配合offload optimizer,能把optimizer states放到CPU上省出一块显存。torch.compile在2.1版本上对LLM的支持还在完善中,有时候反而会踩奇怪的bug,建议等后续版本或者先试试用Inductor后端。另外一个小技巧:把gradient checkpointing的granularity调细一点,比如每层都checkpoint而不是默认的每几层,虽然会慢点但能再挤点显存出来。
24G跑8B LoRA按理说不会爆的,你2k长度算正常范围,但3090的显存瓶颈可能在激活值上。gradient checkpointing确实开了,不过建议检查下是不是把输入padding到统一长度了,很多教程默认pad到最大长度,那样显存会浪费在无效token上。Flash Attention值得试试,它不光省显存还能加速,而且现在PyTorch 2.1以上版本可以直接用torch.backends.cuda.enable_flash_sdp()开启,不用改模型结构。DeepSpeed ZeRO-3对单卡来说有点重量级,配置不当反而会引入额外通信开销,不如先试试ZeRO-2或者干脆只用offload optimizer。torch.compile在微调场景下提升有限,而且容易跟Peft的LoRA层有兼容性问题,建议先别开。另外可以检查下是否开了fp16混合精度,3090对bf16支持一般,用fp16的话记得加个gradient scaler,不然loss可能直接nan。最后一个小技巧:把输入按实际长度动态batch,用DataLoader的collate_fn做截断,别用固定长度。
我最近刚踩过这个坑,8B模型在24G显存上跑微调确实挺极限的。你可以先试试把序列长度裁到1024,如果业务允许的话,显存能省不少。另外DeepSpeed ZeRO-3确实有用,配置其实没那么复杂,可以参考HuggingFace的官方教程,照着来基本能跑通。Flash Attention也建议加上,能省显存还提速,torch.compile可以先不急,有时候反而不稳定。
24G跑8B LoRA确实挺极限的,2k序列长度加梯度检查点还爆,大概率是优化器状态和中间激活在作祟。DeepSpeed ZeRO-3能分摊优化器显存,Flash Attention主要省attention那块,两个组合效果挺明显,配置其实网上有现成模板,照着改个config就行。torch.compile对训练提速有帮助,但显存节省不明显,建议先把ZeRO和Flash Attention搞定。另外检查下是不是把模型参数全量加载了,8B模型光权重就要16G,LoRA一定要确保只有adapters在更新。
24G跑8B LoRA按理说够的,2k长度确实有点吃紧,试试把max_length显式设成1024或者用梯度累积凑batch size,能降瞬时显存峰值。Flash Attention对长序列提升很大,torch.compile在2.1上可能不稳定,建议先升到2.3以上再试。ZeRO-3配置其实照着官方模板改几行就行,注意offload参数别开太大,不然反而慢。
Flash Attention基本是必开的,能省个30%显存,torch compile对长序列效果明显但第一次编译很慢。
3090 24G跑8B LoRA按理说没那么容易爆,2k长度确实有点吃紧,但更可能是attention的计算开销把显存撑爆了。Flash Attention对长序列效果很明显,配置没你想的那么复杂,Llama官方repo里就有现成的调用示例,直接pip装包然后改下model config就行。DeepSpeed ZeRO-3虽然能省,但小规模单卡跑反而可能因为通信开销拖慢速度,建议先试试Flash Attention和torch.compile,后者我试过在A100上能省15%左右显存,PyTorch 2.1兼容性还行。另外检查下是不是把optimizer states也塞进显存了,用bitsandbytes的8-bit Adam能再省几G。
老实说24G跑8B LoRA 2k长度确实挺极限的,你提到的几个方向都对,但顺序得调一下——Flash Attention优先级最高,我自己的经验是装了它之后峰值显存能降30%左右,而且配置其实没想象中复杂,pip装完改个modeling文件里的attention实现就行。DeepSpeed ZeRO-3虽然省显存但跟LoRA有时候会有兼容性问题,尤其你用的是3090这种卡,ZeRO-2可能更稳一些,记得把offload设到cpu。torch.compile的话,我试过在微调场景下收益不大,反而可能因为动态图导致编译时间很长,不如先把精力放在梯度累计上,batch size=1的时候累计个8步效果差不多。另外检查一下是不是把优化器状态也塞进显存了,用bitsandbytes的8-bit Adam能再省点。还有个容易忽略的点——你的数据加载器有没有做tokenized预缓存?每次训练都重新分词会很慢但显存倒是还好。如果实在不行,可以考虑把输入截断到1536 tokens,很多开源模型对长序列的注意力计算其实没那么高效。最后建议你先用huggingface的trainer加个memory_metrics参数跑一下,看看具体哪块在吃显存,对症下药比瞎调参数强。
3090 24G跑8B LoRA确实容易爆,2k长度加上gradient checkpointing还OOM,大概率是优化器状态吃得太猛了,试试把adamw的betas调低一点或者换8bit adam。DeepSpeed ZeRO-3配置起来其实没那么吓人,官方有现成的config模板,改几个参数就能跑,我上周刚在3090上把7B的序列拉到4k,显存占用反而降了。torch.compile对LLM微调收益不大,反而可能因为图编译炸掉,先别折腾那个。另外检查下是不是开了全参数微调的某些hook,LoRA只插适配器的话不该这么吃显存。
显存爆和2k序列长度关系很大,flash attention基本是必装的,能省不少。torch.compile对显存帮助不大,建议先试deepspeed zero-3或offload。
同款配置踩过这坑,2k长度下8B模型attention占大头,Flash Attention必须安排上,能省30%左右显存。另外torch.compile建议试,对推理和训练都有提升,但注意先升级到PyTorch 2.2+,2.1对llama支持不够稳。你LoRA的target modules只设了q_proj和v_proj吗?加上k_proj和o_proj有时候能减少中间激活开销。还有个小技巧,把pad到2的幂次长度,比如2048,偶尔能救回来一点。DeepSpeed ZeRO-3配置确实麻烦,可以先试试ZeRO-2加offload optimizer,省显存效果也够明显。
我之前也卡在过这,24G跑8B LoRA按理说够的,问题大概率出在2k序列长度上,激活值吃显存比想象中凶。建议你先试试把max_seq_len临时砍到512看看还爆不爆,如果正常了就得考虑Flash Attention,其实装起来没你想的复杂,直接pip装flash-attn然后改一行代码就行。ZeRO-3对LoRA来说有点重,8B参数用ZeRO-2或者干脆offload到CPU更省事。torch.compile我试过,对显存帮助不大,主要提速度,建议等你把显存问题搞定再说。
我之前也卡在这步,8B模型2k长度确实容易爆,你试试把序列长度先砍到1k看看,如果能跑起来基本就是长度问题。Flash Attention值得配一下,代码改动很小但省显存效果明显,比ZeRO-3好上手多了。另外torch.compile对显存帮助不大,主要是提速,别指望它省内存。还有个没提到的点:LoRA的target modules别全加,挑着几层加能省不少。
我之前也卡在这上面好久,2k长度加8B确实挺吃紧的,你试试把LoRA的target modules换成只训attention的q和v,能省不少。Flash Attention对长序列提升很明显,值得折腾一下,torch.compile在2.1上收益不大,建议先放着。另外别忘了把输入padding到固定长度,动态padding有时候反而会触发额外的显存碎片。
2k长度确实吃显存,Flash Attention能救,ZeRO-3配LoRA有点杀鸡用牛刀。
你这配置跑8B LoRA按理说不该OOM,重点大概率是序列长度——2k tokens对attention来说显存是平方级增长,建议先试下把max_length砍到1k看能不能跑通,确认瓶颈在哪。Flash Attention确实值得搞,能省不少显存而且配置不算难,找个现成脚本改改就行;DeepSpeed ZeRO-3对这种单卡场景帮助不大,反而增加通信开销,不用优先考虑。torch.compile对显存优化作用有限,但能提点速度,等稳定跑通后再试不迟。另外检查下是不是把lora模块的target_modules设得太宽,或者用了全量微调的默认优化器状态,这俩也常是隐性显存杀手。
跟你一样配置,之前也被24G卡得怀疑人生。2k长度确实是大头,我后来换了Flash Attention直接省了快5G,建议先从这个入手,比DeepSpeed简单多了。torch.compile对显存优化其实有限,主要是加速,别指望它解决OOM。另外可以试试把LoRA的r值调小到8,或者用gradient accumulation模拟更大batch,但关键是检查一下是不是把optimizer states也塞进显存了,用AdamW的8-bit版能再挤出一块空间。
我之前也卡在这过,2k输入长度确实挺吃显存的,你试试把max_seq_len临时砍到1k跑一次,能跑通就说明是序列长度的问题。Flash Attention真的建议装一下,你那24G卡开FA之后还能把batch提上去,比调DeepSpeed省事多了。torch.compile对显存帮助不大,主要是提速,别指望它救OOM。另外检查下你有没有把labels设成-100的部分都正确mask掉,有时候padding没处理干净也会白白多占显存。
你提到2k tokens这个点很关键,llama的attention是二次方增长的,24G跑8B其实很极限,可以先试试把序列长度砍到1k看看,如果能跑就说明问题出在这。Flash Attention确实能省不少显存,而且现在transformers里直接调就行,不用太担心配置复杂度,跟着官方文档走基本不会踩坑。DeepSpeed ZeRO-3对单卡来说收益不大,主要是多卡场景,你单卡3090的话优先级可以放低点。torch.compile对显存优化帮助有限,但能提升训练速度,建议等基本流程跑通后再试。还有个容易忽略的点是optimizer状态,可以换8bit adamw,能再省出几个G。
24G跑8B全参微调确实憋屈,但LoRA+gradient checkpointing还爆,大概率不是batch size的锅,2k序列长度对Llama来说真的不小,attention那部分内存是跟长度平方走的。你可以先试试把序列截到1k看还爆不爆,能跑就是长度问题,然后去开Flash Attention,那个其实装起来没那么玄乎,pip装个包改两行配置就行。DeepSpeed ZeRO-3对单卡LoRA收益不大,主要是给多卡用的,别急着上。torch.compile对显存优化帮助有限,但能提速,等你能跑通再折腾。另外检查下是不是忘了关gradient computation或者没设padding策略,有时候是数据预处理的小坑。