1. 问题背景:全量微调7B模型,我连显存都塞不下
上周接到一个任务:把Llama-2-7B基座模型微调成能回答中文医疗问答的助手。第一反应是直接全量微调,结果一跑就崩——单卡A100 40G勉强够,但公司只有RTX 3090,24G显存,全量微调连模型权重都加载不完(7B float16权重需要14G,加上优化器状态直接爆显存)。
换方案:冻结基座,只训练Adapter。最初想用PEFT的LoRA,但LoRA仍然需要加载完整模型权重,24G显存勉强能跑,可一旦batch size调大或者序列长度超过512,直接OOM。最后决定用QLoRA——把基座模型量化到4bit,再注入LoRA适配器。这样基座权重只需4G左右,加上LoRA参数和激活值,8G显存就能跑起来。
实际效果比预期好,但坑也不少。下面按步骤记录整个过程。
2. 环境与版本
- GPU: NVIDIA RTX 3090 24G
- CUDA: 11.8
- Python: 3.10.12
- transformers: 4.35.2
- peft: 0.7.1
- bitsandbytes: 0.41.3
- datasets: 2.15.0
- accelerate: 0.24.1
- 基座模型: meta-llama/Llama-2-7b-chat-hf
特别注意:bitsandbytes必须和CUDA版本匹配,否则会报CUDA SETUP: ERROR。建议直接用pip安装预编译版本,别自己编译,我在这上面浪费了2小时。
3. 方案设计:QLoRA的配置与参数选择
QLoRA的核心思路是:基座模型用4bit NF4量化存储,LoRA适配器用float16训练。关键配置如下:
from transformers import BitsAndBytesConfig
import torch
# 4bit量化配置
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4", # NF4量化,比fp4精度高
bnb_4bit_use_double_quant=True, # 双量化,进一步压缩
bnb_4bit_compute_dtype=torch.float16 # 计算时反量化为fp16
)
# LoRA配置
from peft import LoraConfig, TaskType
lora_config = LoraConfig(
task_type=TaskType.CAUSAL_LM,
r=16, # LoRA秩
lora_alpha=32, # 缩放系数,一般设为r的2倍
target_modules=["q_proj", "k_proj", "v_proj", "o_proj"],
lora_dropout=0.05,
bias="none"
)
这里有个关键点:target_modules必须匹配Llama-2的模块名。Llama-2的注意力层叫q_proj/k_proj/v_proj/o_proj,不是query/key/value。如果你用别的模型(比如Qwen),模块名是c_attn,照抄会报错。用model.state_dict().keys()先看一眼最保险。
4. 核心实现:数据准备到训练完整流程
4.1 数据准备
我用了公开的cmrc2018中文医疗问答数据集,清洗后得到约2万条数据。格式统一为:
问题:患者出现发热、咳嗽、咳痰,血常规提示白细胞升高,应首先考虑什么诊断?
回答:细菌性肺炎。建议完善痰培养、胸片检查,经验性使用抗生素。
数据加载和预处理代码:
from datasets import load_dataset
from transformers import AutoTokenizer
dataset = load_dataset("json", data_files="medical_qa.jsonl")
tokenizer = AutoTokenizer.from_pretrained(
"meta-llama/Llama-2-7b-chat-hf",
padding_side="right",
trust_remote_code=True
)
# 关键:加padding token,否则batch训练会报错
tokenizer.add_special_tokens({"pad_token": ""})
model.resize_token_embeddings(len(tokenizer))
def preprocess(example):
prompt = f"问题:{example['question']}\n回答:"
full_text = prompt + example["answer"]
tokenized = tokenizer(
full_text,
max_length=512,
truncation=True,
padding="max_length"
)
# 只计算answer部分的loss,问句部分不参与训练
labels = tokenized["input_ids"].copy()
prompt_len = len(tokenizer(prompt)["input_ids"])
labels[:prompt_len] = -100 # -100是CrossEntropyLoss的默认ignore_index
tokenized["labels"] = labels
return tokenized
dataset = dataset.map(preprocess, remove_columns=dataset.column_names)
这里有个血泪教训:我一开始没有把prompt部分的label设为-100,结果模型把「问题:」这几个字也学会了预测,生成的时候会把问题重复一遍才给答案。
4.2 训练配置
用HuggingFace的TrainingArguments,关键参数如下:
from transformers import TrainingArguments, Trainer
training_args = TrainingArguments(
output_dir="./qlora_medical",
num_train_epochs=3,
per_device_train_batch_size=8,
gradient_accumulation_steps=4,
gradient_checkpointing=True, # 用计算换显存,降低激活值占用
logging_steps=50,
save_steps=500,
learning_rate=2e-4,
warmup_ratio=0.03,
lr_scheduler_type="cosine",
fp16=True,
report_to="tensorboard"
)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=dataset["train"],
data_collator=DataCollatorForLanguageModeling(tokenizer, mlm=False),
)
trainer.train()
batch size=8 + grad accumulation=4,等效batch=32。这个配置在RTX 3090上显存峰值约8.2G,训练速度约3.5步/秒(序列长度512)。
4.3 训练过程中的Loss曲线
训练3个epoch,总共约7500步,loss曲线如下:
- 初始loss:1.82
- 第500步:0.93
- 第1500步:0.51
- 第3000步:0.34
- 第5000步:0.27
- 第7500步(结束):0.21
曲线整体平滑下降,没有出现loss spike。如果你看到loss突然暴涨,大概率是学习率太大或数据里有异常长序列。我用的是2e-4,如果你用1e-4会更稳,但收敛会慢20%左右。
5. 踩坑与优化:量化感知训练的三个坑
坑1:4bit量化后训练不收敛
第一次跑QLoRA时,loss在0.8附近震荡不下降。排查了半天,发现是bnb_4bit_compute_dtype设成了torch.float32——虽然显存占用没变,但计算精度和后续的fp16训练不匹配,导致梯度更新异常。改成torch.float16后问题消失。
坑2:中文分词器截断
Llama-2的tokenizer对中文支持很差,一个汉字会被切成多个token。512的最大长度实际只能覆盖约200个汉字。我的解决方式是:在preprocess里对回答部分单独截断,保证回答完整:
answer = example["answer"]
if len(tokenizer.encode(answer)) > 300:
# 按字符粗截,保证语义完整
answer = answer[:200] # 200个汉字大约对应300+ token
坑3:eval时显存爆掉
训练没问题,但一跑eval就OOM。原因是Trainer默认会在eval时计算完整hidden state,加上4bit基座的反量化开销。解决:设置evaluation_strategy="steps"并在TrainingArguments里加eval_accumulation_steps=2,减少eval时的中间变量累积。
6. 推理效果对比:LoRA微调 vs 基座模型
微调完成后,用PEFT加载模型:
from peft import PeftModel
base_model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-2-7b-chat-hf",
quantization_config=bnb_config,
device_map="auto"
)
model = PeftModel.from_pretrained(base_model, "./qlora_medical/checkpoint-7500")
model.eval()
测试样本:
- 输入:患者60岁男性,有长期吸烟史,近期出现痰中带血,CT提示右肺上叶结节,最可能的诊断是什么?
- 基座模型输出:患者有长期吸烟史,痰中带血,CT提示结节,诊断需要考虑肺癌。建议进行进一步检查,如痰细胞学检查、纤维支气管镜检查,必要时行CT引导下肺穿刺活检。
- LoRA微调输出:根据患者长期吸烟史、痰中带血及CT提示右肺上叶结节,高度怀疑周围型肺癌。需进一步完善增强CT、痰脱落细胞学检查及经皮肺穿刺活检以明确病理类型。若确诊,需进行临床分期,首选手术治疗。
微调后模型明显更聚焦,直接给出了「周围型肺癌」的明确判断,而且给出了「穿刺活检」和「分期」的具体建议,比基座模型的泛泛而谈实用得多。
量化对比指标(在200条测试集上):
| 指标 | 基座模型 | LoRA微调后 |
|---|---|---|
| BLEU-4 | 0.21 | 0.67 |
| ROUGE-L | 0.31 | 0.78 |
| 医学实体F1 | 0.43 | 0.81 |
| 生成平均长度 | 87字 | 132字 |
7. 总结与建议
QLoRA确实能以极低的显存成本微调大模型,但需要注意:
- 量化配置别乱调:NF4+double quantization+fp16 compute是最稳妥的组合,改任何一个都可能掉精度。
- loss要盯紧:如果loss不降,先检查
compute_dtype,再检查学习率。QLoRA一般用1e-4到3e-4,太高必炸。 - 数据质量比模型重要:我清洗数据时去掉了2000多条含错别字的样本,BLEU从0.58直接跳到0.67。
- 推理时也要用量化:微调完直接保存LoRA adapter(约40MB),推理时再加载基座模型,这样部署时显存依然只有8G左右。
最后留个思考:LoRA的秩r=16是我试出来的最佳值,r=8时收敛慢但泛化更好,r=32时容易过拟合。如果你的任务数据量小于1万,建议r=8起步。数据量超过5万,r=32可能更好。这个需要自己试,没有标准答案。