1. 项目背景:为什么我放弃全参微调
上周接到一个医疗问答任务,需要让7B模型理解病历摘要并回答诊断依据。老板的原话是:“用ChatGLM3-6B微调一下,很快吧?”——很快?全参微调6B模型需要84GB显存,我们组唯一的A100还在跑另一个项目,剩下只有一张3090(24GB)。
查了一圈方案:PEFT库的LoRA把可训练参数压缩到0.4%,显存降到14GB,但还是超。直到看到QLoRA论文——用4bit量化基座模型+分页优化器,能把7B的微调显存压到6GB。这数字让我从椅子上坐直了。
2. 环境与版本清单
先交代环境,避免读者踩坑:
torch==2.1.2+cu118
transformers==4.36.2
peft==0.7.1
bitsandbytes==0.41.3
datasets==2.15.0
accelerate==0.25.0
特别提醒:bitsandbytes版本必须≥0.39,否则load_in_4bit参数不生效。我一开始用的0.38,直接报CUDA error: no kernel image available,折腾了半小时。
3. 方案设计:量化+LoRA的组合拳
核心思路:基座模型冻结并4bit量化,只训练注入的LoRA低秩矩阵。这样做有两个好处:
1. 显存占用主要来自量化后的基座参数+激活值,大幅降低;
2. 训练参数从60亿降到1200万,用AdamW也能扛住。
具体配置:
# 量化配置
quantization_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4", # NF4量化类型
bnb_4bit_use_double_quant=True, # 双重量化
bnb_4bit_compute_dtype=torch.float16 # 计算类型用FP16
)
# LoRA配置
lora_config = LoraConfig(
r=8, # 低秩矩阵维度
lora_alpha=16, # 缩放因子
target_modules=["query_key_value"], # ChatGLM3的注意力投影层
lora_dropout=0.05,
bias="none",
task_type="CAUSAL_LM"
)
这里有个关键点:target_modules必须匹配模型实际层名。ChatGLM3用的是query_key_value,如果是LLaMA则是q_proj, v_proj。可以通过model.named_modules()查看。
4. 核心实现:从数据到训练
4.1 数据准备:把病历变成指令对
数据来自公开的CMID医学数据集,我清洗后构造为指令格式:
def format_example(example):
"""把原始病历转成指令微调格式"""
return {
"input": f"病历:{example['history']}\n请回答:患者最可能的诊断是什么?",
"output": example["diagnosis"]
}
# 用tokenizer处理,注意padding策略
def preprocess_function(examples):
inputs = [f"### 指令:\n{example}\n### 回答:\n"
for example in examples["input"]]
model_inputs = tokenizer(
inputs,
max_length=512,
truncation=True,
padding="max_length"
)
# 把输出拼接到后面,训练时计算loss
with tokenizer.as_target_tokenizer():
labels = tokenizer(
examples["output"],
max_length=128,
truncation=True
)
model_inputs["labels"] = labels["input_ids"]
return model_inputs
注意:ChatGLM3的tokenizer不需要设置pad_token为eos_token(LLaMA需要),它自带``。
4.2 训练启动:一个函数搞定
from transformers import TrainingArguments, Trainer
training_args = TrainingArguments(
output_dir="./chatglm3-6b-lora-medical",
per_device_train_batch_size=2, # 3090上实测2是极限
gradient_accumulation_steps=4, # 等效batch_size=8
learning_rate=2e-4, # LoRA常用1e-4~5e-4
num_train_epochs=3,
logging_steps=50,
save_steps=500,
fp16=True, # 混合精度关键
optim="paged_adamw_8bit", # QLoRA专属优化器
lr_scheduler_type="cosine",
warmup_ratio=0.03,
report_to=["tensorboard"]
)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=dataset["train"],
eval_dataset=dataset["val"],
data_collator=data_collator,
tokenizer=tokenizer
)
trainer.train()
坑1:optim必须用paged_adamw_8bit,这是QLoRA论文的关键创新。如果用默认的AdamW,会在反向传播时爆显存——因为优化器状态需要额外内存。
坑2:gradient_checkpointing要在创建模型后、训练前开启:
model.gradient_checkpointing_enable()
但注意,这会降低约20%的训练速度。
5. 踩坑记录与优化
训练过程中我盯着nvidia-smi,发现显存峰值稳定在6.2GB,比预想低——因为4bit量化+FP16混合精度,激活值占大头。
loss曲线走势如下:
- Step 0-200:loss从1.89快速下降到0.82,这是LoRA在适应新任务分布;
- Step 200-600:下降变缓,0.82→0.51,出现小幅震荡(lr=2e-4偏高导致);
- Step 600-1200:稳定下降到0.38,接近收敛。
我发现两个有效的小优化:
1. warmup_ratio从0.03提升到0.06后,前期震荡明显减小;
2. 把lora_dropout从0.05降到0.01,在验证集上BLEU反而涨了1.2——因为医疗领域数据量小(仅8000条),过高的dropout反而损害知识保留。
6. 推理效果:数字说话
微调前我用原版ChatGLM3-6B跑测试集(500条病历),微调后同样数据测试:
| 指标 | 原版模型 | 微调后 | 提升幅度 |
|---|---|---|---|
| BLEU-4 | 13.8 | 22.4 | +62.3% |
| ROUGE-L | 18.2 | 27.6 | +51.6% |
| 医疗术语准确率 | 41% | 76% | 人工抽检50条 |
但发现了灾难性遗忘:用通用测试集(CMRC2018)评估,微调后模型通用阅读理解F1从64.8降至57.3。因为LoRA权重在特定领域过拟合。解决方案:采用多任务混合训练——将10%通用数据混入医疗数据,F1回升到62.1%,医疗指标仅降2%。
推理阶段的坑:LoRA权重和基座模型需要合并才能部署。直接用model.generate()会因量化模型推理慢且精度损失。我的做法:
from peft import PeftModel
# 加载量化基座
base_model = AutoModelForCausalLM.from_pretrained(
"chatglm3-6b",
load_in_4bit=True,
device_map="auto"
)
# 加载LoRA权重
model = PeftModel.from_pretrained(base_model, "./output/checkpoint-1200")
# 合并权重并转为FP16
merged_model = model.merge_and_unload()
merged_model = merged_model.half()
merged_model.save_pretrained("./merged_model_fp16")
合并后模型大小约12GB(FP16),单张3090可跑,推理速度比4bit量化快1.8倍。
7. 总结:QLoRA的适用边界
这次实践验证了QLoRA在小显存上微调大模型的可能性,但也提醒几个注意事项:
1. 适合数据量<10万的领域适配,如果数据量大,建议用8bit或直接全参;
2. 必须做通用能力评估,避免灾难性遗忘影响线上效果;
3. rank的选择:我试了8/16/32,在医疗场景rank=8和16差别不大(BLEU差0.3),但rank=32显存需求跳涨到8.1GB。
如果手头只有消费级显卡又想调7B模型,QLoRA是目前性价比最高的方案。但如果你有A100/H100,建议直接全参微调或多卡LoRA,省去量化的精度损失。技术选型永远是场景驱动的,别为了炫技强行上量化。
最后留个问题:如果你把LoRA的target_modules扩展到MLP层,效果会提升还是损害? 我自己试了加dense层后医疗指标提升2.8%,但显存涨了700MB。欢迎评论区讨论。