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")
五、评估和迭代
| 指标 | 评估方法 |
|---|---|
| 输出质量 | 人工评分 |
| 安全性 | 红队测试 |
| 指令遵循 | 自动化测试集 |
| 领域准确性 | 专家审核 |
六、常见问题
- 过拟合:数据集太小或训练轮数过多
- 灾难性遗忘:模型忘记了预训练学到的知识
- 格式不一致:训练数据格式不统一