最近在试着用LoRA微调一个7B的底座模型(做代码生成),单卡A100 40G,用的transformers+peft。batch size设成1,gradient accumulation调到8,序列长度512,按理说显存应该够,但跑不到两步就OOM。我看日志里显存占用一直在涨,怀疑是不是保存中间激活或者优化器状态的问题?另外我用了fp16=True但没开gradient checkpointing,会不会是这里的原因?求有经验的大佬指点一下,或者有没有推荐的稳定配置组合?先谢过了。
LoRA微调7B模型显存总爆,是batch size问题还是我配置错了?
全部回复
共 48 条fp16开着但没开gradient checkpointing,7B模型在40G上其实挺悬的,激活值累积起来很吃显存。你可以先试试把gradient checkpointing打开,batch size保持1,序列长度降到512以下,应该能稳住。另外检查下是不是dataloader里没设drop_last,最后一批数据形状不一致也可能导致显存异常增长。我之前跑类似配置,把optimizer换成adamw_8bit,显存能再省一截。
fp16开着但没开gradient checkpointing,7B模型跑512长度其实挺悬的,激活值在反向传播时会吃爆显存,而且显存一直涨更像峰值溢出而不是泄漏。你可以先开一下gradient checkpointing,显存能省一半左右,batch size保持1就行,accumulation不影响显存。另外注意peft里lora的target_modules别设太多,默认的q_proj,v_proj就够了,我之前全加上也爆过。如果还不行,试试用bitsandbytes的4bit量化加载底座,配lora跑起来很稳,我这边24G卡都这么玩。
fp16开着但没开gradient checkpointing,7B模型序列512其实激活值挺吃显存的,尤其你gradient accumulation设8,等效batch size其实不小,log里显存一直涨大概率就是激活缓存没释放。我建议先把gradient checkpointing打开,显存能省一半以上,batch size可以试着提到2或者4,反而比accumulation更稳。另外检查下是不是把model.enable_input_require_grads()漏了,peft有时候会因为这个多占缓存。我之前用40G跑7B,fp16+checkpointing+batch size 2,序列1024都没爆过,你可以参考下。
fp16开了但没开gradient checkpointing基本等于白省,7B模型光中间激活就能吃满30G,你batch size=1也扛不住。建议把gradient checkpointing打开,显存能掉一半,另外确认下是不是transformers版本里默认把past_key_values缓存也加进去了,那个也会涨。我自己的配置是batch size=1、grad_accum=16、seq_len=512,开checkpointing后A100 40G稳跑,实在不行就把seq_len降到384,代码生成任务影响不大。
fp16开了但没开gradient checkpointing,7B模型在这个配置下确实很容易爆,激活值在反向传播时会占掉一大块。你可以先试试开gradient checkpointing,batch size保持1不动,显存占用一般能降个30%以上。另外注意一下transformers版本,有些老版本对peft的显存释放有bug,会导致你看到的“占用一直涨”的现象。稳定配置的话,我一般用batch size=1+gradient checkpointing+fp16,序列长度512跑7B完全没问题,实在不行就把max_length降到384,效果差不了太多。
fp16没配gradient checkpointing等于白省,开一下能砍掉一半激活显存。
fp16开着但没开gradient checkpointing确实挺伤的,7B模型激活值在512长度下累积起来很夸张,尤其你batch size=1但梯度累积8,显存峰值会按单步计算,激活值不会因为累积而分摊。我之前跑13B遇到过类似情况,把gradient checkpointing打开后显存直接掉了快40%,建议先试这个,成本只是慢一点。另外看看是不是transformers版本和peft的兼容问题,之前有过版本不匹配导致显存泄漏的坑,升级或对齐一下说不定就好了。你优化器用的AdamW吧?如果没开8bit优化器,那部分内存也占不少,可以换bitsandbytes的优化器试试。
fp16开了但没开gradient checkpointing,这基本就是主要问题了。7B模型即使LoRA,反向传播时中间激活值在512序列长度下也会吃掉大量显存,尤其代码生成任务往往batch内token长度不均,padding会进一步放大占用。你看到显存持续上涨而不是瞬间爆掉,很可能就是激活值累积加上偶尔的长序列样本触顶。建议先开gradient checkpointing,代价是训练速度慢约30%,但显存能降一半以上;另外检查一下peft的target_modules,如果默认全量调attention的qkv,其实可以只调q和v,省一点参数也省显存。还有个小技巧,把fp16换成bf16试试,A100对bf16支持更好,数值稳定性也好,有时能避免fp16下loss异常导致的缓存抖动。稳定配置的话,我一般7B+LoRA用batch size 2,gradient accumulation 4,序列长度1024,开checkpointing,显存占用大概稳定在30G上下。你可以先跑一个step看峰值,别急着看整体loss,峰值不过就基本稳了。
fp16开了但没开gradient checkpointing,7B模型序列512其实激活值还是占不少,尤其LoRA虽然只训练小权重但base model的前向计算一点没省。我之前试过类似配置,把gradient checkpointing打开后显存直接降了快10G,batch size还能往上提。另外你显存一直涨这个现象,更像是缓存没清或者某个地方有泄漏,建议先看看是不是dataloader那边num_workers太多把内存也吃上来了。配置的话我目前是batch size 2加gradient checkpointing加8倍累积,稳得很,你可以试试。
fp16开着但没开gradient checkpointing,7B模型序列512其实激活值也挺可观的,尤其代码生成任务attention计算密集,显存慢慢涨大概率是激活缓存堆的。你可以先试着把gradient checkpointing打开,显存能省不少,batch size暂时不用动。另外确认下是不是peft的target_modules设置太宽了,有时候默认把全部linear层都lora化,反而比预期吃显存。我之前跑类似配置,开checkpointing后显存从36G降到22G左右,你可以参考下。
fp16开了但没开gradient checkpointing,这基本就是主因了,7B模型即使LoRA,激活值在512长度下也很吃显存,尤其代码任务batch内token密集。你可以先把gradient checkpointing打开,显存能省接近一半,同时把fp16换成bf16试试,A100对bf16支持更好,数值稳定性也强。另外确认下是不是dataloader里有啥东西在累积显存碎片,比如pin_memory或者自定义collate,可以换个简单的数据集跑跑看,排除数据侧问题。我自己的配置是bs=2+grad_accum=16+checkpointing,峰值显存稳定在35G左右,你可以参考下。
fp16开了但没开gradient checkpointing,这基本就是主因了。7B模型即使LoRA,反向传播时中间激活值在512序列长度下也会吃掉大量显存,尤其你batch size=1但accumulation steps=8,实际等效batch是8,但显存峰值是按单batch算的,所以问题不在accumulation。我建议你先把gradient checkpointing打开,显存能省将近一半,代价是训练速度慢20%-30%,但稳定很多。另外检查一下transformers版本,有些老版本对peft的显存优化有bug,升级到最新版可能直接解决。还有个细节,你如果用了fp16=True,记得同时设置fp16_opt_level=“O1”或者用bf16(如果卡支持),A100对bf16支持很好,能进一步降低显存压力。我自己的经验是,7B模型配LoRA,A100 40G开checkpointing后,batch size能跑到4,序列长度1024也不爆,你可以试试这个配置。如果还不行,就看看是不是数据加载时num_workers开太多导致CPU内存溢出,有时候OOM日志不一定只指显存。
fp16开着但不开gradient checkpointing,7B模型在40G上确实容易爆,尤其序列长度512时激活值占了不少。你试试把gradient checkpointing打开,显存能省下一大截,batch size和gradient accumulation不用动。另外确认下transformers版本和peft的lora配置,target modules别选太多,rank设8或16就够了。我之前用类似配置跑过,开checkpointing后稳定在30G左右,你可以先试这个组合。
你这个配置一看就是没开gradient checkpointing,7B模型即使LoRA,激活值在512长度下也挺吃显存的,尤其A100 40G跑fp16其实没想象中那么宽裕。我之前用类似配置,开了checkpointing之后显存直接降了快一半,batch size还能往上提。另外你确认下是不是transformers版本太新,有些默认行为改了,比如把输入embedding也算了梯度,建议把model.gradient_checkpointing_enable()加上试试,顺便把optimizer换adamw_8bit。如果还炸,就把序列长度砍到256,代码生成任务一般也够用。
fp16开着但没开gradient checkpointing,7B模型在40G上跑512长度确实容易卡在激活值上,尤其是代码生成这种序列里token分布不均匀的,显存会波动很厉害。我建议你先把checkpointing打开,显存能省一半左右,batch size暂时不用动。另外可以留意下是不是peft的默认lora dropout或者target modules设置太激进,导致中间tensor被复制了多份。我之前用类似配置跑过CodeLlama,开checkpointing后稳定多了,实在不行就试试把序列长度降到384,先跑通再说。
没开gradient checkpointing基本就是主因了,7B模型即使LoRA,激活值在512长度下也很夸张,尤其A100 40G跑fp16,中间变量累积起来很容易爆。你可以先打开gradient checkpointing试试,显存占用能掉一大截,代价是慢个20%左右,但稳定很多。另外fp16训练时,optimizer state虽然比fp32省一半,但Adam的momentum和variance还是占着fp32精度,这部分固定开销不小,你gradient accumulation设8其实不影响峰值显存,因为它是累积梯度而不是放大batch,所以别指望靠这个省显存。我怀疑你日志里显存一直涨,可能是某个hook或者缓存没清理,比如transformers的past_key_values在生成时容易累积,但训练时不该这样,建议你监控一下每步的torch.cuda.max_memory_allocated,对比峰值而不是当前占用。还有个坑是peft默认把lora dropout设0.1,虽然参数少,但dropout在forward时也会生成随机mask,显存占用跟序列长度成正比,可以试着降到0.05。如果还不行,就把序列长度砍到256,代码生成任务其实短序列也能跑,或者换8bit优化器像bitsandbytes的AdamW8bit,能再省2-3G。最后检查下transformers版本,有些老版本对flash attention支持不好,会额外申请临时buffer,升级到最新版再试。
fp16开着但没开gradient checkpointing,这基本就是主因了,7B模型即使LoRA,激活值在512序列下也相当吃显存,尤其batch accumulation会让峰值叠加。你可以先试试把gradient checkpointing打开,显存能省一半左右,代价就是慢点,但总比OOM强。另外确认下peft的target_modules是不是只设了q和v,如果全量改attention层,LoRA占的显存也会明显变大。我之前跑类似配置,batch size 2加gradient checkpointing就没再爆过,你可以参考下。
看描述大概率就是activation堆积的问题,7B模型就算LoRA,fp16下开gradient checkpointing和不开的显存差距能有2-3倍。你batch size已经是1了,再调低没意义,建议先开gradient checkpointing试试,同时把optimizer换成adamw_8bit,显存占用应该会明显降下来。另外你确认一下transformers版本,老版本对7B的attention缓存管理有bug,也会导致显存缓慢增长。我自己的配置是batch size 1 + grad acc 16 + checkpointing + bf16,A100跑7B训练稳得很。
fp16开着但没开gradient checkpointing,7B模型跑512长度确实容易爆,激活值在反向传播时累积起来很夸张。你可以先试试把gradient checkpointing打开,显存占用能降不少,代价就是慢一点,但对40G来说应该够用了。另外我怀疑你那个“显存一直涨”可能是transformers版本里缓存的问题,升级到最新版或者设一下max_memory试试。我之前用7B+LoRA,batch size 1+gradient checkpointing,序列长度1024都没爆过,你参考下。
fp16开着但没开gradient checkpointing,这基本就是显存刺客了,7B模型即使LoRA,激活值在512长度下也占不少,尤其代码生成任务序列实际利用率高。你可以先试试开gradient checkpointing,batch size保持1,应该能直接解决,代价就是慢个20%左右。另外确认下是不是用的最新版peft,老版本偶尔有显存泄漏的bug,我上次就是更新后就好了。如果还不行,把optimizer换成adamw_8bit或者adafactor,能省一大截。