最近想试试微调LLaMA-3-8B,按照网上的教程用LoRA在单卡A100上跑,batch size设了1,结果刚加载模型就OOM了,显存直接飙到70多G。我查了资料说用bitsandbytes做4bit量化能降显存,但装完量化后加载模型报错说“不支持的架构”,一脸懵……
有没有大佬指点一下,具体应该怎么配置量化参数?或者有没有更省显存的方法(比如gradient checkpointing、混合精度之类的)?我主要想做文本分类微调,数据集不大(约5000条),但卡在环境配置这一步了,求靠谱的实操经验!🙏
新手求教:用PyTorch跑LLaMA-3微调,显存爆炸怎么优化?
全部回复
共 156 条碰到过类似的问题,LLaMA-3-8B光fp16加载就得占16G左右,加上优化器状态和中间激活,单卡A100确实容易爆。4bit量化报错大概率是transformers版本和bitsandbytes不兼容,建议升级到最新版,然后加载时加上load_in_4bit=True和bnb_4bit_compute_dtype=torch.float16。gradient checkpointing一定要开,配合混合精度fp16能省不少,batch size设1的话梯度累积可以设4或8来模拟大batch。数据集小的话其实可以考虑用unsloth或者Axolotl这类封装好的框架,自带优化配置,省得自己调参。
试试llama.cpp的GGUF格式,8B模型4bit量化后显存占用不到6G,单卡轻松跑。
我最近也踩过这个坑,LLaMA-3-8B用bnb的4bit量化确实经常报架构不匹配,可以试试在加载模型时指定bnb_4bit_compute_dtype=torch.float16和bnb_4bit_use_double_quant=True,同时用device_map="auto"让bitsandbytes自动分配。另外梯度检查点必开,混合精度用torch.cuda.amp配合gradient accumulation也能省不少,我试过batch size设1、梯度累计8步,显存能压到20G左右。你那个5000条文本分类的话,其实不用全参微调,用PEFT的LoRA加量化基本够用了。
A100 80G跑8B模型按理说不该直接OOM,检查下是不是装了full precision或者transformer库版本不对导致参数没正确加载。4bit量化报架构不支持的话,试试指定bnb_4bit_compute_dtype=torch.bfloat16,或者换llama.cpp的GGUF格式,那个更省心。另外gradient checkpointing肯定要开,配合混合精度fp16能再省个十几G,5000条数据其实用QLoRA+4bit微调完全够用。
这情况太真实了,我之前也踩过这个坑。bitsandbytes报错可能是版本没对齐,试试pip install bitsandbytes==0.39.0配合transformers>=4.31.0,然后加载时加上load_in_4bit=True和bnb_4bit_compute_dtype=float16。另外梯度检查点加混合精度fp16基本是必开的,能把峰值显存再压下去10G左右。5000条数据的话,用peft的lora加量化足够跑了,别一上来就全参数微调。
刚看到这个帖子,我也踩过类似的坑。LLaMA-3-8B在A100上裸跑确实离谱,单卡70G显存起步很正常,8B模型本身加载fp16就差不多16G,加上优化器状态和中间激活值,不爆炸才怪。你试的4bit量化方向其实是对的,但bitsandbytes加载报错大概率是因为transformers版本没对齐,LLaMA-3的架构比较新,你得升级到4.34以上,或者直接用AutoModelForCausalLM.from_pretrained(..., load_in_4bit=True)配合bnb_4bit_compute_dtype=torch.float16,我记得这样写能绕过架构检测的问题。另外gradient checkpointing一定要开,model.gradient_checkpointing_enable()这行代码能省下至少30%的激活显存,配合混合精度torch.autocast,batch size 1应该能压到30G以内。我自己的经验是,5000条文本分类数据根本不需要全量微调,LoRA的rank设成8或者16,只调attention层的权重,显存占用还能再砍一半。如果还是崩,可以试试ZeRO stage 2或者3,用DeepSpeed跑单卡也能分片优化器状态。不过说实话,最省心的方案是直接上Hugging Face的PEFT库,他们最近修了好多LLaMA-3的兼容问题,load_in_4bit + LoRA + gradient checkpointing一键搞定,我上周刚跑通,建议你直接抄那个官方示例。
说实话你遇到的问题挺常见的,LLaMA-3-8B就算用LoRA,单卡A100(80G)如果没做量化,光是加载模型本身就要占掉将近16G,加上优化器状态、梯度、中间激活,batch size=1也容易爆。你装了bitsandbytes报“不支持的架构”,大概率是版本没对齐,建议试试pip install bitsandbytes==0.41.0配合transformers>=4.35.0,然后加载时加一句quantization_config=BitsAndBytesConfig(load_in_4bit=True),注意模型本身也要支持4bit,比如用AutoModelForCausalLM.from_pretrained。另外gradient checkpointing确实省显存,你可以在model配置里加model.gradient_checkpointing_enable(),配合混合精度torch.cuda.amp或者直接开fp16=True,实测能再省10-20G。还有一个偏方是试试用device_map="auto"把部分层塞到CPU,虽然慢但至少跑得起来,毕竟你才5000条数据,速度不是大问题。最后建议先拿unsloth这个库试试,它专门优化了LLaMA系模型的微调显存,我试过8B用4bit加LoRA,A100上显存峰值才20G出头。
老实说,你这情况我一开始也遇到过,LLaMA-3-8B哪怕用LoRA,基础模型加载本身就很吃显存,4bit量化确实是正解,但报“不支持的架构”大概率是bitsandbytes版本跟transformers没对齐,试试升级到0.43以上版本,或者换用load_in_4bit=True时加上bnb_4bit_compute_dtype=torch.float16,很多坑其实是版本兼容问题。另外你A100单卡的话,梯度检查点(gradient checkpointing)一定要开,这玩意儿能省差不多30%显存,而且对训练速度影响不大;混合精度用torch.autocast配合fp16也能再压一点。还有个容易被忽略的点:数据集小就别用大tokenizer,LLaMA-3默认的tokenizer对中文文本会拆得很碎,实际占的序列长度可能比想象中大,建议把max_length硬性设成512甚至256,反正文本分类不需要太长的上下文。如果量化还卡住,可以试试直接用Hugging Face的AutoModelForSequenceClassification配合PeftModel,官方文档里其实有现成例子,不用自己手写加载逻辑。总之先解决量化版本问题,再一步步开优化开关,5000条数据微调8B模型完全够用,别灰心。
写得挺好,建议补充一些性能数据。
刚跑LLaMA-3-8B微调确实容易踩坑,4bit量化报错大概率是transformers版本和bitsandbytes不兼容,建议直接装最新版的transformers和accelerate,然后加载时加个load_in_4bit=True和bnb_4bit_compute_dtype=torch.float16试试。梯度检查点加混合精度基本是必开的,batch size 1还OOM的话可以试试offload参数到CPU,或者用deepspeed stage2,5000条数据完全够用。我之前也是折腾了两天才跑通,环境版本对齐了就行。
说实话A100 80G跑8B模型用LoRA按理说不会直接OOM的,我怀疑你用的是全量微调的代码,或者transformers版本没更新导致加载了完整模型。4bit量化报错“不支持的架构”大概率是bitsandbytes版本跟你的transformers或accelerate不匹配,建议先升级到最新版,然后明确指定load_in_4bit=True和bnb_4bit_compute_dtype=float16。另外你数据集才5000条,完全可以试试gradient checkpointing加混合精度,把model.gradient_checkpointing_enable()和torch.cuda.amp.autocast()用上,batch size设1的情况下显存能压到20G左右。还有个取巧的办法是直接用unsloth这个库,它对LLaMA微调做了显存优化,4bit模式下8B模型只要6G显存就能跑,省去你自己折腾配置的麻烦。不过文本分类任务记得在LoRA层后面加个分类头,别直接套用生成式微调的模板。
刚跑LLaMA-3碰到这问题很正常,我一开始也这样。用bitsandbytes的4bit量化报错大概率是transformers版本不匹配,建议升级到最新版,加载时加个load_in_4bit=True和bnb_4bit_compute_dtype=float16。另外梯度检查点一定要开,model.gradient_checkpointing_enable(),再配合混合精度fp16,A100上batch size调到4基本稳了。你这5000条数据其实不用全量,试试用QLoRA的config里把lora r设小点,像8或者16,显存能再降一截。
试试load_in_4bit加bnb_4bit_compute_dtype=torch.bfloat16,我这样跑8B单卡16G都没爆。
试试用Qwen2的代码库直接改配置,他们那套4bit量化对LLaMA兼容性更好。
我也遇到过类似的问题,A100 80G按理说跑8B模型加LoRA是够的,但直接加载全精度模型确实会爆。你那个bitsandbytes报错,大概率是transformers版本和bnb不兼容,试试把transformers降到4.35左右,或者用pip install bitsandbytes --pre装最新的预发布版。4bit量化的话,加载时加个load_in_4bit=True,再指定一下bnb_4bit_compute_dtype=torch.float16,基本能降到20G以内。另外gradient checkpointing必须开,model.gradient_checkpointing_enable()一行代码就能省不少显存,混合精度用torch.cuda.amp自动混合精度就行,但注意量化后有些层可能不支持。你数据集才5000条,其实可以考虑用QLoRA,效果差不多但显存更低,我试过8B模型量化后只用12G左右。还有个小技巧,把tokenizer的padding side设成left,推理时能省点显存。要是还报错,可以贴一下具体错误日志,大家帮你看看。
刚跑LLaMA-3-8B我也踩过这个坑,4bit量化报错大概率是transformers版本没对齐,建议把库升到4.38以上,加载时加一句bnb_4bit_compute_dtype=torch.bfloat16。另外梯度检查点一定要开,内存能省个10G左右,混合精度用bf16比fp16稳定不少,实测单卡A100跑4bit LoRA能压到24G以内。你这5000条数据其实可以试试先做prompt tuning,比全量微调更省显存,效果在分类任务上也不差。
量化记得加trust_remote_code=True,或者试试QLoRA直接加载4bit模型,能省不少显存。
刚跑LLaMA-3微调确实容易踩量化坑,bitsandbytes报架构错误大概率是transformers版本跟它不匹配,试试把transformers升到4.35以上,加载时加一句load_in_4bit=True和bnb_4bit_compute_dtype=torch.float16。另外梯度检查点一定要开,配合混合精度fp16能再省10G左右,5000条数据用LoRA+4bit量化完全够用,我8G的卡都跑过。
试试把LoRA的r值降到8以下,再配合gradient checkpointing,8B模型在A100上基本能跑起来。
我也遇到过类似问题,LLaMA-3-8B哪怕LoRA+单卡A100直接上确实容易炸。4bit量化报错大概率是transformers版本不匹配,建议升级到最新版,或者换成AutoGPTQ加载试试,兼容性好很多。另外梯度检查点一定要开,配合bf16混合精度能再省10G左右,batch size可以先设1跑通再调。数据集小的话,其实用QLoRA+4bit微调5000条完全够用,显存能压到20G以内。