1. 问题背景:为什么不用全参微调而要折腾QLoRA
上周接到一个医疗问答场景的需求,要在一个7B模型上做领域适配。leader上来就说“直接全参微调吧,卡够用”。我看了眼账单——全参微调Llama-3-8B需要约60GB显存,而A100 80G按小时计费,一个epoch要跑40分钟。更要命的是,医疗数据只有128条,全参微调大概率过拟合。
试了LoRA(秩r=8),显存降到32GB,但训练loss在1.8左右震荡,怎么都降不下去。换成QLoRA(4-bit NF4量化+双LoRA)后,显存直接砍到18GB,loss能稳定降到1.2。这篇文章就记录这个从“能跑”到“跑得好”的过程。
2. 环境与版本:这些坑都是版本不匹配引出来的
先交代环境,Python 3.10.14,PyTorch 2.2.2+cu118,transformers 4.40.0,peft 0.10.0,bitsandbytes 0.43.1,trl 0.8.6。重点说下bitsandbytes,0.43.0之前有个bug会导致4-bit反量化出错,loss直接变NaN。如果你用0.39.0,那恭喜你,会撞上CUDA 11.8的兼容性问题。
pip install torch==2.2.2 transformers==4.40.0 peft==0.10.0 bitsandbytes==0.43.1 trl==0.8.6
数据集是我自己整理的医疗问答对,只有128条,格式如下:
{"instruction": "患者出现胸痛、呼吸困难,应考虑哪些疾病?", "output": "胸痛伴呼吸困难需警惕急性心肌梗死、肺栓塞、主动脉夹层等急症,建议立即行心电图、心肌酶谱、D-二聚体检查..."}
3. 方案设计:为什么是4-bit NF4 + 双LoRA + 冻结embedding
先说结论:最终采用的配置是bitsandbytes的NF4量化(不是FP4),LoRA只挂在Q和V矩阵上(不是全矩阵),r=16, alpha=32, dropout=0.1。冻结了embedding和lm_head层。
为什么冻结embedding?这是实验出来的。第一次跑的时候没冻结,128条数据训3个epoch后,模型开始胡言乱语,把“胸痛”和“头痛”混为一谈。看了下embedding的梯度范数,发现它占了总梯度的40%以上,说明模型在疯狂调整词向量来死记硬背128条数据。冻结后,loss虽然下降慢了点,但泛化明显变好。
4. 核心实现:QLoRA训练完整代码
先放模型加载和量化配置:
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training
# 4-bit NF4量化配置
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4", # 用NF4不用FP4
bnb_4bit_use_double_quant=True, # 双量化,能省1-2GB
bnb_4bit_compute_dtype=torch.bfloat16
)
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3-8B",
quantization_config=bnb_config,
device_map="auto",
torch_dtype=torch.bfloat16,
trust_remote_code=True
)
# 冻结embedding和lm_head
for param in model.model.embed_tokens.parameters():
param.requires_grad = False
for param in model.lm_head.parameters():
param.requires_grad = False
model = prepare_model_for_kbit_training(model)
# LoRA配置:只挂Q和V
lora_config = LoraConfig(
r=16,
lora_alpha=32,
lora_dropout=0.1,
bias="none",
task_type="CAUSAL_LM",
target_modules=["q_proj", "v_proj"]
)
model = get_peft_model(model, lora_config)
model.print_trainable_parameters() # 输出:trainable params: 4,194,304 || all params: 8,030,244,864
训练配置用trl的SFTTrainer,这里有个关键参数max_seq_length,我设了1024。别设太大,否则会触发KV cache的显存爆炸。
from trl import SFTTrainer
from transformers import TrainingArguments
training_args = TrainingArguments(
output_dir="./qlora_medical",
per_device_train_batch_size=4,
gradient_accumulation_steps=4, # 等效batch size=16
num_train_epochs=3,
learning_rate=2e-4,
lr_scheduler_type="cosine",
warmup_ratio=0.1,
logging_steps=5,
save_strategy="epoch",
fp16=False,
bf16=True, # A100上必须用bf16,fp16会掉精度
gradient_checkpointing=True,
optim="paged_adamw_8bit"
)
trainer = SFTTrainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
tokenizer=tokenizer,
max_seq_length=1024,
dataset_text_field="text",
packing=False
)
trainer.train()
5. 踩坑与优化:三个loss曲线的血泪教训
第一个坑:loss震荡不收敛。第一次用LoRA r=4,learning_rate=1e-4,跑了200步,loss在1.7-1.9之间来回跳。排查发现是LoRA的alpha设太小(alpha=8),导致更新幅度和r不匹配。按经验,alpha应该设为r的2倍(即alpha=32),这样初始化时LoRA分支的输出接近零,但梯度更新时能保持有效幅度。改完loss稳定下降到1.3。
第二个坑:eval_loss先降后升。训练到第2个epoch时,train_loss还在降,但eval_loss开始反弹,典型过拟合。128条数据训3个epoch本来就容易过。解决方案:把dropout从0.05提到0.1,同时在SFTTrainer里加了dataset_text_field的拼接方式,把instruction和output用\n\n### Response:\n连接,这样模型能学习到更稳定的格式。
第三个坑:生成时重复输出。微调完做推理,发现模型回答“胸痛”时会连续输出5遍。查了下是temperature=0.7时采样过随机。把do_sample设为False,改用greedy decoding,问题解决。但这也暴露了模型对答案的置信度不够,后来在推理时加了repetition_penalty=1.1,效果更好。
最终训练loss曲线(每5步记录一次):step 0-50从2.1快速降到1.4,step 50-100缓慢降到1.25,step 100-150(第2个epoch开始)降到1.15,之后基本稳定在1.12-1.18之间。eval_loss在第120步左右达到最低0.98,之后略微上升到1.05。
6. 效果数据:量化对比Q-LoRA vs LoRA vs 原始模型
我在50条手工标注的测试集上做了三组对比:原始模型(zero-shot)、LoRA微调(r=8, fp16)、QLoRA微调(r=16, nf4)。评价指标用ROUGE-L和指令遵循准确率(人工判断回答是否覆盖所有关键诊断点)。
| 模型 | ROUGE-L | 指令遵循准确率 | 显存占用 | 训练耗时 |
|---|---|---|---|---|
| Llama-3-8B原始 | 12.6 | 31% | - | - |
| LoRA r=8 | 24.1 | 64% | 32GB | 52min |
| QLoRA r=16 | 28.4 | 79% | 18GB | 38min |
QLoRA反而比普通LoRA效果好的原因:r=16的秩比r=8能学到更多领域知识,同时4-bit量化本身带有轻微的正则化效果(噪声注入),抑制了过拟合。这个结论在128条小数据集上尤其明显。
7. 总结与后续优化方向
QLoRA在小数据集(128条)上的表现超出预期,核心收益来自三点:4-bit量化省下的显存让batch size和序列长度可以开得更大;冻结embedding强制模型用已有词向量组合来回答;r=16的高秩配置在小数据上反而比低秩更稳。
后续打算试的方向:1)用LoRA的target_modules扩展到全部attention层,看效果是否有提升;2)将alpha从32提到64,配合更低的learning_rate(1e-4)看能否进一步降低loss;3)在推理时用vLLM做batch推理,吞吐量能到原始HuggingFace的3倍。最后提醒一句,QLoRA虽然省显存,但训练速度比全参微调慢约20%(量化反算开销),如果不是显存瓶颈,别盲目上量化。