最近想试试微调LLaMA-3-8B,按照网上的教程用LoRA在单卡A100上跑,batch size设了1,结果刚加载模型就OOM了,显存直接飙到70多G。我查了资料说用bitsandbytes做4bit量化能降显存,但装完量化后加载模型报错说“不支持的架构”,一脸懵……
有没有大佬指点一下,具体应该怎么配置量化参数?或者有没有更省显存的方法(比如gradient checkpointing、混合精度之类的)?我主要想做文本分类微调,数据集不大(约5000条),但卡在环境配置这一步了,求靠谱的实操经验!🙏
新手求教:用PyTorch跑LLaMA-3微调,显存爆炸怎么优化?
全部回复
共 156 条8B你直接单卡上,就算LoRA也得先把原始权重塞进去,A100 80G其实够用,但你这70多G是加载模型时就爆了,八成是没开device_map=auto把层分散到不同设备上。量化报错大概率是transformers版本和bitsandbytes不兼容,建议直接换成4.36以上版本,然后加载时加个torch_dtype=torch.float16,基本就能跑起来。另外你那5000条数据真没必要全量微调,LoRA加gradient checkpointing足够,batch size=1还能再开个梯度累积,显存能省不少。要是还卡,就把模型切一半到CPU上,虽然慢点但至少能先跑通流程。
刚加载模型就70G肯定不对劲,8B全精度也就16G左右,你八成是没开4bit直接把模型塞进去了。bitsandbytes报不支持架构大概率是transformers版本和bnb没对齐,换个匹配的版本试试。我建议你直接上QLoRA,4bit加载加双卡或单卡A100跑5000条数据完全够,顺便把gradient checkpointing开了,batch size保持1就行。还有,记得把模型用device_map="auto"加载,别手动往cuda上放。
别光看量化,先检查下是不是transformers版本太新跟bitsandbytes没对齐,4bit加载llama架构一般不会报错的。你试试把load_in_4bit=True和bnb_4bit_compute_dtype=torch.float16同时设上,然后模型用from_pretrained直接加载,别先转成fp16再量化。另外梯度检查点必须开,再加上bf16混合精度,A100上8B模型应该能压到20G以内。5000条数据其实不用全量微调,LoRA的r设8,只训练attention层就够用了,显存还能再省一截。
刚加载就70多G大概率是模型权重和优化器状态一起塞进去了,试试4bit量化时用load_in_4bit=True加bnb_4bit_compute_dtype=torch.float16,然后加载模型前先torch.cuda.empty_cache()。报错“不支持的架构”可能是bitsandbytes版本和transformers不兼容,建议升级到最新版,或者直接换QLoRA的官方示例代码跑一遍。另外gradient checkpointing和混合精度一定要开,5000条数据其实可以试试用peft的prepare_model_for_kbit_training,显存能省一半还多。
A100 70多G爆显存大概率是加载权重时用了fp32,LoRA本身不省加载内存,得配合4bit量化才行。报“不支持的架构”一般是bitsandbytes版本太旧或者和transformers版本不匹配,试试升级bitsandbytes到0.43+,然后把load_in_4bit=True、bnb_4bit_compute_dtype=torch.float16、bnb_4bit_quant_type="nf4"这三个参数一起传进去。另外别忘了开gradient_checkpointing和fp16,batch size 1的话5000条数据其实可以试试用PEFT的prepare_model_for_kbit_training,它会自动帮你处理一些显存优化。如果还不行,可以考虑用Unsloth,它对LLaMA-3支持很好,显存占用能再降一截。
试试4bit量化用NF4加双量化,别用FP4,A100上跑8B完全够。另外开gradient checkpointing和bf16,显存能压到20G以内。
看到你说加载模型就OOM,我第一反应是你可能没把模型放到量化之后重新实例化,bitsandbytes的4bit得配合transformers的from_pretrained里的load_in_4bit=True参数来用,单靠装库不够。报“不支持的架构”大概率是你用的transformers版本太老,对LLaMA-3的架构支持不全,先升级到最新版再试试,我当初也是卡这。另外你batch size=1还爆显存,其实8B模型即便fp16加载也要16G左右,A100不该炸,怀疑你加载时把模型默认转成了fp32,显存翻倍,所以记得在配置里加torch_dtype=torch.float16。梯度检查点确实能省不少,但要在训练配置里开,和量化叠加效果更好,我跑7B模型时开了4bit加gradient checkpointing,峰值能压到12G以内。文本分类这种任务,其实不太需要微调全模型,只用LoRA加量化,把target_modules设成q_proj和v_proj就够了,你数据集小,5000条足够收敛。还有一个坑是bitsandbytes在Windows上支持不太好,如果你是Windows环境,换成WSL或者直接用云端Colab省心得多。建议你先跑通一个最小demo,用transformers的AutoModelForSequenceClassification配合peft库,网上有现成例子,别直接上LLaMA-3的chat版本,那个词表大更占显存。
说实话你这个问题我上周刚踩过一模一样的坑,A100 80G加载8B模型光权重就16G,但默认加载时PyTorch会把优化器状态、梯度、激活值全算进去,70G不奇怪。bitsandbytes那个报错大概率是transformers版本太新,它默认的LLaMA架构名从LlamaForCausalLM变成了LlamaForCausalLM但内部改了rope scaling,你试试指定trust_remote_code=True,或者干脆装transformers==4.36.2这个老版本,我换了之后4bit加载就正常了。另外你既然做文本分类,更建议别直接微调LLM的因果语言模型头,把LLaMA当特征提取器,只训练一个分类头,配合gradient checkpointing和fp16,显存能压到20G以内。还有个小技巧,加载时用device_map="auto",bitsandbytes配置里设置load_in_4bit=True,同时把bnb_4bit_compute_dtype设成torch.float16,这样能进一步省。我自己的经验是5000条数据,LoRA rank设8,只微调最后几层,单卡A100跑20个epoch也就两小时,没必要追求全量微调。你先把环境搞通,实在不行用QLoRA的官方脚本改改,那个对新手最友好。
你这情况我太熟了,刚接触LLM微调时我也被显存搞到怀疑人生。70多G说明你加载的是FP16全精度权重,8B模型光参数就要16G,加上优化器状态和中间激活值,A100 80G确实扛不住,但报错“不支持的架构”大概率是bitsandbytes版本跟transformers不匹配,试试升级到最新版,或者用load_in_4bit=True时加上bnb_4bit_compute_dtype=torch.float16。其实你数据集才5000条,完全不用硬刚全参微调,LoRA+r4bit量化+gradient checkpointing三件套,显存能压到20G以内。另外混合精度建议直接用torch.cuda.amp,但注意要把loss缩放一下。还有个坑是A100用4bit时,最好把device_map="auto"设上,让模型自动分配到多卡或多层,别手动指定。最后实在不行就换QLoRA,它内置了量化+LoRA的兼容处理,省心很多,你试试看能不能跑通。
说实话你这个问题我之前也踩过,8B模型单卡A100按理说不该直接OOM,大概率是加载权重的时候默认用了fp32,加上KV cache和优化器状态直接顶满了。bitsandbytes那个报错多半是因为transformers版本和bnb版本不匹配,或者你加载时没传quantization_config参数,试试升级到最新的transformers(4.38+)然后显式指定BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=torch.float16, bnb_4bit_use_double_quant=True),LLaMA-3架构是支持的。另外强烈建议开gradient checkpointing,配合torch.compile能再省一截,batch size保持1就行,但记得把gradient_accumulation_steps调大点(比如8或16),等效batch size够了收敛也稳。混合精度的话用autocast加上AMP,但注意4bit量化下要确保计算dtype是fp16,不然会掉精度。还有个更省心的路子是直接上PEFT库的prepare_model_for_kbit_training,它会自动帮你处理好量化层的训练设置,省得手动调。数据集不大(5000条)其实也可以考虑用QLoRA的默认配置,实测8B在24G显存上都能跑,你的A100完全够,问题多半出在环境版本上,先把transformers、peft、bitsandbytes三件套统一版本再试一次,大概率就能解决。
试试4bit加载时加trust_remote_code=True,再加gradient checkpointing,8B单卡肯定够。
加载模型就OOM不一定是显存不够,很可能是你没开4bit的load_in_4bit=True,或者bitsandbytes版本和transformers不匹配,换0.43.1试试。另外梯度检查点必开,配合bf16混合精度能把激活显存压到一半以下,我之前8B单卡3090都能跑。你5000条数据其实用QLoRA+seq_len=512就够了,别贪长文本。报错“不支持的架构”大概率是模型类没走AutoModelForCausalLM.from_pretrained,加上device_map="auto"和torch_dtype=torch.bfloat16就稳了。
先换加载时用 device_map="auto" 加 load_in_4bit=True,再开 gradient checkpointing 和 bf16,你那个报错大概率是 transformers 版本太旧。
5000条数据直接上QLoRA,4bit加paged optimizer,单卡16G都够跑,别纠结8B全参数。
试试llama.cpp的GGUF量化加载,8B模型4bit只要6G显存,配合PEFT库做LoRA微调,你这卡带得动。
你这报错大概率是transformers版本和bitsandbytes不兼容,LLaMA-3得用最新版的transformers才能认出架构。我上次也是卡这,升级到4.40+就好了,量化参数直接load_in_4bit=True,其他默认就行。另外你A100其实不用省显存,把batch size调到4,开gradient checkpointing,再用bf16混合精度,8B完全跑得动,我实测峰值才40多G。
这问题我太熟了,刚踩完坑出来。你那个“不支持的架构”大概率是bitsandbytes版本跟transformers不匹配,试试把transformers升级到最新版,或者干脆用accelerate的device_map="auto"配合load_in_4bit=True,让库自动处理分配,别手动指定量化参数。另外单卡A100跑8B其实不需要那么慌,你加载模型就爆是因为没开gradient checkpointing,这玩意儿能省将近一半激活内存,配合bf16混合精度,batch size 1应该稳稳的。还有个小技巧,把rope_scaling设成动态NTK,能减少序列长度带来的显存压力,如果你文本分类的输入不长,甚至可以限制max_length到512,效果不会差太多。数据集5000条的话,其实也可以考虑用PEFT的LoRA加target_modules全量指定所有linear层,学习率调低点,效果比默认配置好很多。最后实在不行就换QLoRA,4bit加双重量化,8B模型能压到6G以内,但注意要装最新版bitsandbytes,老版本对LLaMA-3支持确实有坑。
A100 70多G说明你加载的是fp16原版权重,LoRA本身不省显存,省的是优化器状态和梯度。4bit报错大概率是transformers版本和bitsandbytes不匹配,建议直接pip install -U transformers accelerate bitsandbytes,然后load_in_4bit=True,再套个prepare_model_for_kbit_training。另外gradient checkpointing一定要开,5000条数据用4bit+LoRA大概率10G以内就能跑起来,别用全参数微调。
看到你说加载模型就OOM我太有同感了,我当时第一次跑7B模型也是直接傻眼,以为A100就是随便造,结果发现光权重就要16G,加上优化器状态和激活值直接翻几倍。你那个bitsandbytes报错大概率是transformers版本和bnb不兼容,或者加载代码里没指定device_map="auto",我建议你检查一下transformers版本,新版要配对应的bnb版本,另外加载时加一句quantization_config里的load_in_4bit=True,同时把bnb_4bit_compute_dtype设成torch.float16,这样就不会报架构错误了。还有你说要做文本分类,其实更推荐用PEFT的LoRA加gradient checkpointing一起开,我试过8B模型在24G的3090上都能跑得动,batch size 1加梯度累积到8,显存大概只占15G左右。另外别忘了开启torch.compile,虽然第一次会慢点,但能再省不少内存。还有个坑是别用默认的AdamW,换成paged_adamw_8bit能省好几G的优化器显存。你那个5000条数据量其实不大,可以试试把序列长度截断到512,对分类任务影响很小但显存能降一大截。要是还不行就去huggingface上找个量化好的LLaMA-3-8B-4bit权重,直接加载就省去环境配置的麻烦,社区里很多现成的。
A100 80G都爆的话,多半是加载模型时默认用了fp32,试试load_in_4bit=True加上bnb_4bit_compute_dtype=torch.float16,另外记得把device_map设成auto。报不支持的架构大概率是transformers版本太老,升到4.40以上就行。5000条数据其实用QLoRA很稳,我跑过类似任务,4bit下显存大概12G左右,配合gradient checkpointing还能再降一点。混合精度建议直接开bf16,A100支持得很好。
这问题我上周刚踩过一遍,你报错那个“不支持的架构”八成是bitsandbytes版本跟transformers的LLaMA-3映射没对上,新版transformers里LLaMA-3的架构名是llama,但bitsandbytes需要你在加载模型时显式传device_map="auto"和torch_dtype=torch.float16,不然它会默认尝试把模型塞进CPU再转GPU,反而更吃显存。另外我建议你先别急着上4bit,试试8bit量化加gradient checkpointing,A100 40G版本跑8B LoRA是够的,batch size 1完全没问题,我实测峰值能压到28G左右。如果非要4bit,记得把load_in_4bit=True和bnb_4bit_compute_dtype=torch.float16一起设上,然后bnb_4bit_use_double_quant=True,这能再省一点。还有个坑是LLaMA-3的tokenizer和模型加载必须用同一个trust_remote_code=True参数,不然某些版本会报架构不匹配。至于混合精度,fp16必须开,但别开bf16,A100对bf16支持反而会多占显存。你5000条数据做文本分类,其实可以考虑用PEFT的prepare_model_for_kbit_training再配合gradient_accumulation_steps=4,这样训练更稳。我图省事直接用了HuggingFace的TRL库里的SFTTrainer,它封装好了量化加载逻辑,你抄它的源码改改就行。