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() 一样 |
一、微调价值与时机

什么时候需要微调?
- ❌ 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)

大模型训练/推理时看到的不是 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 的输入,毫无意义

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 ✅)
五、损失函数
🟡 【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]))八、本阶段文件索引
| 优先级 | 文件 | 路径 |
|---|---|---|
| 🔥 P0 | 02_sft_train.py | 3.代码/fine_tune_proj/02_sft_train.py |
| 🟡 P1 | 05_sft_demo.ipynb | 3.代码/fine_tune_proj/05_sft_demo.ipynb |