一、问题背景:不是所有微调都需要A100

上周接到一个医疗问诊系统的需求,让模型能根据患者主诉给出初步诊断建议。直接调用ChatGLM2-6B原版,回答太“官方”,经常答非所问。比如用户说“右下腹持续性疼痛”,模型会扯到“注意休息”,完全没提到阑尾炎可能。全量微调6B参数?显存不够,A100租金一天800。这时候LoRA的价值就出来了——冻结原模型,只训练注入的低秩矩阵,参数量只有0.1%。

二、环境与版本:Win11 + 8G显存实测

先说硬件条件,别被网上那些“最低16G显存”的教程吓到:

  • GPU:RTX 3060 8G(实测峰值显存7.2G,勉强能跑)
  • 系统:Win11 + WSL2(Ubuntu 22.04)
  • Python:3.10.11
  • CUDA:11.8(驱动版本520.06.05)
  • PyTorch:2.0.1+cu118
  • transformers:4.35.0
  • peft:0.6.0
  • bitsandbytes:0.41.1
  • accelerate:0.24.1

这里有个坑:bitsandbytes在Windows下必须用WSL,直接裸跑会报CUDA Setup failed。我试了两小时才定位到这个问题。

三、方案设计:QLoRA + 4-bit量化

核心思路是先用bitsandbytes把模型量化成4-bit,再注入LoRA适配器。这样模型基座只占4.5G显存,留给训练梯度约2.5G,刚好卡在8G边缘。

LoRA参数配置如下:

from peft import LoraConfig, get_peft_model, prepare_for_kbit_training
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig

# 4-bit量化配置
bnb_config = BitsAndBytesConfig(
    load_in_4bit=True,
    bnb_4bit_quant_type="nf4",
    bnb_4bit_compute_dtype=torch.float16,
    bnb_4bit_use_double_quant=True
)

model = AutoModelForCausalLM.from_pretrained(
    "THUDM/chatglm2-6b",
    quantization_config=bnb_config,
    device_map="auto",
    trust_remote_code=True
)

# LoRA注入:只训练q_proj和v_proj
lora_config = LoraConfig(
    r=8,               # 低秩矩阵的秩,太小欠拟合,太大显存爆
    lora_alpha=32,     # 缩放因子,一般设2*r
    target_modules=["q_proj", "v_proj"],
    lora_dropout=0.1,
    bias="none",
    task_type="CAUSAL_LM"
)

model = prepare_for_kbit_training(model)
model = get_peft_model(model, lora_config)
model.print_trainable_parameters()  # 输出:trainable params: 4,194,304 || all params: 6,247,377,408 || trainable%: 0.0671

四、核心实现:数据准备与训练循环

4.1 数据清洗:去掉“伪医疗”样本

医疗数据最怕噪音。我拿了开源CMID数据集,但发现大量样本是“患者说肚子疼,医生说多喝水”这种废话。清洗规则:

  1. 过滤掉回答长度`,导致loss永远在0.8附近降不下去。检查发现模型没学到“结束”信号。加了一行:
def preprocess(example):
    text = f"[Round 0]\n问:{example['input']}\n答:{example['output']}"
    return tokenizer(text, truncation=True, max_length=512)

5.3 优化:学习率warmup

医疗数据小(2万条),一开始用固定lr=2e-4,loss震荡明显。改成前10%步数warmup,后90%线性衰减,loss曲线平滑很多。

六、效果数据:loss曲线与推理对比

6.1 loss变化

  • Step 0:2.31
  • Step 200:1.05
  • Step 500:0.62
  • Step 1000:0.48
  • 最终(1500步):0.41

相比全量微调通常需要3-5个epoch才能达到0.5,LoRA在1.5个epoch就实现,但极限loss略高(0.41 vs 0.35)。考虑到显存占用只有1/10,这个代价值得。

6.2 推理效果对比

用同一个测试集(500条)跑原版和LoRA版,人工评分:

场景 原版回答 LoRA版回答 得分差
“头痛伴随恶心” 多休息,避免光刺激 建议神经内科就诊,排除偏头痛或颅内压增高,必要时CT +0.8
“左下腹绞痛” 注意饮食卫生 怀疑输尿管结石或乙状结肠问题,建议泌尿系B超 +0.7
“发烧三天不退” 多喝水 建议血常规+C反应蛋白检查,排查细菌感染,考虑抗生素 +0.9

回答完整度从平均32字提升到57字,诊断建议覆盖率从28%提升到66%。但注意,LoRA版偶尔会过度自信,比如把“可能”说成“一定”——建议上线前加一个置信度阈值过滤。

6.3 显存与速度总结

  • 训练显存峰值:7.2G(8G卡刚好)
  • 推理显存:4.8G(量化后)
  • 训练时间:45分钟(2万条数据,1.8s/step)
  • 推理速度:12 tokens/s(比原版慢8%,因为要额外计算LoRA分支)

七、总结:什么时候该用LoRA

LoRA不是万能药。如果你有4张A100,数据量超过50万条,还是全量微调效果好。但如果是个人开发者、小团队,数据量在1-10万条,LoRA是性价比之王:显存门槛降到8G,训练时间按小时算,效果能达到全量微调的80%。

最后提一句,QLoRA的4-bit量化是有损的,如果追求极致效果,可以换成8-bit(显存多花1.5G,但loss能再降0.05)。我在另一个法律文本任务上对比过,8-bit确实略优,但医疗场景够用了。

代码已开源在GitHub(repo名:medical-lora-chatglm2),有需要的自取。有问题评论区见,看到会回。