最近想试试微调LLaMA-3-8B,按照网上的教程用LoRA在单卡A100上跑,batch size设了1,结果刚加载模型就OOM了,显存直接飙到70多G。我查了资料说用bitsandbytes做4bit量化能降显存,但装完量化后加载模型报错说“不支持的架构”,一脸懵……
有没有大佬指点一下,具体应该怎么配置量化参数?或者有没有更省显存的方法(比如gradient checkpointing、混合精度之类的)?我主要想做文本分类微调,数据集不大(约5000条),但卡在环境配置这一步了,求靠谱的实操经验!🙏
新手求教:用PyTorch跑LLaMA-3微调,显存爆炸怎么优化?
全部回复
共 156 条检查下transformers版本,新版才支持LLaMA-3的4bit,另外开gradient checkpointing能再省一截。
试试把加载模型的代码改成load_in_4bit=True加bnb_4bit_compute_dtype=torch.float16,之前我也是这问题,升级库就解决了。
8B单卡其实不用量化,先开gradient checkpointing加bf16,batch size再砍半试试。
刚看完你的描述,我第一反应是“A100 70多G显存爆了”这个点太真实了,8B模型全精度加载本来就得16G左右,但LoRA按说只训adapters,问题多半出在你加载基座模型时默认用了fp32,而且没开gradient checkpointing。4bit量化报“不支持的架构”大概率是bitsandbytes版本跟transformers版本不匹配,或者你加载时没传device_map='auto',试试把bitsandbytes升到0.43以上,然后加载时加一句load_in_4bit=True,再加上bnb_4bit_compute_dtype=torch.float16,基本能解决。另外你数据集才5000条,其实不用纠结全参微调,直接上QLoRA,把lora_r设成8,lora_alpha设成16,target_modules里把q_proj和v_proj都加上,显存能压到12G左右。还有个小技巧,训练时把num_workers设成0,有时候DataLoader的预取也会偷偷吃显存。要是还卡,就开梯度累积,batch_size设1但accumulation_steps设4,效果跟batch_size=4差不多,但显存占用能再降一截。最后别忘了一件事:加载分词器时把padding_side='left'设上,文本分类微调时这个细节能省不少显存。
A100 80G都爆的话,八成是加载时用了float32,试试load_in_4bit=True加bnb_4bit_compute_dtype=torch.float16,另外quantization_config里记得指定trust_remote_code=True,LLaMA-3的架构有点特殊。你那个报错大概率是transformers版本太旧,升到4.40+应该能解决。5000条数据其实不用全量微调,用LoRA加gradient checkpointing,batch size 1再加梯度累积,显存能压到20G以内。
看到你说刚加载模型就70多G,我猜你八成是没开4bit就直接硬怼了,LLaMA-3-8B的fp16权重本身就占16G,加上优化器状态和中间激活值,单卡A100确实扛不住。bitsandbytes报“不支持的架构”多半是版本没对齐,你检查下transformers和bnb的版本,新版transformers已经内置了BitsAndBytesConfig,直接load_in_4bit=True,再配合torch_dtype=torch.float16,基本能压到6-8G显存。不过你既然做文本分类这种小任务,其实不用全量微调,用PEFT的LoRA再加gradient checkpointing,batch size开到4-8都没问题,显存占用大概10G出头。另外,你那5000条数据建议先看下类别分布,如果分类数不多,试试把模型改成AutoModelForSequenceClassification后再套LoRA,别在生成模型上硬改,能省不少事。还有个小坑,4bit下最好把bnb_4bit_compute_dtype设成fp16,不然CPU offload会很慢,你如果还卡在环境阶段,可以贴下transformers和torch版本,我帮你看看是不是兼容性冲突。
量化加载报错大概率是transformers版本和bitsandbytes不匹配,升到最新版然后model_kwargs里加load_in_4bit试试,A100跑8B四比特再加梯度检查点肯定够。
之前也踩过这个坑,4bit量化加载报错大概率是transformers版本和bitsandbytes不匹配,试试把transformers升到4.35以上,然后量化配置里加上trust_remote_code=True,LLaMA架构就能识别了。另外你batch size=1还爆显存有点不正常,先确认下是不是把梯度和优化器状态也算进去了,开gradient checkpointing能省一半,混合精度bf16在A100上很稳。5000条数据其实不用上LoRA,直接冻住前面层,只训最后几层分类头,显存能压在20G以内,速度还快。
你这报错八成是transformers版本太老,跟bitsandbytes的4bit不兼容,换个0.42以上的版本基本能解决。另外A100跑8B用LoRA其实不该爆,先确认下是不是把全量参数都加载了,记得用load_in_4bit=True同时把llm_int8_enable_fp32_cpu_offload加上。我上次跑类似任务,batch size=1加gradient checkpointing,峰值也就20G出头,你试试把flash attention也打开,能省不少。实在不行就换QLoRA,5000条数据效果差不了多少。
报错大概率是transformers版本太旧,升级到最新版再试下,4bit加gradient checkpointing基本就能跑起来。
刚踩过类似的坑,A100 80G跑8B全量加载本来就得60多G,LoRA只省了梯度那部分,模型权重还是占满。4bit报错大概率是transformers版本和bitsandbytes不匹配,建议直接升到最新版,然后加载时加一句quantization_config=BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=torch.float16),别用默认的nf4试试。另外5000条数据真不用上8B,换个7B或更小的模型,开gradient checkpointing加fp16,batch size设4都没压力,跑起来也快很多。
8B全参加载本来就吃显存,4bit要用transformers的BitsAndBytesConfig,别直接用bitsandbytes的api。
8B全量加载本来就吃显存,LoRA只是省了训练时的梯度,模型权重该占还得占。你报错大概率是bitsandbytes版本和transformers不匹配,试试升级到最新版,然后加载时加一句quantization_config=BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=torch.float16),顺便把device_map="auto"加上。另外可以开gradient checkpointing,再把fp16=True打开,A100上4bit+LoRA跑8B应该能把峰值压到20G以内,5000条文本分类数据完全够用。
我之前也踩过这个坑,4bit量化报错大概率是transformers版本跟bitsandbytes不兼容,试试升级transformers到4.40以上,然后加载时加一句quantization_config的配置,别用默认参数。另外你单卡A100其实跑8B不算难,开gradient checkpointing和bf16就能省不少,batch size=1都OOM多半是上下文长度设太长了,把max_seq_len砍到512试试,5000条文本分类数据根本用不着长上下文。
顺便说下,如果量化还是搞不定,可以退一步用PEFT的LoRA加torch_dtype=float16直接跑,我上次在4090上这么干过,显存占用大概20G左右,A100绰绰有余,而且速度比4bit还快一点。你那个报错把完整traceback贴出来可能更好排查,要不先试试最简单的fp16+gradient checkpointing组合?
A100 80G跑8B原模型确实紧,但你这70G有点怪,LoRA本身不该吃这么多,检查下是不是没开gradient checkpointing,那个能省将近一半。4bit报错大概率是transformers版本和bitsandbytes不匹配,试试升级到最新版,加载时加load_in_4bit=True,再配个bnb_4bit_compute_dtype=torch.float16。5000条数据其实可以考虑QLoRA+deepspeed stage3,或者干脆用PEFT的prepare_model_for_kbit_training,我上次跑7B就是这么搞定的,显存能压到20G以内。
8B全参加载本来就要60多G,你光看显存占用高其实不奇怪,关键是你LoRA那层配置可能没吃到量化红利。试试用transformers的load_in_4bit=True加bnb_4bit_compute_dtype=torch.float16,然后显存不够就把max_memory设成CPU offload,另外把gradient_checkpointing打开,5000条数据真不用那么大的batch,gradient_accumulation_steps凑一下就行。
刚踩过一模一样的坑,4bit报错大概率是transformers版本和bitsandbytes不兼容,建议把transformers升到4.40+,然后加载时用quantization_config=BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=torch.float16, bnb_4bit_quant_type="nf4"),另外记得先model = prepare_model_for_kbit_training(model)再套LoRA。其实5000条数据做分类,8B有点杀鸡用牛刀了,试试DeBERTa-v3-large或者Llama-3-8B的蒸馏版,显存压力能小一半,训练速度还快。
4bit得用QLoRA,别直接load_in_4bit,试试transformers源码里LLaMA的量化映射,顺便开gradient_checkpointing。
试试Unsloth框架,自带4bit优化,直接改几行代码就能跑,显存能压到20G以内。
这问题我太熟了,刚踩完坑出来。你那个bitsandbytes报错大概率是版本不匹配,transformers和bnb得配套,尤其LLaMA-3的架构名在最新版里改过,你试试把transformers升到4.40+,然后加载时用load_in_4bit=True,同时把bnb_4bit_compute_dtype设成torch.float16,这样能压到6G左右。不过说实话,你A100都70G显存了,跑8B模型原版都能塞下,LoRA加batch size 1理论上不该炸,我怀疑你加载时没开low_cpu_mem_usage=True,模型权重会重复拷贝一份,直接翻倍。另外强烈建议开gradient checkpointing,虽然会慢一点但显存能省30%,再把torch.compile打开,速度能拉回来。至于混合精度,AMP对LoRA微调效果不错,但注意别用bf16在A100上,有些算子会回退到fp32反而更吃显存。你5000条数据做分类,其实可以试试QLoRA加8bit优化器,效果和全量微调差距很小,但显存能压到10G以内。还有个野路子,如果只是实验,用device_map="auto"让模型分片到CPU和GPU,速度慢点但绝不会OOM,数据集小的话训练时间也就多个半小时。最后检查下是不是CUDA版本和PyTorch不匹配,有时这也会导致显存分配异常。
看到你这个报错我太有同感了,上周刚踩完同一个坑。bitsandbytes那个“不支持的架构”大概率是因为transformers版本和bnb版本不匹配,或者加载时没传quantization_config参数,我建议你直接升级transformers到4.40以上,然后用BitsAndBytesConfig显式指定load_in_4bit=True、bnb_4bit_quant_type="nf4"和bnb_4bit_compute_dtype=torch.float16,这样基本能稳。另外你单卡A100其实不用太慌,8B模型4bit后权重只占5G左右,但激活值还是会涨,所以gradient checkpointing一定要开,配合torch.cuda.amp混合精度,batch size 1跑起来应该没问题。还有个冷门技巧,把padding_side设为left,文本分类任务里能省不少无效计算。如果还嫌显存紧,可以试试flash_attention,A100上能再压一截。你数据量5000条其实很小,甚至可以试下冻结所有层只训练分类头,显存直接掉到10G以内,效果也不一定差。装环境建议直接用pip install -q -U bitsandbytes transformers accelerate,别用conda源,容易版本冲突。跑通后可以看看loss和验证集表现,有问题随时交流。