微调与训练入门
大模型微调(Fine-tuning)的基础概念、方法和实践指南。
什么是微调?
微调是在预训练模型基础上,使用特定领域数据进行二次训练,使模型更适应特定任务。
flowchart LR
A[预训练模型] --> B[+ 领域数据]
B --> C[微调训练]
C --> D[专用模型]微调的本质:
预训练模型已经学会了语言的通用知识,微调是在此基础上”教”它特定领域的知识或特定的行为模式。这比从头训练高效得多。
微调能解决什么问题?
| 问题 | 说明 | 微调效果 |
|---|---|---|
| 格式不一致 | 模型输出格式不稳定 | ⭐⭐⭐⭐⭐ |
| 风格不匹配 | 语气、措辞不符合要求 | ⭐⭐⭐⭐⭐ |
| 领域知识缺乏 | 专业术语理解不准 | ⭐⭐⭐ |
| 推理成本高 | 大模型太贵 | ⭐⭐⭐⭐ |
| 响应速度慢 | 需要更快响应 | ⭐⭐⭐⭐ |
微调 vs Prompt Engineering
| 方面 | Prompt Engineering | 微调 |
|---|---|---|
| 成本 | 低(无训练成本) | 高(需要 GPU 和数据) |
| 效果 | 通用场景好 | 特定任务更好 |
| 灵活性 | 高(随时修改) | 低(需要重新训练) |
| 数据需求 | 无 | 需要训练数据 |
| 适用场景 | 快速验证、通用任务 | 生产部署、特定任务 |
| 迭代速度 | 快(分钟级) | 慢(小时/天级) |
决策流程:
flowchart TD
A[需求] --> B{Prompt 能解决?}
B -->|是| C[使用 Prompt Engineering]
B -->|否| D{有足够数据?}
D -->|否| E[收集数据或用 RAG]
D -->|是| F{需要专业知识?}
F -->|是| G[先试 RAG]
F -->|否| H[微调]何时需要微调?
| 场景 | 是否微调 | 原因 |
|---|---|---|
| 通用问答 | ❌ Prompt 即可 | 模型已经很擅长 |
| 特定格式输出 | ✅ 微调效果好 | 格式一致性要求高 |
| 专业领域知识 | ⚠️ 先试 RAG | 微调不擅长注入新知识 |
| 特定风格/语气 | ✅ 微调效果好 | 风格是可学习的模式 |
| 降低推理成本 | ✅ 小模型微调 | 小模型+微调可接近大模型效果 |
微调方法
flowchart TD
A[微调方法] --> B[全量微调]
A --> C[参数高效微调 PEFT]
C --> D[LoRA]
C --> E[QLoRA]
C --> F[Prefix Tuning]
C --> G[Adapter]方法对比
| 方法 | 训练参数 | 显存需求 | 效果 | 适用场景 |
|---|---|---|---|---|
| 全量微调 | 100% | 极高 | 最好 | 资源充足、追求最佳效果 |
| LoRA | ~1% | 低 | 接近全量 | 大多数场景推荐 |
| QLoRA | ~1% | 更低 | 接近 LoRA | 显存受限 |
| Prefix Tuning | <1% | 最低 | 一般 | 极端资源受限 |
LoRA 原理
LoRA(Low-Rank Adaptation)是目前最流行的参数高效微调方法。
核心思想:
不直接修改原始权重,而是学习一个”增量”。这个增量用两个小矩阵的乘积表示(低秩分解)。
flowchart LR
A[输入] --> B[原始权重 W]
A --> C[低秩矩阵 A×B]
B --> D[+]
C --> D
D --> E[输出]为什么有效?
- 原始权重 W 是 d×d 的大矩阵(如 4096×4096)
- LoRA 学习 A(d×r) 和 B(r×d),r 通常是 8-64
- 参数量从 d² 降到 2dr,减少 99%+
- 研究表明,微调的”有效维度”很低,低秩足够
LoRA 只训练低秩分解矩阵,大幅减少参数量。
数据准备
数据质量是微调成功的关键。“垃圾进,垃圾出”在微调中尤为明显。
数据格式
OpenAI 格式
OpenAI 微调 API 使用的格式,也是最通用的格式:
{"messages": [
{"role": "system", "content": "你是客服助手"},
{"role": "user", "content": "如何退款?"},
{"role": "assistant", "content": "退款流程如下..."}
]}
Alpaca 格式
开源社区常用的格式:
{
"instruction": "翻译成英文",
"input": "你好世界",
"output": "Hello World"
}
ShareGPT 格式
多轮对话格式:
{
"conversations": [
{"from": "human", "value": "你好"},
{"from": "gpt", "value": "你好!有什么可以帮助你的?"},
{"from": "human", "value": "今天天气怎么样?"},
{"from": "gpt", "value": "抱歉,我无法获取实时天气信息。"}
]
}
数据质量要求
| 要求 | 说明 | 重要性 |
|---|---|---|
| 数量 | 通常 1000+ 条 | ⭐⭐⭐ |
| 质量 | 准确、一致 | ⭐⭐⭐⭐⭐ |
| 多样性 | 覆盖各种情况 | ⭐⭐⭐⭐ |
| 格式 | 统一规范 | ⭐⭐⭐⭐ |
数据量参考:
| 数据量 | 效果 | 适用场景 |
|---|---|---|
| 100-500 | 初步效果 | 验证可行性 |
| 1000-5000 | 较好效果 | 简单任务 |
| 5000-10000 | 良好效果 | 复杂任务 |
| 10000+ | 最佳效果 | 生产部署 |
数据清洗
def clean_data(examples):
cleaned = []
for ex in examples:
# 去除空白
if not ex["input"].strip() or not ex["output"].strip():
continue
# 长度限制
if len(ex["input"]) > 2000 or len(ex["output"]) > 2000:
continue
# 去重
# ...
cleaned.append(ex)
return cleaned
云端微调
OpenAI Fine-tuning
# 1. 上传数据
file = client.files.create(
file=open("train.jsonl", "rb"),
purpose="fine-tune"
)
# 2. 创建微调任务
job = client.fine_tuning.jobs.create(
training_file=file.id,
model="gpt-4o-mini-2024-07-18"
)
# 3. 查看状态
status = client.fine_tuning.jobs.retrieve(job.id)
# 4. 使用微调模型
response = client.chat.completions.create(
model="ft:gpt-4o-mini:org::xxx",
messages=[...]
)
支持微调的模型
| 厂商 | 模型 | 说明 |
|---|---|---|
| OpenAI | gpt-4o-mini | 推荐 |
| OpenAI | gpt-3.5-turbo | 经济 |
| 通义千问 | qwen 系列 | 国内 |
| 智谱 | glm 系列 | 国内 |
本地微调
环境准备
# 基础依赖
pip install torch transformers datasets
pip install peft accelerate bitsandbytes
# 可选
pip install wandb # 实验追踪
LoRA 微调示例
from transformers import AutoModelForCausalLM, AutoTokenizer
from peft import LoraConfig, get_peft_model
# 加载模型
model = AutoModelForCausalLM.from_pretrained("model_name")
tokenizer = AutoTokenizer.from_pretrained("model_name")
# LoRA 配置
lora_config = LoraConfig(
r=8, # 秩
lora_alpha=32, # 缩放因子
target_modules=["q_proj", "v_proj"],
lora_dropout=0.1,
task_type="CAUSAL_LM"
)
# 应用 LoRA
model = get_peft_model(model, lora_config)
model.print_trainable_parameters()
# 输出: trainable params: 0.1% of total
训练配置
from transformers import TrainingArguments
training_args = TrainingArguments(
output_dir="./output",
num_train_epochs=3,
per_device_train_batch_size=4,
gradient_accumulation_steps=4,
learning_rate=2e-4,
warmup_steps=100,
logging_steps=10,
save_steps=500,
fp16=True, # 混合精度
)
硬件需求
显存估算
| 模型大小 | 全量微调 | LoRA | QLoRA |
|---|---|---|---|
| 7B | 60GB+ | 16GB | 8GB |
| 13B | 120GB+ | 32GB | 16GB |
| 70B | 600GB+ | 160GB | 48GB |
推荐配置
| 场景 | GPU | 显存 |
|---|---|---|
| 学习实验 | RTX 3090/4090 | 24GB |
| 小模型微调 | A100 40GB | 40GB |
| 大模型微调 | A100 80GB × N | 80GB+ |
评估与验证
评估是微调成功的关键,需要在训练过程中和训练后进行多维度评估。
评估指标
| 指标 | 说明 | 用途 |
|---|---|---|
| Loss | 训练损失,越低越好 | 监控训练过程 |
| Perplexity | 困惑度,越低越好 | 衡量语言建模能力 |
| 任务指标 | 准确率、F1、BLEU 等 | 衡量任务效果 |
| 人工评估 | 质量、相关性、流畅度 | 最终验收 |
验证方法
flowchart TD
A[微调模型] --> B[自动评估]
A --> C[人工评估]
B --> D[测试集指标]
B --> E[基准测试]
C --> F[A/B 对比]
C --> G[专家评审]评估流程:
def evaluate_model(model, test_data):
results = {
"loss": [],
"accuracy": [],
"samples": []
}
for sample in test_data:
# 生成输出
output = model.generate(sample["input"])
# 计算指标
results["accuracy"].append(
output == sample["expected"]
)
# 保存样本用于人工评估
results["samples"].append({
"input": sample["input"],
"expected": sample["expected"],
"actual": output
})
return results
过拟合检测
| 现象 | 说明 | 解决方案 |
|---|---|---|
| 训练 Loss 低,验证 Loss 高 | 典型过拟合 | 增加数据、正则化、早停 |
| 只会回答训练数据中的问题 | 记忆化 | 增加数据多样性 |
| 通用能力下降 | 灾难性遗忘 | 混合通用数据训练 |
| 输出格式过于固定 | 过度拟合格式 | 增加格式变化 |
# 早停策略
from transformers import EarlyStoppingCallback
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_data,
eval_dataset=eval_data,
callbacks=[EarlyStoppingCallback(early_stopping_patience=3)]
)
最佳实践
1. 从小规模开始
循序渐进,逐步扩大规模:
100 条数据 → 验证流程是否正确
1000 条数据 → 观察初步效果
5000 条数据 → 调优超参数
10000+ 条数据 → 生产质量
2. 数据质量优先
| 优先级 | 说明 | 原因 |
|---|---|---|
| 1 | 数据准确性 | 错误数据会被学习 |
| 2 | 数据多样性 | 避免过拟合 |
| 3 | 数据数量 | 量变引起质变 |
数据质量检查清单:
- 输入输出配对正确
- 无重复数据
- 格式统一
- 无敏感信息
- 覆盖各种边界情况
3. 保留基础能力
# 混合通用数据,防止灾难性遗忘
domain_data = load_domain_data() # 领域数据
general_data = load_general_data() # 通用数据
# 按比例混合
train_data = domain_data + random.sample(general_data, len(domain_data) // 10)
random.shuffle(train_data)
4. 版本管理
| 记录内容 | 说明 | 工具 |
|---|---|---|
| 数据版本 | 训练数据快照 | DVC, Git LFS |
| 超参数 | 学习率、epoch 等 | MLflow, W&B |
| 模型权重 | checkpoint | Hugging Face Hub |
| 评估结果 | 各项指标 | MLflow, W&B |
# 使用 Weights & Biases 追踪实验
import wandb
wandb.init(project="my-finetune")
wandb.config.update({
"learning_rate": 2e-4,
"epochs": 3,
"batch_size": 4,
"lora_r": 8
})
# 训练过程中记录
wandb.log({"loss": loss, "eval_accuracy": accuracy})
5. 超参数调优
关键超参数:
| 参数 | 建议范围 | 说明 |
|---|---|---|
| learning_rate | 1e-5 ~ 5e-4 | LoRA 可以用较大学习率 |
| epochs | 1-5 | 过多容易过拟合 |
| batch_size | 4-32 | 受显存限制 |
| lora_r | 8-64 | 越大效果越好但参数越多 |
| lora_alpha | 16-64 | 通常设为 2×r |
常见问题
Q: 微调后效果变差?
| 原因 | 解决 |
|---|---|
| 数据质量差 | 清洗数据 |
| 过拟合 | 减少 epoch、增加数据 |
| 学习率太高 | 降低学习率 |
| 灾难性遗忘 | 混合通用数据 |
Q: 显存不足?
# 使用 QLoRA(4-bit 量化 + LoRA)
from transformers import BitsAndBytesConfig
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.float16
)
model = AutoModelForCausalLM.from_pretrained(
model_name,
quantization_config=bnb_config
)
Q: 训练太慢?
| 优化 | 效果 |
|---|---|
| 混合精度 (fp16/bf16) | 2x 加速 |
| 梯度累积 | 模拟大 batch |
| 多 GPU | 线性加速 |
| Flash Attention | 显著加速 |