后训练深度解析:LLM 从基座到助手的能力锻造
6788 字
34 分钟
后训练深度解析:LLM 从基座到助手的能力锻造
1. 引言:预训练是学会”说话”,后训练是学会”做事”
1.1 后训练的地位
大语言模型的能力分为两层:
预训练(Pre-training): 学习目标:预测下一个 token 训练数据:互联网规模的文本(trillions of tokens) 学到的能力: → 语法、语义、世界知识 → 基本的推理能力 → 涌现的 in-context learning
问题: → 模型会生成"互联网上最可能的文本" → 但这不一定是"用户想要的回答" → 缺乏指令遵循、偏好对齐
后训练(Post-Training): 学习目标:学会"按照人类意图"行动 训练数据:人类标注的指令-回答对、偏好数据 学到的能力: → 指令遵循(Follow instructions) → 偏好对齐(Helpful, Harmless, Honest) → 特定任务能力(代码、数学、推理) → 对话能力(Chat)1.2 后训练的历史演进
后训练技术演进时间线:
2020 - InstructGPT (OpenAI): → 3 阶段: SFT → RM → PPO → 开创了 RLHF 后训练范式 → 确立了"Helpful, Harmless, Honest" 对齐目标
2022 - FLAN (Google): → 开创了 SFT 为核心的后训练路线 → 证明纯 SFT 也能达到良好对齐
2022 - Constitutional AI (Anthropic): → 用 AI 反馈替代人类反馈 → 降低标注成本
2023 - DPO (Stanford): → 绕过 Reward Model 的简化对齐 → 离线直接优化
2024 - GRPO (DeepSeek): → 去掉 Critic 的高效对齐 → 推理能力的突破
2024 - KTO (Khazron): → 超越DPO的偏好对齐 → 从 Kahneman-Tversky 心理学理论出发 → 只需正负样本,无需配对
2025 - 融合时代: → 多阶段流水线成为标准 → SFT + DPO + GRPO 组合使用 → 领域自适应后训练(医疗/法律/金融)1.3 后训练的核心挑战
后训练的三大核心挑战:
挑战 1: 数据质量 vs 数量 → 数据太少:泛化能力差 → 数据太多:噪声累积、能力遗忘 → 解决:质量过滤 + 课程学习
挑战 2: 能力 vs 对齐的权衡 → 过度对齐 → 过于"安全" → 缺乏创造性 → 对齐不足 → 有害/偏见输出 → 解决:多阶段对齐 + 混合比例
挑战 3: 评估 vs 能力的差距 → Benchmark 不能完全反映用户体验 → Goodhart's Law: 指标被优化后就不再是好的指标 → 解决:多维度评估 + 人类评估1.4 本系列文章关联
| 文章 | 关联 |
|---|---|
| RLHF 深度解析 | RLHF/DPO/GRPO/KTO 详细原理 |
| GRPO 深度解析 | DeepSeek 的高效对齐算法 |
| DeepSeek-R1 | 推理能力后训练的最佳实践 |
| DPO 深度解析 | 直接偏好优化的理论分析 |
| 知识蒸馏 | 后训练数据的蒸馏方法 |
2. 后训练全景图
2.1 完整训练流水线
现代 LLM 的后训练通常采用多阶段流水线:
┌─────────────────────────────────────────────────────────────────┐│ LLM 后训练完整流水线 │├─────────────────────────────────────────────────────────────────┤│ ││ 基座模型 (Pre-trained) ││ │ ││ ▼ ││ 阶段 0: 领域自适应 (可选) ││ ├─ 继续在领域数据上预训练 ││ ├─ 目的: 注入领域知识 (医疗/法律/金融) ││ └─ 风险: 灾难性遗忘 ││ │ ││ ▼ ││ 阶段 1: 监督微调 (SFT) ││ ├─ 核心: 学习指令-回答格式 ││ ├─ 数据: 人工标注 / LLM 生成 ││ ├─ 目标: 最小化交叉熵 ││ └─ 关键: 数据质量 > 数据数量 ││ │ ││ ▼ ││ 阶段 2: 偏好对齐 (Alignment) ││ ├─ 方法 1: RLHF (PPO) ││ ├─ 方法 2: DPO / IPO ││ ├─ 方法 3: GRPO (DeepSeek) ││ ├─ 方法 4: KTO ││ └─ 目标: 学习人类偏好 ││ │ ││ ▼ ││ 阶段 3: 特殊能力强化 (可选) ││ ├─ 数学推理: GRPO + 规则奖励 ││ ├─ 代码生成: SFT + 执行反馈 ││ └─ 安全对齐: 红队测试 + 对抗训练 ││ │ ││ ▼ ││ 助手模型 (Aligned Assistant) ││ │└─────────────────────────────────────────────────────────────────┘2.2 各阶段数据量级
# 各阶段数据量经验值(以 7B 模型为例)POST_TRAINING_DATA_SCALES = { # 领域自适应(继续预训练) "domain_adaptation": { "scale": "1B - 10B tokens", "duration": "1-4 epochs", "lr": "1e-5 to 5e-5", "purpose": "注入领域知识", },
# SFT "sft": { "scale": "10K - 100K samples", "duration": "1-3 epochs", # 注意是 samples 而非 tokens! "lr": "1e-6 to 5e-6", "purpose": "学习指令格式", },
# RLHF/DPO 对齐 "alignment": { "scale": "10K - 100K preference pairs", "duration": "1-2 epochs", "lr": "1e-7 to 1e-6", "purpose": "学习人类偏好", },
# 特殊能力强化 "specialization": { "math": { "scale": "10K - 100K math problems", "method": "GRPO + rule reward", }, "code": { "scale": "10K - 50K coding problems", "method": "SFT + execution feedback", }, },}2.3 关键决策点
def post_training_decision_tree(): """ 后训练策略选择决策树 """
# Q1: 基座模型质量如何? # 基座强 → 可跳过领域自适应 # 基座弱 → 需要领域自适应
# Q2: 对齐目标是什么?
# 目标 1: 通用助手 # → SFT + DPO/GRPO # → 偏好数据: Helpfulness 为主
# 目标 2: 特定领域专家 # → 领域自适应 + SFT + RLHF # → 需要领域专家标注
# 目标 3: 推理能力 # → SFT (CoT 数据) + GRPO + 规则奖励 # → DeepSeek-R1 路线
# Q3: 计算资源?
# 资源充足: # → Full RLHF (PPO) # → 可负担更大的模型和更多的训练
# 资源有限: # → DPO / GRPO # → 离线对齐,无需价值网络3. 监督微调(SFT)
3.1 SFT 的本质
SFT(Supervised Fine-Tuning)的核心是让模型学习”如何回答问题”。
def sft_loss(model, batch): """ SFT 损失函数
核心:标准的语言模型交叉熵损失 L = -Σ y_i · log P(y_i | x, y_{<i})
唯一不同于预训练的地方: → 只在 response 上计算 loss → prompt 的 tokens 不参与 loss 计算 """ prompt_ids = batch["prompt_ids"] # (B, L_p) response_ids = batch["response_ids"] # (B, L_r)
# 拼接 input_ids = torch.cat([prompt_ids, response_ids], dim=1)
# 前向传播 outputs = model(input_ids) logits = outputs.logits # (B, L, V)
# 只在 response 部分计算 loss # 位置: 从 len(prompt_ids) 开始,到末尾 response_logits = logits[:, prompt_ids.size(1):, :]
# 目标: response_ids 右移一位 labels = response_ids
# 交叉熵损失 loss = F.cross_entropy( response_logits.reshape(-1, response_logits.size(-1)), labels.reshape(-1), reduction='mean' )
return loss
# ⚠️ SFT 的常见错误:把 prompt 也算进 lossdef sft_loss_WRONG(model, batch): """ 错误示范:在 prompt 上也计算 loss
这会导致: → 模型过度记忆 prompt 的格式 → 对不同 prompt 的泛化能力下降 → 训练效率降低(prompt 不需要学习) """ input_ids = batch["input_ids"]
outputs = model(input_ids) logits = outputs.logits
# ❌ 错误:整个序列都算 loss loss = F.cross_entropy( logits[:, :-1].reshape(-1, logits.size(-1)), input_ids[:, 1:].reshape(-1), )
return loss3.2 SFT 数据格式
SFT 数据的格式设计直接影响模型学习效果:
# 主流格式 1: ChatML (Mistral, Qwen)CHATML_FORMAT = """<|im_start|>system{system_prompt}<|im_end|><|im_start|>user{user_message}<|im_end|><|im_start|>assistant{assistant_response}<|im_end|>"""
# 主流格式 2: Llama 3 (Meta)LLAMA3_FORMAT = """<|begin_of_text|><|start_header_id|>system<|end_header_id|>
{system_prompt}<|eot_id|><|start_header_id|>user<|end_header_id|>
{user_message}<|eot_id|><|start_header_id|>assistant<|end_header_id|>
{assistant_response}<|eot_id|>"""
# 主流格式 3: DeepSeek (无特殊标记,用 special tokens)DEEPSEEK_FORMAT = """{special_token}你是 DeepSeek Chat<|special_token|>
Human: {user_message}
Assistant: {assistant_response}"""
def format_sft_sample(sample, format_type="chatml"): """ 格式化 SFT 样本 """ if format_type == "chatml": return CHATML_FORMAT.format(**sample) elif format_type == "llama3": return LLAMA3_FORMAT.format(**sample) elif format_type == "deepseek": return DEEPSEEK_FORMAT.format(**sample) else: raise ValueError(f"Unknown format: {format_type}")
class SFTDataset(Dataset): """ SFT 数据集 """ def __init__(self, data, tokenizer, max_length=4096): self.data = data self.tokenizer = tokenizer self.max_length = max_length
def __getitem__(self, idx): sample = self.data[idx]
# 格式化为对话 text = format_sft_sample(sample)
# Tokenize encoding = self.tokenizer( text, max_length=self.max_length, truncation=True, padding="max_length", return_tensors="pt", )
return { "input_ids": encoding["input_ids"].squeeze(0), "attention_mask": encoding["attention_mask"].squeeze(0), "labels": encoding["input_ids"].squeeze(0).clone(), }
def collate_fn(self, batch): """ ⚠️ 关键:在 collate 中 mask 掉 prompt 的 loss """ # 这需要在格式化时就记录 prompt 长度 # 或者用特殊 token 的 id 来识别 for i, item in enumerate(batch): prompt_len = self.get_prompt_length(item) # prompt 部分的 label 设为 -100(忽略) item["labels"][:prompt_len] = -100
return self.default_collate(batch)3.3 SFT 数据构建
class SFTDataBuilder: """ SFT 数据构建器 """
def __init__(self, base_model, tokenizer): self.base_model = base_model self.tokenizer = tokenizer
def build_from_human_annotation(self, annotations): """ 方法 1: 人类标注 最可靠,但成本最高 """ return [self.validate_sample(a) for a in annotations]
def build_from_llm_generation(self, prompts, strategy="self"): """ 方法 2: LLM 生成
策略: - self: 用自己生成 (Llama-Factory 风格) - teacher: 用大模型生成 (蒸馏风格) - mix: 混合 """ if strategy == "self": return self.self_generate(prompts) elif strategy == "teacher": return self.teacher_generate(prompts) elif strategy == "mix": return self.mix_generate(prompts)
def self_generate(self, prompts): """ Self-Generation: 用当前模型生成回答 (Llama-Factory 的做法) """ data = []
for prompt in prompts: # 生成回答 response = self.base_model.generate( prompt, max_tokens=2048, temperature=0.8, do_sample=True, )
# 质量过滤 if self.quality_check(prompt, response): data.append({ "prompt": prompt, "response": response, "source": "self_generated", })
return data
def quality_check(self, prompt, response): """ 质量检查 """ # 1. 长度检查 if len(response.split()) < 10: return False if len(response.split()) > 3000: return False
# 2. 格式检查 if response.startswith("I'm sorry") and len(response) < 50: return False
# 3. 困惑度检查 ppl = compute_perplexity(response, self.base_model) if ppl > 100: # 困惑度太高说明是乱码 return False
return True
def build_from_feedback(self, prompts, feedback_fn): """ 方法 3: 带反馈的迭代生成
典型应用: 用 reward model 筛选 """ data = []
for prompt in prompts: # 生成多个候选 candidates = [ self.base_model.generate(prompt, temperature=t) for t in [0.5, 0.7, 0.9, 1.1] ]
# 用反馈函数评分 scores = [feedback_fn(prompt, c) for c in candidates] best_idx = scores.index(max(scores))
if scores[best_idx] > 0.8: data.append({ "prompt": prompt, "response": candidates[best_idx], "score": scores[best_idx], })
return data3.4 SFT 最佳实践
# SFT 超参数推荐值SFT_HYPERPARAMS = { # 学习率 "learning_rate": { "7B_model": 1e-5, # ~10x 小于预训练 "13B_model": 5e-6, "70B_model": 2e-6, "note": "通常用余弦衰减 + warmup", },
# Epochs "num_epochs": { "high_quality_data": 1, # 高质量数据,1 epoch 足够 "medium_quality_data": 2, # 中等质量,2 epochs "low_quality_data": 3, # 低质量数据需要更多 epochs "warning": "太多 epochs 会导致过拟合和遗忘", },
# Batch size "batch_size": { "effective_batch_size": 128, # 梯度累积后的有效 batch "per_device": 4, # 单卡 batch "gradient_accumulation_steps": 32, },
# 其他 "weight_decay": 0.01, "warmup_ratio": 0.03, # 3% 的 steps 用于 warmup "lr_scheduler": "cosine", "max_grad_norm": 1.0,}SFT 的黄金法则
数据质量 > 数据数量 > 模型大小 > 训练时长
10K 条高质量人工标注数据的训练效果 > 100K 条低质量自动生成数据
4. 偏好对齐:RLHF/DPO/GRPO/KTO
4.1 对齐的核心问题
偏好对齐要解决的是:如何让模型学习”什么样的回答是好的”。
SFT vs 对齐的本质区别:
SFT: 学习目标: P(回答|问题) 训练数据: (问题, 最佳回答) 学习内容: 给定问题,应该回答什么
对齐: 学习目标: P(偏好回答|问题, 回答A, 回答B) 训练数据: (问题, 回答A, 回答B, 偏好) 学习内容: 给定问题,哪个回答更好
为什么对齐必要? → 同一个问题可能有多个"正确"回答 → SFT 只学了其中一个 → 对齐让模型学会区分回答的优劣 → 提升泛化能力(即使没见过的 prompt 也知道什么是好回答)4.2 RLHF(人类反馈强化学习)
RLHF 是后训练对齐的经典方法:
class RLHFPipeline: """ RLHF 三阶段流水线
阶段 1: SFT 阶段 2: Reward Model 阶段 3: PPO """
def __init__(self, sft_model, reference_model): self.sft_model = sft_model self.reference_model = reference_model self.reward_model = None self.value_model = None
def stage2_train_reward_model(self, preference_data): """ 阶段 2: 训练 Reward Model
数据格式: (prompt, chosen_response, rejected_response) 目标: 学习 ELO/Bradley-Terry 偏好模型 """ # 从 SFT 模型初始化(比随机初始化好) self.reward_model = copy.deepcopy(self.sft_model) # 替换输出层为 1 维(reward 分数) self.reward_model.replace_head(1)
# 训练损失 loss = self.preference_loss(preference_data)
return self.reward_model
def preference_loss(self, batch): """ Reward Model 损失函数
Bradley-Terry 模型: P(prefer A) = σ(r_A - r_B)
损失: -log P(prefer chosen) """ chosen_rewards = self.reward_model(batch["prompt"], batch["chosen"]) rejected_rewards = self.reward_model(batch["prompt"], batch["rejected"])
# 偏好概率 probs = torch.sigmoid(chosen_rewards - rejected_rewards)
# 损失 = -log P(prefer chosen) loss = -torch.log(probs + 1e-8).mean()
return loss
def stage3_ppo(self, prompts, rm, value_model): """ 阶段 3: PPO 对齐
目标函数: L = E[r_θ(x,y)] - β · KL(π_θ || π_ref)
其中 r_θ 是 learned reward function """ self.value_model = value_model
for step in range(num_ppo_steps): # 1. Rollout: 用当前策略生成回答 rollouts = self.rollout(prompts)
# 2. 计算 reward rewards = rm.score(prompts, rollouts)
# 3. 计算 advantage (GAE) advantages = self.compute_advantages(rewards)
# 4. PPO 更新 self.ppo_update(rollouts, advantages)
return self.policy_model4.3 DPO(直接偏好优化)
DPO 绕过 Reward Model,直接优化偏好:
class DPO: """ Direct Preference Optimization (Rafailov et al., 2023)
核心思想: 绕过 Reward Model,直接用偏好数据优化 Policy
DPO 损失函数: L = -E_{(x, y_w, y_l) ~ D}[log σ( β · log(π_θ(y_w|x) / π_ref(y_w|x)) - β · log(π_θ(y_l|x) / π_ref(y_l|x)) )]
等价于同时优化: 1. 最大化偏好回答的概率 2. 最小化拒绝回答的概率 3. 隐式地保持 KL 约束 """
def __init__(self, policy_model, reference_model, beta=0.1): self.policy = policy_model self.reference = reference_model self.beta = beta # KL 系数
def dpo_loss(self, batch): """ DPO 损失 """ prompt = batch["prompt"] # (B,) chosen = batch["chosen_response"] # (B,) rejected = batch["rejected_response"] # (B,)
# 计算 log probs # π_θ(y|x) - π_ref(y|x) chosen_logps = self.get_log_probs(self.policy, prompt, chosen) rejected_logps = self.get_log_probs(self.policy, prompt, rejected)
ref_chosen_logps = self.get_log_probs(self.reference, prompt, chosen) ref_rejected_logps = self.get_log_probs(self.reference, prompt, rejected)
# 优势 chosen_rewards = self.beta * (chosen_logps - ref_chosen_logps) rejected_rewards = self.beta * (rejected_logps - ref_rejected_logps)
# DPO 损失 # L = -log σ(chosen_rewards - rejected_rewards) loss = -F.logsigmoid(chosen_rewards - rejected_rewards).mean()
# 可选:加入 SFT 损失防止灾难性遗忘 sft_loss = -chosen_logps.mean()
return loss + 0.1 * sft_loss
def get_log_probs(self, model, prompts, responses): """ 计算 log P(y|x) """ # 拼接 prompt 和 response inputs = tokenize_concat(prompts, responses, tokenizer)
outputs = model(**inputs) logits = outputs.logits
# 提取 response 部分的 log prob log_probs = gather_log_probs(logits, inputs, responses)
return log_probs4.4 GRPO(组相对策略优化)
DeepSeek 提出的高效对齐方法:
class GRPO: """ Group Relative Policy Optimization (DeepSeek)
核心思想: 用组内相对排名替代 Value 网络 大幅降低显存和计算开销
GRPO 优势估计: A_i = (r_i - μ) / σ
其中 r_i 是第 i 个回答的 reward μ, σ 是组内均值和标准差 """
def __init__(self, policy, reference, reward_fn, group_size=8): self.policy = policy self.reference = reference self.reward_fn = reward_fn self.group_size = group_size
def grpo_loss(self, prompts): """ GRPO 单步 """ # 1. 采样 G 个回答 responses, old_logps = self.rollout(prompts, self.group_size)
# 2. 评分 rewards = [self.reward_fn(p, r) for p, r in zip(prompts, responses)]
# 3. 计算组内归一化优势 advantages = self.compute_advantages(rewards, self.group_size)
# 4. 计算 KL 散度(与 reference) ref_logps = self.reference_score(prompts, responses) kl = (old_logps - ref_logps).mean()
# 5. PPO-style clip ratio = torch.exp(old_logps - old_logps.detach()) clipped_ratio = torch.clamp(ratio, 1 - 0.2, 1 + 0.2)
policy_loss = -torch.min(ratio * advantages, clipped_ratio * advantages).mean()
# 6. 总损失 loss = policy_loss + 0.04 * kl
return loss
def compute_advantages(self, rewards, group_size): """ 组内归一化优势 """ rewards = rewards.view(-1, group_size) # (num_prompts, G)
mean_r = rewards.mean(dim=-1, keepdim=True) std_r = rewards.std(dim=-1, keepdim=True) + 1e-8
advantages = (rewards - mean_r) / std_r
return advantages.flatten()4.5 KTO(Kahneman-Tversky 优化)
从心理学理论出发的对齐方法:
class KTO: """ Kahneman-Tversky Optimization (Ethayarajh et al., 2024)
核心思想: DPO 需要偏好对 (chosen, rejected) KTO 只需要正负样本 (desirable, undesirable) 更易于收集数据
损失函数基于前景理论: L = -log σ(β · (v_desirable - v_undesirable))
其中 v(y) = log P(y|x) - β · log P_ref(y|x) """
def __init__(self, policy, reference, beta=0.1): self.policy = policy self.reference = reference self.beta = beta
def kto_loss(self, batch): """ KTO 损失 """ desirable_responses = batch["desirable"] # 正样本 undesirable_responses = batch["undesirable"] # 负样本 prompts = batch["prompt"]
# 计算 utility v_desirable = self.utility(prompts, desirable_responses) v_undesirable = self.utility(prompts, undesirable_responses)
# KTO 损失 # 目标: v_desirable > v_undesirable loss = -F.logsigmoid(self.beta * (v_desirable - v_undesirable)).mean()
return loss
def utility(self, prompts, responses): """ Utility 函数
v(y) = log P(y|x) - β · log P_ref(y|x)
第一项: 鼓励高概率回答 第二项: KL 正则,防止偏离太远 """ logps = self.policy.get_log_probs(prompts, responses) ref_logps = self.reference.get_log_probs(prompts, responses)
return logps - self.beta * ref_logps4.6 对齐方法对比
| 维度 | PPO (RLHF) | DPO | GRPO | KTO |
|---|---|---|---|---|
| 需要 RM | ✅ | ❌ | ✅/❌* | ❌ |
| 需要 Critic | ✅ | ❌ | ❌ | ❌ |
| 数据格式 | (prompt, reward) | (prompt, y_w, y_l) | (prompt, r) | (prompt, y) |
| 显存占用 | ~4× | ~2× | ~3× | ~2× |
| 训练稳定性 | 中 | 高 | 高 | 高 |
| 偏好一致性 | 高 | 中 | 高 | 高 |
| 计算效率 | 低 | 高 | 中 | 高 |
| 超参数敏感性 | 高 | 中 | 中 | 低 |
*GRPO 可用规则奖励替代 RM
5. 多阶段训练流水线
5.1 标准流水线:SFT + 对齐
def standard_post_training_pipeline( base_model, sft_data, preference_data, reward_model_data=None,): """ 标准后训练流水线
流程: 1. SFT: 学习指令格式 2. (可选) Reward Model: 如果用 RLHF/GRPO 3. 对齐: DPO/GRPO/PPO """
# === 阶段 1: SFT === print("Stage 1: Supervised Fine-Tuning") sft_model = base_model.clone() sft_model = train_sft(sft_model, sft_data, epochs=2, lr=1e-5)
# === 阶段 2: Reward Model (如果需要) === if preference_data and not use_dpo: print("Stage 2: Training Reward Model") reward_model = train_reward_model(sft_model, preference_data) else: reward_model = None
# === 阶段 3: 对齐 === print("Stage 3: Alignment") if method == "dpo": aligned_model = train_dpo(sft_model, preference_data) elif method == "grpo": aligned_model = train_grpo(sft_model, reward_model, prompts) elif method == "kto": aligned_model = train_kto(sft_model, preference_data)
return aligned_model5.2 领域自适应流水线
def domain_adaptation_pipeline( base_model, domain_data, general_data, preference_data,): """ 领域自适应后训练
典型场景: 医疗、法律、金融等专业领域 """
# === 阶段 0: 领域继续预训练 === print("Stage 0: Domain Continuation Pre-training") domain_model = continue_pretrain( base_model, domain_data, # 领域文本(无标签) lr=5e-5, epochs=1, )
# === 阶段 1: 领域 SFT === print("Stage 1: Domain SFT") domain_sft_data = build_domain_sft_data(domain_data) domain_model = train_sft( domain_model, domain_sft_data, epochs=2, lr=1e-5, )
# === 阶段 2: 混合对齐 === print("Stage 2: Hybrid Alignment") # 混合领域偏好和通用偏好 mixed_preference = mix_preference_data( domain_preference_data, general_preference_data, ratio=0.3, # 30% 通用,保持通用能力 ) aligned_model = train_dpo(domain_model, mixed_preference)
return aligned_model5.3 推理能力流水线(DeepSeek-R1 风格)
def reasoning_post_training_pipeline( base_model, math_problems, code_problems,): """ 推理能力后训练流水线 典型案例: DeepSeek-R1 """
# === 阶段 1: 冷启动 SFT === print("Stage 1: Cold Start SFT") # 加入少量高质量 CoT 数据 cold_start_data = build_cold_start_data(math_problems, with_reasoning=True) model = train_sft(base_model, cold_start_data, epochs=1)
# === 阶段 2: GRPO for Reasoning === print("Stage 2: GRPO with Rule Rewards") model = grpo_training( model, math_problems, reward_fn=rule_based_math_reward, # 规则奖励 group_size=16, beta=0.04, entropy_coef=0.02, )
# === 阶段 3: RFT (Reinforced Fine-Tuning) === print("Stage 3: Reinforced Fine-Tuning") # 用 GRPO 生成的数据做 SFT reasoning_data = generate_reasoning_data(model, math_problems) filtered_data = filter_by_reward(reasoning_data, threshold=0.9) model = train_sft(model, filtered_data, epochs=2)
# === 阶段 4: 拒绝采样 + 扩展 === print("Stage 4: Rejection Sampling + Expansion") all_data = concat(math_problems, code_problems, general_tasks) expanded_data = rejection_sampling(model, all_data) model = train_sft(model, expanded_data, epochs=1)
return model5.4 安全对齐流水线
def safety_alignment_pipeline( base_model, helpful_data, harmful_data,): """ 安全对齐流水线 目标: Helpfulness + Harmlessness """
# === 阶段 1: Helpfulness SFT === print("Stage 1: Helpfulness SFT") helpful_model = train_sft(base_model, helpful_data)
# === 阶段 2: Safety Alignment === print("Stage 2: Safety Alignment") # 方法: DPO + 特殊损失
def safety_preference_fn(prompt, response): """ 安全偏好数据构建 """ is_safe = evaluate_safety(response) is_helpful = evaluate_helpfulness(response)
if is_safe and is_helpful: return ("helpful", "safe") # chosen=helpful, rejected=safe elif is_safe and not is_helpful: return ("safe", "helpful") elif not is_safe: # 有害回答必须被拒绝 return None
safety_preference = build_safety_preference(harmful_data, safety_preference_fn) safe_model = train_dpo(helpful_model, safety_preference)
# === 阶段 3: 红队测试 === print("Stage 3: Red Teaming") adversarial_prompts = generate_adversarial_prompts() safe_model = adversarial_training(safe_model, adversarial_prompts)
return safe_model6. 数据工程
6.1 数据质量过滤
class DataQualityFilter: """ SFT 数据质量过滤器 """
def __init__(self, judge_model=None): self.judge = judge_model
def filter(self, dataset): """ 多维度质量过滤 """ filtered = []
for sample in tqdm(dataset): scores = self.score_sample(sample)
if self.passes_threshold(scores): sample["quality_score"] = scores["overall"] filtered.append(sample)
return filtered
def score_sample(self, sample): """ 多维度评分 """ scores = {}
# 1. 格式正确性 scores["format"] = self.check_format(sample)
# 2. 长度合理性 scores["length"] = self.check_length(sample)
# 3. 语义相关性 scores["relevance"] = self.check_relevance(sample)
# 4. 毒性/有害内容 scores["safety"] = self.check_safety(sample)
# 5. (可选) LLM 评分 if self.judge: scores["llm_judge"] = self.llm_judge(sample)
# 综合评分 scores["overall"] = self.weighted_sum(scores)
return scores
def llm_judge(self, sample): """ 用 LLM 评判质量 """ prompt = f""" 请评估以下对话的质量,评分 1-5:
问题: {sample['prompt']} 回答: {sample['response']}
评分标准: 1. 准确性:回答是否正确 2. 完整性:回答是否全面 3. 清晰度:回答是否易于理解 4. 有用性:回答是否有帮助
评分: [1-5] """
response = self.judge.generate(prompt) score = extract_score(response)
return score
def passes_threshold(self, scores): """ 判断是否通过阈值 """ return ( scores["format"] > 0.8 and scores["length"] > 0.5 and scores["relevance"] > 0.7 and scores["safety"] > 0.95 )6.2 数据去重
class DataDeduplicator: """ 数据去重 """
def deduplicate_by_exact_match(self, dataset): """ 精确匹配去重 """ seen = set() unique = []
for sample in dataset: key = self.make_key(sample)
if key not in seen: seen.add(key) unique.append(sample)
return unique
def deduplicate_by_similarity(self, dataset, threshold=0.9): """ 语义相似去重 """ unique = [] embeddings = self.get_embeddings(dataset)
for i, sample in enumerate(dataset): is_duplicate = False
for existing in unique: j = len(unique) sim = cosine_similarity(embeddings[i], embeddings[j])
if sim > threshold: is_duplicate = True break
if not is_duplicate: unique.append(sample)
return unique6.3 课程学习(Curriculum Learning)
class CurriculumScheduler: """ 课程学习调度器 在训练过程中从易到难安排数据 """
def __init__(self, difficulty_fn): self.get_difficulty = difficulty_fn
def create_curriculum(self, dataset, num_stages=3): """ 创建课程
将数据按难度分为多个阶段 训练时从易到难逐步引入 """ # 按难度排序 scored = [(self.get_difficulty(s), s) for s in dataset] scored.sort(key=lambda x: x[0])
# 划分阶段 stage_size = len(scored) // num_stages curriculum = []
for i in range(num_stages): start = i * stage_size end = (i + 1) * stage_size if i < num_stages - 1 else len(scored)
curriculum.append({ "stage": i, "data": [s for _, s in scored[start:end]], "difficulty_range": (scored[start][0], scored[end-1][0]), })
return curriculum
def get_training_data(self, curriculum, current_step, total_steps): """ 根据当前步数返回训练数据
早期: 只用简单数据 后期: 逐步引入复杂数据 """ # 计算当前处于哪个阶段 progress = current_step / total_steps num_stages = len(curriculum) current_stage = min(int(progress * num_stages), num_stages - 1)
# 收集前面所有阶段的数据 training_data = [] for i in range(current_stage + 1): training_data.extend(curriculum[i]["data"])
return training_data
def difficulty_fn_math(self, sample): """ 数学问题难度评估 """ problem = sample["prompt"] response = sample["response"]
# 1. 问题复杂度 num_steps = response.count("Step") complexity_score = min(num_steps / 10, 1.0)
# 2. 答案正确性(如果有 ground truth) if "ground_truth" in sample: correct = sample["response"].endswith(sample["ground_truth"]) correctness_score = 1.0 if correct else 0.0 else: correctness_score = 0.5
# 3. 公式密度 formula_density = response.count("$") / len(response.split())
# 综合难度 difficulty = ( 0.5 * complexity_score + 0.3 * (1 - correctness_score) + # 正确答案反而简单 0.2 * formula_density )
return difficulty7. 超参数调优
7.1 关键超参数
| 阶段 | 超参数 | 推荐范围 | 说明 |
|---|---|---|---|
| SFT | 学习率 | 1e-6 ~ 2e-5 | 模型越大越小 |
| SFT | Epochs | 1 ~ 3 | 高质量数据 1-2 轮 |
| SFT | Batch Size | 4 ~ 16(per device) | 梯度累积到 128-256 |
| SFT | Warmup Ratio | 0.01 ~ 0.05 | 3% warmup 最常见 |
| 对齐 | KL 系数 β | 0.01 ~ 0.3 | DPO: 0.1, GRPO: 0.04 |
| 对齐 | Clip Ratio ε | 0.1 ~ 0.3 | 默认 0.2 |
| 对齐 | 学习率 | 5e-7 ~ 2e-6 | 对齐比 SFT 更小 |
| 对齐 | Entropy Coef | 0 ~ 0.02 | 防止策略坍缩 |
7.2 学习率调度
class LRScheduler: """ 学习率调度策略 """
def get_cosine_schedule_with_warmup( self, optimizer, num_warmup_steps, num_training_steps, min_lr_ratio=0.1, ): """ 余弦退火 + Warmup
学习率曲线:
lr ↑ ┌────── │ /│ ╲ │ / │ ╲ │ / │ ╲ │/ │ ╲ └────┴──────────┴→ steps warmup cosine decay """ def lr_lambda(current_step): if current_step < num_warmup_steps: # Linear warmup return float(current_step) / float(max(1, num_warmup_steps))
# Cosine decay progress = float(current_step - num_warmup_steps) / float( max(1, num_training_steps - num_warmup_steps) )
return max(min_lr_ratio, 0.5 * (1.0 + math.cos(math.pi * progress)))
return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)7.3 超参敏感度分析
┌──────────────────────────────────────────────────────────┐│ 后训练超参敏感度 │├──────────────────────────────────────────────────────────┤│ ││ 高度敏感 (需精细调优): ││ ├─ 学习率(±50% 性能差异) ││ ├─ KL 系数 β(±30% 差异) ││ └─ 熵系数(可能导致坍缩) ││ ││ 中度敏感 (有默认值可用): ││ ├─ Batch Size(影响收敛速度,不影响最终性能) ││ ├─ Warmup Ratio(影响初期稳定性) ││ └─ Epochs(过拟合/欠拟合) ││ ││ 低敏感 (影响小): ││ ├─ Weight Decay(默认 0.01 即可) ││ ├─ Dropout(微调时通常不用) ││ └─ Attention Dropout(通常 0) ││ │└──────────────────────────────────────────────────────────┘8. 评估体系
8.1 评估维度
class AlignmentEvaluator: """ 对齐效果评估 """
EVALUATION_DIMENSIONS = { # 1. 任务能力 "task_capability": { "math": ["MATH", "GSM8K", "AIME"], "code": ["HumanEval", "MBPP", "LiveCodeBench"], "reasoning": ["BBH", "ARC-Challenge"], "knowledge": ["MMLU", "TriviaQA"], },
# 2. 对齐质量 "alignment_quality": { "helpfulness": ["MT-Bench", "AlpacaEval"], "safety": ["ToxiGen", "HarmBench"], "honesty": ["TruthfulQA"], },
# 3. 格式遵循 "format_adherence": { "length": "输出长度是否合适", "structure": "是否遵循指定格式", "completeness": "是否完整回答问题", },
# 4. 鲁棒性 "robustness": { "adversarial": "对抗样本攻击成功率", "jailbreak": "越狱尝试成功率", "persistence": "多次询问是否保持一致", }, }
def evaluate(self, model, test_sets): """ 综合评估 """ results = {}
for dimension, benchmarks in self.EVALUATION_DIMENSIONS.items(): results[dimension] = {}
if dimension == "task_capability": for task, task_benchmarks in benchmarks.items(): results[dimension][task] = self.evaluate_task_capability( model, task_benchmarks )
elif dimension == "alignment_quality": for aspect, aspect_benchmarks in benchmarks.items(): results[dimension][aspect] = self.evaluate_alignment( model, aspect_benchmarks )
return results8.2 评估方法
def benchmark_evaluation(model, benchmark_name): """ Benchmark 评估 """ if benchmark_name == "MMLU": return evaluate_mmmlu(model) elif benchmark_name == "HumanEval": return evaluate_humaneval(model, execution=True) elif benchmark_name == "TruthfulQA": return evaluate_truthfulqa(model) # ...
def human_evaluation(samples, num_humans=10): """ 人类评估(黄金标准) """ criteria = [ "Helpfulness: 回答是否有帮助", "Accuracy: 回答是否正确", "Completeness: 回答是否完整", "Safety: 回答是否安全", "Coherence: 回答是否连贯", ]
ratings = {c: [] for c in criteria}
for human in range(num_humans): for sample in samples: for criterion in criteria: rating = human_rate(sample, criterion) ratings[criterion].append(rating)
# 汇总 results = {c: np.mean(ratings[c]) for c in criteria}
return results
def llm_as_judge(prompts, responses): """ LLM 评判(可扩展的人类评估替代) """ judge_prompt = """ 请评估以下回答的质量,评分 1-10。
问题: {prompt} 回答: {response}
评分标准: - 准确性 (1-10) - 有帮助性 (1-10) - 完整性 (1-10)
请以 JSON 格式输出评分。 """
scores = [] for prompt, response in zip(prompts, responses): formatted_prompt = judge_prompt.format(prompt=prompt, response=response) score_json = llm_judge.generate(formatted_prompt) scores.append(json.loads(score_json))
return scores9. 工程实践
9.1 分布式训练
# DeepSpeed ZeRO 配置示例DEEPSPEED_CONFIG = { "train_batch_size": "auto", "train_micro_batch_size_per_gpu": "auto", "gradient_accumulation_steps": "auto", "gradient_clipping": 1.0, "zero_optimization": { "stage": 2, # ZeRO-2: 优化器分片 "offload_optimizer": { "device": "cpu", # CPU offload 节省显存 }, "stage3_max_reuse_distance": 10000, "stage3_param_persistence_threshold": 1e5, }, "fp16": { "enabled": "auto", "loss_scale_window": 100, }, "bf16": { "enabled": True, # BF16 比 FP16 更稳定 }, "gradient_checkpointing": { "enabled": True, # 节省 60% 显存 "checkpoint_ratio": 0.5, },}
# FSDP 配置示例FSDP_CONFIG = { "sharding_strategy": "FULL_SHARD", # 梯度 + 模型分片 "cpu_offload": True, # CPU offload "backward_prefetch": "backward_pre", "mixed_precision": { "param_dtype": "bfloat16", "reduce_dtype": "float32", "buffer_dtype": "bfloat16", }, "activation_checkpointing": { "num_checkpoints": 4, },}9.2 常见问题与解决
┌────────────────────────────────┬──────────────────────────────────────────┐│ 问题 │ 解决 │├────────────────────────────────┼──────────────────────────────────────────┤│ SFT 过拟合 │ 减少 epochs、增加数据量、降低学习率 ││ 对齐后通用能力下降 │ 增加通用数据比例、减少对齐轮数 ││ Reward Hacking │ 混合奖励、多样性正则、人类评估 ││ 策略坍缩(熵降为 0) │ 加入熵正则、增大 KL 系数 ││ 回答过于冗长/简短 │ 长度奖励/惩罚、过滤异常长度数据 ││ 对特定 prompt 过拟合 │ 数据增强、prompt 多样化 ││ 多轮对话能力下降 │ 加入多轮数据、contex window 扩展 ││ 领域能力与通用能力冲突 │ 课程学习、两阶段对齐 │└────────────────────────────────┴──────────────────────────────────────────┘10. 主流模型后训练实践
10.1 DeepSeek-R1
DEEPSEEK_R1_POST_TRAINING = { "base_model": "DeepSeek-V3", "stages": [ { "name": "Cold Start SFT", "data_size": "5,000 samples", "method": "SFT", "purpose": "格式引导", }, { "name": "GRPO for Math", "data_size": "200K math problems", "method": "GRPO", "reward": "rule-based (accuracy)", "group_size": 16, "entropy_coef": 0.02, }, { "name": "RFT", "data_size": "600K samples", "method": "SFT", "purpose": "压缩推理能力", }, { "name": "Rejection Sampling", "data_size": "800K samples", "method": "filter + SFT", }, { "name": "GRPO for All", "data_size": "mixed", "method": "GRPO", "reward": "mixed (rule + RM + safety)", }, ],}10.2 Llama 3
LLAMA3_POST_TRAINING = { "base_model": "Llama 3 Base", "stages": [ { "name": "Quality Filtering", "data_size": "10M samples → 2M samples", "method": "LLM-based filtering", }, { "name": "SFT", "data_size": "9M samples", "epochs": 2, "lr": "2e-5", }, { "name": "DPO Alignment", "data_size": "1M preference pairs", "beta": 0.1, }, ], "key_insight": "数据质量过滤是性能的关键",}10.3 GPT-4o
GPT4O_POST_TRAINING = { "stages": [ "SFT on curated data", "PPO RLHF", "HPOT (Hybrid Preference Optimization)", "Red teaming + Safety alignment", ], "notes": [ "闭源模型,细节未公开", "推测使用多阶段 + 大量人类评估", "多模态能力是训练重点", ],}11. 总结
11.1 后训练的核心原则
┌─────────────────────────────────────────────────────────────┐│ 后训练的五大核心原则 │├─────────────────────────────────────────────────────────────┤│ ││ 1. 数据质量 > 数据数量 ││ → 10K 高质量 > 100K 低质量 ││ → 质量过滤是投入产出比最高的工作 ││ ││ 2. 多阶段优于单阶段 ││ → SFT → 对齐 → 特殊能力 逐步推进 ││ → 每阶段专注一个目标 ││ ││ 3. 对齐是双刃剑 ││ → 过度对齐 → 过于保守,缺乏创造性 ││ → 对齐不足 → 有害/偏见输出 ││ → 平衡点需要持续探索 ││ ││ 4. 评估驱动开发 ││ → Benchmark 是必要但不充分的 ││ → 必须有人类评估作为最终标准 ││ ││ 5. 领域适应需要谨慎 ││ → 领域自适应可能导致通用能力遗忘 ││ → 混合数据 + 两阶段对齐是推荐方案 ││ │└─────────────────────────────────────────────────────────────┘11.2 方法选择指南
后训练方法选择决策树:
目标是什么? │ ├── 通用助手(Helpful & Harmless) │ │ │ └── SFT + DPO (或 GRPO) │ ├── 特定领域专家 │ │ │ ├── 有领域数据 │ │ └── 领域自适应 + SFT + DPO │ │ │ └── 无领域数据 │ └── SFT + 对齐 │ ├── 推理能力 │ │ │ └── DeepSeek-R1 路线: │ SFT(CoT) → GRPO(规则奖励) → RFT │ └── Code 模型 │ └── SFT(exec feedback) + GRPO(rule reward)推荐阅读
- InstructGPT(Ouyang et al., 2022)—— RLHF 开山之作
- FLAN(Wei et al., 2022)—— SFT 为核心的对齐
- Constitutional AI(Bai et al., 2022)—— AI 反馈对齐
- DPO(Rafailov et al., 2023)—— 直接偏好优化
- GRPO(DeepSeek, 2024)—— 高效对齐算法
- KTO(Ethayarajh et al., 2024)—— 心理学对齐
- DeepSeek-R1(DeepSeek, 2025)—— 推理后训练最佳实践
参考资料
- Ouyang, L., et al. (2022). “Training language models to follow instructions with human feedback.” NeurIPS.
- Wei, J., et al. (2022). “Finetuned language models are zero-shot learners.” ICLR.
- Bai, Y., et al. (2022). “Constitutional AI: Harmlessness from AI Feedback.” arXiv.
- Rafailov, R., et al. (2023). “Direct Preference Optimization: Your Language Model is Secretly a Reward Model.” NeurIPS.
- Schulman, J., et al. (2017). “Proximal Policy Optimization Algorithms.” arXiv.
- DeepSeek-AI. (2025). “DeepSeek-R1: Incentivizing Reasoning Capability in LLMs via Reinforcement Learning.” arXiv.
- Ethayarajh, K., et al. (2024). “KTO: Kahneman-Tversky Optimization.” arXiv.
- Shao, Z., et al. (2024). “DeepSeekMath: Pushing the Limits of Mathematical Reasoning in Open Language Models.” arXiv.
- Touvron, H., et al. (2023). “Llama 2: Open Foundation and Fine-Tuned Chat Models.” arXiv.
- Meta AI. (2024). “The Llama 3 Herd of Models.” arXiv.
- Ji, Y., et al. (2024). “LLM as Dataset Generator: Quality-Diversity Trade-offs.” arXiv.
- Gudibande, A., et al. (2023). “The False Promise of Imitating Proprietary LLMs.” arXiv.
文章分享
如果这篇文章对你有帮助,欢迎分享给更多人!
后训练深度解析:LLM 从基座到助手的能力锻造
https://aiattnstudio.link/posts/post-training/
