监督微调 SFT:从基础模型到任务专家的关键一步
8037 字
40 分钟
监督微调 SFT:从基础模型到任务专家的关键一步
1. 引言:从预训练到应用
1.1 一个生动的比喻
如果说预训练是让学生读完 9 年义务教育,那么:
预训练 (Pre-training): 学生读完了所有教科书 = 模型学会了语言、世界知识、推理 = 能接龙,但不知道怎么"做事"
监督微调 (SFT): "基于以下指令完成任务" = 教学生特定的工作技能 = 模型学会回答问题、对话、完成任务 = 拿到"岗位证书"1.2 为什么需要 SFT?
预训练后的模型虽然强大,但存在三个关键问题:
| 问题 | 表现 | SFT 的作用 |
|---|---|---|
| 不知道怎么用 | 接龙不会回答 | 学会遵循指令 |
| 任务格式不对 | 输出含糊不清 | 学习格式化输出 |
| 角色错位 | 继续”接龙”而非”服务” | 建立助手身份 |
1.2.1 没有 SFT 的模型行为
用户: 翻译成中文: Hello World
预训练模型可能回复:"翻译成中文这个任务,Hello World 翻译成中文是... 继续:世界 Hello 翻译成..."
(它在"接龙",没有真正完成任务)1.2.2 经过 SFT 的模型行为
用户: 翻译成中文: Hello World
SFT 后模型回复:"你好世界"(直接、精准完成任务)1.3 SFT 在 LLM 训练流程中的位置
完整 LLM 训练流程:─────────────────────────────────────1. 预训练 (Pre-training) ↓ 海量无标注文本 = 基础模型 (Base Model) ↓2. 监督微调 (SFT) ← 本章 ↓ 高质量指令-回答对 = 指令模型 (Instruct Model) ↓3. 偏好对齐 (RLHF/DPO) ↓ 人类偏好数据 = 对齐模型 (Chat Model)1.4 SFT 的简史
2017: "Fine-Tuning" 概念兴起 │ 用于 BERT 等编码器 │ 适配下游任务 │2020: GPT-3 in-context learning │ 发现无需微调即可做任务 │ 但仍需要 SFT 用于生产 │2021: FLAN / T5 指令微调 │ Instruction Tuning 概念 │ 多任务指令提升泛化 │2022: InstructGPT (OpenAI) │ GPT-3 + SFT + RLHF 流程 │ ChatGPT 的前身 │2023: Alpaca / Vicuna / WizardLM │ 开源指令微调爆发 │2024: 多轮对话、Tool-Use SFT │ 复杂的指令格式 │2025+: 高质量小数据集 SFT2. SFT 基础概念
2.1 SFT 的本质
SFT = 高质量 (instruction, response) 数据 + 监督学习
# 训练样本example = { "instruction": "解释什么是机器学习", "response": "机器学习是一种让计算机...",}
# 转换为对话格式training_text = f"""用户: {example['instruction']}助手: {example['response']}"""
# 监督学习: 教会模型在看到 instruction 时生成 response2.2 三种 SFT 范式
2.2.1 Task-Specific Fine-Tuning (任务特定微调)
针对单一任务的微调数据: 单一任务的标注数据目的: 在该任务上达到最优例子: 训练 BERT 做情感分类2.2.2 Instruction Tuning (指令微调)
用多种任务的指令数据微调数据: 多个任务的 (指令, 回答) 对目的: 学会"理解指令"例子: FLAN, Natural Instructions2.2.3 Chat Fine-Tuning (对话微调)
用多轮对话数据微调数据: 多轮对话数据目的: 学会"对话"例子: RLHF 第一阶段本文主要讨论后两种。
2.3 Token 角色
# 训练时,每个 token 都有特定的角色TOKEN_ROLES = { 'instruction': '用户输入, 不计算损失', 'response': '助手回答, 计算损失', 'system': '系统提示, 不计算损失', 'padding': '填充, 不计算损失',}
# 损失掩码示意:# [INST] 你好 [SEP] 你好, 我是助手 [/INST] [EOS] [PAD] [PAD]# 计算损失: ____ _____ ____ _____ 好好 我是助手 [/INST] [EOS] ___ ___# ^^^^^^^^^^^^^^^^^^^^# 只对 response 计算 loss2.4 损失函数
与预训练相同的交叉熵损失,但只对 response 计算:
def sft_loss(model, batch): """SFT 的损失函数""" input_ids = batch['input_ids'] labels = batch['labels'] # -100 表示忽略
logits = model(input_ids)
# 标准语言建模损失,但只对 response 部分计算 loss = F.cross_entropy( logits.view(-1, logits.size(-1)), labels.view(-1), ignore_index=-100 # 关键!忽略非 response 位置 )
return loss核心思想:
3. 数据工程:SFT 的灵魂
3.1 数据的重要性
“SFT 的效果 ≈ 模型质量 + 数据质量” 高质量的 10K 数据胜过低质量的 100K 数据
Llama 2 论文的发现:
- 数据质量 比数据量重要 10 倍
- 数据多样性 比单一来源重要
- 数据清洗 是生命线
3.2 数据规模
训练规模 数据规模 训练时长─────────────────────────────────────7B 模型 10K-100K 数小时-1天13B 模型 20K-200K 数小时-2天70B 模型 50K-500K 1-3 天100B+ 100K+ 3+ 天经验公式:
高质量 SFT 数据量: 1万 ~ 10 万 条样本 (优秀模型)消耗 token 数: 数据量 × 平均长度3.3 数据来源
DATA_SOURCES = { 'human_annotated': { 'examples': ['', 'HelpSteer', 'OASST'], 'pros': '质量最高', 'cons': '成本高、规模有限', 'best_for': '关键领域', }, 'self_instruct': { 'examples': ['Alpaca', 'WizardLM'], 'pros': '易于扩展', 'cons': '可能继承错误', 'best_for': '通用指令', }, 'synthesized': { 'examples': ['Evol-Instruct', 'OpenHermes'], 'pros': '无限规模', 'cons': '可能单一样本', 'best_for': '特定任务', }, 'distilled': { 'examples': ['从 GPT-4 蒸馏'], 'pros': '质量高', 'cons': '成本/合规', 'best_for': '需要高质量', }, 'real_conversations': { 'examples': ['ShareGPT', 'WildChat'], 'pros': '真实分布', 'cons': '噪声多', 'best_for': '对话微调', },}3.4 经典数据集
3.4.1 主流公开数据集
DATASETS = { # 早期 'Alpaca': { 'size': '52K', 'year': 2023, 'feature': '首次开源 SFT 数据集', }, 'Dolly': { 'size': '15K', 'feature': '人工标注', },
# 进阶 'ShareGPT': { 'size': '70K', 'feature': '真实对话', }, 'OpenAssistant': { 'size': '160K', 'feature': '多轮对话', }, 'WizardLM': { 'size': '100K', 'feature': 'Evol-Instruct', },
# 现代 'OpenHermes-2.5': { 'size': '1M', 'feature': '高质量混合', }, 'UltraChat': { 'size': '1.5M', 'feature': '覆盖多种任务', }, 'FLAN-v2': { 'size': '20M+', 'feature': '多任务指令', }, 'Tulu-3': { 'size': '~1M', 'feature': '现代高质量', },}3.5 数据格式
3.5.1 基础格式
{ "instruction": "把下面句子翻译成英文: 你好", "input": "", "output": "Hello",}3.5.2 对话格式
{ "messages": [ {"role": "user", "content": "你好"}, {"role": "assistant", "content": "你好, 我是助手"}, {"role": "user", "content": "你会做什么?"}, {"role": "assistant", "content": "我可以帮你..."} ]}3.5.3 Alpaca 格式
{ "instruction": "任务描述", "input": "(可选)补充输入", "output": "期望输出",}3.5.4 ShareGPT 格式
{ "conversations": [ {"from": "human", "value": "用户消息"}, {"from": "gpt", "value": "助手回复"} ]}3.5.5 OpenAI 格式
{ "messages": [ {"role": "system", "content": "你是助手"}, {"role": "user", "content": "..."}, {"role": "assistant", "content": "..."} ]}3.6 Chat Template 设计
# Llama-2 Chat TemplateLLAMA2_TEMPLATE = """<s>[INST] <<SYS>>{system_message}<</SYS>>
{user_message_1} [/INST] {assistant_message_1} </s><s>[INST] {user_message_2} [/INST]"""
# Llama-3 Chat TemplateLLAMA3_TEMPLATE = """<|begin_of_text|><|start_header_id|>system<|end_header_id|>
{system}<|eot_id|><|start_header_id|>user<|end_header_id|>
{user}<|eot_id|><|start_header_id|>assistant<|end_header_id|>
{assistant}<|eot_id|>"""
# Mistral / Mixtral Chat TemplateMISTRAL_TEMPLATE = """<s>[INST] {system}\n\n{user} [/INST]{assistant}</s>"""
# ChatML Format (通用)CHATML_TEMPLATE = """<|im_start|>system{system}<|im_end|><|im_start|>user{user}<|im_end|><|im_start|>assistant{assistant}<|im_end|>"""
# QwenQWEN_TEMPLATE = """<|im_start|>system{system}<|im_end|><|im_start|>user{user}<|im_end|><|im_start|>assistant{assistant}<|im_end|>"""3.7 数据清洗
class DataCleaner: """数据清洗流程"""
def clean(self, dataset): # 1. 去重 dataset = self.deduplicate(dataset)
# 2. 长度过滤 dataset = self.filter_by_length(dataset, min=10, max=2048)
# 3. 质量过滤 dataset = self.filter_quality(dataset)
# 4. PII 移除 dataset = self.remove_pii(dataset)
# 5. 格式验证 dataset = self.validate_format(dataset)
# 6. 平衡采样 dataset = self.balance_sampling(dataset)
return dataset
def filter_quality(self, examples): """质量过滤规则""" filtered = [] for ex in examples: if self._passes_quality(ex): filtered.append(ex) return filtered
def _passes_quality(self, ex): # 检查规则: # 1. 长度合理 (10-2048 字符) # 2. 没有重复句子 # 3. 没有 placeholder # 4. 没有乱码 # 5. 不是纯模板 text = ex['output']
if len(text) < 10: return False if text.count(text[:10]) > 2: # 重复内容 return False if '...' in text.lower() and len(text) < 50: # 过短 return False
return True3.8 数据增强
3.8.1 Self-Instruct
class SelfInstructGenerator: """Self-Instruct 数据生成"""
def generate(self, seed_tasks, n=10000): """从少量种子生成大量数据""" new_tasks = []
for _ in range(n // 20): # 每批 20 个 # 1. 从已有任务采样 6-8 个作为示例 examples = random.sample(seed_tasks, k=random.randint(6, 8))
# 2. 让 LLM 生成新任务 prompt = self._build_prompt(examples) new_task = self.llm.generate(prompt)
# 3. 验证质量 if self._is_valid(new_task): new_tasks.append(new_task)
return new_tasks
# Self-Instruct 的 prompt 示例SELF_INSTRUCT_PROMPT = """你被要求生成 20 个多样化的任务。这些任务将用于微调语言模型。
任务类别: 写作、编程、数学、问答...
示例任务:1. 任务: ..., 输出: ...2. ...
生成新任务:"""3.8.2 Evol-Instruct
# WizardLM 的进化指令EVOLUTION_PROMPTS = { 'add_constraints': '在原指令上添加约束', 'deepen': '将指令深化为更复杂的问题', 'concretize': '将抽象问题具体化', 'increase_reasoning': '需要更多推理步骤', 'complicate_input': '输入更复杂', 'breadth': '探索新的话题',}
def evolve_instruction(instruction, evolution_type): """根据进化类型改写指令""" prompt = f""" 原指令: {instruction} 进化类型: {evolution_type}
生成新的、改进的指令: """ return llm.generate(prompt)3.8.3 数据蒸馏
class DataDistiller: """从强大模型蒸馏数据"""
def distill_from_gpt4(self, instructions, batch_size=100): """用 GPT-4 生成回答作为 SFT 数据""" results = []
for batch in chunks(instructions, batch_size): # 1. 并发调用 GPT-4 responses = self.gpt4.generate_batch(batch)
# 2. 配对 (instruction, response) for instr, resp in zip(batch, responses): results.append({ 'instruction': instr, 'output': resp })
return results
# 注意事项:# - 蒸馏后要做后处理# - 检查有害内容# - 不要过拟合老师模型的风格3.9 数据集构建完整流程
原始数据 ↓[1] 格式统一 ↓[2] 去重 ├── 精确去重 └── 模糊去重 (MinHash) ↓[3] 质量过滤 ├── 长度过滤 ├── 语言识别 ├── 关键词过滤 └── 质量模型评分 ↓[4] 内容审核 ├── NSFW 过滤 ├── PII 移除 └── 偏差检测 ↓[5] 数据配比 ├── 任务平衡 └── 多样性约束 ↓[6] 分割 ├── Train (90%) ├── Dev (5%) └── Test (5%) ↓最终 SFT 数据集4. SFT 训练流程
4.1 训练循环
class SFTTrainer: """SFT 训练器"""
def __init__(self, model, tokenizer, config): self.model = model self.tokenizer = tokenizer self.config = config
def train(self, dataset): # 1. 应用 chat template dataset = self._apply_chat_template(dataset)
# 2. 标记 response 部分(用于损失掩码) dataset = self._mark_response_tokens(dataset)
# 3. 创建 DataLoader dataloader = DataLoader( dataset, batch_size=self.config.batch_size, shuffle=True, collate_fn=self._collate )
# 4. 训练循环 for epoch in range(self.config.epochs): for batch in dataloader: # 计算损失 loss = self._compute_loss(batch)
# 反向传播 loss.backward()
# 优化器步骤 self.optimizer.step() self.optimizer.zero_grad()
# 学习率调度 self.scheduler.step()
def _compute_loss(self, batch): # 只对 response 计算损失 input_ids = batch['input_ids'] labels = batch['labels'] # response 部分保持原值,其他位置是 -100
outputs = self.model(input_ids=input_ids, labels=labels) return outputs.loss4.2 损失掩码实现
def mask_instruction_tokens(input_ids, tokenizer): """为损失函数掩码掉 instruction 部分"""
# 假设格式: [INST] instruction [/INST] response </s>
# 找到 response 开始位置 # 例如,找 [/INST] 之后的位置 response_start = find_token_sequence(input_ids, tokenizer.encode("[/INST]")) + len(tokenizer.encode("[/INST]"))
# 创建 labels: -100 表示忽略损失 labels = input_ids.clone() labels[:response_start] = -100 # 忽略 instruction
# 也忽略 padding labels[labels == tokenizer.pad_token_id] = -100
return labels4.3 DataCollator 实现
class SFTCollator: """SFT 数据整理器"""
def __init__(self, tokenizer, max_length=2048): self.tokenizer = tokenizer self.max_length = max_length
def __call__(self, batch): # 1. 提取文本 texts = [self.format_example(ex) for ex in batch]
# 2. Tokenize tokenized = self.tokenizer( texts, truncation=True, max_length=self.max_length, padding=True, return_tensors='pt', )
# 3. 创建损失掩码 labels = tokenized['input_ids'].clone()
# 找每个样本的 response 开始位置 for i, text in enumerate(texts): response_start = self._find_response_start(text) labels[i, :response_start] = -100
# 4. 忽略 padding labels[labels == self.tokenizer.pad_token_id] = -100
return { 'input_ids': tokenized['input_ids'], 'attention_mask': tokenized['attention_mask'], 'labels': labels, }5. 训练参数与最佳实践
5.1 关键超参数
HYPERPARAMETER_SETTINGS = { 'learning_rate': { 'full_finetuning': '1e-5 到 5e-5', 'LoRA': '1e-4 到 5e-4', 'common': '比预训练小 10-100 倍', 'best_practice': '2e-5 (lr), cosine schedule', }, 'epochs': { 'common': '2-3 个 epoch', 'warning': '太多 epoch 容易过拟合', 'best_practice': '在 val loss 上升时停止', }, 'batch_size': { 'small_model': '32-128', 'large_model': '1-16 (per GPU)', 'gradient_accumulation': '通常 4-16', 'best_practice': '等效 batch size 128-512', }, 'max_length': { 'short': '512-1024', 'medium': '2048', 'long': '4096-8192', 'best_practice': '根据数据分布决定', }, 'warmup_ratio': { 'common': '0.03-0.1', 'best_practice': '0.03', }, 'weight_decay': { 'common': '0.0-0.1', 'best_practice': '0.0 (与预训练相同)', },}5.2 学习率调度
def get_lr_scheduler(optimizer, num_training_steps, warmup_ratio=0.03): """SFT 常用学习率调度"""
warmup_steps = int(num_training_steps * warmup_ratio)
# Cosine 退火 (主流) scheduler = get_cosine_schedule_with_warmup( optimizer, num_warmup_steps=warmup_steps, num_training_steps=num_training_steps, )
return scheduler
# 学习率曲线示例LR_CURVE = """LR^| ___________| / \___| / \____| / \___| / \____+--------------------------→ Steps ^ warmup"""5.3 Batch Size 的选择
BATCH_SIZE_GUIDELINES = """总 batch_size = per_device_batch × grad_accum × n_gpus
7B 模型 (单 GPU 24GB): per_device_batch=2-4, grad_accum=16, total=32-64
70B 模型 (8 GPU): per_device_batch=1, grad_accum=8-16, total=8-16
推荐: total batch size = 64-128
经验公式: GPU 内存 ≈ 2 × per_device_batch × seq_len × 4 bytes × num_layers (粗略估算,FP16/BF16)"""5.4 多轮对话的处理
class MultiTurnCollator: """多轮对话 SFT 数据整理器"""
def format_conversation(self, messages, tokenizer): # 累积对话 text = "" assistant_ranges = [] # (start, end) of assistant tokens
for msg in messages: # 添加对话块 block_start = len(tokenizer.encode(text)) text += f"{msg['role']}: {msg['content']}\n" block_end = len(tokenizer.encode(text))
# 记录 assistant 消息范围 if msg['role'] == 'assistant': assistant_ranges.append((block_start, block_end))
# Tokenize tokens = tokenizer.encode(text)
# 创建 labels: assistant 部分保留,其他部分 -100 labels = [-100] * len(tokens) for start, end in assistant_ranges: labels[start:end] = tokens[start:end]
return tokens, labels
# 注意:使用 mask 而不是简单的"只看最后一句"# 这样可以让模型学会所有轮的生成6. PEFT:参数高效微调
6.1 为什么需要 PEFT?
全参数微调的问题:
70B 模型全参数微调:- 模型参数: 140 GB (FP16)- 优化器状态: 280 GB (Adam)- 激活值: 数十 GB- 总计: 数百 GB,>100 GPU
对单家公司、研究者不现实!6.2 LoRA
Low-Rank Adaptation:
6.2.1 LoRA 原理
class LoRALinear(nn.Module): """LoRA 包装的线性层"""
def __init__(self, original_linear, r=8, alpha=16): super().__init__() self.original = original_linear self.original.weight.requires_grad = False # 冻结原参数
in_dim = original_linear.in_features out_dim = original_linear.out_features
# LoRA 矩阵 self.lora_A = nn.Parameter(torch.randn(r, in_dim)) self.lora_B = nn.Parameter(torch.zeros(out_dim, r))
# 缩放 self.scaling = alpha / r
def forward(self, x): # 原始输出 + LoRA 更新 original_out = self.original(x) lora_out = (x @ self.lora_A.T @ self.lora_B.T) * self.scaling return original_out + lora_out数学表达:
原始权重 ,LoRA 增加低秩更新:
其中 , , 秩 。
6.2.2 LoRA 参数分析
# 假设: d=4096, k=4096, r=8ORIGINAL_PARAMS = 4096 * 4096 # 16.78MLORA_PARAMS = 4096 * 8 + 8 * 4096 # 65KRATIO = LORA_PARAMS / ORIGINAL_PARAMS # 0.39%
# 但训练中:# - 原始权重: 冻结 (无梯度)# - LoRA 权重: 训练# - 优化器状态: 只有 LoRA
LORA_TRAINABLE_TOTAL = 0.39% # 65K / 16.78M6.2.3 LoRA 实战配置
LORA_CONFIG = { 'r': 8, # rank, 推荐 8-64 'lora_alpha': 16, # 缩放, alpha=2*r 通常效果好 'target_modules': [ # 应用到哪些层 'q_proj', # 注意: attention 投影 'k_proj', 'v_proj', 'o_proj', 'gate_proj', # FFN 'up_proj', 'down_proj', ], 'lora_dropout': 0.05, 'bias': 'none', # 不训练 bias 'task_type': 'CAUSAL_LM',}6.3 QLoRA
Quantized LoRA - 在 4-bit 量化基础上做 LoRA:
class QLoRALayer: """QLoRA: 4-bit 量化 + LoRA"""
def __init__(self): # 1. 将基础模型量化为 4-bit (NF4) self.base_weight = quantize_4bit(original_weight) # ~4 GB
# 2. LoRA 参数保持 FP16 self.lora_A = nn.Parameter(torch.randn(8, 4096)) self.lora_B = nn.Parameter(torch.randn(4096, 8))
# 3. 分页优化器 (Paged Optimizer) # 4. 双量化 # 5. 嵌套量化
def forward(self, x): # 反量化 base 做计算 base_out = x @ dequantize(self.base_weight).T lora_out = lora_computation(x, self.lora_A, self.lora_B) return base_out + lora_out
# 内存对比 (65B 模型):# FP16 全量: 130 GB# LoRA: 130 GB base + 100 MB LoRA = 130 GB# QLoRA: 65 GB (4-bit) + 100 MB LoRA = 65 GB (省一半!)6.4 其他 PEFT 方法
6.4.1 Adapter
class Adapter(nn.Module): """Adapter 模块"""
def __init__(self, dim, bottleneck_dim=64): super().__init__() self.down = nn.Linear(dim, bottleneck_dim) self.up = nn.Linear(bottleneck_dim, dim) self.act = nn.ReLU()
def forward(self, x): # 残差连接 return x + self.up(self.act(self.down(x)))6.4.2 Prefix Tuning
class PrefixTuning(nn.Module): """在每层前面添加可训练 prefix"""
def __init__(self, prefix_length=20, n_layers=32, dim=4096): # prefix_tokens: [n_layers, 2, prefix_length, dim] # 2 = key + value self.prefix_tokens = nn.Parameter( torch.randn(n_layers, 2, prefix_length, dim) )
def forward(self, x): # 将 prefix_tokens 拼接到 K, V # ... pass6.4.3 IA³
class IA3(nn.Module): """IA³ - 元素级缩放"""
def __init__(self): # 每个层一个可学习的缩放向量 self.l = nn.Parameter(torch.ones(dim))
def forward(self, x): # 元素级乘法 return x * self.l6.5 PEFT 方法对比
| 方法 | 训练参数 | 内存占用 | 效果 | 推理延迟 |
|---|---|---|---|---|
| Full FT | 100% | 极高 | 基准 | 无 |
| LoRA | 0.1-5% | 中 | 95-99% | 无 |
| QLoRA | 0.1-5% | 低 | 90-95% | 略高 |
| Adapter | 2-10% | 中 | 90-95% | 略高 |
| Prefix Tuning | <1% | 低 | 85-90% | 略高 |
| IA³ | <0.5% | 极低 | 85-95% | 无 |
7. 训练框架与工具
7.1 Hugging Face TRL
from trl import SFTTrainer, SFTConfig
# 配置 SFTconfig = SFTConfig( output_dir="./output", num_train_epochs=3, per_device_train_batch_size=4, gradient_accumulation_steps=4, learning_rate=2e-5, warmup_ratio=0.03, lr_scheduler_type="cosine", max_seq_length=2048, logging_steps=10, save_steps=500, bf16=True, gradient_checkpointing=True,)
# 初始化 Trainertrainer = SFTTrainer( model=model, args=config, train_dataset=dataset, processing_class=tokenizer,)
# 开始训练trainer.train()7.2 Hugging Face Transformers
from transformers import Trainer, TrainingArguments
training_args = TrainingArguments( output_dir="./output", num_train_epochs=3, per_device_train_batch_size=2, gradient_accumulation_steps=8, learning_rate=2e-5, bf16=True, gradient_checkpointing=True, logging_steps=10, save_strategy="steps", save_steps=500,)
trainer = Trainer( model=model, args=training_args, train_dataset=train_dataset, eval_dataset=eval_dataset, data_collator=collator,)
trainer.train()7.3 PEFT 集成
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training
# 准备模型 (可选 QLoRA)model = prepare_model_for_kbit_training( model, use_gradient_checkpointing=True,)
# LoRA 配置lora_config = LoraConfig( r=16, lora_alpha=32, target_modules=[ "q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj", ], lora_dropout=0.05, bias="none", task_type="CAUSAL_LM",)
# 包装模型model = get_peft_model(model, lora_config)
# 打印可训练参数model.print_trainable_parameters()# trainable params: 16,384,000 || all params: 6,738,415,616 || trainable%: 0.24%
# 正常训练trainer = SFTTrainer( model=model, args=training_args, train_dataset=dataset,)7.4 Llama Factory
# 一站式 SFT 框架# 支持 100+ 模型 + 多种 SFT 模式
# 命令行llamafactory-cli train \ --model_name_or_path meta-llama/Llama-3-8B-Instruct \ --template llama3 \ --dataset alpaca_zh \ --output_dir ./output \ --finetuning_type lora \ --lora_rank 16 \ --per_device_train_batch_size 2 \ --gradient_accumulation_steps 8 \ --learning_rate 2e-4 \ --num_train_epochs 38. 训练技巧与最佳实践
8.1 显存优化
MEMORY_OPTIMIZATIONS = { 'mixed_precision': { 'method': 'BF16 混合精度', 'saves': '~50%', 'code': 'bf16=True', }, 'gradient_checkpointing': { 'method': '激活值不保存, 反向时重算', 'saves': '~30-50%', 'code': 'gradient_checkpointing=True', 'cost': '+30% 计算', }, 'lora': { 'method': '参数高效微调', 'saves': '~70% (LLM 部分)', }, 'qlora': { 'method': '4-bit 量化 + LoRA', 'saves': '~80%', }, 'deepspeed_zeRO3': { 'method': '完全分片数据并行', 'saves': '线性扩展', 'code': 'deepspeed_stage=3', }, 'flash_attention': { 'method': 'IO 优化的注意力', 'saves': '长序列显著', 'code': 'attn_implementation="flash_attention_2"', },}8.2 训练稳定性
STABILITY_TIPS = """1. 学习率不要太大 - 推荐 2e-5 (全量) / 2e-4 (LoRA) - 比预训练小 10-100 倍
2. 监控 loss 和 grad_norm - grad_norm > 5: 警告 - grad_norm > 10: 干预
3. 检查 NaN - 添加: torch.autograd.set_detect_anomaly(True) - 跳过 NaN 批次
4. 慢启动 - Warmup ratio 0.03 - 不要从高 LR 开始
5. 定期评估 - val loss 上升时停止 - 不要等到计划结束"""8.3 数据配比与训练时长
TRAINING_DURATION_GUIDELINES = """数据 1万条: 1-3 epochs数据 5万条: 1-2 epochs数据 10万条: 1 epoch数据 50万条+: 通常不够, 需要更长时间
经验:- 数据少时: 多 epoch (但注意过拟合)- 数据多时: 少 epoch (通常 1-2)- 学习率与 epoch 数相关"""8.4 Epoch vs Steps
# 选择标准EPOCH_VS_STEPS = """按 epoch (遍历整个数据集):- 数据集固定不变- 简单直观- 适合中小数据集
按 step (固定步数):- 数据集持续扩展- 适合大数据集- 训练时间可控
注意: 总更新次数 = steps × batch_size 这个数量级应该和模型规模匹配: - 7B: 数千到数万个更新 - 70B: 更少(更稳定的 SFT)"""8.5 评估 Checkpoint
CHECKPOINT_STRATEGY = """1. 频繁保存 checkpoint save_steps=500 (10K 数据集每 epoch 约 10 次)
2. 评估每个 checkpoint eval_steps=500 (与 save_steps 同步)
3. 选择最佳 checkpoint 基于 dev set eval_loss 或下游任务分数
4. 防止过拟合 - 监控 val loss - 早停 (early stopping)"""8.6 多任务 SFT
class MultiTaskSFTDataset: """多任务 SFT"""
def __init__(self, datasets, mixing_weights): """ datasets: 多个任务的 SFT 数据集 mixing_weights: 每个任务的比例 (和为 1) """ self.datasets = datasets self.weights = mixing_weights
def __iter__(self): # 多任务混合采样 while True: task_idx = random.choices( range(len(self.datasets)), weights=self.weights, )[0]
dataset = self.datasets[task_idx] if dataset: yield next(iter(dataset))
# 典型配比:TYPICAL_MIX = { 'general_chat': 0.40, # 通用对话 'code': 0.15, # 代码 'math': 0.10, # 数学 'reasoning': 0.15, # 推理 'creative': 0.10, # 创意写作 'qa': 0.10, # 问答}9. SFT 后效果评估
9.1 评估方法
class SFTEvaluator: """SFT 评估器"""
def evaluate(self, model, eval_set): results = {}
# 1. 困惑度 (Perplexity) results['ppl'] = self.compute_perplexity(model, eval_set)
# 2. 任务准确率 results['accuracy'] = self.compute_accuracy(model, eval_set)
# 3. 人类偏好对比 results['win_rate'] = self.human_eval_comparison(model)
# 4. 安全性 results['safety'] = self.safety_eval(model)
return results
def compute_perplexity(self, model, eval_set): total_loss = 0 total_tokens = 0
for example in eval_set: # 只看 response 部分 tokens = example['response_tokens']
with torch.no_grad(): logits = model(example['input_ids'])
loss = F.cross_entropy( logits.view(-1, logits.size(-1)), tokens.view(-1), )
total_loss += loss.item() * len(tokens) total_tokens += len(tokens)
avg_loss = total_loss / total_tokens perplexity = math.exp(avg_loss) return perplexity9.2 自动化评估
# 使用 GPT-4 作为评判def gpt4_as_judge(model_output, reference): prompt = f""" 参考答案: {reference} 模型输出: {model_output}
评估模型输出的: 1. 准确性 (1-5) 2. 帮助性 (1-5) 3. 安全性 (1-5) 4. 一致性 (1-5)
JSON 格式: """ return json.loads(gpt4.generate(prompt))
# Alpaca Eval: 自动评估# MT-Bench: 多任务评估# LMSys Arena: 人类偏好对比9.3 关键指标
SFT_METRICS = { # 模型能力 'task_accuracy': '任务准确率', 'perplexity': '困惑度 (越低越好)', 'bleu/rouge': '与参考答案的相似度',
# 指令遵循 'instruction_following': '指令遵循率', 'format_compliance': '格式合规率',
# 安全性 'toxicity_score': '有害性分数', 'bias_score': '偏差分数',
# 用户体验 'helpfulness': '有用性', 'harmlessness': '无害性', 'honesty': '诚实性',
# 业务指标 'user_satisfaction': '用户满意度', 'task_completion_rate': '任务完成率',}9.4 评估陷阱
EVALUATION_PITFALLS = """1. 数据泄漏 - 测试集出现在训练集 - 使用专门的去重流程
2. 评估偏差 - 单一任务的分数不能代表整体 - 多基准综合评估
3. 表面指标 - 只看 BLEU 不可靠 - 结合人类评估
4. 基准过拟合 - 不要只针对某个基准优化 - 关注泛化能力
5. 能力遗忘 - 评估是否保留了原模型能力 - 与原始模型对比"""10. 应对 SFT 的常见问题
10.1 模型学不会
TROUBLESHOOTING_NO_LEARNING = """可能原因与解决:
1. 学习率太低 - 尝试 5e-5 到 1e-4
2. 数据质量差 - 检查: 答案真的正确吗? - 改善: 重新标注
3. 数据太简单 - 评估: 模型 base 已经会这些 - 改善: 增加难度
4. 输入长度限制 - 数据被截断 - 改善: 检查 max_length
5. 模型太小 - 7B 学不会复杂任务 - 改善: 增加模型规模"""10.2 模型过度模仿
TROUBLESHOOTING_OVERFITTING = """症状: 模型只会复读训练数据
1. 早期过拟合 - epoch 过多 - 改善: 减少 epoch (1-2)
2. 数据太相似 - 多样性不足 - 改善: 增加数据多样性
3. 复制响应 - 模型学到了逐字记忆 - 改善: 增加数据量
4. 评估信号停滞 - train loss 下降, val loss 上升 - 改善: 早期停止"""10.3 灾难性遗忘
TROUBLESHOOTING_CATASTROPHIC_FORGETTING = """症状: SFT 后, 模型在新任务上变好, 但原能力下降
1. 混入原始数据 - 10-20% 原始预训练数据混到 SFT
2. 使用 EWC (Elastic Weight Consolidation) - 保护重要参数
3. 降低学习率 - 避免大幅改动
4. LoRA 微调 - 限制可训练参数范围
5. 评估检查 - 使用通用基准 (MMLU, Hellaswag) - 确认原能力保留"""10.4 格式不一致
TROUBLESHOOTING_FORMAT_INCONSISTENCY = """症状: 模型有时输出正确格式, 有时不正确
1. 增加格式样本 - 用多个样本训练同一格式
2. 强化 chat template - 训练和推理用相同模板
3. 添加格式 marker - 使用特殊 token 标识格式
4. 输出层规范化 - 添加后处理"""11. SFT vs RLHF vs DPO
11.1 三种对齐方式对比
LLM 训练阶段:─────────────────────────────────────────────────Base Model → SFT → RLHF → Chat Model ↓ (可选) ↓ DPO─────────────────────────────────────────────────
对比:┌───────────┬───────────┬────────────┬──────────┐│ │ SFT │ RLHF │ DPO │├───────────┼───────────┼────────────┼──────────┤│ 数据 │ 指令-回答 │ 偏好对 │ 偏好对 ││ 训练目标 │ 拟合回答 │ 学习奖励 │ 直接对齐 ││ 复杂度 │ 简单 │ 复杂 │ 中等 ││ 稳定性 │ 高 │ 波动 │ 中等 ││ 效果 │ 基础 │ 强 │ 强 ││ 成本 │ 低 │ 高 │ 低 │└───────────┴───────────┴────────────┴──────────┘11.2 何时选择 SFT
WHEN_TO_USE_SFT = """应该用 SFT 当:✓ 主要需求是任务格式✓ 数据 (instruction, response) 充足✓ 不需要过强的对齐✓ 资源有限 (RLHF 需要更多)
后续可以加:- DPO 进一步对齐- Online RL 在线强化
不适合 SFT 当:✗ 数据中答案质量差✗ 需要复杂推理对齐✗ 价值观精确对齐"""11.3 SFT + DPO 流程
class SFTThenDPOPipeline: """SFT + DPO 组合"""
def __init__(self): # 1. 预训练模型 self.base = load_base_model()
# 2. SFT 阶段 self.sft_model = self.sft_train(self.base, sft_data)
# 3. DPO 阶段 self.dpo_model = self.dpo_train(self.sft_model, preference_data)
return self.dpo_model
def sft_train(self, model, data): """SFT 训练""" trainer = SFTTrainer(model, data) trainer.train(epochs=3, lr=2e-5) return trainer.model
def dpo_train(self, model, preferences): """DPO 训练""" trainer = DPOTrainer(model, preferences) trainer.train(epochs=1, lr=5e-7) # DPO 用小 LR return trainer.model12. 高级 SFT 技术
12.1 Constitutional AI (CAI)
class CAIDataGenerator: """Anthropic 的 Constitutional AI"""
def generate_sft_data(self, base_sft_data, constitution): """基于原则生成改进数据""" improved = []
for example in base_sft_data: # 1. 评估是否符合 constitution critique = self.llm.evaluate_against( example['response'], constitution, )
# 2. 基于 critique 改进 revised = self.llm.revise_response( example['response'], critique, )
improved.append({ 'instruction': example['instruction'], 'output': revised, })
return improved
# Constitution 示例:CONSTITUTION = [ "请勿生成有害内容", "请保持中立", "请避免偏见", "不要冒充身份",]12.2 Iterative SFT
class IterativeSFT: """迭代 SFT"""
def __init__(self): # 多轮 SFT 与评估 pass
def train_iteratively(self, max_iterations=5): current_model = load_base_model()
for i in range(max_iterations): # 1. 用当前模型生成数据 new_data = self.generate_data(current_model)
# 2. 人工或自动筛选 quality_data = self.filter_quality(new_data)
# 3. 继续 SFT current_model = self.continue_sft( current_model, quality_data, )
# 4. 评估 score = self.evaluate(current_model)
if score > THRESHOLD: break
return current_model12.3 Rejection Sampling Fine-Tuning
class RSFT: """拒绝采样微调"""
def train(self, samples_per_prompt=10): # 1. 为每个 instruction 生成多个 response results = []
for example in self.dataset: candidates = [] for _ in range(samples_per_prompt): response = self.model.generate(example['instruction']) candidates.append(response)
# 2. 用奖励模型评分 scored = self.reward_model.score(candidates)
# 3. 选择最佳响应 best_idx = max(range(len(scored)), key=lambda i: scored[i]) best_response = candidates[best_idx]
results.append({ 'instruction': example['instruction'], 'output': best_response, })
# 4. 在精选数据上训练 return self.train_on(results)
# Llama 2 / WizardLM 都用过这种方法12.4 Online SFT
class OnlineSFT: """在线 SFT - 边训练边生成"""
def train(self): for prompt_batch in self.dataloader: # 1. 当前模型生成 response responses = self.model.generate(prompt_batch)
# 2. 评估质量 rewards = self.reward_model(responses)
# 3. 用高质量响应 SFT loss = self.compute_sft_loss(prompt_batch, responses) loss.backward() self.optimizer.step()
# 注意: 避免 reward hacking # 注意: 避免分布漂移13. 实战案例分析
13.1 案例 1: 客服助手 SFT
# 场景: 电商客服助手# 数据来源: 历史人工客服对话 (脱敏)
class CustomerServiceSFT: def __init__(self): # 准备数据 self.data = self.prepare_data()
# 配置训练 self.config = SFTConfig( model_name='llama-3-8b-instruct', template='llama3', output_dir='./customer-service-sft', epochs=3, batch_size=4, lora_rank=16, learning_rate=2e-5, )
def prepare_data(self): # 1. 收集: 100,000 真实对话 (脱敏) # 2. 清洗: 移除 PII # 3. 标注: 标注意图、情感 # 4. 格式化: 应用 chat template # 5. 分割: 95/2.5/2.5 train/val/test ...
def train(self): trainer = SFTTrainer(self.config) trainer.train(self.data)
# 评估 results = trainer.evaluate() return results13.2 案例 2: 代码助手 SFT
# 场景: 企业内部代码助手# 数据: (instruction, code) 对
class CodeAssistantSFT: def __init__(self): # 准备代码相关数据 # MagiCoder, WizardCoder 等数据集 ...
def train(self): # 关键: # 1. 代码块要有正确语法 # 2. 包含解释 # 3. 多语言 # 4. 多种任务 (生成、解释、修复)13.3 案例 3: 医学诊断助手 SFT
# 场景: 医学专业知识# 数据: 医学问答数据 + 安全审查
class MedicalSFT: def __init__(self): # 强烈注意: # 1. 数据需要医学专家审阅 # 2. 加入安全约束 # 3. 不替代医生明确强调 ...
def safety_filter(self, output): # 确保模型: # - 不替代医生 # - 强调专业建议 # - 不诊断严重疾病 ...14. SFT 的成本与规模
14.1 训练成本估算
模型规模 方法 GPU 数 数据 时长 估算成本──────────────────────────────────────────────────────────7B Full FT 4-8 50K 数天 $0.5K-2K7B LoRA 1-2 50K 数小时 $0.1K-0.5K13B Full FT 8-16 100K 数天 $1K-5K13B LoRA 2-4 100K 数小时 $0.5K-1K70B QLoRA 4-8 50K 数天 $5K-20K14.2 推理部署成本
INFERENCE_COST_OPTIMIZATION = """1. 合并 LoRA - LoRA 训练后合并到基础模型 - 推理时无额外开销
2. 量化部署 - SFT 后量化到 INT4/INT8 - 大幅减小模型大小
3. 蒸馏 - 用大模型 SFT 数据训小模型 - 性能损失较小
4. vLLM/TGI 部署 - 高吞吐推理 - 批处理优化"""15. 未来发展方向
15.1 SFT 的趋势
SFT_FUTURE_TRENDS = """1. 更高效的数据利用 - 质量 > 数量 - 用更少数据训更好模型
2. 自动化数据生成 - 自我进化 - 弱监督信号
3. 多模态 SFT - 视觉-语言 SFT - 视频 SFT
4. 在线 SFT - 持续学习 - 实时适应
5. Agent SFT - 工具使用 SFT - 多轮 Agent 训练
6. 价值对齐 - 直接在 SFT 中加入价值观"""15.2 新兴技术
EMERGING_TECHNIQUES = """1. Direct Preference Optimization (DPO) - 直接用偏好优化 - 不需要奖励模型
2. Self-Rewarding SFT - 模型评估自己 - 自我提升
3. Constitutional AI - 用原则引导训练 - 避免昂贵的人类标注
4. Synthetic Data - 用 LLM 生成训练数据 - 飞轮自我增强"""16. 核心概念总结
16.1 SFT 公式化表达
- SFT 损失函数:
- LoRA 数学:
- QLoRA 内存:
- 采样效率:
16.2 SFT 决策树
需要微调模型吗?├── 是│ ├── 全参数还是 PEFT?│ │ ├── 全参数 (效果好, 资源多)│ │ └── PEFT (资源有限)│ │ ├── LoRA (主流)│ │ └── QLoRA (更省资源)│ ├── 是否有偏好数据?│ │ ├── 没有 → SFT 后发布│ │ └── 有 → SFT + DPO│ └── 是否需要持续对齐?│ ├── 是 → 在线学习│ └── 否 → 训练完成即结束└── 否 → 考虑 prompting / RAG16.3 关键原则
SFT_PRINCIPLES = """1. 数据是灵魂 - 质量 > 数量 - 多样性 > 单一 - 真实 > 完美
2. 模型对齐 - Format 一致 (train/eval/inference) - Chat template 标准化
3. 评估驱动 - 不靠感觉 - 用数据说话
4. 安全第一 - 数据审查 - 输出过滤 - 人类监督
5. 持续改进 - 监控用户反馈 - 迭代数据集"""17. 实战建议
17.1 决策清单
DECISION_CHECKLIST = """□ 数据准备 □ 数据来源选择 □ 数据规模 (10K-100K) □ 数据质量评分 □ 去重 (精确/模糊) □ 格式化 (chat template) □ 划分 (train/val/test)
□ 模型选择 □ Base 模型选择 □ 参数规模 (7B/13B/70B) □ 训练方法 (Full/LoRA/QLoRA) □ 训练目标模块
□ 训练配置 □ 学习率 (2e-5 full, 2e-4 LoRA) □ Epoch 数 (1-3) □ Batch size (等效 64-128) □ 学习率调度 (cosine + warmup) □ 优化器 (AdamW) □ 精度 (BF16)
□ 训练后评估 □ Perplexity □ 任务准确率 □ 安全性 □ 泛化能力
□ 部署准备 □ 模型合并 (LoRA) □ 量化 (INT4/8) □ 推理测试 □ 性能监控"""17.2 常见失败与解决
COMMON_FAILURES = """1. 模型输出格式混乱 → 检查 chat template 一致性
2. 模型失去通用能力 → 加入 10-20% 原始数据
3. 模型过度重复训练数据 → 减少 epoch, 增加数据多样性
4. 模型"胡说八道" → 数据质量太差, 重新标注
5. 训练 loss 不降 → 学习率太低, 改成 5e-5
6. 训练 loss spike → 降低学习率, 增加 warmup
7. 显存 OOM → 启用 gradient checkpointing + LoRA → 减小 batch size"""17.3 推荐学习路径
LEARNING_ROADMAP = """1. 入门 (1 周) - 理解 SFT 的作用 - 用 TRL 在 1K 数据 SFT - 观察效果
2. 进阶 (2-4 周) - 学习 LoRA/QLoRA - 设计 prompt 模板 - 数据工程
3. 高级 (1-2 月) - 多任务 SFT - SFT + DPO - 大规模数据处理
4. 专家 (3 月+) - 在线学习 - 持续改进流程 - 生产部署"""17.4 推荐资源
RESOURCES = """论文:- "LLaMA: Open Foundation and Fine-Tuned Chat Models" (Llama 2)- "Training language models to follow instructions with human feedback" (InstructGPT)- "Self-Instruct: Aligning Language Models with Self-Generated Instructions"- "LoRA: Low-Rank Adaptation of Large Language Models"- "QLoRA: Efficient Finetuning of Quantized LLMs"
框架/工具:- Hugging Face Transformers / TRL / PEFT- Llama Factory- Unsloth (快速 SFT)- axolotl- Liger Kernel
数据集:- HuggingFace H4 (用于训练)- OpenAssistant- ShareGPT- Tulu-3- OpenHermes-2.5
教程:- Hugging Face SFT 教程- "Instruction Tuning for LLMs" 课程"""18. 总结与展望
18.1 SFT 的核心价值
SFT 是把”无所不知”的基础模型变成”能做事”的助手的关键步骤:
- 任务对齐:让模型学会做特定任务
- 格式对齐:让模型遵循统一格式
- 角色对齐:让模型扮演助手角色
- 风格对齐:让模型有特定风格
18.2 SFT 的局限
SFT_LIMITATIONS = """1. 数据上限 - 永远不能超过训练数据 - 数据偏差传递
2. 表面学习 - 学习的是 token 模式 - 不是真正的"理解"
3. 对齐困难 - 不能完全对齐价值观 - 需要额外的 RLHF/DPO
4. 能力约束 - 受限于 base model 能力 - 孱弱的 base 训不出强助手"""18.3 未来愿景
LLM 训练的下一阶段:┌────────────────────────────────────────────┐│ ││ · 自动化数据工程 ││ · 自我对齐 (RLAIF) ││ · 多模态原生 SFT ││ · Agent 原生 SFT ││ · 持续在线学习 ││ · 终身学习 ││ │└────────────────────────────────────────────┘18.4 结语
SFT 是现代 LLM 流程中最成熟、最实用的环节。
掌握 SFT 你可以:
- 把任何基础模型变成有用的助手
- 用 LoRA 在消费级 GPU 微调大模型
- 构建领域专用的 AI 系统
“SFT 是从基础模型到产品模型的桥梁。”
组合的力量:
- 预训练 → 打基础
- SFT → 学技能
- RLHF/DPO → 价值观对齐
每一阶段都很重要,但 SFT 是从”知道”到”做到”的关键一步。
参考资料
- Wei, J., et al. (2022). “Finetuned Language Models Are Zero-Shot Learners.” ICLR.
- Sanh, V., et al. (2022). “Multitask Prompted Training Enables Zero-Shot Task Generalization.” ICLR.
- Ouyang, L., et al. (2022). “Training Language Models to Follow Instructions with Human Feedback.” NeurIPS.
- Touvron, H., et al. (2023). “Llama 2: Open Foundation and Fine-Tuned Chat Models.” arXiv.
- Hu, E., et al. (2022). “LoRA: Low-Rank Adaptation of Large Language Models.” ICLR.
- Dettmers, T., et al. (2023). “QLoRA: Efficient Finetuning of Quantized LLMs.” NeurIPS.
- Wang, Y., et al. (2023). “Self-Instruct: Aligning Language Models with Self-Generated Instructions.” ACL.
- Xu, C., et al. (2023). “WizardLM: Empowering Large Language Models to Follow Complex Instructions.” arXiv.
- Conover, M., et al. (2023). “Free Dolly: Introducing the World’s First Truly Open Instruction-Tuned LLM.”
- Radford, A., et al. (2019). “Language Models are Unsupervised Multitask Learners.” OpenAI.
- Liu, P., et al. (2023). “Pre-train, Prompt, and Predict: A Systematic Survey of Prompting Methods.” ACL.
- Liu, X., et al. (2024). “Instruction-Following Evaluation for Large Language Models.” ICLR.
- Biderman, S., et al. (2023). “Pythia: A Suite for Analyzing Large Language Models Across Training Time and Scale.” arXiv.
- Taori, R., et al. (2023). “Stanford Alpaca: An Instruction-following LLaMA Model.”
- Chiang, W., et al. (2023). “Vicuna: An Open-Source Chatbot Impressing GPT-4 with 90%* ChatGPT Quality.”
- Rafailov, R., et al. (2023). “Direct Preference Optimization: Your Language Model is Secretly a Reward Model.” NeurIPS.
- Bai, Y., et al. (2022). “Constitutional AI: Harmlessness from AI Feedback.” arXiv.
- Houlsby, N., et al. (2019). “Parameter-Efficient Transfer Learning for NLP.” ICML.
- Li, X., et al. (2021). “Prefix-Tuning: Optimizing Continuous Prompts for Generation.” ACL.
- Ansell, J., et al. (2024). “Unsloth: 2-5x faster LoRA fine-tuning.” GitHub.
文章分享
如果这篇文章对你有帮助,欢迎分享给更多人!
监督微调 SFT:从基础模型到任务专家的关键一步
https://aiattnstudio.link/posts/supervised-finetuning/
