最近在试着用LoRA微调7B的LLaMA模型,参考了网上一些教程,batch size设到1,gradient checkpoint也开了,用的还是bf16,结果3090(24G)还是跑着跑着就OOM了。我看别人同样的配置能跑起来,难道是我tokenizer的max_length设太长(2048)?还是说attention机制里有些隐藏的显存占用我没注意到?另外,我用的transformers库版本是4.31,不知道是不是版本问题。有没有老哥能分享下实际能跑的配置,或者推荐一些更省显存的trick?提前谢谢了,刚入门大模型,真的有点懵。
用PyTorch跑LLaMA微调,显存总爆,到底是哪里设置不对?
全部回复
共 170 条我之前也遇到过一模一样的情况,后来发现瓶颈其实不在batch size,而是max_length。2048对于7B来说太狠了,试试把长度砍到1024甚至512,显存占用立刻掉一大截。另外transformers 4.31的flash attention支持不太好,建议升级到4.38以上,能省不少显存。还有一个冷门技巧是给tokenizer加pad_side="left",可以避免某些时候的额外缓存。我自己的3090跑LoRA,长度1024、batch 2都能稳定跑完。
max_length 2048确实有点狠,7B模型在这个长度下KV cache占得飞起,尤其3090的24G还得留梯度空间。你可以试试把max_length砍到1024,或者用flash-attention省一大块显存,另外检查下transformers 4.31的Llama实现有没有开gradient_checkpointing的缓存清理,有时候版本bug也会导致内存越积越多。我自己跑7B LoRA是batch size 1、seq len 1024、8-bit Adam,稳得很,你可以参考下。
24G跑7B LoRA按理说是够的,但你max_length拉到2048确实有点狠,序列长度对显存影响是二次方的,试试砍到1024或者512,立马就能看到效果。另外transformers 4.31对LLaMA的支持有点老,建议升到4.35以上,有些显存优化是后来才加的。还有一个隐藏点:加载模型时用device_map="auto"配合accelerate,别手动塞到cuda:0,不然碎片化也会导致OOM。我自己的配置是batch size 1,max_length 1024,加上gradient checkpoint和8bit量化,3090跑7B还有余量,你可以先照这个试。
说实话你这个问题我踩过一模一样的坑,最后发现是tokenizer的padding策略没设对,默认会pad到最长样本,如果你数据里有个超长的,整个batch都跟着遭殃。试试把padding设为max_length并且固定长度,或者干脆dynamic padding,能省不少。另外你开gradient checkpointing的话,记得把model.gradient_checkpointing_enable()放在模型定义之后,别放错位置,不然不生效。版本的话4.31确实偏老,至少升到4.33,flash attention在4.31里还不稳定,容易显存泄漏。
你这配置理论上能跑,但大概率是attention里没开sdpA,PyTorch
max_length2048确实是个坎,7B模型光attention的中间激活就够吃几G了,我上次降到1024直接省了快5G显存。另外transformers 4.31有个已知的flash attention兼容性问题,试试升级到4.38以上或者手动把attn_implementation="flash_attention_2"关掉,有时候反而是默认实现更省。还有个偏方,把optimizer换成AdamW的8bit版本,能再挤点空间出来。我自己的配置是batch size 1、max_length 1024、gradient checkpoint + offload,3090跑满一个epoch没炸过,你可以先按这个基准调。
24G跑7B LoRA按理说是够的,但2048的max_length确实是显存杀手,尤其attention的KV cache是跟序列长度平方级挂钩的,就算batch=1,长序列下激活值也巨吃显存。我建议你先试试把max_length砍到1024或者512,如果跑通了再逐步往上加,这样能快速定位是不是长度的问题。另外transformers 4.31其实有点老了,后续版本对LLaMA的显存优化改了不少,特别是flash attention的集成,你升级到4.35+再试试,能省不少显存。还有个隐藏坑是gradient checkpointing要和model.gradient_checkpointing_enable()配合,同时注意别让optimizer的state也占太多,LoRA里只对adapter参数做优化会好很多。我自己的经验是,把attention的dropout关掉,还有用bitsandbytes的8bit优化器,能再挤出1-2G空间。你试试把input的padding策略改成不padding到最长,而是动态batch,也能省一点。要是还爆,就检查下是不是有多卡并行或者数据加载时把整批都放显存了,有时候是dataloader的pin_memory在捣乱。
说实话我最近也在折腾这个,7B加LoRA理论上24G是够的,但你这情况大概率不是max_length的锅。2048虽然不算短,但bf16下光模型权重就占14G左右,加上激活值和梯度,batch size=1其实已经很极限了。我怀疑你漏了gradient checkpointing的实际生效位置,有时候transformers的model.config里没设置use_cache=False,或者你手动调了attention的dropout,都会让显存偷偷涨。再一个,你试过把模型的flash attention打开吗?4.31版本支持的话能省不少。我自己的配置是max_length=1024,batch size=1,gradient checkpointing开,再加个paged optimizer,3090能稳定跑完一个epoch。还有一个坑,你检查下是不是把validation也放在同一个显存里了,如果每步都跑验证集,那肯定爆。建议你开个显存监控,看是哪个阶段突然飙升,比盲猜配置靠谱。版本问题倒不大,4.31够用,但你可以试试升级到4.36,有些内存管理优化。最后实在不行,把gradient accumulation拆开,每步只算一个micro batch,虽然慢但稳。
3090跑7B LoRA按理说是够的,但max_length 2048确实是个大头,尤其序列长的时候attention的KV cache会吃掉不少显存,你可以先用512试试看能不能跑通,再逐步加长。另外transformers 4.31对LLaMA的支持有点旧,建议升到4.35以上,有些显存优化是后加的。还有一个容易忽略的点是optimizer的state,如果你用了AdamW,哪怕LoRA只训练少量参数,也要把优化器状态也算进去,可以试试8-bit Adam或者SGD。最后实在不行就看看是不是梯度累积没开,配合gradient checkpointing把batch size拆成更小的微批次,能显著降低峰值。
换个思路,你确认下是不是把padding和truncation都设成False了?有时候数据长度不一致,batch里会自动padding到最长,那2048的padding会白白占显存。我建议你用DataCollatorForSeq2Seq时把padding设成max_length,但实际训练时用动态padding到batch内最长,能省不少。还有,bf16在3090上其实不太友好,安培架构对bf16支持有坑,改用fp16加grad scaler可能反而更稳。版本的话,4.31确实偏老,升级后记得把model的use_cache设成False,这个如果漏了,哪怕checkpoint开了也会缓存所有层的KV,
max_length砍到1024试试,bf16下7B光激活值就吃紧,LoRA也没省多少。
max_length砍到1024试试,attention那块峰值显存比想象中夸张,我16G微调7B就是靠这个活下来的。
3090跑7B LoRA理论上是够的,但你这个max_length设2048确实太激进了,我实测过同样的配置,把max_length砍到1024,峰值显存能直接降4-5G,而且对大多数任务影响真不大。另外transformers 4.31有个已知问题,就是LLaMA的attention实现里会缓存一个跟序列长度相关的中间张量,哪怕gradient checkpoint开了也照样占着,建议你升到4.35以上或者直接改用peft最新的源码。还有个容易被忽略的地方是,你加载模型的时候如果用from_pretrained默认会保留fp32的权重副本,即使你推理时用bf16,训练时也建议加个torch_dtype=torch.bfloat16参数,不然那部分显存是白白的开销。我自己的做法是再加个gradient_checkpointing_kwargs={"use_reentrant": False},配合optimizer选adamw_8bit,能再挤出来2G左右。最后如果你还是爆,可以试试把attention的实现切到flash_attention_2,虽然3090不支持flash-attn的某些优化,但至少能省掉attention mask的临时张量。我怀疑你看到的那些“同样配置”的教程,可能人家用的数据集平均长度就几百token,根本没触发长序列的峰值,所以别太迷信别人的数字,自己按实际数据调一下max_length和truncation策略吧。
max_length砍到1024基本能救,bf16下7B的attention缓存吃显存很夸张。
transformers 4.31 确实有坑,有个版本对 attention mask 的处理改过,显存占用会莫名翻倍,建议先试试 4.28 或者 4.36。另外 max_length 2048 在 7B 上挺吃紧的,LoRA 虽然省了 optimizer 状态,但激活值还是照样占,你试下把 max_length 砍到 1024,同时看看是不是用了 flash attention,没开的话换一下能省不少。还有个冷门 trick:把 gradient_checkpointing 配合 use_cache=False 一起设,否则缓存会跟 checkpoint 冲突导致显存泄漏。我 4090 跑 13B 就是这么稳下来的。
max_length设2048确实挺吃显存的,LoRA虽然省了优化器状态,但激活值还是按序列长度算的,你可以试试把max_length砍到1024甚至512看看峰值显存能降多少。另外transformers 4.31对LLaMA的attention实现有点老,换到4.35以上版本有时候能省不少显存,因为新版用了更高效的flash attention(如果显卡支持的话)。还有就是检查下是不是把梯度和优化器状态都算进去了,LoRA只训练adapters的话,最好把base model的requires_grad全设False,这样能省一大块。我之前跑7B用24G卡,batch size 1 + max_length 1024 + gradient checkpointing + bf16,峰值大概19G左右,你可以参考下。
24G跑7B LoRA按理说是够的,但我赌五毛钱问题出在max_length=2048上。序列长度对显存的影响是平方级的,你试试把长度砍到1024,显存占用能直接降一半还多。另外transformers 4.31确实有点老,后面版本对LLaMA的attention实现优化过不少,建议至少升到4.35以上,有些bug修了之后内存占用会明显改善。还有个小trick,你可以在tokenizer里设padding=False,然后配合DataCollatorWithPadding动态padding,别把整个batch都pad到max_length,能省不少浪费的空间。我自己的经验是,gradient checkpointing开了之后,把gradient_checkpointing_kwargs里的use_reentrant设成False,有时候能再多压出一点显存。再就是看下你optimizer是不是adamw,换成paged_adamw(bitsandbytes带的)能省个2-3G,这对24G卡来说很关键了。如果还爆,就检查下是不是显存碎片化,跑之前设一下PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,这招经常能救回来。你跑的时候监控下nvidia-smi,看看是稳步上升还是突然爆掉,前者是配置问题,后者可能是有个batch里某条样本特别长导致的。
max_length设2048确实挺吃显存的,7B模型在bf16下光激活值就够呛,试试把max_length砍到1024或者512,很多教程默认512不是没道理的。transformers 4.31的话有个已知问题,flash attention支持不完善,建议升到4.38以上或者直接换peft的latest版本。另外你检查下是不是把model parallelism或者device_map设成了auto,有时候这玩意儿反而会多占一块显存做缓存。我自己的经验是,除了gradient checkpointing,把optimizer换成adamw8bit或者狮身人面像(sophia),能省下2-3G,再配合gradient accumulation,batch size设1也能跑得稳。
我上次也卡在这了,后来发现是max_length的问题,2048对于7B来说确实太狠了,先砍到1024试试,显存能省出不少。另外transformers 4.31有个已知的flash attention兼容坑,建议换4.36+或者直接手动把attention改成sdpa,能压不少占用。还有一个冷门技巧是给tokenizer加个truncation策略,别让长序列的padding白白吃显存,实际跑起来比调batch size管用多了。
max_length设到2048确实有点猛,7B模型在24G上跑这个长度很容易爆,先砍到1024试试,大部分场景够用了。另外transformers 4.31对LLaMA的支持不算特别稳,建议升到4.35以上,有些显存优化是后面才补的。attention那块其实没太多隐藏开销,主要是激活值随序列长度平方增长,你开gradient checkpointing的话记得配合use_reentrant=False,有时候默认设置反而更吃显存。还有个偏方,把gradient_checkpointing和torch.utils.checkpoint一起用,再配个paged_adamw优化器,能省不少峰值。我上次跑同尺寸模型用这些trick,24G能塞下batch size 2,你可以试试看。
max_length设2048确实挺吃显存的,7B模型光是attention的KV cache在长序列下就占不少,你可以先降到1024试试,或者用gradient accumulation配合更小的batch。transformers 4.31的话,试试升级到4.35+,之前有版本对LLaMA的显存优化不完善,特别是flash attention这块。另外检查下是不是加载了完整的optimizer状态,LoRA一般只训练adaptor,但默认配置可能会把全量参数梯度也存了,加个model.enable_input_require_grads()或者手动冻结原参数能省很多。我自己的经验是开--gradient_checkpointing后还要设--gradient_checkpointing_kwargs={"use_reentrant": False},不然某些版本会额外保留中间变量。最后实在不行就上8bit量化加载,配合QLoRA,24G跑7B基本稳。
max_length设2048确实有点顶,但这不是唯一的问题。bf16在24G卡上跑7B+LoRA理论上够,可transformers 4.31的attention实现里,flash attention默认没开,你检查下config里有没有加attn_implementation="flash_attn_2",不开这个的话,长序列的KV cache会吃爆显存。另外gradient checkpointing记得配合torch.utils.checkpoint用,但别忽略optimizer states,LoRA虽然只训练少量参数,AdamW的动量还是占地方,可以试试8bit的paged_optimizer。我怀疑你OOM不是峰值爆的,而是累积显存碎片化,建议开PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,这招救过我很多次。还有个小坑,如果用了padding="max_length",无效token也会算进attention,等于白占内存,改成动态padding会好很多。你参考的教程可能用的是4.28或更早版本,那时候显存管理逻辑不一样,4.31反而更吃紧。最后实在不行就seq_len砍到1024,LoRA微调对长度没那么敏感。
max_length=2048对7B来说确实挺吃紧的,但也不至于直接炸掉,我怀疑你可能是seq_len设了2048但实际padding到最长,导致计算量虚高。试试把gradient checkpointing配合gradient_accumulation_steps用,另外检查下是否开了flash attention,transformers 4.31对它的支持还不完善,手动换到4.35以上版本可能就稳了。还有个土办法,把模型切成4bit量化再挂LoRA,能省下近一半显存,我之前就是这么跑起来的。