最近在试着用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上确实容易翻车,激活值累积起来很要命。建议先把这个打开试试,显存占用能降不少,代价就是慢一点。另外你观察下是不是max_length设了但实际padding没处理好,有时候输入长度波动会导致显存峰值忽高忽低。我之前跑类似配置,batch size=1+gradient checkpointing+fp16,序列长度1024都能稳在35G左右,你可以参考下。
看到你提到显存一直在涨而不是直接爆掉,这其实是很典型的症状。fp16只省了张量存储,但中间激活值默认还是fp32,你序列512加batch 1按说激活不该占太多,不过我怀疑你用的7B模型可能本身KV cache就吃掉了不少显存,加上peft的lora参数虽然小,但反向传播时梯度也需要额外空间。gradient checkpointing大概率是主因,没开的话激活值全部保留,7B模型每层都存一份,跑两步累积下来直接炸很正常,你开一下试试,代价就是慢个30%左右但显存能降一半。另外你注意下transformers版本,有些老版本对fp16和lora的叠加有bug,会额外分配优化器状态,建议更新到最新版。如果还不行,可以试试把gradient accumulation改成4,然后序列长度降到384,代码生成任务其实不需要太长的上下文。我之前调过类似配置,A100 40G跑7B lora,batch 2加gradient checkpointing加fp16,峰值大概在30G左右,你可以参考一下。
大概率是gradient checkpointing没开的问题,LoRA虽然省了大部分参数梯度,但中间激活值照样吃满显存,尤其序列长度512加7B模型,不开checkpointing很容易爆。你可以先开gradient checkpointing试试,batch size保持1,显存占用应该能降一半以上。另外fp16混合精度最好配合torch.cuda.amp用,单纯设fp16=True有时候反而会多留一些缓存碎片。如果还不行,看看是不是transformers版本里use_cache没关,生成时缓存会一直累积导致显存涨。
fp16开了但没开gradient checkpointing,这基本就是显存爆掉的直接原因。7B模型即使LoRA只训练 adapter 参数,forward 过程里base model的中间激活值还是会全部存在显存里,序列长度512、batch 1的情况下,激活大概占6-8G,但加上优化器状态和梯度累积的中间缓存,峰值很容易冲到40G以上。我自己的经验是,开gradient checkpointing能省掉一大半激活显存,代价是慢个20%左右,但至少能跑起来。另外你可以把optimizer换成AdamW的8bit版本,或者直接用paged_adamw_8bit,能再省几个G。还有个坑是transformers加载模型时默认会缓存所有层的hidden state,你可以在modeling代码里把output_hidden_states关掉,或者手动清理一下每步的中间变量。如果还不行,试试把序列长度降到256看是不是稳定,先确认是不是长度导致的峰值问题。配置上我建议fp16+gradient checkpointing+8bit adamw,batch 1,accumulation 8,这个组合在40G上跑7B LoRA是稳的,至少我同参数跑CodeLlama没爆过。
我之前也踩过一模一样的坑,尤其7B在40G上跑LoRA,batch size=1还爆显存大概率不是配置问题,而是你猜的那个方向——gradient checkpointing没开。transformers的模型默认会缓存所有中间激活,序列长度512虽然不长,但7B的层数深,激活值累积起来非常吓人,加上fp16只是减半了参数和梯度的内存,激活值该占多少还是多少。你把gradient checkpointing打开,内存能掉三分之一到一半,这是最直接有效的解法。另外优化器状态也得盯一下,LoRA虽然只训练adaptor参数,但如果你用了AdamW,它的状态还是按全量参数尺寸算的,除非你显式指定了只优化LoRA参数,否则等于白省。还有个我后来发现的坑是peft的默认实现可能把base model的梯度也保留了,建议在training_args里加上remove_unused_columns=False,同时确认model.gradient_checkpointing_enable()真的生效了。如果还不行,可以试试把fp16换成bf16,A100对bf16支持更好,有时候数值稳定性反而能减少显存碎片。我自己的稳定组合是batch size=1、gradient checkpointing开、learning rate 1e-4、LoRA rank=8,跑13B都没再爆过,你可以参考下。
fp16开着但没开gradient checkpointing,这基本就是主因了,7B模型就算LoRA,激活值在512长度下也能吃好几个G,而且你gradient accumulation设8,反向传播时梯度累积本身不会额外占显存,但如果你没开checkpointing,中间变量全攒着,跑两步爆很正常。我建议先把gradient checkpointing打开,显存能省一半左右,batch size可以保持1,accumulation调到16,序列长度如果数据允许降到384也行。另外检查下transformers版本,老版本对fp16的显存优化有bug,升级到4.38以上试试,我之前遇到过类似问题,换了版本就好了。
开gradient checkpointing吧,显存能省一半,fp16不加这个7B很容易炸。
fp16开着但没开gradient checkpointing的话,7B模型光中间激活就能吃掉十几个G,你batch size=1加梯度累积8其实等效batch没变,显存峰值还是没降下来。建议先把checkpointing开了,能省一半以上,另外优化器状态可以用8bit adam或者adamw-torch+foreach试试。我跑7B一般序列512、batch1、acc16,开checkpointing后峰值大概22G左右,你参考下。