LLM训练三阶段-预训练-SFT-RLHF
> **本文适合谁**
LLM 训练三阶段:预训练、SFT 与 RLHF
本文适合谁
想理解"为什么模型会拒绝某些请求"、"微调能改变什么不能改变什么"、"为什么 System Prompt 这么有效"的开发者。理解训练过程,就能理解模型行为的根源。
LLM 的训练分三个阶段,分工清晰:
- 预训练(Pre-training):学知识
- 监督微调(SFT,Supervised Fine-Tuning):学格式
- 强化学习对齐(RLHF/DPO):学价值观
理解这三个阶段,才能理解模型为什么能回答问题、会遵循指令、会拒绝有害请求。
1.1 第一阶段:预训练(Pre-training)
1.1.1 为什么"预测下一个词"能学会一切:自监督学习的魔力
预训练的核心思想是自监督学习:不需要任何人工标注,直接从文本本身构造学习信号。这个想法看起来简单,但背后有深刻的逻辑。
图 6.12:LLM 训练三阶段流程——预训练给能力,SFT 给格式,RLHF 给价值观
预测下一个词,要求模型必须学会很多能力:语法(错误的语法会让下一个词的预测变差)、语义(不理解词义无法预测合理的续写)、常识("苹果从树上落下来,因为__"——要预测"重力",模型需要知道物理知识)、推理(数学、逻辑问题的模式在大量文本中反复出现,预测下一个词时需要推理)。
所有这些能力,都隐藏在"预测下一个词"这一个目标中。互联网上几乎所有人类知识都以文本形式存在,通过海量文本的预测训练,模型把这些知识全部编码进了神经网络的权重。
1.1.2 Next-Token Prediction:最简单的目标函数
预训练的训练目标极其简单——给定前面所有的 Token,预测下一个 Token。这个目标函数称为语言模型损失(Language Modeling Loss),即负对数似然(模型对正确答案的预测概率越高,损失越低):
$$\mathcal{L} = -\sum_{t=1}^{T} \log P(x_t \mid x_1, x_2, \ldots, x_{t-1})$$
import torch
import torch.nn as nn
# 预训练损失函数的本质:交叉熵损失
# 模型输出每个位置的 logits,目标是预测下一个 token
loss_fn = nn.CrossEntropyLoss()
# logits: [batch_size, seq_len, vocab_size]
# labels: [batch_size, seq_len],即将输入整体左移一位
def compute_pretrain_loss(logits, input_ids):
# 输入序列向右移动一位作为标签
# 即:给定 token[0..n-1],预测 token[1..n]
shift_logits = logits[:, :-1, :].contiguous() # 去掉最后一个位置的预测
shift_labels = input_ids[:, 1:].contiguous() # 去掉第一个 token 作为标签
# 展平后计算交叉熵
loss = loss_fn(
shift_logits.view(-1, shift_logits.size(-1)),
shift_labels.view(-1)
)
return loss
这个目标函数看似简单,却迫使模型学会了语法、常识、推理、代码——因为这些知识全部隐藏在"预测下一个词"的统计规律中。
1.1.3 预训练数据来源
| 数据集 | 内容 | 规模 | 特点 |
|---|---|---|---|
| Common Crawl | 网页爬取文本 | ~PB 级 | 最大但噪声最多,需过滤 |
| The Pile | 22 种来源混合 | 825 GB | 高质量多样性 |
| Books(BookCorpus、Gutenberg) | 书籍文本 | 数百 GB | 长文依赖、叙事结构 |
| Wikipedia | 百科全书 | ~20 GB | 高质量、结构化知识 |
| GitHub | 代码仓库 | 数百 GB | 代码和注释 |
| ArXiv / PubMed | 学术论文 | 数十 GB | 专业推理能力 |
| StackExchange | 问答对话 | ~30 GB | 问答格式 |
数据质量远比数量重要。Llama 3(Meta开源的大语言模型,是当前最流行的开源LLM系列之一)的预训练数据经过严格的语言识别、去重、质量过滤,最终 15 万亿 Token 中每条都经过多轮筛选。
1.1.4 Scaling Law:算力、数据、模型三要素
OpenAI 和 DeepMind 的研究发现,模型性能与训练算力呈幂律关系。Chinchilla Scaling Law(由 DeepMind 提出的规模最优配比研究)给出了最优配比:
对于给定的计算预算 $C$(FLOPs),最优模型大小 $N$(参数量)和训练 Token 数 $D$ 应等比例扩展,最优比例约为 $D \approx 20N$(每个参数对应约 20 个训练 Token)。
实践结论: 用同样的算力,训练一个更小的模型但使用更多数据,比训练一个大模型但数据不足效果更好。Llama 3 8B 在 15T Token 上训练,性能超过早期在 1T Token 上训练的 70B 模型。
1.2 第二阶段:监督微调(SFT)
1.2.1 为什么不能直接用基础模型:对话格式的缺失
预训练完成后,模型已经拥有了大量知识,但它不会"对话"。给它输入 "法国的首都是哪里?" 它可能会续写 "这道题考察的是……" 而不是直接回答 "巴黎"。
原因在于:预训练语料里有各种各样的文本——百科词条、新闻、论文、小说、问答帖子……它见过"问题 + 答案"的格式,但也见过无数其他格式。它不知道自己现在应该扮演"助手"的角色,给出直接的答案。
SFT 解决的就是这个问题:用"指令 → 回答"格式的数据对,明确教模型在面对指令时应该如何回应。SFT 的任务就是让模型理解"用户提问→助手回答"这种对话格式。
1.2.2 指令数据格式
Alpaca 格式(单轮指令):
{
"instruction": "将下面的句子翻译成英文",
"input": "人工智能正在改变世界",
"output": "Artificial intelligence is changing the world"
}
ShareGPT 格式(多轮对话):
{
"conversations": [
{"from": "human", "value": "Python 中如何读取文件?"},
{"from": "gpt", "value": "可以使用 open() 函数...\n```python\nwith open('file.txt', 'r') as f:\n content = f.read()\n```"},
{"from": "human", "value": "如果文件不存在会怎样?"},
{"from": "gpt", "value": "会抛出 FileNotFoundError 异常,建议用 try-except 处理..."}
]
}
1.2.3 Chat Template:格式转换的桥梁
不同模型使用不同的对话模板将多轮对话序列化为一维 Token 序列。错误的模板会导致性能大幅下降。
from transformers import AutoTokenizer
# 以 Llama 3 为例,展示 chat_template 的作用
tokenizer = AutoTokenizer.from_pretrained("meta-llama/Meta-Llama-3-8B-Instruct")
messages = [
{"role": "system", "content": "你是一个有帮助的 AI 助手。"},
{"role": "user", "content": "什么是梯度下降?"},
{"role": "assistant", "content": "梯度下降是一种优化算法..."},
{"role": "user", "content": "它有什么变体?"}
]
# apply_chat_template 将消息列表转换为模型实际接收的文本格式
# tokenize=False 先看原始文本,了解模型"看到"的是什么
formatted = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True # 在末尾添加 <|start_header_id|>assistant<|end_header_id|>
)
print(formatted)
# 输出类似:
# <|begin_of_text|><|start_header_id|>system<|end_header_id|>
# 你是一个有帮助的 AI 助手。<|eot_id|>
# <|start_header_id|>user<|end_header_id|>
# 什么是梯度下降?<|eot_id|>
# <|start_header_id|>assistant<|end_header_id|>
# 梯度下降是一种优化算法...<|eot_id|>
# <|start_header_id|>user<|end_header_id|>
# 它有什么变体?<|eot_id|>
# <|start_header_id|>assistant<|end_header_id|>
SFT 的训练损失只计算在助手回复部分,用户输入部分的损失被 mask 掉——模型学习的是"如何回答",不是"如何提问"。
1.3 第三阶段:RLHF(基于人类反馈的强化学习)
1.3.1 为什么需要 RLHF:SFT 之后仍然存在的问题
SFT 之后,模型能够对话了。但存在一个新问题:它没有办法判断什么是"好的"回答。
对于同一个问题,"你这个傻瓜自己去查"和"这个问题的答案是……"都是有效的回答,SFT 只教了格式。模型不知道第二个回答比第一个好很多,不知道详细解释比简单否定更有价值,不知道诚实承认不确定比自信给出错误答案更好。
这就是 RLHF 要解决的问题:教模型什么是"好的"。
SFT 后的模型虽然会对话,但对"什么是好的回答"没有判断力。RLHF 通过引入人类偏好信号,让模型学习人类价值观。
1.3.2 RLHF 三步骤
步骤 1:收集人类偏好数据
对同一个问题,模型生成多个回答,人工标注员选择哪个更好:
问题:如何快速减肥?
回答 A:建议通过极端节食,每天只吃 500 卡路里...
回答 B:健康减重需要均衡饮食和规律运动...
标注:B 优于 A
步骤 2:训练奖励模型(Reward Model)
奖励模型是 RLHF 中的关键组件。它的作用是把人类的判断力"固化"成一个可以自动评分的系统,这样才能在训练中自动给模型的回答打分,不需要每次都让人工标注员参与。
奖励模型是一个独立的语言模型,输入(问题+回答),输出一个标量分数。训练目标是让"被偏好的回答"得分高于"被拒绝的回答":
import torch
import torch.nn as nn
class RewardModel(nn.Module):
def __init__(self, base_model):
super().__init__()
self.model = base_model
# 在语言模型顶部加一个线性层,将 hidden_state 映射为标量分数
self.value_head = nn.Linear(base_model.config.hidden_size, 1)
def forward(self, input_ids, attention_mask):
outputs = self.model(input_ids, attention_mask=attention_mask)
# 取最后一个 token 的 hidden state 作为整个序列的表示
last_hidden = outputs.last_hidden_state[:, -1, :]
reward = self.value_head(last_hidden).squeeze(-1)
return reward
# 奖励模型的训练损失:确保 chosen 的分数高于 rejected
def reward_loss(chosen_reward, rejected_reward):
# 使用 log-sigmoid 损失(来自 Bradley-Terry 偏好模型)
return -torch.log(torch.sigmoid(chosen_reward - rejected_reward)).mean()
步骤 3:PPO 强化学习微调
PPO(Proximal Policy Optimization,近端策略优化)的直觉:让语言模型生成回答,奖励模型打分,根据得分更新语言模型参数。关键约束是不能偏离 SFT 模型太远(KL 散度惩罚——KL 散度是衡量两个概率分布差异的指标,用于限制模型每次更新幅度不能太大),否则模型会"奖励黑客"——生成高分但无意义的文本。
1.3.3 对齐税(Alignment Tax)
RLHF 对齐会带来轻微的能力下降,称为"对齐税"(Alignment Tax)。例如,对齐后的模型在某些基准测试(如数学推理)上的分数可能略低于 SFT 模型,但在用户满意度上显著提升。这是安全性与能力之间真实存在的权衡。
1.4 DPO:直接偏好优化
1.4.1 为什么需要更简单的对齐方法
PPO 需要同时维护四个模型(策略模型、参考模型、奖励模型、价值模型),训练过程复杂,超参数敏感,工程实现门槛很高。对于大多数团队来说,PPO 的实现复杂度是一个巨大的障碍。
DPO 通过一个数学洞察简化了整个过程:奖励模型其实可以隐式地用策略模型和参考模型的概率比来表示。这意味着不需要独立训练一个奖励模型,可以直接把偏好数据转化为分类问题,大幅降低实现复杂度。
PPO 的实现复杂,需要同时维护 4 个模型(策略模型、参考模型、奖励模型、价值模型),训练不稳定。2023 年提出的 DPO(Direct Preference Optimization) 绕过了奖励模型,直接从偏好数据优化语言模型。
1.4.2 DPO 偏好数据格式
{
"prompt": "如何处理 Python 中的异常?",
"chosen": "使用 try-except 块是处理异常的标准方式。建议捕获具体异常类型而非裸 except...",
"rejected": "用 try except 就行了,随便写写"
}
1.4.3 为什么 DPO 比 PPO 更稳定
# DPO 损失函数的核心思路(简化版)
import torch
import torch.nn.functional as F
def dpo_loss(policy_chosen_logprob, policy_rejected_logprob,
ref_chosen_logprob, ref_rejected_logprob, beta=0.1):
"""
policy_*: 当前训练的策略模型对 chosen/rejected 的对数概率
ref_*: 固定的参考模型(SFT 模型)对 chosen/rejected 的对数概率
beta: 控制偏离参考模型的程度,越小越保守
DPO 的关键洞见:奖励模型实际上可以用策略模型和参考模型的对数概率比来隐式表示,
从而将 RL 问题转化为一个分类问题。
"""
# 计算隐式奖励差值
pi_log_ratio = policy_chosen_logprob - policy_rejected_logprob
ref_log_ratio = ref_chosen_logprob - ref_rejected_logprob
# 使 chosen 回答的相对概率高于 rejected 回答
loss = -F.logsigmoid(beta * (pi_log_ratio - ref_log_ratio)).mean()
return loss
| 对比维度 | PPO | DPO |
|---|---|---|
| 实现复杂度 | 高(4 个模型) | 低(2 个模型) |
| 训练稳定性 | 较差(超参敏感) | 较好 |
| 计算开销 | 高 | 低约 50% |
| 效果 | RLHF 天花板更高 | 中等规模任务效果相当 |
| 典型应用 | InstructGPT、ChatGPT | Llama 3、Zephyr |
1.5 Constitutional AI(Anthropic 方法)
Anthropic 提出了 Constitutional AI(CAI),通过一组明确的"宪法原则"(如"不应提供危险信息"、"应承认不确定性")来替代人工标注偏好数据。
流程:
- 用 SFT 模型生成回答
- 让模型根据宪法原则批判自己的回答("这个回答是否有害?")
- 让模型基于批判重写回答
- 用重写前后的对比数据训练 RLAIF(AI 反馈强化学习,替代人工反馈)
这使得 Anthropic 能以更少的人工标注成本实现高质量对齐,Claude 系列模型正是基于此方法训练。
1.6 对应用开发者的意义
理解三阶段训练直接影响使用决策:
为什么微调(Fine-tuning)能改变模型行为:微调本质上是在预训练基础上继续执行 SFT,用你的数据覆盖模型原有的"格式偏好"。微调不能增加训练数据中没有的知识(应使用 RAG),但可以改变模型的回答风格、领域倾向、格式偏好。
为什么不同 API 有不同的"安全级别":基础模型(Base)几乎没有安全过滤;指令模型(Instruct)有 SFT 但 RLHF 较少;对齐模型(Claude、ChatGPT)经过完整三阶段训练,拒绝有害请求。
为什么 System Prompt 能显著影响行为:RLHF 训练中大量数据包含 System Prompt,模型学会了遵守 System Prompt 中的指令,这是提示工程(Prompt Engineering)有效的根本原因。
1.7 小结
| 阶段 | 数据类型 | 训练目标 | 结果 |
|---|---|---|---|
| 预训练 | 万亿级原始文本 | Next-Token Prediction | 基础知识与语言能力 |
| SFT | 数万至数百万条指令对话 | 监督学习(只算助手部分损失) | 遵循指令的对话能力 |
| RLHF/DPO | 人类偏好比较数据 | 强化学习 / 隐式偏好分类 | 符合人类价值观的回答 |
没有预训练,模型没有知识基础;没有 SFT,模型不会对话;没有 RLHF/DPO,模型不可信任。三个阶段各司其职,缺哪个都不行。
作为应用开发者,你不需要自己跑这套流程,但知道这三个阶段的存在,能帮你理解模型为什么会有某些固执的偏好,以及微调能改变什么、不能改变什么。