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_tokeneos_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()

坑1optim必须用paged_adamw_8bit,这是QLoRA论文的关键创新。如果用默认的AdamW,会在反向传播时爆显存——因为优化器状态需要额外内存。

坑2gradient_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。欢迎评论区讨论。