一、为什么不用全参数微调?—— 资源与效果的平衡
我手头有一个垂直领域的问答需求:医疗诊断辅助,需要模型理解专业术语并给出结构化的回答。直接调用API成本高、延迟不可控,而且无法定制输出格式。全参数微调Qwen2-7B需要约56GB显存(FP16),单卡A100 80G勉强能跑,但实际项目中只有RTX 4090(24G)可用。这时候LoRA(Low-Rank Adaptation)是唯一的选择。
LoRA的核心思想是:冻结原始权重矩阵W,在W旁边插入两个低秩矩阵A和B(秩r远小于d),训练时只更新A和B。参数量缩减为原始参数的1/1000到1/10000。QLoRA更进一步,将基座模型量化为4-bit NormalFloat,进一步降低显存占用。
我最终选择的是QLoRA + 4-bit量化,单卡24G显存完全够用,batch size还能开到8。
二、环境与版本 —— 一个稳定的组合
硬件:单张RTX 4090 24G,CPU Intel i9-13900K,内存64GB。
软件版本(非常重要,版本不一致会遇到各种玄学问题):
Python 3.10.12
torch 2.1.2+cu121
transformers 4.36.2
peft 0.7.1
bitsandbytes 0.41.3
accelerate 0.25.0
datasets 2.16.1
trl 0.7.10
特别注意:bitsandbytes的0.41.3版本修复了4-bit训练时的一个梯度bug。如果低于0.41.0,会随机出现loss不收敛或NaN。
安装命令(建议新建conda环境):
conda create -n lora_finetune python=3.10 -y
conda activate lora_finetune
pip install torch==2.1.2 --index-url https://download.pytorch.org/whl/cu121
pip install transformers==4.36.2 peft==0.7.1 bitsandbytes==0.41.3 accelerate==0.25.0 datasets==2.16.1 trl==0.7.10
三、数据准备 —— 30万条领域问答的结构化清洗
我的数据来源是公开的医疗QA数据集(中文),原始数据包含大量噪声:重复问题、空答案、格式不统一。清洗流程分三步:
- 去重:基于问题文本的SimHash,剔除相似度大于0.85的重复项,从50万降到30万。
- 格式标准化:所有答案统一为结构化JSON格式,包含“症状”、“诊断”、“建议”三个字段。
- 模板化:使用ChatML格式(Qwen模型的官方模板)封装对话。
最终每条数据格式如下:
system
你是一个专业的医疗助手。请根据用户描述提供症状分析、可能诊断和就医建议。
user
{用户问题}
assistant
{结构化答案}
使用datasets库加载:
from datasets import Dataset, load_dataset
import json
# 假设数据是jsonl格式,每行一个对话
def load_and_format(data_path):
with open(data_path, 'r', encoding='utf-8') as f:
lines = [json.loads(line) for line in f]
# 格式化函数
def format_example(example):
system_msg = "你是一个专业的医疗助手。请根据用户描述提供症状分析、可能诊断和就医建议。"
user_msg = example['question']
assistant_msg = json.dumps(example['answer'], ensure_ascii=False)
# ChatML格式
text = f"system\n{system_msg}\nuser\n{user_msg}\nassistant\n{assistant_msg}"
return {'text': text}
dataset = Dataset.from_list(lines)
dataset = dataset.map(format_example, remove_columns=dataset.column_names)
return dataset
train_dataset = load_and_format('medical_qa_30w.jsonl')
print(f"训练数据量: {len(train_dataset)}") # 输出: 300000
关键点:``后的换行符必须保留,否则tokenizer的eos token会拼接错误。
四、核心实现 —— LoRA配置、训练与loss曲线
4.1 模型加载与量化
使用bitsandbytes的4-bit量化加载Qwen2-7B:
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training
model_name = "Qwen/Qwen2-7B"
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_use_double_quant=True # 双重量化,进一步节省显存
)
model = AutoModelForCausalLM.from_pretrained(
model_name,
quantization_config=bnb_config,
device_map="auto",
trust_remote_code=True
)
tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
tokenizer.pad_token = tokenizer.eos_token # Qwen没有pad token,用eos代替
# 准备k-bit训练:冻结所有参数,仅保留lm_head可训练(实际不需要,但保留以防万一)
model = prepare_model_for_kbit_training(model)
4.2 LoRA参数配置
LoRA的关键超参数:rank r(秩)、alpha(缩放因子)、target_modules(目标模块)。
lora_config = LoraConfig(
r=16, # 秩,实验发现8-32之间效果差异不大,16是平衡点
lora_alpha=32, # 缩放因子,通常设置为r的2倍
target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],
lora_dropout=0.1, # dropout防止过拟合,小数据集建议0.1
bias="none",
task_type="CAUSAL_LM",
modules_to_save=["lm_head"] # 保留lm_head可训练,稳定输出分布
)
model = get_peft_model(model, lora_config)
model.print_trainable_parameters()
# 输出: trainable params: 7,077,888 || all params: 7,073,686,528 || trainable%: 0.1000
注意:target_modules我加上了gate_proj和up_proj(MLP层),实验表明比只调attention层效果好5%左右(在验证集loss上)。
4.3 训练配置与启动
使用transformers的Trainer,配置如下:
from transformers import TrainingArguments, Trainer
from trl import DataCollatorForCompletionOnlyLM
# 数据整理器:只计算assistant部分的loss,忽略system和user
response_template = "assistant\n"
collator = DataCollatorForCompletionOnlyLM(
response_template=response_template,
tokenizer=tokenizer,
mlm=False
)
training_args = TrainingArguments(
output_dir="./qwen2-lora-medical",
num_train_epochs=3,
per_device_train_batch_size=8, # 4090 24G刚好
gradient_accumulation_steps=4, # 等效batch size = 8*4 = 32
learning_rate=2e-4, # LoRA通常用1e-4到5e-4,比全参数微调大
lr_scheduler_type="cosine",
warmup_ratio=0.03,
logging_steps=10,
save_steps=500,
evaluation_strategy="steps",
eval_steps=500,
save_total_limit=2,
fp16=True, # 混合精度训练
gradient_checkpointing=True, # 梯度检查点,节省显存的关键
report_to="tensorboard",
dataloader_num_workers=4,
)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=eval_dataset, # 另外准备的验证集,约5000条
data_collator=collator,
tokenizer=tokenizer,
)
trainer.train()
4.4 Loss曲线分析
训练过程中记录的loss曲线(使用TensorBoard查看):
- 前100步:loss从6.2快速下降到3.1,模型在学习回答格式和基础语法。
- 100-1000步:loss从3.1缓慢下降到2.4,开始学习领域知识。
- 1000-3000步:loss从2.4下降到2.1,进入稳定收敛期。
- 3000步之后:loss在2.0附近震荡,出现过拟合风险(验证集loss不再下降)。
最佳checkpoint在2500步左右(第2个epoch结束),验证集loss = 2.03。
如果跑满3个epoch(约4500步),验证集loss反而上升到2.12,说明过拟合了。
训练过程中的显存峰值:约22.5GB(batch size=8 + gradient checkpoint + 4-bit量化)。
单epoch耗时:约3.5小时(30万条数据,每条平均512 tokens)。
五、踩坑与优化 —— 那些让loss突然爆炸的问题
5.1 序列长度截断的陷阱
初始时我设置max_seq_length=2048,但发现loss下降极慢,且验证集loss始终在3.5以上。排查后发现:数据中最长的样本有1800个token,但大部分集中在500-800。直接截断到2048没问题,但截断策略很重要。默认是截断尾巴(truncation_side='right'),但Qwen的ChatML格式中,重要的回答在结尾。改成truncation_side='left'后,保留完整的assistant部分,loss立刻降到2.5以下。
5.2 学习率与rank值的匹配
尝试过rank=8, 16, 32。rank=8时效果最差(验证集loss高0.3),rank=32时训练显存多占2G但效果提升有限。最终选择rank=16。
学习率方面:LoRA的推荐学习率是1e-4到5e-4。我测试了1e-4和5e-4,发现5e-4会导致loss在第200步附近出现一次尖峰(跳到4.5),然后回落。最终使用2e-4,曲线最平滑。
5.3 梯度检查点导致的训练速度下降
开启gradient_checkpointing=True后,训练速度从4步/秒降到2步/秒,但显存占用从23G降到18G。为了稳定运行(避免OOM),这个代价可以接受。如果显存充足(比如A100 80G),建议关闭以加快训练。
六、推理效果对比 —— 微调前后的差异
使用同样的prompt测试微调前后的模型表现。
测试问题:“我最近总是头晕,伴随恶心,血压正常,请问可能是什么原因?”
微调前(基座模型):
根据您的描述,头晕和恶心可能与多种因素有关,包括低血糖、颈椎问题、内耳疾病等。建议您去医院做进一步检查。
输出比较宽泛,没有针对性。
微调后(LoRA模型):
{
"症状分析": "患者主诉头晕伴恶心,血压正常,排除高血压相关头晕。可能原因包括:1. 颈椎病(颈性眩晕)——长时间低头工作常见;2. 前庭系统疾病——如耳石症、梅尼埃病;3. 低血糖或贫血——需结合血糖和血常规检查。",
"可能诊断": "1. 颈性眩晕(可能性较大,占60%);2. 耳石症(20%);3. 其他(20%)。",
"就医建议": "1. 建议挂神经内科或耳鼻喉科;2. 做颈椎MRI和耳石诱发试验;3. 注意休息,避免突然转头。"
}
输出格式严格遵循了JSON结构,且诊断更加细化,给出了概率和具体检查建议。这是30万条领域数据微调带来的效果:模型学会了输出结构化回答,且内容更贴近医疗场景。
七、总结
| 项目 | 参数/数据 |
|---|---|
| 基座模型 | Qwen2-7B (4-bit量化) |
| 训练参数量 | 7.08M (0.1%) |
| 训练数据 | 30万条医疗QA |
| 单卡显存 | 22.5GB (RTX 4090) |
| 训练时间 | 约10小时 (3 epoch) |
| 最佳验证loss | 2.03 |
| 推理格式 | 结构化JSON |
LoRA/QLoRA的核心优势在于:用单卡24G显存就能微调7B模型,且效果显著。但需要注意几个关键点:1)数据的格式必须对齐模型的模板;2)截断策略要保住回答部分;3)LoRA的rank和学习率需要调参,不是越大越好。
如果你手头有领域数据,并且资源有限,LoRA是目前最实用的微调方案。下一步我计划尝试DoRA(Weight-Decomposed Low-Rank Adaptation),据说在效果上比LoRA有进一步提升,但实现复杂度也更高。欢迎有经验的同学交流。
博客作者:一个正在和loss作斗争的算法工程师
日期:2024年5月
项目代码:已整理到GitHub,关键词“qwen2-lora-medical”