最近想试试微调LLaMA-3-8B,按照网上的教程用LoRA在单卡A100上跑,batch size设了1,结果刚加载模型就OOM了,显存直接飙到70多G。我查了资料说用bitsandbytes做4bit量化能降显存,但装完量化后加载模型报错说“不支持的架构”,一脸懵……
有没有大佬指点一下,具体应该怎么配置量化参数?或者有没有更省显存的方法(比如gradient checkpointing、混合精度之类的)?我主要想做文本分类微调,数据集不大(约5000条),但卡在环境配置这一步了,求靠谱的实操经验!🙏
新手求教:用PyTorch跑LLaMA-3微调,显存爆炸怎么优化?
全部回复
共 156 条试试加载时加device_map="auto"配合load_in_4bit=True,记得先升级bitsandbytes到最新版,老版本不认LLaMA架构。
8B全量加载本身就吃显存,你试试transformers里load_in_4bit配好bnb_config,别直接用bitsandbytes默认参数。
这问题我当初也踩过,A100 80G看着挺大,但加载8B模型原始权重就要16G,算上优化器状态和激活值,LoRA的显存大头反而不在可训练参数上。你那个70多G应该是把整个模型都塞进显存了,bitsandbytes报“不支持的架构”大概率是transformers版本太老,得升到4.30以上,并且加载时要用load_in_4bit=True配合BitsAndBytesConfig里设bnb_4bit_compute_dtype=torch.float16,同时记得给模型套一层prepare_model_for_kbit_training。另外你说的gradient checkpointing和混合精度其实比量化更优先,开起来能把激活值显存砍掉一半以上,加上LoRA的r=8、alpha=16,batch size=1在A100上应该能压到20G以内。还有个野路子是直接上QLoRA,它把4bit量化、双重量化和分页优化器都集成好了,HuggingFace的TRL库有现成示例,你那个5000条文本分类任务完全够用。如果还爆,就把gradient_accumulation_steps开到8,实际batch size不变但显存曲线平缓很多。最后提醒下,量化后微调完要保存merge_and_unload()合并权重,不然推理时还得挂着量化配置,挺烦的。
试试llama-factory吧,配置好量化参数直接跑,还有CPU offload兜底,你这情况换它准能救。
试试先开gradient checkpointing+bf16,A100吃8B没理由爆,多半是默认fp32在作怪。
A100 80G跑8B居然也爆,大概率是加载时没走量化,4bit得用bitsandbytes配transformers的load_in_4bit=True,然后传个bnb_4bit_compute_dtype=torch.float16试试,你报错那个不支持的架构八成是版本没对齐。另外梯度检查点肯定要开,LoRA的target_modules记得只选q_proj和v_proj,别全加上,5000条数据其实用QLoRA跑个几轮就够了,没必要上全参数。
你这个问题我上周刚踩过,A100跑8B用LoRA其实不用4bit也能塞下,关键是把加载时的torch_dtype设成torch.float16,再加gradient checkpointing和gradient_accumulation_steps,显存基本能压到30G以内。bitsandbytes那个报错大概率是transformers版本太老,升级到4.40+再试试,或者直接用QLoRA的官方demo配置。你5000条数据做分类的话,其实可以考虑PEFT的IA3或者Adapter,比LoRA更省。另外加载模型时记得用low_cpu_mem_usage=True,能省不少临时显存。
加载模型就OOM大概率是bitsandbytes版本跟transformers不匹配,LLaMA-3的架构名没被识别,试试把bitsandbytes升到0.43以上,同时transformers版本别太旧。4bit量化参数里load_in_4bit加bnb_4bit_compute_dtype=torch.float16,再加bnb_4bit_use_double_quant=True,基本能压到10G以内。你数据集才5000条,其实用QLoRA加gradient checkpointing就够,batch size可以保持1,梯度累积设8。另外记得加载模型时把device_map设成auto,不然容易把权重全塞进单卡显存。
看到你说刚加载模型就70G,我猜你可能没开4bit加载,直接fp16硬怼的。bitsandbytes报“不支持的架构”大概率是版本和transformers不匹配,你试试把transformers升到4.35以上,然后load_in_4bit=True,同时指定bnb_4bit_compute_dtype=torch.float16,这样模型权重压到4bit后显存基本能控制在10G以内。不过光量化还不够,你训练时务必要开gradient_checkpointing,这个能把激活显存砍掉一大半,另外把optimizer换成paged_adamw_8bit,能省下优化器状态的内存。5000条数据其实用LoRA加量化完全够跑,但建议你把LoRA的target_modules限定在q_proj和v_proj上,别全模块都挂,这样显存压力会小很多。还有个小技巧,训练时把input_ids和attention_mask直接放到cuda:0上,别让模型自己搬数据,有时候会自动多复制一份。最后提醒下,A100上如果还嫌不够,可以把batch size设为1但梯度累积步数设成8,效果一样且更稳。你检查下bitsandbytes是不是装成CPU版了,这个坑我踩过,pip list里看下版本号带不带cu后缀。
建议先试试4bit的NF4量化加双卡offload,A100跑8B不该爆,八成是bitsandbytes版本和transformers不匹配。
说实话你这个报错我也踩过,bitsandbytes对LLaMA-3的支持得看transformers版本,太新或太旧都会报“不支持的架构”,建议直接换成4.38到4.40之间,然后加载时用load_in_4bit=True加bnb_4bit_compute_dtype=torch.float16试试。另外A100跑8B其实不用量化也能塞下,关键是开gradient checkpointing和把optimizer换成AdamW 8bit,batch size=1加梯度累积到16,显存能压到40G以内。5000条数据做分类任务的话,其实用LoRA加全量微调中间层的方案就够了,别碰那些花里胡哨的量化配置。
试试4bit得用QLoRA那套,加载时加device_map="auto"再配个trust_remote_code=True,A100跑8B轻轻松松。
把模型塞进cpu再load_state_dict,配合gradient_checkpointing和bf16,5000条数据真不用硬扛70G显存。
4bit得用QLoRA那套,别直接load_in_4bit,再把gradient checkpointing和bf16开起来,A100跑8B轻轻松松。
看到你说加载模型就OOM,大概率是bitsandbytes版本和transformers不匹配,LLaMA-3需要比较新的transformers(4.40+)才认架构,建议直接升级到最新版再试。另外4bit量化建议用load_in_4bit=True加bnb_4bit_compute_dtype=torch.float16,同时开gradient checkpointing,A100上8B全参加载加LoRA其实能压到15G以内。还有个小技巧,把attn_implementation="flash_attention_2"加上,显存和速度都会好很多,你5000条数据其实用QLoRA跑几轮就够了。
我前几天刚在A100上跑通LLaMA-3-8B的LoRA微调,你那个报错大概率是transformers版本和bitsandbytes不匹配,试试升级到最新版,然后加载时加上quantization_config=BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=torch.float16),基本能压到20G出头。另外gradient checkpointing一定要开,再配合torch.compile,batch size可以提到4,反正5000条数据跑20个epoch也就一晚上,别着急。
先开gradient checkpointing和bf16,再把LoRA的target_modules设全点,4bit用QLoRA的NF4格式别用FP4。