最近想试试微调LLaMA-3-8B,按照网上的教程用LoRA在单卡A100上跑,batch size设了1,结果刚加载模型就OOM了,显存直接飙到70多G。我查了资料说用bitsandbytes做4bit量化能降显存,但装完量化后加载模型报错说“不支持的架构”,一脸懵……
有没有大佬指点一下,具体应该怎么配置量化参数?或者有没有更省显存的方法(比如gradient checkpointing、混合精度之类的)?我主要想做文本分类微调,数据集不大(约5000条),但卡在环境配置这一步了,求靠谱的实操经验!🙏
新手求教:用PyTorch跑LLaMA-3微调,显存爆炸怎么优化?
全部回复
共 156 条学到了,感谢分享!
我最近也在折腾这个,4bit报错大概率是transformers版本太老,不支持LLaMA-3的架构映射,升级到4.40以上基本能解决。你A100都70G爆了,纯属加载原始权重没开量化,建议直接load_in_4bit加bnb_4bit_compute_dtype=torch.float16,显存能压到15G以内。另外gradient checkpointing和混合精度必开,尤其你的数据集才5000条,LoRA+4bit完全够用,别用全参数微调。
A100 70多G确实不对劲,你八成是直接加载了fp16权重没开device_map,llama3官方模型默认就会把整层都塞进显存。4bit报错大概率是transformers版本太老,bitsandbytes得配合accelerate一起更新到最新才行。你试试load_in_4bit=True加上bnb_4bit_compute_dtype=torch.float16,然后model.enable_input_require_grads(),基本能压到12G以内。另外你这数据量用LoRA其实没必要上8B,试试Llama-3-8B的蒸馏版或者直接换Qwen2-1.5B,跑起来会舒服很多,微调效果也不差。
刚加载模型就70多G,说明你加载的是fp16原版,LoRA本身不省显存,省的是优化器状态和梯度。4bit报错大概率是transformers版本和bitsandbytes不兼容,换个0.42以上的版本试试,加载时加load_in_4bit=True,同时把bnb_4bit_compute_dtype设成torch.float16。另外gradient checkpointing必须开,A100上8B模型4bit+checkpointing+LoRA,显存能压到12G以内,5000条数据跑分类完全够用。
我上次也踩过这个坑,后来发现torch.compile也能省不少,但跟bitsandbytes有冲突,建议先跑通再优化。你那个不支持的架构报错,可以试试先加载meta设备再替换模块,或者直接用QLoRA的官方脚本改,别自己拼。混合精度建议fp16,bf16在某些卡上会慢。
我之前也踩过这个坑,8B模型直接加载光权重就得16G,加上优化器状态和激活值,70G不奇怪。你那个bitsandbytes报错大概率是transformers版本和模型config没匹配上,试试升级transformers到4.40+,然后加载时明确传quantization_config,别用默认的。另外梯度检查点真的香,能省一半激活内存,配合4bit量化加batch size 1,我实际跑下来峰值也就20G出头。你数据集才5000条,其实可以考虑用QLoRA的paged_adamw优化器,显存不够时它会自动溢出到CPU,慢是慢点但至少不崩。
4bit得用QLoRA那套,别直接load_in_4bit,再开gradient_checkpointing和bf16,5k条数据小batch完全够。
这问题太典型了,我上周刚踩完同一个坑。你那个bitsandbytes报错大概率是transformers版本和bnb的兼容性问题,LLaMA-3的架构比较新,得用transformers>=4.40,然后bitsandbytes最好升到0.43以上,加载时用load_in_4bit=True加上bnb_4bit_compute_dtype=torch.float16,别用默认的float32,不然计算图还是吃满显存。另外你光看加载模型就70G,其实LoRA本身只占一点点,大头是优化器状态和中间激活值,建议把gradient checkpointing打开,这个能省掉大部分激活显存,代价就是训练慢个20%左右,但绝对值得。混合精度的话用torch.cuda.amp的autocast就行,配合4bit量化,8B模型应该能压到12G以内。还有个隐藏技巧是device_map="auto",让accelerate自动切分层到不同设备,但单卡上意义不大。最后你那5000条数据其实可以试试把per_device_train_batch_size先设成1,然后gradient_accumulation_steps设4,效果等同batch size 4但显存只吃一份,先跑通再调参数。环境这块别硬刚,直接建个新的conda环境装peft+transformers+accelerate+bitsandbytes的latest版本,成功率最高。
A100 80G都爆的话大概率是加载时没开低精度,model = AutoModelForCausalLM.from_pretrained(..., torch_dtype=torch.float16) 这步很容易漏,加上bf16能省一半。bitsandbytes那个报错多半是transformers版本和bnb不兼容,建议直接装最新版,然后加载时写load_in_4bit=True,同时指定device_map="auto",另外quantization_config里记得设bnb_4bit_compute_dtype=torch.float16,不然默认float32照样爆。5000条数据其实用不着全参数微调,LoRA r=8就够,梯度检查点开一下能再省30%左右,我上次跑7B就是这么过来的。
试下transformers加载时传device_map="auto"加load_in_4bit=True,bitsandbytes报架构错误多半是版本没对齐,LLaMA-3需要transformers>=4.40和bnb>=0.43。另外你数据集才5000条,其实用QLoRA把lora_r设成8,加gradient_checkpointing和bf16,显存能压到12G左右,A100绰绰有余。我前两天刚跑通类似任务,关键是加载后先冻结原模型再配peft,别直接改config。
这问题我上周刚踩过,A100 80G跑8B全量加载确实会爆,但LoRA不至于70G,你八成是没开8bit或4bit的load_in_4bit=True,bitsandbytes报错大概率是transformers版本和模型类不匹配,试试升级transformers到4.36+,或者直接用AutoModelForCausalLM.from_pretrained(..., quantization_config=BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=torch.float16))。另外gradient checkpointing必须开,再加个torch.cuda.amp自动混合精度,5000条数据其实用QLoRA+LoRA rank=8就够,显存能压到15G以内,别一上来就追求满血。
A100 80G按理说跑8B LoRA是够的,你OOM大概率是加载基座模型时峰值太高,试试先load模型再套LoRA,别一上来就整全套。量化报错大概率是transformers版本和bitsandbytes不匹配,建议直接上peft的官方文档看支持的架构列表,或者换bnb的4bit加上torch_dtype=torch.float16参数试试。另外你这个数据量做分类任务,其实可以考虑freeze所有层只训分类头,配合gradient checkpointing和混合精度,显存能压到20G以内。
8B全精度加载就要16G,A100 70G不该爆,八成是LoRA配置里target_modules写错了导致全参微调。
试试transformers的load_in_4bit加device_map="auto",再配合gradient_checkpointing,5000条数据绝对够用。
我当初也踩过这个坑,A100 80G加载8B原模型本来就快满了,LoRA省的是训练时的梯度显存,加载时该爆还是爆。你那个bitsandbytes报错大概率是transformers版本太老,升级到4.40以上然后加载时直接传load_in_4bit=True,再配合bnb_4bit_compute_dtype=torch.float16试试,应该能压到10G以内。另外gradient checkpointing记得开,配合混合精度能再省不少,5000条数据其实用QLoRA完全够,没必要全参数微调。
8B模型用LoRA还爆显存,大概率是加载基座模型时没开低精度,默认fp32直接就把显存吃满了。你先把torch_dtype=torch.float16加上,光这一步就能省一半,然后再开gradient checkpointing,LoRA的target_modules别全选,只挑q_proj和v_proj,显存能压到20G左右。bitsandbytes那个报错,可能是transformers版本和bnb不兼容,建议换个4.40以上的transformers,然后加载时明确指定quantization_config里的quant_method="bitsandbytes"和bnb_4bit_quant_type="nf4",别用fp4。另外你5000条数据做分类,其实用LLaMA-3-8B都有点大材小用,可以试试Mistral-7B或者更小的模型,收敛更快,调参成本也低。如果非要坚持LLaMA,建议把sequence length限制在512以内,文本分类用不到长上下文,padding和truncation都设好,显存还能再降一截。还有个野路子,用DeepSpeed的ZeRO-3配合CPU offload,单卡也能跑,但速度会慢不少,你自己权衡一下。
LLaMA-3用4bit得改trust_remote_code=True,A100跑8B其实bf16+gradient checkpointing就够了,5000条数据LoRA完全够。
4bit量化得用BitsAndBytesConfig指定load_in_4bit=True,再配torch_dtype=torch.float16试试,另外梯度检查点必须开,5000条数据其实用QLoRA完全够。
检查下bitsandbytes版本和transformers是否匹配,新版要传bnb_4bit_compute_dtype=torch.float16,另外记得开gradient checkpointing,显存能再省一半。
8B裸跑本来就要70G左右,你直接上4bit是对的,报错大概率是transformers版本太老不认LLaMA-3的架构,升级到4.40以上再试。微调5000条数据其实不用上LoRA,用PEFT的IA3或者直接冻结embedding只训分类头,A100 40G都够。另外gradient checkpointing记得开,配合bf16能把峰值再砍一半。实在不行就换QLoRA,load_in_4bit加nf4量化,模型加载完才6G,训练也就10G出头。
8B全参加载本来就要60多G,A100单卡确实紧张,但LoRA不该爆成这样,大概率是transformers版本太新或者bitsandbytes没适配LLaMA-3的架构。你试试把load_in_4bit=True和bnb_4bit_compute_dtype=torch.float16写进BitsAndBytesConfig,然后模型加载时加个device_map="auto",基本能压到12G以内。另外gradient checkpointing必须开,配合fp16混合精度,5000条数据完全够跑。实在不行就换QLoRA,效果差别不大,环境问题会少很多。
说实话你这情况挺典型的,A100 80G看着大但8B全参加载就得占16G左右,LoRA+4bit其实完全能压到10G以内。你报错大概率是transformers版本太旧,不支持LLaMA-3的架构映射,把transformers和accelerate都升到最新版,然后bitsandbytes加载时加上trust_remote_code=True试试。另外gradient checkpointing一定要开,配合bf16混合精度能再省不少,5000条数据其实用QLoRA跑两三个epoch完全够。对了,如果你只是做文本分类,不如直接试下PEFT库里的prefix tuning,显存占用比LoRA还低,省得折腾量化。