最近在尝试用LoRA微调一个7B的模型做代码补全,用的是一张24G的卡。但发现只要把序列长度拉到2048以上,batch size调到2就OOM了。我查了下,很多人说LoRA很省显存,但我实际跑起来感觉跟全量微调也没差太多?我用的transformers + PEFT,是不是我哪里配置不对?还是说7B模型本来就这样,想跑长上下文就得靠梯度累积或者换更大的卡?另外想问下,有没有什么技巧能减少激活内存的占用,比如gradient checkpointing能有多大帮助?现在卡在数据准备这一步,序列长度上不去感觉效果也出不来,有点迷茫,求有经验的大佬指点一下。
用LoRA微调7B模型,显存只够跑2k上下文,有什么办法吗?
全部回复
共 13 条24G跑7B LoRA长序列确实紧,你试试把gradient checkpointing打开,激活内存能省一半以上,再加上梯度累积,batch size设成1基本能稳。另外PEFT的target_modules别全选,只挑attention的q和v,能再省一截。还有个小技巧,用flash attention替代原生的,2k到4k的上下文显存开销差异会小很多。我之前跑代码模型也是这个配置,开checkpointing后8k都能勉强塞进去,但训练速度会慢点,你可以权衡下。
gradient checkpointing必须开,显存能省一半多,但速度会慢不少。另外检查下是不是忘关gradient accumulation了,还有attention的flash实现能再挤点空间。
24G跑7B LoRA长序列确实紧,但你这情况不太正常,LoRA省的是优化器状态和梯度,激活内存该占还是占。gradient checkpointing必须开,能省一半以上激活显存,代价是慢30%左右,配合梯度累积把batch先降到1试试。另外建议用FlashAttention-2,能把注意力内存从O(n²)降到O(n),2k到4k序列这个优化特别明显。你如果用的是旧版transformers,建议升级下,新版对PEFT的兼容性和显存分配都优化过。我同配置跑过CodeLlama 7B,4k上下文batch 1加梯度累积8是没问题的。
24G跑7B LoRA长序列确实紧,你这不是配置问题,7B的激活内存本来就很吃显存。gradient checkpointing一定要开,能省接近一半,但速度会慢不少,配合梯度累积到8或者16基本是标配了。另外可以试试把LoRA的target modules收敛到只调attention的q和v,别碰mlp,能省一点是一点。还有个偏方,用8bit优化器加混合精度,能再挤出一两千token的余量。不过说实话,想上4096以上序列,要么换80G卡,要么用序列并行或DeepSpeed ZeRO-3拆层,但后者配置起来挺折腾的。
24G跑7B+LoRA,2k上下文batch 2 OOM太正常了,你以为省显存省的是优化器状态和梯度,但激活值该占多少还是多少,这块大头省不掉。gradient checkpointing建议直接开,能省差不多一半激活内存,代价是慢个20%-30%,但至少能让你把batch加上去或者序列拉长。另外你试试用unsloth或者flash-attention,光这两个就能把内存占用再压一截,尤其flash-attn对长序列效果很明显。我个人体感7B想跑4k以上,24G卡基本就是极限了,要么接受2k+梯度累积,要么就得考虑量化到4bit再加LoRA,实在不行只能换卡。
gradient checkpointing能省不少,配合梯度累积试试,24G跑7B长序列确实紧巴。
试试8bit优化器加flash attention,序列长度能再提一截。
gradient checkpointing必须开,24G跑7B长序列不开这个基本没戏,开了之后激活内存能砍掉一半以上,代价就是训练慢个20%左右但完全能接受。另外你说的对比全量微调感觉没差太多,可能是你LoRA的r设太高了,或者你把target modules全勾上了,试试只微调q_proj和v_proj,r=8到16就够代码补全这种任务用了。还有个小技巧是attention里用sliding window或者把position embedding换成ALiBi,能直接降显存,不过要改模型结构比较麻烦。我自己的经验是2048长度batch size 1加上8步梯度累积效果其实也不差,别太纠结单步batch。
gradient checkpointing必须开,能省将近一半显存,再加梯度累积就能跑长序列了。
24G跑7B LoRA长序列确实紧张,你这情况挺正常的。gradient checkpointing必须开,能省差不多一半激活显存,但训练会慢个20%-30%,值得换。另外把batch size降到1,配合梯度累积到4或8,效果一样但显存压力小很多。还有个小技巧,检查下是否用了flash attention,能省不少内存。序列长度卡在2048可以先跑起来看效果,代码补全未必非要超长上下文。
开gradient checkpointing能省不少激活显存,但计算会慢点,先试试把batch降到1配梯度累积跑起来再说。
试试flash attention加梯度检查点,激活内存能降不少,我24G跑4k没问题。
LoRA省的是优化器状态和权重梯度,激活内存该占还是占,所以长序列下OOM很正常。gradient checkpointing能帮不少忙,大概能降30-50%激活内存,但会慢个20%左右,建议先开上试试。另外可以看看flash attention,对长上下文挺友好的,还有就是把batch size压到1配合梯度累积,效果基本不受影响。24G跑7B的2k以上确实紧,实在不行考虑QLoRA量化一下基座,能再挤出不少空间。
24G跑7B的LoRA确实有点紧,尤其序列一长激活值涨得飞快,LoRA省的是优化器状态和梯度那部分,激活内存该占还是占。gradient checkpointing能省不少,大概能砍一半多激活,但速度会慢个20-30%,建议先开上试试。另外flash attention 2一定要用,对长序列帮助很大,再配合bf16训练,2k应该能稳住。实在不行就gradient accumulation拉满,batch size压到1,效果不会差太多。