最近在试着用LoRA微调7B的LLaMA模型,参考了网上一些教程,batch size设到1,gradient checkpoint也开了,用的还是bf16,结果3090(24G)还是跑着跑着就OOM了。我看别人同样的配置能跑起来,难道是我tokenizer的max_length设太长(2048)?还是说attention机制里有些隐藏的显存占用我没注意到?另外,我用的transformers库版本是4.31,不知道是不是版本问题。有没有老哥能分享下实际能跑的配置,或者推荐一些更省显存的trick?提前谢谢了,刚入门大模型,真的有点懵。
用PyTorch跑LLaMA微调,显存总爆,到底是哪里设置不对?
全部回复
共 170 条我之前也被这个坑过,问题多半出在max_length上,2048对7B来说确实太狠了,尤其3090只有24G,建议先砍到1024试试,显存能省一大截。另外transformers 4.31的attention实现比较老,可以升到4.40以上,用flash attention能明显降低峰值占用。还有个小技巧是开启gradient_accumulation_steps,虽然不省显存但能让你用更小batch稳定训练,配合bf16的话记得把torch.backends.cuda.matmul.allow_tf32也打开。我自己的配置是max_length=1024,batch=2,梯度累积8步,跑7B完全没问题,你试试看。
之前我也在3090上卡了好久,最后发现是max_length的锅,2048对7B来说确实太激进了,降到1024直接稳了。另外你可以试试unsloth这个库,同样的LoRA能省差不多一半显存,我换了之后就没再爆过。transformers 4.31有点老,新版对attention的显存优化了不少,建议升到4.38以上。还有个小trick,把gradient_accumulation_steps设成4,同时配合torch.compile,体感上比单纯开gradient checkpoint更稳。
我最近也踩过类似的坑,最后发现是max_length在作祟。2048对于7B来说确实太激进了,尤其你还在用LoRA,输入序列长度直接决定激活内存的峰值,试试512或者768,显存占用能掉一大截。另外transformers 4.31有个已知问题,就是gradient checkpointing和bf16在某些算子组合下不会真正释放中间张量,建议升级到4.35以上,或者直接换用peft的官方示例代码,那里面隐含设置了torch.nn.functional.scaled_dot_product_attention,能省不少显存。还有个冷门trick,把attention的dropout设成0.0,虽然影响不大,但多少能挤点空间。还有就是你检查下是不是把eval mode的模型也塞进显存了,有时候验证阶段会额外占2-3G。我最后是max_length=1024,batch=1,gradient checkpoint=True,再加一个8bit的AdamW优化器,24G跑起来还剩4G左右。你试试看吧,这种问题基本都是细节堆出来的。
24G跑7B LoRA按理说真够用了,问题大概率出在max_length=2048上。你想想,序列长度对显存的影响是二次方的,2048和1024差着四倍,很多人跑7B微调实际都压到1024或者更短。另外transformers 4.31对LLaMA的attention实现确实有点老,建议升级到4.35以上,那边对SDPA的支持更好,能省不少显存。我自己的经验是,除了gradient checkpoint,还可以把gradient_accumulation_steps调大一点,batch size保持1没问题,但配合4步累积效果一样。还有个容易被忽略的点,optimizer的动量状态也吃显存,你可以试试用paged_adamw_8bit,或者直接切到adamw_8bit,能挤出几个G。如果还爆,就把模型并行或者offload开一下,虽然慢点但至少不会崩。对了,确认下你是不是在训练循环里忘记调用model.train(),有时候eval模式的缓存也会占着不放。最后,你可以开个nvidia-smi实时盯着,看峰值到底出现在前向还是反向,这样能更精准定位。
24G跑7B LoRA按理说够用,但max_length=2048确实挺吃显存的,你可以先砍到1024试试,很多教程其实没明说这点。另外transformers 4.31对LLaMA的attention实现有点老,建议升到4.35+,配合最新的peft版本,显存占用能明显降下来。还有个冷门trick是把gradient_checkpointing的use_reentrant设为False,有时候能省出几个G。我上次用同样配置跑通,是把LoRA的r降到8,target_modules只选q和v,你参考下。
3090跑7B LoRA其实挺极限的,24G看着够用但实际峰值很容易被attention的中间激活值吃满,尤其max_length拉到2048时,序列长度对显存是平方级增长,试试把长度砍到1024或者512,效果可能差别不大但显存会宽裕很多。另外你transformers 4.31有点旧了,4.38之后对LLaMA的attention实现做了不少优化,内存占用能省下10%左右,建议升个版。还有个容易忽略的点是LoRA的target_modules,如果默认把所有linear层都加了,那反向传播时梯度显存也会翻倍,我一般只改q_proj和v_proj,省下的显存能撑住更大的batch。gradient checkpointing开了的话,注意别和torch.compile一起用,俩有冲突反而会爆。另外可以试试把optimizer换成AdamW8bit或者Adafactor,优化器状态占的显存能砍掉一大半。最后实在不行就上序列并行或者offload到CPU,虽然慢点但至少不崩。你那个“别人能跑”的教程,可能人家用的8bit量化加载基座模型,3090上这招特别管用。
max_length设2048确实有点激进,7B模型在24G卡上跑2048长度基本就是极限边缘了,试试把max_length砍到1024或者512,显存压力立刻小一个档次。另外transformers 4.31有点老了,换到4.38以上版本,有些显存优化是后加的,比如flash attention的集成方式变了。还有一个容易忽略的点,就是optimizer的state,AdamW的momentum也吃显存,你可以试下8-bit optimizer,比如bitsandbytes的AdamW8bit,能省不少。我自己的经验是LoRA的r值先设8,target_modules别全勾,只改q_proj和v_proj,跑起来稳很多。
max_length=2048确实是个坎,LoRA虽然省了主干显存,但激活值还是按全量序列长度算的,你试试把max_length砍到1024,显存能掉一大截。transformers 4.31有个已知的attention mask分配问题,旧版本反而更省,建议升到4.38+或者干脆降到4.28。另外检查下是不是把gradient_checkpointing开在了model.enable_input_require_grads之前,顺序错了等于没开。还有个冷门技巧:用unsloth的优化版LoRA,同样配置能多塞一倍batch。我跑7B时是max_length=1024,batch=2,grad_accum=4,峰值才18G,你可以参考下。
3090跑7B LoRA按理说24G是够的,但max_length设2048确实太狠了,我平时用1024都嫌多,你可以先砍到512试试,显存直接少一大截。另外transformers 4.31有点老,有些版本的attention实现会额外吃显存,建议升到4.38以上或者直接换flash-attention,效果立竿见影。还有个偏方,把gradient checkpointing和优化器分片一起开,比如用accelerate的cpu offload,虽然慢点但至少不爆。我上次跑同配置是batch size 2,max_length 768,峰值大概18G,你可以参考下。
max_length降到512试试,我调参时发现这玩意对显存影响比想象中大得多。
24G跑7B LoRA按理说够用,你试着把max_length砍到1024,长文本场景下显存占用是二次方涨的,2048跟1024差出来好几G。另外transformers 4.31有个已知bug,flash attention和gradient checkpoint一起开的时候会重复缓存激活值,升到4.35+或者干脆手动把attn_implementation设为eager试试。我自己的经验是再加个--gradient_checkpointing_kwargs use_reentrant=False,能再省一截。你跑的时候留意下nvidia-smi,如果显存是稳步涨而不是突然爆,多半是数据加载那边缓存没清,加个dataloader pin_memory=False看看。
max_length设2048确实有点顶,7B模型长序列下KV cache吃显存很猛,先砍到1024试试。另外transformers 4.31对LLaMA的attention实现比较老,建议升到4.35以上,flash attention能省不少。我自己的经验是3090跑7B LoRA,batch 1加gradient checkpoint,序列长度1024,再加个8bit优化器,基本能稳在21G左右。还有个小坑,别忘了把model的pad_token_id设好,不然偶尔会多算几个token的梯度。
max_length砍到1024试试,bf16下显存占用跟序列长度几乎线性涨,2048确实容易爆。
max_length设到2048确实是主要嫌疑,但更关键的是你bf16下KV cache的占用没算进去。7B模型光参数就14G,激活值加上KV cache在2048长度下峰值能到20G以上,24G的卡不爆才怪。我试过把max_length砍到1024,同样配置能跑,但损失不小。有个trick是开gradient checkpointing的同时把attention的dropout关掉,能省点显存。另外transformers 4.31确实有点老,换4.36以上版本,LlamaForCausalLM会默认用flash attention 2,显存占用直接少三分之一。你参考的教程是不是用的8卡环境?单卡跑7B的LoRA本来就极限,建议把batch size设成0试试(用gradient accumulation模拟),或者干脆用QLoRA量化到4bit,这样能稳很多。最后检查一下你是不是把模型加载到CPU再搬到GPU的,中间过程会多占一份显存。
我之前也遇到过一模一样的情况,后来发现是max_length的问题,2048对于7B来说确实太吃显存了,我降到1024之后立马稳了。另外你可以试试用unsloth这个库,它对显存优化做得特别好,同样的配置能省出将近一半的占用。transformers版本建议升到4.36以上,旧版本的一些算子分配确实有坑,我之前4.31也是各种爆,换了就好了。还有一个冷门的trick是关掉flash attention,虽然慢点但能省不少临时显存,你可以交叉验证下到底哪一步是瓶颈。
我上次也卡在这,后来发现是max_length的锅,2048对7B来说太狠了,降到1024立马稳了。另外你可以试试gradient_accumulation_steps配合小batch,还有attention里加个sdpa或者flash-attention,显存能省不少。transformers版本的话,4.31确实有点老,升到4.38以上有些内存优化,但别升太新,有些API会变。要是还爆,就看看是不是模型加载时把torch_dtype设成auto了,有时候这个也会偷偷多吃显存。
我之前也卡在这过,你试试把max_length砍到1024,LoRA的target_modules换成只改q和v,显存能省出不少。另外transformers 4.31确实有点老,升到4.38+对LLaMA的显存优化明显,特别是attention那块。实在不行就上8bit优化器,3090跑7B用bitsandbytes的bnb_8bit_adam,我这边峰值能降到18G左右。
你检查下是不是把padding开着了,默认的pad_token会强制把整个batch拉到最长,哪怕你设了batch size=1也白搭。我一般直接把padding设为False,然后配合unpad_inputs那个flash attention的变体,24G跑7B还挺宽裕的。版本的话至少升到4.35,之前有个显存泄漏的bug在4.31还挺严重。
gradient checkpoint开了但没生效吧?记得要把model.gradient_checkpointing_enable()放在模型加载之后,然后input的requires_grad也得设True。我之前就是栽在这,开了等于没开。另外你把gradient_accumulation_steps设成8,实际batch还是1,但能摊平loss波动,显存也会稳一点。max_length其实影响不大,主要看你序列真实长度,别硬塞2048的pad就行。
我最近也踩过类似的坑,而且我最后发现主因还真不是max_length或gradient checkpoint,而是transformers 4.31对LLaMA的attention实现有bug,会在某些条件下把缓存张量重复分配。你试试把库版本升到4.35以上,或者干脆换用peft的官方示例代码跑一遍,它自带flash attention的兼容补丁,显存能直接省掉将近一半。另外,bf16在3090上其实不如fp16稳,因为3090的bf16计算效率是阉割过的,而且某些算子会临时转回fp32,额外吃显存。你那个2048的max_length确实偏大,7B模型就算LoRA,序列长度2048的激活值也够呛,建议先砍到1024,等跑通再往上加。还有个容易被忽略的是optimizer state——你如果用的AdamW,LoRA参数虽然少,但主模型梯度还是会被保留,试试在LoRA配置里把bias设成none,同时把训练参数里requires_grad=True的部分尽量只留在adapter层。我最后是把batch size降到1,gradient accumulation设8,然后显存峰值稳定在19G左右,你可以参考这个组合。
说到这个我太有共鸣了,上周刚在4090上踩完一模一样的坑。你试试把max_length砍到1024,大多数教程为了省事都默认2048,但实际训练时序列长度对显存的影响是二次方的,7B模型在24G卡上撑死也就能吃下1400左右的长度,这个差距比什么batch size都狠。另外检查下你数据处理时有没有做padding策略,如果整个batch按最长样本对齐,哪怕只有一条长的也会把显存顶爆,改成动态padding或者按长度分桶能省出好几G。transformers 4.31确实有点老,LoRA相关的优化在4.36以后才比较完善,建议升级到4.40+,顺带把peft也更新下,有些版本兼容问题会莫名其妙多占显存。还有个冷门trick,把attention的dropout在推理阶段关掉,训练时也调成0,能省一点是一点。最关键的还是建议你装个nvidia-smi盯着看显存曲线,爆之前通常有个突然飙升的峰值,八成是某个中间激活值没释放,可以用torch.cuda.empty_cache()在每步后手动清一下缓存。如果还不行,试试把模型切分到CPU offload,只用LoRA的adapter层上GPU,速度慢点但至少能跑通,先验证loss在降再说优化的事。
max_length砍到1024试试,很多教程没提这个但显存差距巨大。另外4.31确实有bug,升到4.35能省不少。