AI 模型微调入门教程:从数据集准备到 LoRA 训练完整指南

通用大模型在特定领域的表现往往不尽如人意。微调(Fine-tuning)让开发者可以用自己的数据将通用模型定制化,大幅提升特定场景的表现。

一、什么是微调

1.1 微调 vs 提示工程 vs RAG

方法 原理 适用场景 成本
提示工程 精心设计输入提示 简单任务 极低
RAG 检索外部知识辅助生成 需要实时知识的场景
微调 在特定数据上继续训练 特定风格/领域优化

1.2 什么时候需要微调

  • 需要模型学习特定领域术语和表达方式
  • 模型输出风格不符合业务要求
  • 提示工程和 RAG 无法满足效果要求
  • 需要降低推理成本(小模型微调后替代大模型)

二、数据集准备

2.1 数据格式

最常见的格式是对话/指令格式:

{
  "messages": [
    {"role": "system", "content": "你是一个专业的客服助手。"},
    {"role": "user", "content": "如何查询我的订单状态?"},
    {"role": "assistant", "content": "您好,您可以在网站右上角点击"我的订单",输入订单号即可查询。"}
  ]
}

2.2 数据质量要求

  • 数量:至少 100-1000 条高质量对话
  • 多样性:覆盖各种场景和边界情况
  • 一致性:风格和格式保持一致
  • 准确性:内容经过人工审核

2.3 数据增强

如果数据集较小,可以通过以下方法增强:

  • 同义词替换
  • 回译(中→英→中)
  • 模板扩写
  • AI 辅助生成

三、微调方法对比

3.1 全参数微调

更新模型所有参数,效果最好但成本最高。

  • 硬件需求:7B 模型至少需要 4×A100 80GB
  • 适用:预算充足、追求极致效果

3.2 LoRA (Low-Rank Adaptation)

LoRA 在模型原有权重旁添加小型可训练矩阵,大幅降低训练成本。

W' = W + BA

其中 W 是冻结的原权重,BA 是低秩可训练矩阵。

  • 硬件需求:7B 模型单张 24GB 显卡即可
  • 适用:大多数微调场景

3.3 QLoRA

QLoRA = 量化 + LoRA,进一步降低硬件需求。

  • 硬件需求:7B 模型单张 12GB 显卡即可
  • 适用:预算有限、实验性项目

四、实战:使用 LoRA 微调 Llama 3

4.1 环境准备

pip install torch transformers datasets peft accelerate bitsandbytes

4.2 加载模型

from transformers import AutoModelForCausalLM, AutoTokenizer
from peft import get_peft_model, LoraConfig, TaskType

model = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3.1-8B")
tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-3.1-8B")

lora_config = LoraConfig(
    r=16,
    lora_alpha=32,
    target_modules=["q_proj", "v_proj"],
    lora_dropout=0.1,
    task_type=TaskType.CAUSAL_LM
)

model = get_peft_model(model, lora_config)
print(model.print_trainable_parameters())

4.3 训练

from transformers import TrainingArguments, Trainer

training_args = TrainingArguments(
    output_dir="./fine-tuned-model",
    num_train_epochs=3,
    per_device_train_batch_size=4,
    gradient_accumulation_steps=4,
    learning_rate=2e-4,
    fp16=True,
    save_steps=500,
    logging_steps=50,
)

trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=dataset,
)

trainer.train()

4.4 合并与导出

from peft import PeftModel

# 加载 LoRA 权重
model = PeftModel.from_pretrained(base_model, "./lora-checkpoint")
# 合并权重
merged_model = model.merge_and_unload()
merged_model.save_pretrained("./final-model")

五、评估和迭代

指标 评估方法
输出质量 人工评分
安全性 红队测试
指令遵循 自动化测试集
领域准确性 专家审核

六、常见问题

  • 过拟合:数据集太小或训练轮数过多
  • 灾难性遗忘:模型忘记了预训练学到的知识
  • 格式不一致:训练数据格式不统一