27届大模型岗面试准备(六):SFT 监督微调深挖——数据构造、损失细节与训练技巧全解析

SFT(Supervised Fine-Tuning)是校招面试里被问得最细的环节之一。原因很简单:多数公司的大模型业务不做预训练,但几乎都做微调——所以面试官会默认你对 SFT 有实操级理解。"SFT 的 loss 和预训练有什么区别""为什么要 mask 掉 prompt 部分""多轮对话怎么组织成训练样本"这三个问题,答不好基本一面就结束了。这篇从数据格式讲到损失实现,再到数据质量工程,把 SFT 的完整知识面串起来。

一、SFT 在训练管线中的位置与本质

完整管线是:预训练(PT)→ 监督微调(SFT)→ 偏好对齐(RLHF/DPO)。base 模型只会"续写",不会"对话"——你问它"中国的首都是哪里?",它可能续写出"这是一道常见的地理题"。SFT 的本质是用少量高质量的(指令,回复)对,教会模型"响应指令"这一行为模式

关键认知(Superficial Alignment Hypothesis,来自 LIMA 论文):模型的知识和能力几乎全部来自预训练,SFT 只是把"以助手口吻响应指令"这个分布激发出来。所以 SFT 数据质量远比数量重要——LIMA 用 1000 条精选数据就达到了不错的对齐效果。这是面试必背结论,但要会辩证补充:复杂任务(数学、代码、多步推理)仍需要足够的数据量与多样性,"1000 条就够"不能绝对化。

二、数据格式:Chat Template 与 Loss Mask

SFT 数据的标准形态是多轮对话,训练前要用 chat template 拼成一条序列。以 ChatML 格式(Qwen 使用)为例:

<|im_start|>system
You are a helpful assistant.<|im_end|>
<|im_start|>user
9.11 和 9.8 哪个大?<|im_end|>
<|im_start|>assistant
9.8 更大。比较小数时先看整数部分相同,再比较十分位:8 > 1,所以 9.8 > 9.11。<|im_end|>

核心考点:loss 只算 assistant 部分。prompt(system + user)的 token 参与前向传播提供上下文,但不计入损失——实现上就是把这些位置的 label 设为 -100(PyTorch 交叉熵的 ignore_index)。为什么?因为我们不希望模型学习"生成用户问题"的分布,只希望它学会"在给定问题下生成回复"。如果不 mask,模型会浪费容量拟合用户输入分布,且训练信号被稀释。

多轮对话的组织有两种做法:

  • 整段拼接、只在各轮 assistant 部分算 loss(主流,一条样本利用所有轮次);
  • 拆成多条样本,每条只保留最后一轮 assistant 算 loss(数据量膨胀,历史轮被重复编码,效率低)。

追问点:为什么 EOS/im_end 一定要算 loss? 因为模型必须学会"什么时候停"。漏了这个细节,推理时模型会停不下来一直生成——这是真实事故高发点,也是面试官爱挖的工程细节。

三、SFT 与预训练的异同对比

维度 预训练(PT) 监督微调(SFT)
目标函数 Next Token Prediction 交叉熵 相同的交叉熵,但仅在 response 部分计算
数据规模 数T token 数万~数百万条对话
数据形态 原始文本 (instruction, response) 结构化对话
学习率 峰值 3e-4 量级 小 1~2 个数量级(1e-5 ~ 2e-5)
训练轮数 <1 epoch(数据不重复) 2~3 epochs
序列组织 文档拼接切块 按对话组织 + loss mask / packing
主要风险 loss spike、数据污染 过拟合、灾难性遗忘、幻觉加重

表里两个点常被追问。为什么 SFT 学习率要小:SFT 数据量小且分布窄,大学习率会迅速过拟合并冲掉预训练学到的通用能力(灾难性遗忘)。为什么 SFT 可能加重幻觉:如果 SFT 数据里包含模型预训练中根本没见过的知识,等于教模型"在不知道的时候也要一本正经地回答",这是 John Schulman 的著名观点——SFT 数据应尽量落在模型已有知识边界内,超出边界的问题应教它说"不知道"。

四、可运行代码:从数据构造到 Loss Mask 的完整演示

下面的代码不依赖 GPU,用 PyTorch 完整演示 SFT 样本的构造:chat template 拼接、tokenize、loss mask 生成,以及带 ignore_index 的损失计算。这段代码的逻辑与 LLaMA-Factory、trl 等主流框架内部实现一致。

import torch
import torch.nn.functional as F

# ---------- 玩具 tokenizer:字符级,仅为演示结构 ----------
class ToyTokenizer:
    def __init__(self):
        self.vocab = {"<pad>": 0, "<im_start>": 1, "<im_end>": 2}
    def encode(self, text):
        ids = []
        for ch in text:
            if ch not in self.vocab:
                self.vocab[ch] = len(self.vocab)
            ids.append(self.vocab[ch])
        return ids
    @property
    def im_start(self): return 1
    @property
    def im_end(self):   return 2

IGNORE_INDEX = -100

def build_sft_sample(tokenizer, messages, max_len=512):
    """把多轮对话拼成 (input_ids, labels),仅 assistant 内容算 loss。"""
    input_ids, labels = [], []
    for msg in messages:
        role_ids = [tokenizer.im_start] + tokenizer.encode(msg["role"] + "\n")
        content_ids = tokenizer.encode(msg["content"])
        end_ids = [tokenizer.im_end]
        input_ids += role_ids + content_ids + end_ids
        if msg["role"] == "assistant":
            # 角色头不算 loss,内容和 <im_end> 都算(模型要学会停止)
            labels += [IGNORE_INDEX] * len(role_ids) + content_ids + end_ids
        else:
            labels += [IGNORE_INDEX] * (len(role_ids) + len(content_ids) + len(end_ids))
    return input_ids[:max_len], labels[:max_len]

def sft_loss(logits, labels):
    """标准的 shift 一位交叉熵:预测下一个 token。"""
    shift_logits = logits[:, :-1, :].contiguous()
    shift_labels = labels[:, 1:].contiguous()
    return F.cross_entropy(
        shift_logits.view(-1, shift_logits.size(-1)),
        shift_labels.view(-1),
        ignore_index=IGNORE_INDEX,
    )

if __name__ == "__main__":
    tok = ToyTokenizer()
    messages = [
        {"role": "system",    "content": "你是一个有用的助手。"},
        {"role": "user",      "content": "9.11和9.8哪个大?"},
        {"role": "assistant", "content": "9.8更大,十分位8>1。"},
        {"role": "user",      "content": "谢谢"},
        {"role": "assistant", "content": "不客气!"},
    ]
    input_ids, labels = build_sft_sample(tok, messages)
    n_total = len(labels)
    n_train = sum(1 for x in labels if x != IGNORE_INDEX)
    print(f"总 token 数: {n_total}, 参与 loss 的 token 数: {n_train} "
          f"(占比 {n_train/n_total:.1%})")

    # 模拟一个随机模型前向,验证 loss 可正常计算
    vocab_size = len(tok.vocab)
    logits = torch.randn(1, n_total, vocab_size)
    loss = sft_loss(logits, torch.tensor([labels]))
    print(f"随机初始化下的 SFT loss: {loss.item():.4f} "
          f"(理论值≈ln({vocab_size})={torch.log(torch.tensor(float(vocab_size))):.4f})")

运行后能看到两个关键输出:参与 loss 的 token 占比(真实项目中这个比例太低说明 prompt 冗长、训练效率差);随机模型的 loss 约等于 ln(vocab_size),这是验证训练代码正确性的经典手段——面试聊到"怎么排查训练代码 bug"时,"检查初始 loss 是否接近 ln(V)"是非常加分的回答。

五、SFT 数据工程:好数据长什么样

数据构造的主流方法:

Self-Instruct 路线:用强模型(GPT-4 级别)从种子任务扩展生成指令与回复,Alpaca 是开山之作。廉价但有天花板——学生模型学到的是教师模型的近似,且容易继承教师的套话与偏见。

进化式增强(WizardLM 的 Evol-Instruct):对已有指令做"深度进化"(加约束、加推理步骤、复杂化)和"广度进化"(换主题),系统性提升指令复杂度分布。

人工精标:贵但质量上限最高,通常用于核心场景(安全、价值观、公司业务数据)。实践中是"合成数据打底 + 人工精标点睛"。

质量过滤的可操作指标:指令多样性(用 embedding 聚类看覆盖度)、回复长度分布(过短的敷衍回复要清掉)、IFD 指标(Instruction Following Difficulty,用模型自身 loss 筛选"有信息量"的样本)、拒答比例控制(拒答样本太多模型会变得过度保守)。

数据配比经验:通用对话、代码、数学、多语言、安全各占一定比例,业务数据不超过 30%——纯业务数据训练会让模型"变笨",通用能力回退。这条在业务落地面试题里几乎必问("给你 5 万条客服数据怎么微调",答案一定要包含"混合通用数据防遗忘")。

六、训练技巧与常见坑

Packing:把多条短样本拼进一个 max_length 序列,配合 attention mask 隔离(或 Flash Attention 的 varlen 接口),可以把训练吞吐提升数倍。追问点:不隔离会怎样?样本间会互相"看见",造成信息泄露,虽然实践中影响常常不大,但严谨做法必须隔离。

NEFTune:给 embedding 加均匀噪声的正则化技巧,一行代码在多个 benchmark 上涨点,面试提到会显得跟进前沿。

学习率与 epoch:2e-5、2~3 个 epoch 是 7B 模型全参 SFT 的常见起点;LoRA 微调学习率要放大到 1e-4 量级(可训练参数少、有效步长需求大)。判断过拟合:验证集 loss 回升、回复开始背诵训练集措辞、多样性下降。

灾难性遗忘的缓解:混入预训练数据(replay)、降低学习率、LoRA 等参数高效方法天然遗忘更少、模型融合(SFT 后与 base 加权平均)。

七、面试答题框架

被问"如何为业务场景做一次高质量 SFT",推荐框架:

  1. 界定目标:明确任务类型与成功指标(人工评估维度 + 自动指标);
  2. 数据侧:业务数据清洗精标 + 开源/合成数据补多样性,控制配比(业务 ≤30%),构造拒答与边界样本;
  3. 格式侧:统一 chat template,检查 loss mask 与 EOS,长样本截断策略;
  4. 训练侧:小学习率 + 2~3 epoch,先跑 LoRA 快速验证数据价值,再决定是否全参;
  5. 评估侧:held-out 业务评测集 + 通用能力回归测试(防遗忘),A/B 上线;
  6. 迭代:badcase 归因(数据缺失 or 能力不足 or 幻觉),针对性补数据。

自检清单:能解释为什么 mask prompt 吗?EOS 为什么要算 loss?初始 loss ≈ ln(V) 的原理?业务微调为什么要混通用数据?这四问过关,SFT 环节就是你的得分点。

下一篇进入对齐算法:RLHF 三阶段与 PPO/DPO 的原理对比与手推。

更多推荐