Skip to content

03 监督微调 SFT

学习理念:SFT 的核心不是"训练一个模型",而是"让已经懂语言的模型学会执行指令"。关键在于 chat template 格式化对话、answer mask 只计算 assistant 回答的 loss、以及完整的训练循环。如果只记一件事:SFT 只对 assistant 的回答计算损失,user 和 system 部分不参与。

海外对标:OpenAI InstructGPT SFT 阶段、HuggingFace TRL SFTTrainer

本节 AI 替代率:~70% | 人工干预率:~30%

角色能力范围
🤖 AI 擅长生成 SFT 训练循环、answer mask 计算、loss 函数
👤 人类需理解为什么只对 assistant 部分算 loss、chat template 的作用、padding + mask 的配合

📌 来源说明:以下内容提取自原始笔记第5章(LLM微调)+ 第6章(SFT实现),代码来自 02_sft_train.py

📖 阅读优先级

等级章节说明
🔴 必须深入二~四、Chat Template + Answer Mask选 Base/Instruct 模型、理解训练数据格式
🟡 理解即可五~六、Loss + 训练循环知道 SFT 只算 assistant 的 loss 即可
🟢 直接跳过七、推理验证和平时用 model.generate() 一样

一、微调价值与时机

SFT训练流程

什么时候需要微调?

  • ❌ Prompt Engineering + RAG 能解决的 → 不需要微调
  • ✅ 模型在目标任务上表现不稳定,且优化 prompt 仍不够
  • ✅ 输出格式/风格有严格规范要求
  • ✅ 使用小模型微调替代大模型调用来降低成本

Base Model vs Instruct Model 选择

起点适用场景数据量难度
Instruct Model ✅ 推荐大多数场景(客服、生成、对话)数千条
Base Model任务高度特殊,Instruct 行为反而干扰数万条

二、SFT 整体流程

模型选择 → 数据准备 → 微调训练 → 模型验证
    │           │           │          │
    ├ Qwen3     ├ chat      ├ loss     ├ 验证集
    ├ LLaMA     ├ template  ├ backward ├ 样例测试
    └ DeepSeek  ├ answer    ├ optimizer └ 业务评估
                └ mask      └ save

三、数据准备(Chat Template)

Chat Template转换

大模型训练/推理时看到的不是 JSON 消息列表,而是一段连续文本。Chat Template 负责把消息列表转换成模型能理解的格式:

消息列表:                          →  Chat Template 转换 →  模型看到的文本:
[
  {"role": "user", "content": "你好"},
  {"role": "assistant", "content": "你好!有什么可以帮助你的吗?"}
]

→ <|im_start|>user\n你好<|im_end|>\n<|im_start|>assistant\n你好!有什么可以帮助你的吗?<|im_end|>

训练数据准备

python
def get_data_ultrachat_200k(config):
    ultrachat_200k_data = datasets.load_dataset("./data/ultrachat_200k")
    train_data = []
    i = 0
    while True:
        data = ultrachat_200k_data["train_sft"][i]["messages"]
        # 插入 system prompt
        data.insert(0, {"role": "system", "content": "You are a helpful assistant."})
        # 使用 tokenizer 自带的 chat template 转换
        input_ids = tokenizer.apply_chat_template(
            data, tokenize=True,
            add_generation_prompt=False,  # SFT 训练不用加生成标记
            truncation=True, max_length=2500)
        train_data.append(input_ids)
        i += 1
        if i == config.train_data_size:
            break
    return train_data

四、Answer Mask

🔥 【P0 必须理解】 SFT 的 loss 只计算 assistant 回答的部分。因为:

  • 我们想让模型学会"在这个 user 输入下,输出正确的 answer"
  • user 和 system 部分是输入条件,不需要模型去预测
  • 如果计算 user 部分的 loss,模型会被鼓励去预测 user 的输入,毫无意义

Answer Mask示意

python
def create_answer_mask(input_ids, tokenizer):
    """从 input_ids 中找出 assistant 回答的部分,设为 1,其余为 0"""
    answer_mask = torch.zeros_like(input_ids)
    eos_token_id = tokenizer.encode('<|im_end|>')[0]

    for idx, ids in enumerate(input_ids):
        eos_position = torch.where(ids == eos_token_id)[0].tolist()
        eos_position = eos_position[1:]                    # 跳过 system 的 eos
        user_ends, assistant_ends = _parse_conversation_turns(eos_position)
        _set_answer_masks(answer_mask[idx], user_ends, assistant_ends)

    return answer_mask

对话结构与 mask 对应:

<|im_start|>system\n...<|im_end|>\n    ← mask=0(system)
<|im_start|>user\n...<|im_end|>\n       ← mask=0(user)
<|im_start|>assistant\n...<|im_end|>\n  ← mask=1(assistant ✅)

Answer Mask处理流程

五、损失函数

🟡 【P1 看注释就行】 SFT 就是最大化"模型在正确答案上的对数概率"——即最大似然估计的负对数形式。

损失计算示意

python
def compute_loss(output_logits, target_labels, assistant_answer_mask):
    # 1. softmax → log 概率
    log_probabilities = torch.log_softmax(output_logits, dim=-1)
    # 2. 取出 target token 位置的对数概率
    gathered_log_probs = torch.gather(log_probabilities, dim=-1,
                                      index=target_labels.unsqueeze(-1))
    # 3. 负对数似然
    token_losses = gathered_log_probs.squeeze(-1) * (-1)
    # 4. 只保留 assistant 部分的 loss
    masked_token_losses = torch.mul(token_losses, assistant_answer_mask)
    # 5. 平均
    valid_token_count = assistant_answer_mask.sum()
    average_loss = masked_token_losses.sum() / valid_token_count
    return average_loss

六、训练循环

🔥 【P0 必须理解】 大模型训练和传统 DL 一样:前向 → loss → 反向 → optimizer。额外多了 padding(因序列长度不同)+ mask 处理。学习率用余弦衰减 + warmup。

python
def train(model, config, tokenizer):
    train_data = get_data_ultrachat_200k(config)
    model.train()
    optimizer = torch.optim.AdamW(model.parameters(), lr=config.max_lr)

    for step in range(total_steps):
        # 1. 取 batch + padding
        batch = train_data[step*config.batch_size:(step+1)*config.batch_size]
        max_len = max(len(ids) for ids in batch)
        padded_ids = [pad(ids, max_len, pad_token_id) for ids in batch]

        batch_input_tensor = torch.tensor(padded_ids)
        model_inputs = batch_input_tensor[:, :-1]   # 输入:[:-1]
        target_labels = batch_input_tensor[:, 1:]    # 标签:[1:]

        # 2. mask 组合:padding mask & assistant mask
        padding_mask = torch.where(target_labels == pad_token_id, 0, 1)
        assistant_answer_mask = create_answer_mask(model_inputs, tokenizer)
        final_loss_mask = assistant_answer_mask & padding_mask

        # 3. 前向 + loss
        model_logits = model(model_inputs).logits
        step_loss = compute_loss(model_logits, target_labels, final_loss_mask)

        # 4. 反向 + 更新
        step_loss.backward()
        optimizer.step()
        optimizer.zero_grad()

七、推理验证

训练完成后,用 model.generate() 测试:

python
inputs = tokenizer.apply_chat_template(
    [{"role": "user", "content": prompt}],
    tokenize=True, add_generation_prompt=True, return_tensors="pt")
output_ids = model.generate(inputs, max_new_tokens=5000)
print(tokenizer.decode(output_ids[0]))

八、本阶段文件索引

优先级文件路径
🔥 P002_sft_train.py3.代码/fine_tune_proj/02_sft_train.py
🟡 P105_sft_demo.ipynb3.代码/fine_tune_proj/05_sft_demo.ipynb

OPC 超级个体实战指南