Skip to content

4.2 SFT 监督微调:教模型学会新任务

承前:从预训练到监督微调

在上一节中,我们建立了对"微调"的整体认知:预训练阶段模型在海量无标注文本上学习语言的基本规律——词汇搭配、语法结构、世界知识——这时的模型像一个博览群书但不会"听话"的学生,它只会接续文本,却不知道如何按人的要求给出回答。微调,就是在这个"通才"基础上,用有标注的数据把它引导成"专才"的过程。我们区分了全参数微调与部分参数微调,也提到微调是连接预训练与对齐的关键桥梁。本节就要走完这座桥梁的第一段——监督微调(SFT),看清楚"给模型出题、对答案"这件事到底是怎么发生的。

什么是 SFT:用"做题对答案"来理解

监督微调(Supervised Fine-Tuning,SFT)的核心思想朴素得近乎直白:给模型一批"题目 + 标准答案"的样本对,让它学会按题目要求作答

可以把预训练想象成一个学生花了几年时间读完了图书馆里所有的书——天文地理、小说散文无所不包。这时如果你问他"请把这句话翻成英文",他可能会继续往下写一段中文小说,而不是执行翻译任务,因为他根本没意识到你在"下指令"。

SFT 做的事情,相当于给这个学生一本习题册:每道题前面写着题目要求(指令),后面写着标准答案。学生做完一页,老师批改一页——如果回答偏离了题目意图,损失值就高,权重就往"贴题"的方向调整。做完整本习题册,学生就学会了"看到题目要先理解要求,再组织回答"这个习惯。

SFT 训练流程:
┌──────────────┐     ┌──────────────────┐     ┌──────────────┐
│  预训练模型    │ ──> │  SFT 训练         │ ──> │  Chat 模型   │
│  (Base Model) │     │  指令-回答对       │     │  (能对话了!)  │
│  只会补全文本  │     │  教模型"对话格式"   │     │  会遵循指令   │
└──────────────┘     └──────────────────┘     └──────────────┘

从数学角度看,SFT 与预训练使用完全相同的损失函数——交叉熵损失(Cross-Entropy Loss),但有一个关键区别:只对"回答部分"计算损失,忽略"指令部分"

SFT 的损失计算策略:
┌─────────────────────────────────────────────────────┐
│  [INST] 解释什么是量子计算 [/INST] 量子计算是...      │
│         ↑── 这部分不计算损失 ──↑  ↑── 只计算这部分的损失 │
│                                                      │
│  只对"回答部分"计算损失,忽略"指令部分"                  │
│  这确保模型学习的是"如何回答",而非"如何提问"            │
└─────────────────────────────────────────────────────┘

为什么要这样设计?想象老师批改作业时,只看学生的解答过程对不对,不会因为题目本身"写得不好"就扣学生的分。模型也一样——指令是"题目",回答才是"解答",损失只在解答上计算,模型才能学到"怎么答题"而不是"怎么出题"。

三种主流指令数据格式

不同模型使用不同的对话模板。理解这些格式是构建数据集的第一步。可以类比为不同的考试答题卡格式——同样的题目,填到不同的表格里。

格式一:Alpaca 格式

斯坦福 Alpaca 项目开创的经典格式,结构简单清晰:

json
{
  "instruction": "将以下句子翻译成英文:今天天气真好。",
  "input": "",
  "output": "The weather is really nice today."
}

使用场景:单轮指令任务,如翻译、分类、生成等。

特点

  • 三字段结构(instruction / input / output)
  • input 字段可选,用于提供上下文信息
  • 简单直观,适合自动化数据生成
格式二:ShareGPT 格式

来源于用户与 ChatGPT 的真实对话分享,天然支持多轮对话:

json
{
  "conversations": [
    {"from": "human", "value": "你好,我想学编程,有什么建议吗?"},
    {"from": "gpt", "value": "当然!我建议从 Python 开始..."},
    {"from": "human", "value": "那学完 Python 基础之后呢?"},
    {"from": "gpt", "value": "接下来可以学数据结构和算法..."}
  ]
}

使用场景:多轮对话、客服系统、教学助手等。

特点

  • 天然支持多轮对话
  • 角色明确(human / gpt / system)
  • 保留真实对话上下文
格式三:ChatML 格式

OpenAI 提出的标准化格式,被许多开源模型(如 Qwen、DeepSeek)采用:

<|im_start|>system
你是一个有用的助手。<|im_end|>
<|im_start|>user
解释什么是机器学习。<|im_end|>
<|im_start|>assistant
机器学习是人工智能的一个分支...<|im_end|>

使用场景:通用对话模型,支持 system prompt。

特点

  • 使用特殊 Token 标记角色边界
  • 支持 system / user / assistant 三种角色
  • 格式严格,便于模型精确解析
格式对比
特性AlpacaShareGPTChatML
多轮对话❌ 不支持✅ 支持✅ 支持
System Prompt❌ 不支持⚠️ 需要扩展✅ 原生支持
复杂度
代表性模型Alpaca, VicunaLLaMA, MistralQwen, DeepSeek
适用场景单轮任务多轮对话通用对话

指令数据集的构建

构建高质量指令数据集是 SFT 中最关键也最耗时的环节。数据集的质量直接决定了微调后模型的能力上限——再好的训练方法也救不了一份糟糕的数据集,这就是社区常说的 "garbage in, garbage out"(垃圾进,垃圾出)。

推荐的构建流程分为四个阶段:

数据构建四个阶段:
┌──────────────┐    ┌──────────────┐    ┌──────────────┐    ┌──────────────┐
│ 1. 数据收集    │ -> │ 2. 格式转换    │ -> │ 3. 质量过滤    │ -> │ 4. 数据增强    │
│ 种子数据       │    │ 统一模板       │    │ 去重去噪       │    │ 扩充多样性     │
│ 爬取/生成      │    │ 角色标注       │    │ 长度过滤       │    │ 难度分层       │
└──────────────┘    └──────────────┘    └──────────────┘    └──────────────┘

第一阶段:数据收集。数据来源有三条主流路径。一是公开数据集,如 Alpaca、BELLE、Firefly 等中文指令数据,可以直接下载使用,适合快速实验。二是人工标注,成本最高但质量最可控,适合对垂直领域有严格要求的场景,如医疗问答、法律咨询。三是模型生成(self-instruct),用 GPT-4 等强模型根据种子指令生成新的指令-回答对,成本低、规模大,但需要严格过滤幻觉与低质内容。

第二阶段:格式转换。收集来的数据格式各异,必须统一到目标模板。比如你决定用 Qwen 模型,那就要把 Alpaca 格式的 instruction/input/output 转成 ChatML 格式。这一步看似机械,却是后续训练能否正确进行的前提——模型只认识它被预训练时用过的那种对话模板。

第三阶段:质量过滤。这是最容易被忽视、却最影响效果的一环。需要做四件事:去重(近似重复的样本只保留一条,避免模型对某些模式过拟合);去噪(剔除回答不完整、答非所问、包含乱码的样本);长度过滤(过短的回答信息量不足,过长的回答容易超长截断,建议回答长度在 50–512 Token 之间);安全审查(剔除有害、偏见、涉政等不安全内容)。

第四阶段:数据增强。在基础数据集上增加多样性:混入不同难度(简单/中等/困难)的样本,覆盖不同任务类型(问答、翻译、代码、推理、创作),并适当加入"拒答样本"(教会模型在不懂时说"我不知道",而不是胡编乱造)。

数据质量检查清单

  • [ ] 回答是否准确、完整?
  • [ ] 格式是否与目标模板一致?
  • [ ] 是否有重复或高度相似的样本?
  • [ ] 回答长度是否合理(不过短也不过长)?
  • [ ] 是否包含有害、偏见或不安全内容?
  • [ ] 覆盖的任务类型是否足够多样?

SFT 训练实战

使用 TRL 库的 SFTTrainer 进行训练是最便捷的方式。它封装了数据处理、模型加载、训练循环等复杂逻辑,让我们能专注于数据和参数本身。下面这段代码是一次完整的 SFT 训练流程,我们对每一行做详细注释。

python
# 从 HuggingFace transformers 库导入分词器、模型和训练参数类
from transformers import AutoTokenizer, AutoModelForCausalLM, TrainingArguments
# 从 TRL 库导入 SFT 训练器和"仅补全部分计算损失"的数据整理器
from trl import SFTTrainer, DataCollatorForCompletionOnlyLM
# 从 datasets 库导入数据集加载函数
from datasets import load_dataset
# 导入 PyTorch,用于指定数据精度
import torch

# ── 第 1 步:加载模型和分词器 ──
# 指定要微调的基座模型,这里用 Qwen2-0.5B(小模型适合实验)
model_name = "Qwen/Qwen2-0.5B"
# 加载分词器:负责把文本切分成 Token ID。trust_remote_code=True 允许执行模型仓库里的自定义代码
tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
# 如果模型没有设置 padding token,用结束符 EOS 兜底,否则 batch 训练会报错
if tokenizer.pad_token is None:
    tokenizer.pad_token = tokenizer.eos_token
# 加载因果语言模型(即"从左到右预测下一个 Token"的标准 LLM)
model = AutoModelForCausalLM.from_pretrained(
    model_name,
    torch_dtype=torch.bfloat16,  # 使用 bfloat16 半精度,显存占用减半
    device_map="auto",           # 自动把模型放到可见的 GPU 上
    trust_remote_code=True,      # 允许执行模型仓库里的自定义代码
)

# ── 第 2 步:加载并格式化数据集 ──
# 从 HuggingFace Hub 加载中文 Alpaca 数据集,只取训练集
dataset = load_dataset("silk-road/alpaca-data-gpt4-chinese", split="train")

# 定义格式转换函数:把 Alpaca 三字段格式拼成 Qwen 使用的 ChatML 格式
def format_chatml(example):
    """将 Alpaca 格式转换为 ChatML 格式"""
    # 如果 input 字段非空,就把指令和输入拼在一起作为用户消息
    if example.get("input") and example["input"].strip():
        user_content = f"{example['instruction']}\n{example['input']}"
    else:
        user_content = example["instruction"]
    # 用 ChatML 特殊 Token 拼出完整的对话文本
    return {
        "text": (
            f"<|im_start|>system\n你是一个有用的助手。<|im_end|>\n"
            f"<|im_start|>user\n{user_content}<|im_end|>\n"
            f"<|im_start|>assistant\n{example['output']}<|im_end|>"
        )
    }

# 用 map 方法批量转换整个数据集
formatted_dataset = dataset.map(format_chatml)

# ── 第 3 步:配置训练参数 ──
training_args = TrainingArguments(
    output_dir="./sft-output",          # 模型和日志的输出目录
    num_train_epochs=3,                  # 训练轮数:整个数据集过 3 遍
    per_device_train_batch_size=4,       # 每张 GPU 一次处理的样本数
    gradient_accumulation_steps=4,       # 梯度累积 4 步,等效 batch=16,缓解显存压力
    learning_rate=2e-5,                  # 学习率:SFT 通常在 1e-5 到 5e-5 之间
    warmup_ratio=0.03,                   # 前 3% 的步数做学习率预热,避免初始震荡
    lr_scheduler_type="cosine",          # 学习率按余弦曲线衰减,后期缓慢下降
    logging_steps=10,                   # 每 10 步打印一次训练日志
    save_strategy="epoch",              # 每个 epoch 结束保存一次 checkpoint
    bf16=True,                           # 开启 bfloat16 混合精度训练
    report_to="none",                    # 不上报到 WandB 等实验平台
)

# ── 第 4 步:创建 SFT Trainer ──
trainer = SFTTrainer(
    model=model,                         # 传入加载好的基座模型
    args=training_args,                  # 传入训练参数
    train_dataset=formatted_dataset,     # 传入格式化后的数据集
    tokenizer=tokenizer,                 # 传入分词器
    max_seq_length=512,                  # 最大序列长度 512 Token(按显存调整)
    dataset_text_field="text",           # 告诉 Trainer 文本数据在 "text" 字段里
)

# ── 第 5 步:开始训练 ──
trainer.train()

# ── 第 6 步:保存模型和分词器 ──
trainer.save_model("./sft-output/final")           # 保存微调后的模型权重
tokenizer.save_pretrained("./sft-output/final")     # 保存分词器,推理时要配套使用

实战练习

练习 1:构建自定义指令数据集

从一个 CSV 文件构建指令微调数据集。

python
import pandas as pd
from datasets import Dataset

# 假设你有一个 CSV 文件,包含 question 和 answer 两列
# 这里我们演示如何从零构建

# 创建示例数据
data = {
    "instruction": [
        "请用一句话总结深度学习的核心思想。",
        "Python 中 list 和 tuple 的区别是什么?",
        "写一个计算斐波那契数列的 Python 函数。",
    ],
    "output": [
        "深度学习通过多层神经网络自动学习数据的层次化特征表示,无需人工设计特征。",
        "list 是可变的(可以增删改元素),tuple 是不可变的(创建后不能修改)。list 用方括号 [],tuple 用圆括号 ()。",
        "def fibonacci(n):\n    a, b = 0, 1\n    for _ in range(n):\n        a, b = b, a + b\n    return a",
    ]
}

df = pd.DataFrame(data)
dataset = Dataset.from_pandas(df)

# 转换为 ChatML 格式
def format_chatml(example):
    return {
        "text": (
            f"<|im_start|>user\n{example['instruction']}<|im_end|>\n"
            f"<|im_start|>assistant\n{example['output']}<|im_end|>"
        )
    }

formatted_dataset = dataset.map(format_chatml)
print(f"数据集大小: {len(formatted_dataset)}")
print(f"\n第一条数据:")
print(formatted_dataset[0]["text"])
练习 2:使用 LLaMA-Factory 进行可视化微调

LLaMA-Factory 是目前最流行的微调框架之一,支持 Web UI 操作,适合不熟悉代码的团队协作场景。

bash
# 安装 LLaMA-Factory
git clone https://github.com/hiyouga/LLaMA-Factory.git
cd LLaMA-Factory
pip install -e ".[torch,metrics]"

# 启动 Web UI
llamafactory-cli webui

# 或者使用命令行进行微调
llamafactory-cli train \
    --model_name_or_path Qwen/Qwen2-0.5B \
    --dataset alpaca_zh \
    --template qwen \
    --finetuning_type lora \
    --output_dir ./qwen2-lora \
    --per_device_train_batch_size 4 \
    --gradient_accumulation_steps 4 \
    --lr_scheduler_type cosine \
    --logging_steps 10 \
    --save_steps 500 \
    --learning_rate 5e-5 \
    --num_train_epochs 3.0 \
    --fp16

常见误区

初学者在实践 SFT 时容易踩几个坑,这里提前预警。

误区一:数据越多越好。 很多人以为指令数据越多效果越好,于是堆上几十万条未经清洗的样本。实际上,5000 条高质量数据的训练效果往往优于 5 万条低质量数据。低质量数据不仅浪费算力,还会引入噪声,让模型学到错误的回答模式。建议先用小而精的数据集跑通流程,再逐步扩充。

误区二:学习率沿用预训练的设置。 预训练阶段的学习率通常在 1e-4 以上,但 SFT 阶段模型已经是一个"成品",过大的学习率会破坏已学到的知识,导致灾难性遗忘(catastrophic forgetting)——模型学会了新任务,却把预训练时学到的常识忘光了。SFT 的学习率一般设为 1e-5 到 5e-5,是预训练的十分之一甚至更小。

误区三:忽略损失屏蔽(loss masking)。 如果不区分"指令部分"和"回答部分",对整段文本统一计算损失,模型会同时学习"怎么出题"和"怎么答题",结果两方面都学不好。正确做法是用 DataCollatorForCompletionOnlyLM 等工具,只在回答部分计算损失。这也是 SFTTrainer 相比手写训练循环的核心价值之一。

误区四:只看训练损失,不验证效果。 训练损失下降不代表模型变好了——它可能只是把训练集背了下来(过拟合)。一定要留出验证集,或者训练后实际跑几条指令看看回答质量。过拟合的典型表现是:训练集上的 loss 很低,但模型对新指令的回答变得僵硬、重复或答非所问。

误区五:格式与模型不匹配。 每个模型家族有自己的对话模板——Qwen 用 ChatML,LLaMA 用 <<SYS>> 标签,Mistral 用 [INST] 标签。如果训练时用的格式和推理时用的格式不一致,模型就会"听不懂"指令。务必确认数据格式、训练模板、推理模板三者统一。

本节小结

要点说明
SFT 本质在预训练模型上用指令-回答对进行监督学习,教会模型对话格式
核心类比像"做题对答案":给题目和标准答案,模型只对答案部分计算损失
损失计算只对"回答"部分计算损失,忽略"指令"部分
格式选择Alpaca 适合单轮任务,ShareGPT 适合多轮对话,ChatML 是通用标准
数据构建四阶段流程:收集→转换→过滤→增强,质量永远优先于数量
训练工具推荐 TRL 的 SFTTrainer 或 LLaMA-Factory 框架
学习率SFT 学习率通常为 1e-5 ~ 5e-5,远小于预训练

启后:从全参数微调到参数高效微调

本节我们走完了 SFT 的完整流程——从指令数据集的构建,到用 SFTTrainer 跑通训练。但你会发现一个问题:上面的代码里我们对模型的全部参数进行了更新。对于一个 7B 参数的模型,全参数 SFT 需要至少 14GB 显存(半精度),如果是 70B 模型则要 140GB 以上——这对大多数开发者和中小团队来说是难以承受的。

有没有办法只更新极少量参数,就能达到接近全参数微调的效果?答案是肯定的,这就是下一节要讲的参数高效微调(PEFT)。PEFT 家族中最著名的成员是 LoRA(Low-Rank Adaptation),它通过在权重矩阵旁挂一个低秩"旁路"来学习增量,训练参数量可降至原来的 1% 以下,而效果却能达到全参数微调的 90% 以上。下一节我们就来揭开 LoRA 的原理与实战。

参考资料

  1. TRL 库 SFTTrainer 官方文档 - HuggingFace 官方 SFT 训练器完整文档
  2. LLaMA-Factory 项目 - 最流行的 LLM 微调框架,支持 100+ 模型
  3. LLM工程师手册--监督微调 - 中文 SFT 实战教程
  4. Stanford Alpaca 项目 - SFT 开山之作,展示了如何用 GPT 生成训练数据
  5. OpenAI ChatML 格式规范 - ChatML 格式的官方说明