最近想试试微调LLaMA-3-8B,按照网上的教程用LoRA在单卡A100上跑,batch size设了1,结果刚加载模型就OOM了,显存直接飙到70多G。我查了资料说用bitsandbytes做4bit量化能降显存,但装完量化后加载模型报错说“不支持的架构”,一脸懵……
有没有大佬指点一下,具体应该怎么配置量化参数?或者有没有更省显存的方法(比如gradient checkpointing、混合精度之类的)?我主要想做文本分类微调,数据集不大(约5000条),但卡在环境配置这一步了,求靠谱的实操经验!🙏
新手求教:用PyTorch跑LLaMA-3微调,显存爆炸怎么优化?
全部回复
共 156 条LLaMA-3-8B全精度加载本来就要16G+,你A100 70多G还OOM大概率是LoRA配置里target_modules没设对,把embedding和lm_head也训了。4bit报错多半是bitsandbytes版本和transformers不匹配,试试pip install -U bitsandbytes,然后加载时记得加trust_remote_code=True。另外gradient checkpointing一定要开,配合fp16能把峰值压到20G以内,5000条数据其实用QLoRA+8bit优化器就够了,batch size可以慢慢加到4。
刚踩过类似的坑,说下我的配置:bitsandbytes加载时加一句load_in_4bit=True,但得配合bnb_4bit_compute_dtype=torch.float16,不然会报架构不兼容。你那个报错大概率是transformers版本和bnb对不上,升级到最新版试试。另外梯度检查点一定要开,配合混合精度能再省不少。5000条数据的话,其实可以考虑用PEFT的prepare_model_for_kbit_training,显存能压到20G左右。
8B模型用LoRA还爆显存,大概率是你在加载基座模型时没开量化,或者bitsandbytes版本跟transformers不匹配。建议先确认transformers版本≥4.30,然后加载模型时明确写load_in_4bit=True,同时指定bnb_4bit_compute_dtype=torch.float16和bnb_4bit_quant_type="nf4",这几个参数缺一不可。另外你说报“不支持的架构”,很可能是bitsandbytes没跟上CUDA版本,直接pip install -U bitsandbytes试试。
如果量化这条路实在走不通,还有个土办法:用gradient checkpointing把激活值缓存关掉,再配合torch.cuda.amp混合精度,能把峰值显存压到40G以内。不过你的A100是80G版本吧?按理说8B模型即使全精度加载也就16G,LoRA训练峰值撑死30G,你70多G肯定是有东西没配对——比如把device_map="auto"漏了,或者LoRA配置里target_modules写成了全量参数。
数据集5000条做分类任务,其实没必要硬啃LLaMA-3,换成Mistral-7B或者Qwen-7B对新手友好得多,生态兼容性更好,量化方案也成熟。真要用LLaMA-3,记得在transformers的加载函数里加上trust_remote_code=True,有些新架构的代码需要动态加载。先跑通一个最小的示例,比如batch size=1,sequence length限到512,确认没问题再慢慢加长度。
讲真你这情况我太懂了,刚入坑的时候我也在A100上被8B模型怼到OOM过。4bit报不支持的架构大概率是transformers版本和bitsandbytes没对齐,LLaMA-3需要最新版的库,你试试把transformers升到4.40以上,然后加载时加个trust_remote_code=True,有时候官方教程漏了这个。另外你就算量化成功了,LoRA在8B上跑5000条数据,batch size=1配合gradient checkpointing应该能压到20G以内,但我觉得你其实可以更激进一点——既然A100有80G,直接上QLoRA加8bit优化器状态,或者干脆用Unsloth那个库,它对LLaMA-3做了专门优化,加载就省一半显存。至于混合精度,bf16一定要开,但别用fp16,8B模型在A100上fp16容易loss不稳。还有个野路子,你把序列长度用max_length=512截断,文本分类根本不需要长上下文,显存能再降一截。最后实在不行就换更小的模型,比如Llama-3-2-3B,精度损失对你这个任务影响真没那么大。
刚踩过类似的坑,A100跑8B确实得抠显存。4bit报错大概率是transformers版本和bitsandbytes不兼容,建议直接装peft和transformers的nightly版,然后加载时用load_in_4bit=True,加上bnb_4bit_compute_dtype=torch.float16,基本能压到10G以内。另外gradient checkpointing一定要开,配合混合精度fp16,batch size就算1也能稳定跑。5000条数据的话,其实也可以考虑用QLoRA的paged optimizer,显存占用会更稳,就是训练速度会慢一点。
或者换个思路,既然数据量不大,试试用unsloth这个库,它对llama系优化得很狠,4bit下8B模型显存能压到6G左右,而且训练速度还快。报错的话直接看它文档里的安装命令,记得先卸载原版transformers再装它的fork。另外文本分类任务其实可以只微调最后一层,不用碰整个decoder,能省超级多显存,效果也够用。
试试4bit的NF4配置再加gradient checkpointing,8B单卡能压到20G左右,量化报错多半是transformers版本太新,降一级就好。
刚加载模型就OOM大概率是加载时没走量化,LoRA本身省的是训练时的优化器显存,模型权重还是16bit占满的。4bit报错一般是bitsandbytes版本跟transformers不匹配,试试装0.43.x的bitsandbytes配transformers 4.40+,加载时直接指定load_in_4bit=True就行。另外你5000条数据做分类其实用8B有点浪费,可以先把max_seq_len砍到512,配合gradient checkpointing和bf16,A100上4bit+LoRA跑起来显存能压到20G以内。
刚加载就OOM大概率是模型权重和优化器状态没吃透显存,试试load_in_4bit加nf4配置,再开gradient checkpointing,5000条数据够用了。
A100 70多G爆显存大概率是加载权重时用了fp32,试试加载时直接指定torch_dtype=torch.float16,能省一半。bitsandbytes报架构不支持,多半是版本和transformers没对齐,升级到最新版的bnb和peft,然后模型加载加个load_in_4bit=True就行。你5000条数据做分类,其实用LLaMA-3-8B有点大材小用,不如先试试7B以下的模型或者直接上量化后的Qwen,省心很多。gradient checkpointing一定要开,配合混合精度,4bit下8B模型显存能压到12G左右,跑起来很稳。
检查下transformers版本,llama3得用4.37以上,bitsandbytes参数加bnb_4bit_compute_dtype=torch.float16试试。
A100 80G不该连8B都装不下啊,你八成是加载时默认float32了,先试试load_in_4bit=True配合bnb_4bit_compute_dtype=torch.float16,bitsandbytes报架构不匹配多半是版本太旧,升级到0.43以上看看。另外5000条数据真没必要上LoRA,直接冻住前面所有层只训最后几层全连接,配合gradient_checkpointing和bf16,显存能压到20G以内,速度还快不少。
刚踩过这个坑,A100 80G跑8B全参加载本来就要60多G,LoRA省的是优化器状态不是模型本体,建议先确认下是不是没开load_in_4bit=True,而且新版transformers要用BitsAndBytesConfig传quant_method参数,老写法会不认架构。实在不行试试QLoRA+gradient checkpointing,5000条数据用4bit够用,另外记得把模型塞到device_map="auto"里,让bitsandbytes自动分配层到CPU。
4bit报错大概率是transformers版本跟bitsandbytes不兼容,换个新版或者指定load_in_4bit的dtype试试。你数据量才5000条,其实不需要全参数微调,冻结embedding和lm_head,用peft做LoRA,再加gradient checkpointing和bf16,8G显存都能跑。顺便把batch size调成1,梯度累积设个8步,效果一样稳。
看到你说加载模型就70G,我猜你八成是直接用fp16加载了原版权重,LoRA其实省的是训练时的梯度显存,不是加载时的权重占用。4bit那个报错大概率是bitsandbytes版本跟transformers不兼容,你试试装最新版的bitsandbytes,然后把模型加载改成load_in_4bit=True、bnb_4bit_compute_dtype=torch.float16,再配合device_map="auto",A100上8B模型4bit加载大概只要6-7G。
另外你那个“不支持的架构”可能是因为没设trust_remote_code=True,LLaMA-3的权重文件里有些自定义代码需要这个参数。如果你急着跑通,我建议先别碰量化,直接开gradient checkpointing加torch.compile,batch size=1的情况下A100理论上能跑起来,只是会慢一点。
文本分类微调其实不需要完整序列长度,你把max_seq_len截到256或者128,显存直接砍半。还有个小技巧,用paged_adamw_8bit优化器能省一些优化器状态显存。数据集才5000条,其实可以考虑用PEFT的prepare_model_for_kbit_training配合gradient_accumulation_steps=8,效果跟大batch差不多,显存压力小很多。
最后建议你直接去Hugging Face看下transformers官方文档里LLaMA-3的示例代码,他们最近更新了4bit加载的兼容性修复,要是还报错就把torch_dtype=torch.float16显式写出来。我上次踩坑就是漏了这个参数导致量化后张量类型对不上。
加载就爆说明是权重本身吃满了,先试试HuggingFace的load_in_4bit=True,配好bnb_4bit_compute_dtype=torch.float16,应该能压到10G以内。
A100 80G都爆的话大概率是LLaMA-3的tokenizer和attention实现跟bitsandbytes的4bit层不兼容,你换个加载方式试试,用load_in_4bit=True配合bnb_4bit_compute_dtype=torch.float16,然后把trust_remote_code=True加上。另外把gradient_checkpointing打开,加上fp16混合精度,5000条数据其实用QDoRA或者IA3这种参数效率更高的方法也够,不一定非要死磕LoRA。
先确认下transformers版本,新版要传quantization_config而不是load_in_4bit,另外梯度检查点必须开,能省不少。
说到4bit报错不支持的架构,大概率是transformers版本和bitsandbytes不匹配,或者LLaMA-3的config里rope_scaling字段没被识别。我之前也踩过这个坑,建议先检查transformers版本是不是4.35以上,然后试试在加载模型时显式指定quantization_config里的trust_remote_code=True,有些第三方实现需要这个。另外你直接上LoRA+4bit的话,其实可以省掉gradient checkpointing,因为量化后显存占用已经很低了,但A100 80G如果还是爆,可能是你在加载时把device_map设成了auto,结果模型被均匀切到多卡上了,反而浪费了单卡内存。我自己的经验是,5000条数据微调文本分类,根本不用全参数微调,你甚至可以只冻住所有层,单独训练一个classification head,配合fp16,batch size可以开到16以上。还有个小技巧,用accelerate库的init_empty_weights先初始化元设备上的模型,再load到GPU,能避免峰值内存。最后实在不行,试试QLoRA的官方实现,它把4bit的nf4格式和双重量化都封装好了,报错率低很多。你数据集小,其实也可以考虑用更小的模型比如7B的Mistral,效果差不了太多但省心。
A100 70多G确实不正常,LoRA本来不该这么吃显存,你八成是加载模型时没把量化参数传到from_pretrained里,试试load_in_4bit=True加bnb_4bit_compute_dtype=torch.float16,另外检查下transformers版本,太老不支持LLaMA-3架构。如果还报错,干脆用QLoRA官方示例改改,或者直接换Unsloth,它对8B模型优化得很狠,我4bit下batch size能开到8。
这问题我上个月刚踩过,A100 80G跑8B全量加载本来就要16G左右权重,但OOM到70多G大概率是transformers版本跟bitsandbytes不兼容,尤其是新版LLaMA架构的rope scaling那块,4bit量化会识别不了。你试试把transformers降到4.36左右,bitsandbytes用0.43.0,然后加载时加上low_cpu_mem_usage=True和torch_dtype=torch.float16,别用默认的float32,这一步能省一半显存。另外你只做文本分类的话,强烈建议别直接微调全模型,用PEFT的target_modules锁定q_proj和v_proj,rank设8,alpha设16,再把gradient_checkpointing打开,batch size甚至能提到4。还有个坑是dataset的padding策略,别用max_length硬截断,用padding="longest"配合DataCollatorWithPadding,不然5000条数据里长样本会把显存撑爆。最后实在不行就上QLoRA,但记得把bnb_4bit_compute_dtype也设成float16,不然反量化那步会莫名多出几个G。我这边跑通后峰值也就35G左右,效果跟全量微调差不了多少。