最近开始尝试用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 条你这配置3090跑8B LoRA按理说不会直接爆,2k长度确实有点吃紧,但更可能是attention计算时中间激活值太大。Flash Attention能省不少显存,而且现在PyTorch 2.2以上直接集成,升级一下就自带支持了,不用额外折腾。ZeRO-3主要是把优化器状态分到多卡或CPU上,单卡其实效果有限,不如先试试gradient checkpointing配合Flash Attention,batch size保持1但调低micro batch size到4或者8,看看能不能跑通。torch.compile对显存优化帮助不大,主要提速,建议先把显存问题解决了再折腾这个。
8B模型就算LoRA+gradient checkpointing,2K长度在24G上确实很极限,建议先把序列长度砍到1K或者512看能不能跑通。DeepSpeed ZeRO-3和Flash Attention确实能救,但ZeRO-3配LoRA有时会有参数同步问题,可以试试ZeRO-2或者用bitsandbytes的4bit量化,简单粗暴。torch.compile对LLM微调提升有限,不用急着折腾,先把显存降下来再说。
24G跑8B LoRA微调确实挺极限的,你提到的几个方向都是正解。序列长度2k token对于Llama 3.1来说已经是显存大头了,Flash Attention能有效降低attention层的显存占用,建议优先配起来,实测能省10%-20%。DeepSpeed ZeRO-3主要是切分优化器状态和梯度,对单卡来说收益其实没多卡那么明显,而且配置起来确实容易踩坑,新手可以先试试ZeRO-2或者干脆不用。torch.compile在微调场景下对显存优化帮助不大,主要是加速计算,反而可能因为编译过程增加临时显存占用。另外一个小技巧:检查一下你的LoRA rank是不是设太高了,8或16通常就够用,ranking偏大也会让adapter权重占不少空间。还有,可以尝试把输入截断到1.5k左右看能不能跑通,很多领域任务其实不需要全部2k context。最后确认一下你用的是不是最新版bitsandbytes,4bit量化配合NF4能把模型压到6-7G,配合LoRA基本是2408的标配玩法了。
24G跑8B LoRA微调确实有点极限,但2k长度下3090爆显存应该不只是序列长度的问题,建议先检查下是否忘了关gradient checkpointing的某些子模块,比如注意一下enable_reentrant=False这个参数,有时候默认设置反而会增加显存占用。Flash Attention对长序列提升很明显,而且现在transformers库直接集成好了,from_pretrained时加attn_implementation="flash_attention_2"就行,不用自己配,能省出4-5G。DeepSpeed ZeRO-3的话,如果没用多卡,单卡意义不大,反而可能因为offload拖慢速度,不如试试ZeRO-2或者直接用bitsandbytes的4bit量化QLoRA,我实测24G跑8B 4bit量化加gradient checkpointing能塞下4k长度的batch size 2。torch.compile在训练场景收益不稳定,尤其对长序列可能编译时间太长,建议先搞定显存再考虑。另外可以看看是不是dataloader的pin_memory=True和num_workers开太高导致缓存占显存,改成pin_memory=False偶尔会有奇效。
建议试试bitsandbytes 4bit量化,能直接把显存压到12G以下,配合Flash Attention效果更好。
24G跑8B LoRA确实有点极限,2k长度加上gradient checkpointing还爆的话,大概率是输入padding或者attention计算本身吃掉了大量临时显存。试试torch.compile?它在PyTorch 2.1上对Llama这类模型效果挺明显的,我自己的经验是能省10%-15%显存,而且编译一次后面跑起来很快。Flash Attention也值得搞一下,虽然配置确实有点绕,但装上后长序列的显存占用直接砍半,你搜一下xformers或者flash-attn的官方安装指南,跟着走一遍其实没那么复杂。DeepSpeed ZeRO-3倒不一定非得用,对单卡场景来说配置成本太高,有时候反而因为offload引入额外开销。另外检查一下你的tokenizer是不是把所有样本都pad到2k了,如果实际长度差异大,用动态padding或者packing技巧能省不少。还有个容易被忽略的点:优化器状态占显存,试试AdamW的8-bit版本或者干脆用SGD+合理的warmup,有时候能省出几G来。
我最近也踩过类似的坑,8B模型2k长度确实很容易爆显存。Flash Attention是必装的,能省20%左右显存,而且配置不复杂,pypi直接装就行。ZeRO-3的话,如果单卡用效果不明显,反而是offload到CPU能解燃眉之急。torch.compile建议先别开,有时反而会多占显存,等调通流程再说。另外可以检查下是不是输入padding太长,把pad到固定长度改成动态batch能省不少。
Flash Attention基本是必开的,能省不少显存,torch.compile对训练收益不大可以先不管。
Flash Attention基本是必装的,能省30%显存,torch.compile对长序列效果也明显。
Flash Attention基本是必装的,能省30%显存,配置也不复杂,pip install一下就完事了。
3090 24G跑8B LoRA 2k长度确实有点极限,我遇到过类似情况,后来把--gradient_accumulation_steps设成8,配合ZeRO-3总算跑起来了,Flash Attention也能省个2-3G显存。torch.compile我试过,编译时间挺长但对显存没啥直接帮助,主要提速度。你先看看是不是attn_implementation没设成flash_attention_2,很多人漏了这一步。
24G跑8B LoRA确实得精打细算,2k tokens长度算是显存大头了,建议先试试gradient checkpointing配合更小的per_device_train_batch_size(甚至1),同时把model并行或者ZeRO-3开起来,官方文档其实有现成config模板。Flash Attention在长序列上效果很明显,主要省的是attention那块的显存,配置起来只要装个包改一行代码就行。torch.compile对推理加速有用,但训练时可能不稳定,建议先搞定基本显存问题再折腾。另外检查下是不是把优化器状态也塞进显存了,用AdamW 8bit或者Adafactor能省不少。
24G跑8B LoRA确实容易爆,2k tokens加上gradient checkpointing还不够的话,试试把输入长度先砍到1k或者512看看是不是模型本身的问题。Flash Attention配置起来其实没那么复杂,装个包改两行代码就行,能省不少显存,DeepSpeed ZeRO-3对单卡意义不大反而麻烦。torch.compile在微调场景下主要是加速不是省显存,建议先解决OOM再考虑这个。
3090 24G跑8B LoRA按说不会爆,2k长度确实挺吃显存,但更可能的是你的LoRA参数或者输入处理里有隐藏的内存泄漏。Flash Attention基本是必装的,能省30%左右显存,而且对精度没影响;DeepSpeed ZeRO-3配置起来虽然麻烦点,但它是解决这种OOM问题的终极方案,建议花点时间搞。torch.compile对推理加速明显,但训练时容易踩坑,建议先把显存问题搞定再试。另外检查下是不是dataloader的num_workers开太高了,有时候多进程也会莫名吞显存。
3090 24G跑8B LoRA加上2k长度确实容易爆,我试过把LoRA的r从16降到8,同时target_modules只选q_proj和v_proj,能省不少。DeepSpeed ZeRO-3配置确实麻烦,但你可以先试试ZeRO-2,性价比高很多。Flash Attention强烈推荐,部署完显存能降30%左右,而且PyTorch 2.2以上原生支持,不用折腾编译。torch.compile对长序列推理有加速,但训练阶段收益不大,我建议先把前两个优化搞定。
8B模型2k长度在24G上确实挺极限的,我猜你的LoRA rank可能设高了,先试试rank=8或者16,同时把target modules减少一些,比如只q和v。DeepSpeed ZeRO-3配Flash Attention确实能救,但新手容易踩坑,建议先只开ZeRO-2和offload optimizer,简单很多,显存能省2-3G。torch.compile对训练加速不明显,主要是推理有用,暂时别折腾它。另外检查下是不是dataloader的num_workers开太多导致CPU内存爆了,有时候不是显存问题。
Flash Attention基本是必装的,能省30%显存,torch.compile对长序列也有奇效。
24G跑8B LoRA其实还算够用,问题大概率出在序列长度上——2k tokens对attention计算量的影响比想象中更大。建议先试试Flash Attention,配置不复杂,把模型加载时加个attn_implementation="flash_attention_2"就行,能省不少显存。DeepSpeed ZeRO-3确实有效但配置稍烦,可以等熟悉了再上;torch.compile对显存节省有限,主要提速度。另外检查下是不是把optimizer states也塞进显存了,用bitsandbytes的8位Adam能再省点。
试试把input长度砍到1k以下,flash attention真的能省不少,配置不难照着教程走就行。
24G跑8B LoRA按理说不会OOM,问题大概率出在序列长度上,2k tokens对attention来说显存消耗是平方增长的。建议先试下flash attention,PyTorch 2.2以上原生支持,改几行代码就能省下30%左右显存,比配DeepSpeed简单多了。torch.compile也有帮助,但得确保你的CUDA版本跟它兼容,不然容易踩坑。另外检查下是不是数据加载时把整个tokenized数据集都塞进显存了,有时候dataloader的pin_memory设置不当也会吃显存。