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确实能以极低的显存成本微调大模型,但需要注意:

  1. 量化配置别乱调:NF4+double quantization+fp16 compute是最稳妥的组合,改任何一个都可能掉精度。
  2. loss要盯紧:如果loss不降,先检查compute_dtype,再检查学习率。QLoRA一般用1e-4到3e-4,太高必炸。
  3. 数据质量比模型重要:我清洗数据时去掉了2000多条含错别字的样本,BLEU从0.58直接跳到0.67。
  4. 推理时也要用量化:微调完直接保存LoRA adapter(约40MB),推理时再加载基座模型,这样部署时显存依然只有8G左右。

最后留个思考:LoRA的秩r=16是我试出来的最佳值,r=8时收敛慢但泛化更好,r=32时容易过拟合。如果你的任务数据量小于1万,建议r=8起步。数据量超过5万,r=32可能更好。这个需要自己试,没有标准答案。