最近开始尝试用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 条你这配置跑8B全参微调确实勉强,但LoRA还爆大概率是序列长度和attention的锅,2k tokens在24G上其实很吃紧。建议先试Flash Attention,配合gradient checkpointing基本能解决大半,ZeRO-3那个配置成本太高,不一定必要。torch.compile对显存帮助不大,主要省的是计算时间,别指望它降峰值。另外检查下是不是把decoder层的model parallel打开了,有时候默认设置会偷偷多占显存。
8B全参微调24G本来就紧,你试试把max_seq_len砍到1024,再开flash attn,基本能稳。
ZeRO-3配LoRA反而容易出问题,torch.compile对显存帮助不大,别折腾。
说实话2k长度在24G上跑8B确实紧张,但也不该连batch=1都爆。你大概率是没开gradient_checkpointing的use_reentrant=False,或者LoRA只套了attention层但把全连接也算了梯度,建议先打印下模型显存占用分布。DeepSpeed ZeRO-3对单卡没啥用,Flash Attention倒是能砍掉一半激活显存,torch.compile建议直接升到2.4再用,2.1版本收益有限还可能踩坑。我上次是换成peft的prepare_model_for_kbit_training加4bit量化才跑通的,你试试把加载改成load_in_4bit=True。
如果你不想量化,还有个偏方:把输入切成两段分两次forward,手动pooling一下hidden state,虽然逻辑上有点取巧但实测能塞进去。ZeRO-3真的别碰,单卡反而会多出通信开销,你重点检查下是不是attn_implementation="sdpa"没设,默认的eager attention在长序列上会吃满显存。另外,torch.compile在2.1上对动态shape支持很烂,跑微调建议先关了,等稳定再开。
序列长度就是元凶,2k tokens在24G上跑8B LoRA本来就吃紧,先试Flash Attention,效果立竿见影。
24G跑8B LoRA还爆,大概率不是batch size的锅,2k序列长度确实是显存杀手。Flash Attention建议优先试,它能直接砍掉attention矩阵的显存占用,效果立竿见影,而且现在transformers里调用很简单,不用手动改模型结构。ZeRO-3我倒觉得先别急着上,配置offload容易踩坑,而且LoRA本身参数少,主要占显存的是激活值,ZeRO对这块帮助有限。torch.compile可以开,但记得用mode=“reduce-overhead”,对长序列的kernel融合挺友好,能省点显存顺带提速。另外检查下是不是把梯度和优化器状态也塞进显存了,LoRA下可以把主模型参数冻结后设requires_grad=False,只用optimizer管理lora参数。还有个容易忽略的点,gradient checkpointing要配合input更长的chunk才有效,试试把2k切成两段1k,或者用packing把短样本拼一起,减少padding浪费。实在不行就换QLoRA,4bit量化后8B大概只要6G,留足余量。
说到2k序列长度这个点,你大概率是踩在激活显存上了。8B模型本身权重用LoRA后占用不大,但2k tokens的前向激活值在24G卡上确实很容易爆,gradient checkpointing能省一部分但会拖慢速度,建议先把checkpointing开着,然后试试把输入截断到1k看下峰值显存掉多少,这样能定位是不是长度的问题。Flash Attention值得折腾一下,尤其你用的是2.1,直接pip装个flash-attn库就行,能显著降低长序列的显存占用,而且对训练速度也有帮助。DeepSpeed ZeRO-3我建议先放一放,单机单卡用不上它,那是多卡才需要的,你现在的瓶颈不在参数存储上。torch.compile对显存优化帮助不大,但能提点速度,不过配上gradient checkpointing有时候会出奇怪的问题,建议先把显存问题解决了再考虑。另外一个小坑,确认下你的optimizer是不是用的AdamW,它会额外存两份模型大小的状态,换成8-bit Adam或者SGD能省不少,不过要调学习率。我之前用类似配置跑Llama-2 7B,最后是Flash Attention加8-bit Adam加梯度累积才稳定在22G左右,你可以参考下这个组合。还有个小技巧,把LoRA的target modules只加到q_proj和v_proj上,别全加上,也能省点显存。
你这情况大概率就是序列长度把激活值撑爆了,2k tokens在8B上真不小。我建议先别上DeepSpeed,把Flash Attention加上试试,能省不少显存,而且跟LoRA不冲突。torch.compile可以开,但注意别跟gradient checkpointing一起用,有时候反而更吃显存。还有个笨办法,把max_length砍到1k,先跑通流程再慢慢加,能省很多排查时间。
24G跑8B LoRA其实挺极限的,你batch size=1还爆,大概率是2k序列长度把激活值撑爆了,尤其没开flash attention的话,QK^T那步直接吃满显存。DeepSpeed ZeRO-3确实能省,但8B模型用ZeRO-3有点杀鸡用牛刀,关键是你要把offload参数和优化器状态到CPU,配合CPU offload才能看到明显效果,不然纯ZeRO-3对单卡训练帮助不大。Flash Attention强烈建议上,不光省显存,训练速度还能快不少,主要是把attention的中间矩阵计算方式改了,显存占用从O(n²)降成O(n),2k长度下能省出好几个G。torch.compile的话,PyTorch 2.1就能用,但要注意它和gradient checkpointing的兼容性,有时候会触发graph break导致收益打折,可以先试试不加checkpointing只开compile,看能不能过。另外检查下你的LoRA是不是把所有线性层都加了,只加qkv和o_proj通常够用,别碰lm_head和embedding,不然会多出不少可训练参数。还有个容易忽略的点,输入长度2k的话,试试在tokenizer里把padding策略改成max_length+truncation,别让batch里出现长度不一致导致显存碎片。我之前用A6000跑7B模型,也遇到过类似问题,最后是Flash Attention+gradient checkpointing+4bit量化过的,3090应该也能这么搞,bnb的nf4量化对显存帮助巨大,虽然精度会略降,但领域微调一般够用。你先从flash attention加起,一步步试,别一下全堆上,出了问题反而不好定位。
跟你一样踩过这个坑,2k长度输入在8B上确实很吃显存,LoRA只省了训练参数但激活值照样占地方。建议先试Flash Attention,配合梯度检查点能把激活内存砍掉一大截,torch.compile对显存帮助不大但能提速,别指望它省内存。ZeRO-3配置确实麻烦,但你这种情况不如直接换更激进的LoRA rank(比如8以下),或者把序列截断到1.5k试试,24G跑8B微调不是不行,关键是激活值优化。另外检查下是不是把优化器状态也塞进显卡了,用AdamW offload到CPU能再挤出几个G。
24G跑8B LoRA还爆显存,这太正常了,问题基本就出在2k序列长度上。你算算,8B模型反推梯度时激活值跟序列长度是二次方关系,哪怕batch=1,2k tokens的激活缓存轻松吃满十几G,加上LoRA的优化器状态和临时变量,OOM一点不意外。Flash Attention确实能救,它把注意力矩阵的显存占用从O(n²)降到O(n),你这情况至少能省4-6G,强烈建议先试这个,配置其实不复杂,就改一下attention实现。DeepSpeed ZeRO-3对单卡意义不大,它主要是多卡分片参数和梯度,你3090跑单机反而可能因为通信开销变慢。torch.compile可以先不开,PyTorch 2.1对这个支持还不完善,Llama某些算子编译后反而容易炸显存,等稳定版再上。另外检查一下是不是LoRA的target modules设得太多,比如把所有linear层都加了,可以只冻结attention里的q和v试试,能少不少显存。还有个小技巧,输入序列可以截断到1k或1.5k做实验,如果还爆就说明不是长度问题,是加载方式有坑。你顺便看看是不是加载模型时用了float32,改成bf16能直接省一半显存。
Flash Attention基本是必开的,ZeRO-3配LoRA反而可能因分片开销更慢,先试试把LoRA只挂q和v再砍下序列长度。
24G跑8B LoRA还爆显存,大概率不是模型权重的问题,而是激活值在2k长度下吃满了。你可以先试试在tokenizer里把max_length从2k砍到1k看能不能跑通,能跑通就说明是序列长度导致的。Flash Attention确实值得装,但更快的方案是换用Unsloth,它针对长序列和LoRA做了极致优化,装上就能省一半显存,不用折腾DeepSpeed。torch.compile对显存帮助不大,主要是提速,而且和某些库不兼容,建议先放一放。还有个常见坑是optimizer的momentum状态也占显存,可以试试8bit AdamW。
24G跑8B LoRA按理说是够的,但2k序列长度确实是个坎,activation占的显存会随序列长度平方级涨。你试试把max_length先砍到512看能不能跑通,能跑就说明是序列长度的问题。Flash Attention值得搞,能把attention那块的内存从O(n²)降到O(n),而且现在transformers库新版本已经内置支持了,不用手动改模型结构。DeepSpeed ZeRO-3对单卡场景其实提升有限,它主要是省下参数和优化器状态的分片内存,但LoRA本身参数量就小,你更该关心的是activation内存,这个ZeRO管不着。torch.compile对显存帮助不大,它主要提速,但有时候能减少点中间变量的峰值,可以试下,不过记得把mode设成reduce-overhead,不然编译时间够你喝杯茶的。还有个容易忽略的点:检查下是不是把eval也塞进训练循环了,eval模式下的gradient checkpointing是不生效的。实在不行就换QLoRA,4bit量化能把base model压到6G左右,省下来的全给activation。
24G跑2k的8B确实紧,先试试把序列截到1k或换4bit QLoRA,立省一大截。
这配置跑8B LoRA按理说24G是够的,问题大概率出在序列长度上。2k tokens的输入对attention来说显存占用是平方级增长的,你试试把max_length临时砍到512看看会不会立刻降下来,先排除这个因素。Flash Attention确实值得装,能省不少显存而且改动很小,PyTorch 2.1的话直接pip装flash-attn就行,不用配DeepSpeed那么麻烦。ZeRO-3对单卡场景其实帮助不大,那是给多卡分布式准备的,你单卡3090开个ZeRO-2或者干脆用offload到CPU反而更实用。torch.compile我试过,对显存优化没啥直接效果,但能提点速度,不过第一次编译会卡很久,别被吓到。还有个容易忽略的坑是LoRA的target_modules,如果你把attention的所有线性层都加了adapter,显存占用会明显上涨,试试只加q和v两层。另外你checkpointing开了的话,记得把input梯度也设成False,有时候默认配置会额外留梯度缓冲。我之前跑同模型用bitsandbytes加载4bit再叠LoRA,序列长度拉到3k都没爆,你可以考虑把基座模型量化一下,效果损失很小。
24G跑8B LoRA按理说是够的,你batch size都1了还爆,大概率不是显存总量问题,而是峰值占用没控制好。2k序列长度确实是个坎,但更关键的是你看下是不是把label也一起塞进模型算了loss,有些教程会让label参与前向,那显存直接翻倍。DeepSpeed ZeRO-3在这种单卡场景其实帮不上太大忙,它主要是解决多卡分片,单卡反而可能因为通信开销变慢,不如先试试ZeRO-2或者干脆关掉offload。Flash Attention值得装,它对长序列的显存优化是数量级的,而且现在很多轮子都内置了,不用自己改模型结构。torch.compile在2.1上对Llama这种动态shape可能收益不大,反而容易踩编译超时的坑,建议先别开。还有个容易忽略的点:optimizer状态,AdamW的momentum和variance在8B模型上本身就吃不少显存,可以试下8bit优化器,比如bitsandbytes的AdamW8bit,能省下2-3G。另外你checkpointing开了,但注意它跟gradient accumulation配合时,如果accumulation step设大了,反向传播的中间激活还是会累积,试着把accumulation step设为1,确认是不是这个在捣鬼。最后实在不行就上QLoRA,4bit量化后8B模型基座也就5G左右,留出大量余量给激活值,基本不会OOM了。
8B全参微调24G本来就紧,2k长度建议先开flash attention,ZeRO-3配合LoRA反而可能拖慢速度。
把位置编码改成RoPE的theta调小点试试,2k长度下能省不少显存,我上次这么搞直接降了30%。
3090 24G跑8B LoRA还爆显存,大概率就是2k序列长度在作祟,这个长度下KV cache能吃掉好几个G。Flash Attention值得先搞,能把显存占用砍一截,而且现在transformers里直接调就行,不用太折腾。DeepSpeed ZeRO-3配置确实麻烦,但你要是愿意试,记得offload optimizer state到CPU,能腾出不少空间。torch.compile对显存帮助不大,主要是提速,别指望它救OOM。另外检查下是不是加载模型时没用bitsandbytes的4bit量化,这步能省一半显存。
显存爆基本就是序列长度+8B基座模型吃掉的,先把max_len砍到1k试试,Flash Attention绝对值得配,能省不少。