一、问题背景:不是所有微调都需要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数据集,但发现大量样本是“患者说肚子疼,医生说多喝水”这种废话。清洗规则:
- 过滤掉回答长度`,导致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),有需要的自取。有问题评论区见,看到会回。