Skip to main content

06 - SFT 指令微调

预训练产出的叫 base 模型。这一篇把它变成能对话的模型,产出 sft.py

前置:05 篇跑完的预训练 checkpoint,01 篇训好的 tokenizer。这一篇不需要多卡,0.5B 做 SFT 单卡就够。

零、开始之前:base 模型为什么不能直接用

0.1 base 模型学会了什么

预训练的目标从头到尾只有一个:猜下一个 token。所以 base 模型学到的是这段文本后面接什么最自然

它读了 10B token 的网页,语法、常识、写作风格都有了。但它不知道「有人在问你问题,你该回答」这回事,因为训练数据里从来没有这种结构。

0.2 拿它当聊天机器人会发生什么

给 base 模型输入「中国的首都是哪里?」,它可能这样接:

中国的首都是哪里?
A. 上海
B. 北京
C. 广州
D. 深圳
答案:B
解析:北京是中华人民共和国的首都……

它没有回答你,它在续写。因为在网页语料里,这样一句话后面最常见的就是选择题选项。

也可能接出这样的:

中国的首都是哪里?这是很多小学生都会问的问题。今天我们就来聊聊……

同样是续写,不是回答。

这就是 base 模型的状态:知识都在,但不知道自己该扮演一个回答问题的角色

0.3 SFT 要解决的是格式问题,不是知识问题

这一点特别容易误解,说清楚:

SFT 教的是行为模式,不是知识。 模型的知识在预训练阶段就基本定型了,SFT 用的那几万条数据,相对 10B token 来说是九牛一毛,不可能塞进什么新知识。

SFT 让模型学会的是:看到 <|im_start|>user 开头的一段,后面跟 <|im_start|>assistant 时,我应该输出一个回答,然后停下来。

这也解释了一个常见现象:SFT 数据里出现的知识错误,模型会照单全收;但想靠 SFT 给模型灌输一个新领域的知识,基本没用。 前者是学格式时顺带记住了内容,后者是想用几万条覆盖百亿 token 建立的分布。

有篇叫 LIMA 的论文把这个观点推到极致,它只用 1000 条精心挑选的数据做 SFT,效果就不错。结论是:SFT 数据的质量远比数量重要,因为你本来就只是在教格式。

0.4 跟预训练的三处不同

预训练SFT
数据组织所有文档首尾拼成一条 token 流按样本组织,要 padding
loss 算在哪每个 token 都算只在 assistant 的回复上算
学习率3e-42e-5,小一个数量级
跑多久一遍 10B token同样的数据跑 2 到 3 轮

这一篇剩下的内容就是把这四行讲清楚。

一、对话模板

1.1 为什么需要模板

模型的输入是一串 token,本身没有「谁在说话」的概念。要让它区分用户和助手,必须在文本里用特殊标记把角色写出来

这套标记方式就是对话模板。用哪套都行,但训练和推理必须用同一套,否则模型认不出来。

1.2 ChatML

我们用 ChatML,现在最通行的一种,Qwen、多数开源模型都用它:

<|im_start|>system
你是一个乐于助人的助手。<|im_end|>
<|im_start|>user
中国的首都是哪里?<|im_end|>
<|im_start|>assistant
北京。<|im_end|>

规则很简单:每一轮是 <|im_start|> + 角色名 + 换行 + 内容 + <|im_end|>

推理的时候,把用户输入按这个格式拼好,最后补上 <|im_start|>assistant\n,模型就会接着往下生成回复,生成到 <|im_end|> 就停。

1.3 special token 现在派上用场了

<|im_start|><|im_end|> 就是 01 篇 4.8 节训 tokenizer 时留的那两个。

当时说「预训练根本用不上,但现在不加,将来加就得重训」,现在到了要用的时候。如果那时候没留,此刻加进去词表会从 32000 变成 32002,embedding 层形状对不上,预训练的权重全废。

代码里加了个断言,就是防这个:

assert self.im_start is not None and self.im_end is not None, (
"tokenizer 里没有 <|im_start|> / <|im_end|>,"
"回 01 篇 4.8 节,这两个必须在训 tokenizer 时就加进 special_tokens"
)

1.4 为什么要把这两个当成单个 token

如果不把 <|im_start|> 注册成 special token,BPE 会把它切成 <|im_start|> 一堆碎片。

坏处有两个。一是浪费,每轮对话多花十几个 token。二是不可靠,模型要学会「这七个碎片连在一起才表示角色开始」,比学一个独立符号难得多,而且用户在正常文本里打出 <|im_start|> 就能伪造角色,这是提示注入的一个口子。

注册成单个 token 之后,它在词表里有独立编号,跟任何正常文本都撞不上。

二、loss mask:这一篇最核心的改动

2.1 不该在用户的话上算 loss

一条 SFT 样本里有三部分:模板标记、用户的问题、助手的回答。

预训练是每个 token 都算 loss。但 SFT 不能这样,只有助手的回答该算

2.2 为什么

在用户的问题上算 loss,等于在教模型「怎么提问」。

后果是模型学会了模仿用户的说话方式,推理时容易自问自答:你问一句,它回答完接着又冒出一个新问题然后自己回答。这是 SFT 最典型的翻车形式之一。

而且这部分 loss 是纯噪声。用户会问什么是不可预测的,让模型去拟合它没有任何意义,只会稀释真正有用的梯度信号。

2.3 怎么做

PyTorch 的 F.cross_entropy 有个 ignore_index 参数,标成这个值的位置直接跳过。约定俗成用 -100

IGNORE = -100

构造 labels 时,非助手回复的位置全填 -100

输入   <|im_start|>user \n 中国的首都是哪里? <|im_end|> \n <|im_start|>assistant \n 北京。 <|im_end|>
labels -100 -100 -100 -100... -100 -100 -100 -100 -100 北京。 <|im_end|>

2.4 代码

def encode(self, messages: list) -> tuple:
ids, labels = [], []
for m in messages:
role = self.tok.encode(m["role"]).ids
content = self.tok.encode(m["content"]).ids

# 头部 <|im_start|>{role}\n 不算 loss
head = [self.im_start] + role + self.nl
ids += head
labels += [IGNORE] * len(head)

body = content + [self.im_end] + self.nl
ids += body
if m["role"] == "assistant":
labels += body # 只有这一支真的填 labels
else:
labels += [IGNORE] * len(body)

return ids, labels

注意 <|im_end|> 是算 loss 的。 这一点很关键:模型必须学会「回答完了要输出结束符」,否则推理时它会一直说下去停不下来。如果把 <|im_end|> 也 mask 掉,训出来的模型会滔滔不绝,这是另一种典型翻车。

2.5 算 loss 时还要错位一格

跟预训练一样,用第 t 个位置的输出预测第 t+1 个 token,所以要错位:

def sft_loss(logits, labels):
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)

预训练时这个错位是在 get_batch 里做的(yx 右移一位),SFT 因为要配合 mask,在算 loss 时做更清楚。

2.6 必须把 mask 打印出来看一眼

这是这一篇最该照做的一条建议。loss mask 写错了不会报错,训练照跑,loss 照降,只是降的是错的东西。

所以写完先可视化验证:

def inspect_masking(dataset, tok, n=2):
for i in range(min(n, len(dataset))):
ids, labels = dataset[i]
kept = [t for t, l in zip(ids, labels) if l != IGNORE]
print("完整输入:", tok.decode(ids)[:300])
print("算 loss 的部分:", tok.decode(kept)[:300])
print(f"占比 {len(kept)}/{len(ids)} = {len(kept) / len(ids):.1%}")
python sft.py --ckpt out/ckpt.pt --data data/sft.jsonl --inspect_only

「算 loss 的部分」打印出来应该只有助手的回答加上结束符,一个字的用户输入都不该有。

顺便看那个占比。正常在 30% 到 60% 之间。明显偏高说明 mask 没生效,明显偏低说明可能把回复也 mask 掉了。

三、数据组织

3.1 不能再首尾拼接

01 篇 5.3 节说过,预训练把所有文档拼成一条流,随机取起点,不管文档边界。SFT 不行。

原因是一条对话必须完整。截断了的样本,模型看到半句问题就要生成回答,学到的是错的模式。而且 loss mask 是按样本结构算的,拼接之后结构就乱了。

所以 SFT 回到常规做法:一条样本一条,不够长的 padding 补齐。

3.2 右侧 padding 配因果注意力,不需要 padding mask

这里有个值得讲清楚的细节。padding 补进去的位置是垃圾数据,直觉上应该用 attention mask 屏蔽掉。但训练时用右侧 padding 的话,其实不需要

理由两条:

真实 token 在前,padding 在后。02 篇的因果掩码保证每个位置只能看到自己和前面,所以真实 token 永远不会注意到后面的 padding。

padding 位置自己的输出确实是垃圾,但它们在 labels 里是 IGNORE,不进 loss。

所以右侧 padding 加因果注意力,padding 天然就被隔离了。

什么时候需要 attention mask:推理时批量生成要用左侧 padding(因为生成是从序列末尾往后接,末尾必须对齐),那时候 padding 在前面,真实 token 会注意到它们,必须显式屏蔽。

代码里 collate 仍然返回了 attn,但训练循环没用它,就是为了让这个区别显式可见。

def collate(batch, pad_id: int):
maxlen = max(len(ids) for ids, _ in batch)
input_ids, labels, attn = [], [], []
for ids, lab in batch:
pad = maxlen - len(ids)
input_ids.append(ids + [pad_id] * pad)
labels.append(lab + [IGNORE] * pad)
attn.append([1] * len(ids) + [0] * pad)
return torch.tensor(input_ids), torch.tensor(labels), torch.tensor(attn)

padding 到本 batch 内的最长,不是 padding 到 max_len。差别很大:按 batch 内最长补,平均只多算百分之几十;按 2048 补,短样本会浪费十几倍算力。

3.3 超长样本直接丢,不要截断

if len(ids) > max_len:
skipped += 1
continue

截断会把助手的回答砍掉一半,模型学到「回答可以不说完」,比丢掉这条样本糟糕得多。

统计跳过了多少条,如果比例很高(超过百分之几),说明 max_len 设小了。

四、指令数据从哪来

4.1 几个常用的开源数据集

以下行数和体积是 2026-08-19 用 HuggingFace datasets-server 实测的:

数据集行数体积许可说明
HuggingFaceH4/ultrachat_200k515,3111,624.0 MBMIT多轮对话,质量高,英文
shibing624/sharegpt_gpt4103,415602.3 MBCC-BY-4.0GPT-4 生成的中文多轮
tatsu-lab/alpaca52,00224.2 MBCC-BY-NC-4.0经典单轮,注意是非商用
llamafactory/alpaca_gpt4_zh42,67727.9 MBApache-2.0Alpaca 的 GPT-4 中文版
GAIR/lima1,000otherLIMA 论文那 1000 条,质量标杆

选的时候注意许可证。tatsu-lab/alpaca 是 CC-BY-NC,只能非商用,很多人没留意就用了。

4.2 要多少条

按 0.3 节说的,SFT 是教格式,几万条足够。

我的建议是从两三万条混合数据起步,中英各半(取决于你 tokenizer 的语言配比)。不够再加,而不是一上来就堆几十万条。

理由是 SFT 阶段调试的循环要快。0.5B 模型两万条数据跑 3 轮,单卡几十分钟就完事,能快速看到效果、快速调整。堆到几十万条,一轮就要好几小时,反馈太慢。

4.3 数据格式

统一成 jsonl,每行一条:

{"messages": [{"role": "user", "content": "中国的首都是哪里?"}, {"role": "assistant", "content": "北京。"}]}

各个开源数据集的原始格式五花八门(Alpaca 是 instruction/input/output,ShareGPT 是 conversations 数组),先写个转换脚本统一成上面这种,后面的代码才好写。

五、超参怎么变

超参预训练SFT为什么
学习率3e-42e-5小一个数量级
epoch相当于不到 1 轮2 到 3 轮数据少,要多看几遍
warmup2%3%数据少,比例稍高
weight decay0.10不需要正则
batch size1M token十几条样本数据量差三个数量级

学习率是这里最要紧的。 预训练是从随机初始化开始学,要大步走。SFT 是在一个已经学好的模型上做微调,目标只是调整输出格式,大学习率会把预训练学到的能力冲掉,这个现象叫灾难性遗忘。表现是模型学会了对话格式,但一问知识就胡说八道。

2e-5 是通行值。不确定的话宁可再小一点,1e-5 也行。

epoch 不要多。 数据只有几万条,跑太多轮会过拟合,模型开始逐字背诵训练数据。一般 2 到 3 轮,边跑边看验证集,效果不涨就停。

六、LoRA

6.1 全参微调的问题

上面讲的都是全参微调,所有 502M 参数都参与更新。对 0.5B 来说这完全可行,02 篇算过静态显存 8 GB。

但如果基座换成 7B、13B,全参微调的显存就吃不消了,而且每个下游任务都要存一份完整的模型权重。

6.2 LoRA 的想法

LoRA 的观察是:微调对权重的改动量 ΔW\Delta W,其实是低秩的。也就是说 ΔW\Delta W 虽然是个大矩阵,但它包含的有效信息很少,可以用两个瘦长矩阵的乘积近似。

于是做法变成:冻结原权重 WW,在旁边加一条支路:

h=Wx+BAx,ARr×din, BRdout×rh = Wx + BAx, \quad A \in \mathbb{R}^{r \times d_{\text{in}}},\ B \in \mathbb{R}^{d_{\text{out}} \times r}

其中 rr 远小于 dind_{\text{in}}doutd_{\text{out}},叫。训练时只更新 AABB

AA 用随机初始化,BB 初始化成全 0,这样训练开始时 BA=0BA = 0,模型输出和原来完全一样,是个平滑的起点。

推理时可以把 BABA 加回 WW,合并成一个矩阵,推理速度和原模型完全一样,没有额外开销。这是 LoRA 相比其他微调方法的一个重要优点。

6.3 参数量算一算

原矩阵是 din×doutd_{\text{in}} \times d_{\text{out}},LoRA 是 r×(din+dout)r \times (d_{\text{in}} + d_{\text{out}})

code/lora_math.py 的实际输出:

加在哪                          r          可训参数      占基座
只 attention (q,k,v,o) 8 1,474,560 0.29%
只 attention (q,k,v,o) 16 2,949,120 0.59%
只 attention (q,k,v,o) 64 11,796,480 2.35%

全部七个 8 3,907,584 0.78%
全部七个 16 7,815,168 1.56%
全部七个 64 31,260,672 6.22%

显存对比(只看优化器状态,AdamW 每参数 12 字节)
全参微调 6.03 GB
只 attention (q,k,v,o) 0.035 GB (r=16)
全部七个 0.094 GB (r=16)

r=16 加在 attention 上,只训 2,949,120 个参数,占基座的 0.59%,优化器状态从 6.03 GB 降到 0.035 GB。

注意最后那句:LoRA 省的是优化器状态和梯度,基座权重仍然要完整放在显存里。 所以它省的不是全部显存,对 0.5B 来说 1 GB 的 bf16 权重还是躲不掉。

6.4 加在哪些矩阵上,秩取多少

原论文只加在 q_projv_proj 上。后来的实践发现加在全部七个投影上效果更好,代价是参数量翻一倍多,但绝对值仍然很小。

秩的选择:

r适用情况
4 到 8任务简单,只是调格式
16 到 32通用起点,大多数情况用这个
64 以上任务和预训练分布差异大,比如换语言、换领域

alpha 参数是缩放系数,实际的更新量是 αrBA\frac{\alpha}{r} BA。通行做法是设 alpha = 2r,改秩的时候等效学习率不用跟着改。

6.5 什么时候用哪个

情况选择
0.5B 这个规模,显存够全参微调,效果上限更高
基座 7B 以上,单卡显存紧LoRA
要为很多任务各存一份LoRA,每份只有几 MB
想让模型学新领域知识都不行,回 0.3 节

对我们这个专题,推荐全参微调。 0.5B 全参跑得动,效果更好,而且流程更简单。LoRA 值得做一次对照实验,看看同样数据下两者差多少,这个对比本身有价值。

实现上不用自己写,peft 库几行就能包上:

from peft import LoraConfig, get_peft_model

cfg = LoraConfig(
r=16, lora_alpha=32, lora_dropout=0.05,
target_modules=["q_proj", "k_proj", "v_proj", "o_proj",
"gate_proj", "up_proj", "down_proj"],
)
model = get_peft_model(model, cfg)
model.print_trainable_parameters()

最后那行会打印可训练参数占比,拿它跟 lora_math.py 算的对一遍,能确认 target_modules 写对了。

七、验收

  • --inspect_only 看过 mask,算 loss 的部分只有助手回复
  • mask 占比在 30% 到 60% 之间
  • <|im_end|> 确认在算 loss 的范围内
  • 超长被跳过的样本比例不高
  • SFT 初始 loss 明显低于 10.37(base 模型已经会说话了,通常在 2 到 4)
  • 训完问一个问题,模型给的是回答不是续写
  • 模型会在回答完之后停下来,不会一直说
  • 不会自问自答
  • 随便问几个常识问题,确认预训练的知识没被冲掉

第 5 条值得说一下。SFT 的初始 loss 不会是 10.37,因为模型不是随机初始化的。如果初始 loss 真的接近 10.37,说明 checkpoint 没加载成功,在拿一个随机模型做 SFT。

第 6 到 8 条是三个典型翻车的对应检查,各自对应 0.2 节、2.4 节、2.2 节讲的问题。

八、常见问题

现象大概率是什么原因
模型自问自答在用户输入上算了 loss,回 2.2 节
模型停不下来一直说<|im_end|> 被 mask 掉了,回 2.4 节
学会对话但知识全忘了学习率太大,从 2e-5 往下调
逐字背诵训练数据epoch 太多,减到 2 轮
初始 loss 就是 10.37checkpoint 没加载上,在训随机模型
断言报没有 <|im_start|>tokenizer 没留 special token,回 01 篇 4.8 节
推理时输出一堆 <|im 碎片special token 没注册成单个 token,回 1.4 节
mask 占比接近 100%mask 完全没生效,labels 直接拿 ids 填了
批量推理结果和单条不一致推理用了右侧 padding,应该左侧,回 3.2 节
先跑通再谈效果

0.5B 的模型做完 SFT,对话能力也就是「能对上话」的程度,别期待它能写代码或者做推理题。这一篇的验收标准是行为对了:回答而不续写、说完会停、不自问自答。至于回答得好不好,那是模型规模决定的,不是 SFT 能解决的。