最近在试着用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模型即使LoRA,反向传播时中间激活值照样是按全参数量级算的,序列512其实不短,A100 40G看着大,但transformers默认的activation缓存方式很吃显存,尤其是代码生成这种任务,attention的中间tensor会很夸张。我建议你先把gradient checkpointing打开,显存占用能直接砍掉一半以上,代价就是训练慢个20%左右,但至少不会OOM。
另外你提到显存一直在涨而不是直接爆,这个很关键。除了激活值,优化器状态和梯度累积的临时buffer也可能有泄漏,但更可能是你设了gradient accumulation=8,这时候每步的梯度累积会额外保存中间梯度,如果没开offload,叠加起来就越来越接近上限。可以把accumulation先降到4,配合checkpointing试试,batch size保持1没问题。
还有个容易忽略的点:peft的LoRA默认不冻结bias和layer norm,这些参数虽然少,但它们的梯度存储也会占显存。你可以显式设置modules_to_save或者直接冻结全部非LoRA参数,这样能再省一点。我自己的稳定配置是7B+LoRA+checkpointing+bf16(如果卡支持),batch size 2,accumulation 4,序列长度不超过1024,40G基本能稳跑。你那个fp16如果是A100,建议直接换bf16,数值稳定性更好,而且有些算子fp16反而更耗显存。先开checkpointing,大概率就解决了。
fp16开着但没开gradient checkpointing,这基本就是显存刺客了,7B模型即使LoRA,激活值在512长度下也吃得很凶。我之前试过类似配置,batch size 1加梯度累积其实不省显存,真正占地方的是前向传播存的那些中间tensor。建议你把gradient checkpointing开了,显存能掉一半还多,训练速度会慢点但至少不会OOM。另外可以查一下transformers版本,有些旧版本对flash attention支持不好,也会导致显存异常。配置上我一般用fp16+gradient checkpointing+batch size 2,A100 40G跑7B很稳。
大概率就是gradient checkpointing没开的事儿,7B模型就算LoRA,激活值在512长度下也挺吃显存的,你fp16只省了模型权重和优化器,中间变量照样爆。我之前用A100跑6.7B,batch size 1加gradient checkpointing,显存能压到20G左右,你试试开了之后把gradient accumulation降到4,应该就稳了。另外顺手把optimizer换成adamw_8bit,能再省一截,代码生成任务微调效果基本不受影响。
fp16开了但没开gradient checkpointing,7B模型加512序列长度其实激活值占得比你想象大得多,尤其代码任务长token下注意力矩阵很吃显存。你可以先试试gradient checkpointing,基本能省一半以上,batch size保持1没问题。另外确认下是不是用的最新版peft,老版本有时会有缓存没释放的问题,显存一直涨很像这个。我上次是加了checkpointing再把优化器换成adamw_8bit才稳住的,你这卡跑7B应该绰绰有余。
fp16开了但没开gradient checkpointing,7B模型序列512按理说40G勉强够,但你日志里显存持续上涨更像是缓存没释放或者activation堆积,建议先开gradient checkpointing试试,显存能掉下来一大截。另外你把gradient accumulation调到8,但batch size=1的话实际等效batch才8,如果数据加载那边没做shuffle或者padding策略不当,也可能有隐性内存碎片。我之前用7B+LoRA训代码模型,A100上稳定配置是batch=2、grad_accum=4、seq_len=1024,加上checkpointing和bf16,峰值大概32G左右,你可以参考下。还有个小细节,peft的target modules别全加上,选q和v就够,能省不少显存。
开gradient checkpointing吧,你这显存八成是被激活值吃掉的,fp16省的那点不够塞牙缝。
fp16开着但没开gradient checkpointing,7B模型加LoRA在40G上确实容易卡在激活值上,尤其序列512不算短。你可以试试把gradient checkpointing打开,显存占用能降不少,速度慢点但稳。另外检查下是不是peft的target modules设置太宽,把太多层都包进去了,LoRA本身不该吃这么多显存。我之前跑类似配置,batch size 1加accumulation 8,开checkpointing后稳定在30G左右,你参考下。
看到你说显存一直在涨而不是直接爆掉,我第一反应就是显存碎片化或者缓存没释放,这个在长序列训练里特别常见。你fp16开了但没开gradient checkpointing,这其实很关键,LoRA虽然只训练适配器,但7B模型的前向激活值还是全量保存的,序列512不算长,但加上8的梯度累积,激活值峰值会叠得很高。我建议你先把gradient checkpointing打开,显存占用能降差不多一半,代价是慢个20%左右,但稳定很多。另外你试试把optimizer换成AdamW的8bit版本,或者直接用paged_adamw_8bit,这个对显存峰值控制特别有效,我上次微调13B就靠它救回来的。还有个细节,transformers里记得设model.config.use_cache=False,不然生成时的KV cache会一直累积,这个很多人会漏。batch size 1加梯度累积8本身没问题,但你可以考虑把累积步数降到4,然后看看是不是某些特殊token导致序列长度动态变化,偶尔超长一次就爆了。我自己的稳定配置是LoRA rank 16加alpha 32,fp16加gradient checkpointing加8bit优化器,40G上跑7B能塞下batch size 2,你可以参考下。如果还不行,就把序列长度砍到256试试,代码生成任务短序列也能学得不错,别死磕512。
gradient checkpointing没开基本就是主因了,7B模型即使LoRA冻结了底座,前向激活值还是得完整算一遍,序列512在A100上虽然不算长,但加上optimizer states和中间缓存,峰值很容易顶爆。你可以先开个gradient_checkpointing=True试试,显存占用至少能砍掉一半,代价是训练慢个20%-30%,但总比OOM强。另外fp16的问题在于,如果模型里有某些层对精度敏感,loss可能震荡,但显存占用不会因此一直涨,所以你那个“显存持续增长”的现象更像是有张量被意外保留,比如output_hidden_states=True这种参数会把每层输出都存下来,或者你用了return_dict_in_place之类的选项。我自己的配置是batch size 1、gradient accumulation 16、seq len 1024、开gradient checkpointing,再配合optim="adamw_torch"和lr=2e-4,A100能稳定跑完几千步。你还可以把peft的target_modules收窄一点,只微调query和value,别动所有linear层,显存也会明显降。最后检查下是不是transformers版本太新和peft有冲突,我之前遇到过一次内存泄漏,更新两个库到最新稳定版就解决了。
gradient checkpointing没开大概率就是主因,7B模型即使LoRA反向传播时激活值也挺吃显存的,尤其序列长度512下每层缓存都叠起来很夸张。你可以先开这个试试,显存占用基本能砍一半多,batch size和gradient accumulation那个组合本身倒是没啥问题。另外fp16=True如果没配fp16_opt_level之类的参数,可能实际没生效,建议确认下训练日志里loss有没有异常跳动,有时候混合精度踩坑也会导致显存曲线不正常。我之前跑类似配置是8卡A100,单卡微调的话还是把max length压到384更稳妥。
fp16开着但没开gradient checkpointing,7B模型seq len 512其实激活值挺吃显存的,尤其你gradient accumulation只是累加梯度,并不会省显存。建议先把gradient checkpointing打开,显存能掉一半左右,batch size保持1就行。另外看看是不是transformers版本太老,有些版本的缓存没释放干净,升级到最新版再试试。我之前用同样配置跑7B都没问题的,你检查下是不是dataloader里num_workers开太多,也会造成额外内存占用。
fp16不配gradient checkpointing等于白开,开一下能省一半显存。
fp16开了但没开gradient checkpointing,7B模型在这个配置下其实还是很容易爆的,尤其序列长度512时激活值占的显存比想象中大得多。我之前用8B模型也遇到过类似情况,把gradient checkpointing打开后显存瞬间降了差不多一半,你可以先试试这个。另外日志里显存一直涨的话,也可能是数据加载或者缓存没清干净,但更大概率是激活值累积的问题。还有个小建议,optimizer state用adamw的话可以试试8bit版本,能再省一截。batch size和梯度累积这个组合本身没问题,主要还是得先解决激活值这块。
gradient checkpointing没开的话,7B模型加LoRA在40G上确实容易爆,你这显存持续上涨大概率是激活值累积的问题,batch size=1但序列长度512也不小。建议先把gradient checkpointing打开,能省不少显存,另外确认下是不是transformers版本和peft的兼容性问题,之前遇到过新版peft默认把gradient checkpointing关了的情况。我自己的配置是batch size=1,gradient accumulation=16,fp16+gradient checkpointing,跑7B代码模型稳定在35G左右,你可以试试。
开gradient checkpointing能省不少,fp16配合这个基本就稳了,试试看。
开gradient checkpointing啊,显存能省一大截,fp16加这个基本就稳了。
开gradient checkpointing能省一大截显存,fp16配合这个基本就稳了,batch size 1加梯度累积没问题。
gradient checkpointing大概率是元凶,7B模型跑LoRA虽然只更新adaptor权重,但前向激活值照样全量存,fp16下序列512+bs1的激活峰值比你想象的高很多。我之前在3090上调13B也遇到过类似情况,开了checkpointing之后显存直接砍半,代价就是训练慢个20%左右,但至少能稳定跑完。另外你gradient accumulation调到8其实是为了模拟大batch,但对显存峰值没用,主要看单步的激活占用,建议把accumulation降回1或者2,先跑通再慢慢加。如果还爆的话检查下transformers版本,有些老版本对peft的显存优化有bug,升到4.38+能好不少。
fp16显存涨大概率是激活没释放,开gradient checkpointing能省一大截,试试看。
fp16开了但梯度检查点没开,这基本就是主因了,7B模型就算LoRA,激活值在512序列下也吃得很凶,尤其代码任务注意力计算更占显存。建议把gradient checkpointing打开,显存能省一半多,batch size甚至可以试着提到2。另外你确认下是不是transformers版本太新,有些版本对peft的显存优化有bug,换个稳定版比如4.38左右试试。我跑7B代码模型一般就是bs1+gc+fp16,稳得很,峰值大概25G左右。