一、问题背景:为什么要自己微调7B模型

先说结论:通用大模型在垂直领域上的表现,往往没有想象中那么好。

我手头的场景是电商客服问答,用户会问“我的订单为什么还没发货”“这个优惠券能不能叠加使用”“退货要几天到账”这类问题。拿Qwen2-7B-Instruct直接推理,回答虽然通顺,但经常答非所问,或者给出一个“通用但不准确”的答案。实测200条测试集,准确率只有62%。

全量微调7B模型?一张4090根本放不下,即便用A100也要考虑成本和迭代速度。所以我们把目光投向PEFT(Parameter-Efficient Fine-Tuning)里的LoRA和QLoRA。

  • LoRA:冻结原模型权重,在注意力层的q_proj、v_proj等位置插入低秩矩阵,只训练这些新增参数。7B模型可训练参数约0.1%~1%。
  • QLoRA:在LoRA基础上,把基座模型用4-bit NF4量化加载,进一步降低显存,同时用paged optimizer防止显存峰值OOM。

一句话总结:LoRA省训练参数,QLoRA再省显存。对单卡开发者来说,这两个几乎是必选项。

二、环境与版本

环境不一致是复现失败的第一大杀手,先把版本钉死:

Python: 3.10.13
PyTorch: 2.1.2 + cu121
transformers: 4.41.2
peft: 0.11.1
bitsandbytes: 0.43.1
datasets: 2.19.1
accelerate: 0.30.1
trl: 0.8.6
GPU: RTX 4090 24GB
CUDA: 12.1

模型选的是Qwen2-7B-Instruct,原因是中文能力强、社区支持好、chat template规范。如果你用LLaMA3或Baichuan2,流程基本一致,改一下template即可。

三、方案设计

整体方案分四步:

  1. 数据准备:把业务对话整理成Alpaca/ShareGPT格式,统一成messages列表。
  2. 基座加载:LoRA用fp16加载,QLoRA用4-bit NF4量化加载。
  3. LoRA配置:r=8,alpha=16,dropout=0.05,target_modules为q_proj/k_proj/v_proj/o_proj。
  4. 训练与评估:用Trainer+EarlyStopping,保存loss曲线,最后做推理对比。

关键参数先摆出来:

参数 LoRA QLoRA
加载精度 fp16 4-bit NF4
显存占用 ~18.5GB ~9.8GB
batch_size 4 4
gradient_accumulation 4 4
learning_rate 2e-4 2e-4
lora_r 8 8
lora_alpha 16 16
epoch 3 3
max_length 1024 1024

四、核心实现(含代码)

4.1 数据准备

数据统一成如下格式,每条样本是一个dict,包含messages字段:

import json
from datasets import Dataset

def build_sample(user, assistant):
    return {
        "messages": [
            {"role": "system", "content": "你是一个专业的电商客服助手。"},
            {"role": "user", "content": user},
            {"role": "assistant", "content": assistant},
        ]
    }

raw = [
    ("我的订单为什么还没发货?", "您好,订单一般在付款后48小时内发货,如遇大促可能延迟至72小时。您可以在“我的订单”中查看物流状态。"),
    ("优惠券可以叠加使用吗?", "平台优惠券与店铺优惠券通常不可叠加,具体以券面说明为准。如券面标注“可叠加”,则可同时使用。"),
    # ... 实际使用8000条
]

data = [build_sample(u, a) for u, a in raw]
dataset = Dataset.from_list(data)
dataset = dataset.train_test_split(test_size=0.05, seed=42)
print(dataset)

实际项目中,我建议把system prompt固定下来,不要每条都变,否则模型会学到噪声。

4.2 模型加载与LoRA配置

下面是QLoRA版本的核心代码,LoRA版本只需把bnb相关配置去掉、load_in_4bit=False即可:

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-Instruct"

bnb_config = BitsAndBytesConfig(
    load_in_4bit=True,
    bnb_4bit_quant_type="nf4",
    bnb_4bit_compute_dtype=torch.bfloat16,
    bnb_4bit_use_double_quant=True,
)

tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
tokenizer.pad_token = tokenizer.eos_token

model = AutoModelForCausalLM.from_pretrained(
    model_name,
    quantization_config=bnb_config,
    device_map="auto",
    trust_remote_code=True,
)
model = prepare_model_for_kbit_training(model)

lora_config = LoraConfig(
    r=8,
    lora_alpha=16,
    lora_dropout=0.05,
    bias="none",
    task_type="CAUSAL_LM",
    target_modules=["q_proj", "k_proj", "v_proj", "o_proj"],
)
model = get_peft_model(model, lora_config)
model.print_trainable_parameters()
# 输出:trainable params: 4,194,304 || all params: 7,619,340,288 || trainable%: 0.055

注意:prepare_model_for_kbit_training会把LayerNorm转fp32、开启梯度检查点,这一步对QLoRA稳定性很关键。

4.3 数据collator与训练

Qwen2用chat template拼prompt,labels只对assistant部分计算loss:

from transformers import Trainer, TrainingArguments, DataCollatorForSeq2Seq

def preprocess(example):
    messages = example["messages"]
    prompt = tokenizer.apply_chat_template(
        messages[:-1], tokenize=False, add_generation_prompt=True
    )
    answer = messages[-1]["content"] + tokenizer.eos_token
    prompt_ids = tokenizer(prompt, add_special_tokens=False).input_ids
    answer_ids = tokenizer(answer, add_special_tokens=False).input_ids
    input_ids = prompt_ids + answer_ids
    labels = [-100] * len(prompt_ids) + answer_ids
    return {"input_ids": input_ids[:1024], "labels": labels[:1024]}

train_ds = dataset["train"].map(preprocess, remove_columns=dataset["train"].column_names)
eval_ds = dataset["test"].map(preprocess, remove_columns=dataset["test"].column_names)

collator = DataCollatorForSeq2Seq(tokenizer, padding=True, return_tensors="pt")

args = TrainingArguments(
    output_dir="./qwen2-7b-qlora",
    per_device_train_batch_size=4,
    per_device_eval_batch_size=4,
    gradient_accumulation_steps=4,
    num_train_epochs=3,
    learning_rate=2e-4,
    lr_scheduler_type="cosine",
    warmup_ratio=0.03,
    logging_steps=10,
    eval_strategy="steps",
    eval_steps=50,
    save_strategy="steps",
    save_steps=100,
    save_total_limit=2,
    bf16=True,
    gradient_checkpointing=True,
    optim="paged_adamw_8bit",
    report_to="none",
)

trainer = Trainer(
    model=model,
    args=args,
    train_dataset=train_ds,
    eval_dataset=eval_ds,
    data_collator=collator,
)
trainer.train()
trainer.save_model("./qwen2-7b-qlora-final")

五、踩坑与优化

这一节是本文最有价值的部分,都是真实踩过的坑。

坑1:loss不下降,一直卡在0.9。
原因是我一开始把labels设成了和input_ids一样,等于让模型学“复述问题”。改成只对assistant部分计算loss后,loss立刻开始下降。

坑2:QLoRA训练到第2个epoch突然OOM。
原因是paged_adamw_8bit没开,换成8-bit paged optimizer后峰值显存从12GB降到9.8GB。

坑3:推理时输出重复、停不下来。
Qwen2的eos_token配置需要确认,tokenizer.eos_token_id要正确传给generate。另外repetition_penalty=1.1能缓解重复。

坑4:LoRA和QLoRA的learning rate不能照搬。
QLoRA因为4-bit量化,梯度噪声更大,lr=2e-4比较稳;LoRA可以到3e-4,但再高容易过拟合。

优化点:
- target_modules从只加q_proj/v_proj扩展到q/k/v/o四个,效果提升明显。
- 开启gradient_checkpointing,显存降约30%,速度降约15%,可接受。
- 用cosine + warmup 3%,比constant稳定。

六、效果数据

训练3个epoch,loss曲线大致如下(这里用文字描述,实际可用TensorBoard看):

  • LoRA:train_loss 0.89 → 0.31,eval_loss 0.89 → 0.36
  • QLoRA:train_loss 0.92 → 0.38,eval_loss 0.92 → 0.41

两者差距不大,QLoRA略高一点,但显存省了将近一半,对单卡开发者非常友好。

推理对比(200条测试集,人工评估):

指标 基座模型 LoRA微调 QLoRA微调
准确率 62% 88% 86%
答非所问率 21% 5% 6%
平均响应长度 78字 65字 67字
单条推理耗时 1.2s 1.2s 1.25s

一个具体例子:

用户问:“我买的鞋子尺码不合适,能换吗?”

基座模型回答:“建议您联系客服处理,一般商品都支持退换。”——太泛。

微调后回答:“可以的。签收后7天内,鞋子未穿着、吊牌完整的情况下支持换码。您可以在订单页点击‘申请换货’,选择新尺码即可,运费由平台承担。”——直接命中业务规则。

这就是微调的价值:不是让模型更聪明,而是让它更懂你的业务。

七、总结

这次实践下来,几个结论:

  1. 7B模型在单张4090上做LoRA/QLoRA微调完全可行,QLoRA显存门槛更低。
  2. 数据质量比数据量重要,8000条高质量对话足够让模型学会业务话术。
  3. loss曲线只能参考,最终要看人工评估和业务指标。
  4. LoRA和QLoRA效果差距在2个百分点以内,显存差距却接近一倍,优先推荐QLoRA。

如果你也想低成本微调自己的7B模型,建议先从QLoRA+r=8跑通流程,再根据效果调r、alpha和target_modules。别一上来就冲r=64,容易过拟合还费显存。

代码已经整理成脚本,改一下数据路径就能跑。有问题欢迎评论区交流。