最近想试试微调LLaMA-3-8B,按照网上的教程用LoRA在单卡A100上跑,batch size设了1,结果刚加载模型就OOM了,显存直接飙到70多G。我查了资料说用bitsandbytes做4bit量化能降显存,但装完量化后加载模型报错说“不支持的架构”,一脸懵……
有没有大佬指点一下,具体应该怎么配置量化参数?或者有没有更省显存的方法(比如gradient checkpointing、混合精度之类的)?我主要想做文本分类微调,数据集不大(约5000条),但卡在环境配置这一步了,求靠谱的实操经验!🙏
新手求教:用PyTorch跑LLaMA-3微调,显存爆炸怎么优化?
全部回复
共 156 条试试把lora的r值调低到8,同时开gradient checkpointing,我32G的卡跑8B都没爆。
你遇到的这个情况还挺典型的,LLaMA-3-8B光加载FP16权重就要16G左右,但A100是40G或80G,按理说单卡应该能跑,问题可能出在缓存和优化器状态上。bitsandbytes报“不支持的架构”,大概率是因为你没匹配对CUDA版本或者bitsandbytes的版本,试试用pip install bitsandbytes --upgrade,或者直接上conda装预编译版。4bit量化确实能压到6-7G,但你要注意LLaMA-3的架构和原始的LLaMA不太一样,加载时得指定trust_remote_code=True。另外gradient checkpointing一定要开,配合混合精度(fp16或bf16)能省一半显存,batch size设1的话,检查下dataloader的num_workers别太高,也会占显存。5000条数据其实不用全量微调,可以考虑用PEFT的LoRA只训练一小部分参数,甚至用IA3更省。如果还是炸,试试DeepSpeed的ZeRO-2或ZeRO-3,单卡也能用,但配置稍微复杂点。最后建议先从4bit量化+LoRA+gradient checkpointing这个组合入手,网上有现成的transformers脚本可以直接改。
我用A100跑过LLaMA-3-8B微调,刚加载模型就70G确实不正常,可能是你加载时没指定device_map或者量化配置没写对。试试用transformers的BitsAndBytesConfig,load_in_4bit设为True,bnb_4bit_compute_dtype设成float16,再加个torch_dtype=auto,一般能压到20G左右。如果还报架构不支持,检查下bitsandbytes版本是不是0.41以上,然后模型名用meta-llama/Meta-Llama-3-8B试试。gradient checkpointing肯定要开,配合混合精度和4bit量化,你那5000条数据单卡A100跑起来完全没问题。
试试HuggingFace的transformers加载时直接传load_in_4bit=True,配bnb_4bit_compute_dtype=float16,比手动改配置稳得多。
8B单卡还做4bit?直接上QLoRA加gradient checkpointing,显存能压到20G以下。
试试用QLoRA,加载时加load_in_4bit和bnb_4bit_compute_dtype=torch.float16,再配合gradient checkpointing,8B单卡稳够。
A100 80G按理说跑8B的LoRA应该够啊,你是不是把模型完整加载进显存了?试试load_in_4bit=True配合bnb_4bit_compute_dtype=torch.float16,再开gradient_checkpointing,显存能压到20G以内。报错不支持的架构大概率是transformers版本太老,升级到最新版再试下。5000条数据的话其实也可以考虑用QLoRA,效果差不太多但省心很多。
你这情况大概率是bitsandbytes版本和transformers没对齐,4bit加载得用from_pretrained里传bitsandbytes_config那个对象,而不是直接改load_in_4bit。另外LoRA其实不用全量加载,可以先试试load_in_8bit加gradient checkpointing,batch size=1的话8B模型在A100上应该能压到30G以内。我上次跑类似任务还开了torch.compile,显存又省一截,不过第一次编译会慢点。你那5000条数据其实可以试试用qlora的脚本,它默认配置就是4bit加paged_adamw,我照着改很少踩坑。
碰到这问题太正常了,LLaMA-3-8B光fp16权重就得16G,加上优化器状态和激活值,单卡A100 80G看着够用,但LoRA其实只是把可训练参数变小了,前向/反向的中间变量照样吃显存,所以batch size=1炸掉不冤。你那个bitsandbytes报“不支持的架构”,八成是transformers版本和bnb没对齐,或者模型加载时没加load_in_4bit=True配合quantization_config,建议直接升到最新版transformers,然后用BitsAndBytesConfig显式指定bnb_4bit_compute_dtype=torch.float16,再配合device_map="auto",基本就能塞进20G以内。另外gradient checkpointing是必须开的,model.gradient_checkpointing_enable()一行代码能省掉大部分激活显存,混合精度用torch.cuda.amp包一下训练循环就行,但注意4bit下最好用bf16而不是fp16,有些卡不支持。还有个更省事的思路,你数据集才5000条,不如直接试试调per_device_train_batch_size=4加gradient_accumulation_steps=8,等效batch够大但显存占用还是小,或者干脆用peft的prepare_model_for_kbit_training函数,它会把量化后的模型自动设置好requires_grad和梯度检查点,省得手动配置。如果量化还是报错,可以先跑个纯LoRA不加量化,8B模型在80G上勉强能跑batch size=1,配合checkpointing大概率能过,先把流程跑通再优化显存。最后提醒一下,bnb对LLaMA-3支持得看具体版本,实在不行换torchao或llama.cpp的量化方案,但核心还是把transformers和peft的版本锁死,别用太旧的。
刚踩过类似的坑,A100 80G按理说跑8B的LoRA不会炸,先检查下是不是加载时把模型参数和梯度都塞显存了,试试model.to_empty()加zero_init,能省不少。4bit报错大概率是transformers版本和bitsandbytes不兼容,建议直接pip install transformers accelerate bitsandbytes --upgrade,然后加载时加load_in_4bit=True,同时设device_map="auto"。另外5000条数据真不用全量微调,开gradient_checkpointing和fp16,batch size调到4应该都能跑,我上次用同样配置微调7B才占28G。
4bit量化得用QLoRA那个分支,普通LoRA不支持量化权重,换load_in_4bit加bnb_4bit_compute_dtype试试。
看到你说70多G我真的一点不意外,8B模型就算fp16权重也得16G,加上优化器状态和激活值,单卡A100想直接硬吃确实够呛。你那个bitsandbytes报“不支持的架构”,八成是transformers版本和bnb的兼容问题,建议先升级transformers到4.40以上,然后加载时明确指定load_in_4bit=True、bnb_4bit_compute_dtype=torch.float16、bnb_4bit_quant_type="nf4"这三个参数,缺一不可。另外你可以试试把device_map="auto"加上,让模型自动分配显存,有时候手动指定device="cuda:0"反而会撞墙。关于省显存,gradient checkpointing真的强烈推荐,开启后大概能省掉30%-40%的激活显存,就一行代码的事,代价是训练慢个20%左右,但对5000条数据来说完全无所谓。混合精度的话,fp16+bf16混着用能再压一点,但A100对bf16支持更好,你可以在TrainingArguments里直接设bf16=True。还有个细节,LoRA的r和alpha别设太大,建议r=8、alpha=16,target_modules也别全选,只挑q_proj和v_proj就够了,这样能显著减少可训练参数量。最后如果还不行,试试把序列长度截断到512,文本分类任务一般用不到那么长的上下文。
查下transformers版本,新版才支持llama3的4bit,顺手把gradient checkpointing打开能再省点。
把load_in_4bit加进去,配好target_modules,再用gradient checkpointing,8B单卡没问题。
8B模型全精度加载本来就要60多G,A100爆显存很正常。你那个bitsandbytes报错大概率是transformers版本太老,升级到4.40以上,加载时直接传load_in_4bit=True,别自己手动配置量化参数。顺便把device_map="auto"加上,让模型自动分配到多个GPU显存。梯度检查点一定要开,配合4bit后显存能压到20G以内。5000条数据其实用QLoRA就够了,学习率调成2e-4,训练时再用paged_adamw优化器,基本能跑通。
这问题我太熟了,刚踩完坑出来。4bit报“不支持的架构”大概率是bitsandbytes版本和transformers没对齐,试试升级到最新的bitsandbytes,然后加载时加一句model_kwargs={"quantization_config": BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=torch.float16, bnb_4bit_use_double_quant=True)},注意别漏了bnb_4bit_quant_type="nf4",这套组合在8B上基本能压到6-7G。不过说实话,你5000条数据做分类微调,其实不用硬刚8B,换LLaMA-3-8B的Qwen2-1.5B或者phi-3-mini可能更舒服,效果差距没那么大,但显存和调试成本直接降一个量级。如果坚持用8B,gradient checkpointing必须开,配合混合精度fp16,batch size 1应该能跑,但建议把max_seq_len截到512以下,不然attention矩阵照样吃满。另外你提到数据量小,可以考虑用PEFT的LoRA加target_modules全开(比如q/k/v/o_proj),rank设8,alpha设16,别贪大,收敛速度反而更快。最后提醒一句,检查下是不是加载了完整模型才转量化,应该在from_pretrained时直接传quantization_config,不要先加载再转,那样峰值显存早就爆了。
A100都爆的话,试试load_in_4bit加bnb_4bit_compute_dtype=float16,配gradient checkpointing,5000条数据够用了。
A100 80G都OOM?那多半是加载时直接用了fp32,LLaMA-3-8B光权重就16G,加上优化器状态和激活值确实吃紧。建议先试试HuggingFace的load_in_4bit=True,配个bnb_4bit_compute_dtype=torch.float16,报错大概率是transformers版本太老,升到4.40以上就行。量化后配合gradient checkpointing,8G显存都能跑,你5000条数据做分类用LoRA完全够。
刚踩过这坑,4bit量化报错大概率是bitsandbytes版本和transformers不匹配,先升级到最新版试试,再不行就换个支持LLaMA-3的加载方式,比如用transformers的load_in_4bit配合device_map="auto"。另外你这单卡A100跑8B其实完全够,开gradient checkpointing加混合精度就能省一半显存,batch size暂时不用动。数据集才5000条的话,其实可以考虑用QLoRA加paged_adamw优化器,实测4bit下能压到20G以内,就是训练速度会慢点。
4bit量化得用QLoRA那个分支,别直接套原版LoRA,A100跑8B真没必要省到这份上。