最近开始尝试用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 token的序列长度上,8B模型光KV cache就能吃掉好几个G,你试试把max_length缩到1k或者用gradient accumulation替代大batch,显存立马就下来了。DeepSpeed ZeRO-3配置确实麻烦,但如果你只是单卡,不如直接上Flash Attention,能省不少缓存,而且torch.compile对训练提速明显,但对显存帮助不大,建议先别折腾。另外记得把optimizer改成8bit AdamW,能再挤出一两个G,我这么调完24G跑8B LoRA稳稳的。
我之前也卡在这过,输入2k确实挺吃显存的,8B模型光激活就占不少。你LoRA只冻了原模型,但梯度还是要过全量,试试把max_length先砍到1k跑通再说。Flash Attention对长序列提升很大,而且现在torch里直接调就行,比DeepSpeed省事多了,ZeRO-3配置容易出问题。torch.compile可以先不开,有时候反而慢,等显存稳了再优化速度。另外检查下是不是把eval也塞进训练循环了,那个很容易忽略。
3090跑8B全参LoRA本来就紧,你这2k长度是显存杀手,先试试把序列截到1k或换8bit量化,能立省5G。
24G跑8B LoRA按理说够用了,你这情况我怀疑是2k序列长度直接把激活值拉爆了,Flash Attention确实能解决大头,但更简单的办法是先试试把max_length砍到1k或者调低rope scaling看还爆不爆。torch.compile对显存优化帮助不大,主要是提速,而且跟DeepSpeed混用容易出幺蛾子。另外检查下是不是把LoRA加到所有linear层了,target_modules收窄点能省不少。实在不行就ZeRO-3加CPU offload,但别开offload optimizer,那玩意儿慢到怀疑人生。
8B全参微调本来就不是24G能玩的,LoRA省的是优化器状态和梯度,但激活值照样吃满,2k长度加batch 1确实会爆。我之前用7B试过,加了gradient checkpointing后把seq len砍到1k才勉强跑起来,你这情况建议先确认是不是显存碎片化,试试环境变量PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128,有时候比上ZeRO管用。Flash Attention对长序列效果很明显,值得花时间配一下,torch.compile在2.1上收益不稳定,别抱太大期望。
我之前跑7B也遇到过类似问题,后来发现大概率是你说的序列长度没错,2k tokens在8B上做LoRA,24G确实很吃紧。DeepSpeed ZeRO-3配置起来是麻烦,但可以先试试Flash Attention,那个改动小收益明显,基本能省30%左右显存。torch.compile我试过,对显存帮助不大,但能加速训练,可以等不爆显存了再开。另外你可以检查下是不是把label也pad到了2k,有些实现会额外占不少内存,改成动态padding能缓解。
24G跑8B LoRA按理说是够的,你这个OOM大概率不是显存容量问题,而是峰值显存被中间激活值吃掉了。2k长度对8B来说不算短,尤其如果开gradient checkpointing还爆,先检查下是不是LoRA的target modules选太多了,或者r设得偏高,这会让可训练参数增多,反而失去省显存的意义。DeepSpeed ZeRO-3确实能救急,但8B模型用ZeRO-3有点杀鸡用牛刀,而且配置offload参数时容易踩坑,建议先试ZeRO-2,配合CPU offload,通常能把峰值压下来。Flash Attention值得搞,尤其你的序列长度是2k,它能把注意力部分的显存从二次方降到线性,效果立竿见影,而且现在flash-attn库安装也不算麻烦,直接pip装预编译版就行。torch.compile我个人经验是能省一部分显存,但收益不如前两个大,而且第一次编译很慢,还会引入一些兼容性问题,建议等前面的优化做完再考虑。另外有个容易被忽略的点:检查下你是否用了padding="max_length",如果每条样本都pad到2k,那batch size=1也等于在算一堆无意义的token,改成动态padding能省不少。最后,3090的24G其实跑8B LoRA加2k输入是可行的,我见过有人用类似配置跑13B,所以大概率是某处设置没到位,建议先开NVIDIA的nsight看看峰值内存分配在哪一层,比盲目调配置更有效。
这配置跑8B LoRA按理说24G是够的,主要问题应该出在序列长度和attention计算上,2k tokens的KV cache占得比想象中多。建议先试下Flash Attention,能省不少显存,而且和PyTorch 2.1兼容性没问题,比DeepSpeed配置简单多了。torch.compile可以开,但优先解决显存问题再说,它主要提速度不降显存。另外检查下是不是把优化器状态和梯度都堆在显存里了,可以考虑用bitsandbytes的8位优化器,效果立竿见影。
2k长度确实吃显存,但8B上LoRA爆24G多半是优化没拉满,先试试卸载优化器到CPU,再开flash-attn看看。
说实话你这个配置爆显存挺正常的,2k序列长度对8B模型来说,就算LoRA+gradient checkpointing,24G也刚好卡在临界线上。建议先试试torch.compile,配合max_autotune模式,有时候能省个20%显存,而且改动最小。DeepSpeed ZeRO-3确实能救急,但配置起来容易踩坑,不如先检查下是不是padding没做对,或者attention mask没传对,这两个坑比显存优化更隐蔽。另外flash attention在3090上收益明显,但记得要装对应CUDA版本的flash-attn库,装错版本反而会报错。
24G跑8B LoRA还爆,八成是序列长度在作祟,2k tokens对attention来说挺吃显存的,建议先试下Flash Attention,那个改造量不大,能省不少。DeepSpeed ZeRO-3配置确实麻烦,但你这个场景其实ZeRO-2就够了,把优化器状态切分一下就能缓解。torch.compile先别急着上,2.1版本对LLM支持一般,有时候反而拖慢,不如检查下是不是梯度累积没开,或者把输入padding到固定长度再截断。我之前也卡在类似地方,最后是换了个更激进的bitsandbytes 4bit量化才跑通的。
之前我也在这卡了好久,后来发现问题不一定在batch size,2k长度对8B来说确实挺吃显存,Flash Attention基本是必开的,能省不少。ZeRO-3配置起来麻烦但效果立竿见影,不过你可以先试试把输入截断到1k看看,很多时候领域数据没那么长。torch.compile在2.1上对LLM微调提升不算明显,反而可能编译时间很久,我建议先把优化重点放在attention和offload上。另外检查下是不是把优化器状态也放显存了,用AdamW的offload能再挤出一块空间。
说实话2k长度在8B上24G确实紧,建议先开flash attention,比DeepSpeed省事多了。
序列长度影响很大,2k tokens配LoRA在24G上确实紧,建议先把max_seq_len砍到1k试试,能省不少显存。
2k长度配8B,24G单卡确实紧,先把flash-attn加上,能省不少。 torch.compile这版本收益不大,先别折腾。
显存爆掉大概率是激活值太大,2k长度确实吃紧,先试试flash attention,能省不少。
torch.compile对显存帮助不大,主要省的是训练时间,ZeRO-3配置复杂但值得搞。
这个情况我太熟了,之前调7B模型的时候也被24G折磨过。你2k tokens的序列长度其实不算夸张,但Llama的attention机制对显存占用是平方级的,所以长度稍微上来点就很容易爆。建议先别急着上DeepSpeed,配置复杂度高不说,ZeRO-3在单卡场景下收益其实有限,更多是给多卡用的。可以先试试把输入截断到1.5k看能不能跑通,排除是不是长度导致的峰值问题。另外Flash Attention真的值得花时间搞一下,能省不少显存,而且现在很多库都封装好了,比如用bitsandbytes的4bit量化配合它,效果立竿见影。torch.compile的话,你那个版本最好升到2.3以上再试,不然可能遇到一些奇怪的兼容问题,而且它对显存优化帮助不大,主要是加速。还有个容易被忽略的坑是optimizer的state,LoRA虽然参数量小,但如果你没冻结原模型,AdamW的momentum照样吃显存,记得用paged_adamw或者干脆换SGD试试,说不定就解决了。要是还不行,就把LoRA的rank从16降到8,或者target_modules只选q和v,牺牲点效果换稳定。
24G跑8B LoRA按理说够的,你batch size都1了还爆,大概率是序列长度在作祟——2k tokens的激活值其实很吃显存,尤其attention部分。Flash Attention能省不少,但更直接的办法是先把max_seq_len临时砍到1k试试,如果显存占用骤降那就实锤了。另外别急着上ZeRO-3,那个更多是为了多卡扩展,单卡上配置复杂收益还小,ZeRO-2或者干脆关掉offload可能更实际。torch.compile对显存优化没直接帮助,但能省点显存碎片,前提是得跟CUDA graph配合,不过你PyTorch 2.1版本有点老,建议先升级到2.3+再试。还有个容易忽略的坑:优化器状态,AdamW的动量项在LoRA下虽然少,但如果你把bnb的4bit量化也开了,记得检查下是不是quantization config里显存分配出了问题。最后实在不行就上Unsloth,那个库对显存抠得很细,能压到10G以内,但要注意它改动了底层kernel,跟某些自定义loss可能不兼容。
序列长度绝对是大头,2k tokens在8B上24G真不够,试试把max_len砍到1k或者换8bit量化,立刻见效。
3090爆显存太正常了,8B模型光权重fp16就16G,2k序列下激活值很容易吃掉剩下的。你LoRA+gradient checkpointing都开了还爆,大概率是attention的中间tensor在作祟,先试试torch.compile,配合cudagraphs能省不少。Flash Attention值得折腾,但先确认是不是输入padding导致的,把padding到最长序列的trick换成右padding或者动态batch,效果立竿见影。DeepSpeed ZeRO-3对单卡反而增加通信开销,不如直接上bitsandbytes的4bit量化QLoRA,24G跑8B绰绰有余。