【大模型】- 指令微调(SFT)

算法

指令微调(SFT)

使语言模型遵循人类指令

类型: 学习 | 语言: Python | 🏷 前置:《指令微调(SFT)》(本系列第 5 篇)

学习目标

  • 理解指令微调的目的和方法
  • 实现一个用于 SFT 的数据集和数据加载器
  • 使用 LoRA 高效地微调大型模型
  • 监控 SFT 训练过程
  • 评估模型的指令遵循能力

为什么需要指令微调

预训练模型学会了预测下一个 token,但不一定会遵循指令。指令微调(SFT)教会模型:

  • 理解指令格式
  • 以有帮助的方式回答问题
  • 遵循特定的行为准则

指令数据格式

指令数据通常有以下格式:

1
2
3
4
5
{
"instruction": "总结以下文本",
"input": "人工智能正在改变世界...",
"output": "人工智能正在对各个领域产生重大影响。"
}

数据集实现

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
from torch.utils.data import Dataset, DataLoader
from transformers import AutoTokenizer

class InstructionDataset(Dataset):
def __init__(self, data, tokenizer, max_length=512):
self.data = data
self.tokenizer = tokenizer
self.max_length = max_length

def __len__(self):
return len(self.data)

def __getitem__(self, idx):
item = self.data[idx]

# 格式化提示
if item["input"]:
prompt = f"### 指令:\n{item['instruction']}\n\n### 输入:\n{item['input']}\n\n### 回答:\n"
else:
prompt = f"### 指令:\n{item['instruction']}\n\n### 回答:\n"

full_text = prompt + item["output"] + self.tokenizer.eos_token

# 分词
encodings = self.tokenizer(
full_text,
max_length=self.max_length,
padding="max_length",
truncation=True,
return_tensors="pt"
)

# 创建标签(只在回答部分计算损失)
input_ids = encodings["input_ids"].squeeze()
labels = input_ids.clone()

# 将指令部分的标签设为 -100
prompt_length = len(self.tokenizer(prompt)["input_ids"])
labels[:prompt_length] = -100

return {
"input_ids": input_ids,
"attention_mask": encodings["attention_mask"].squeeze(),
"labels": labels
}

使用 LoRA 微调

LoRA(低秩适应)是一种参数高效的微调方法:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
from peft import LoraConfig, get_peft_model, TaskType

def apply_lora(model):
lora_config = LoraConfig(
task_type=TaskType.CAUSAL_LM,
inference_mode=False,
r=8, # LoRA 秩
lora_alpha=32, # 缩放因子
lora_dropout=0.1,
target_modules=["q_proj", "v_proj"] # 要适配的层
)

model = get_peft_model(model, lora_config)
model.print_trainable_parameters()
# trainable params: 4,718,592 || all params: 124,439,808 || 3.79%

return model

训练循环

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
from transformers import get_linear_schedule_with_warmup

def train_sft(model, train_dataloader, val_dataloader, config):
model = apply_lora(model)
model = model.to(config["device"])

optimizer = torch.optim.AdamW(
model.parameters(),
lr=config["lr"],
weight_decay=config["weight_decay"]
)

# 学习率调度
num_training_steps = len(train_dataloader) * config["epochs"]
scheduler = get_linear_schedule_with_warmup(
optimizer,
num_warmup_steps=config["warmup_steps"],
num_training_steps=num_training_steps
)

for epoch in range(config["epochs"]):
model.train()
total_loss = 0

for batch in train_dataloader:
# 移到设备
input_ids = batch["input_ids"].to(config["device"])
attention_mask = batch["attention_mask"].to(config["device"])
labels = batch["labels"].to(config["device"])

# 前向传播
outputs = model(
input_ids=input_ids,
attention_mask=attention_mask,
labels=labels
)

loss = outputs.loss

# 反向传播
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
scheduler.step()
optimizer.zero_grad()

total_loss += loss.item()

# 验证
val_loss = evaluate(model, val_dataloader, config)

print(f"Epoch {epoch + 1}:")
print(f" Train Loss: {total_loss / len(train_dataloader):.4f}")
print(f" Val Loss: {val_loss:.4f}")

return model

评估

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
def evaluate(model, dataloader, config):
model.eval()
total_loss = 0

with torch.no_grad():
for batch in dataloader:
input_ids = batch["input_ids"].to(config["device"])
attention_mask = batch["attention_mask"].to(config["device"])
labels = batch["labels"].to(config["device"])

outputs = model(
input_ids=input_ids,
attention_mask=attention_mask,
labels=labels
)

total_loss += outputs.loss.item()

return total_loss / len(dataloader)

def generate_response(model, tokenizer, instruction, input_text=""):
model.eval()

if input_text:
prompt = f"### 指令:\n{instruction}\n\n### 输入:\n{input_text}\n\n### 回答:\n"
else:
prompt = f"### 指令:\n{instruction}\n\n### 回答:\n"

inputs = tokenizer(prompt, return_tensors="pt").to(model.device)

with torch.no_grad():
outputs = model.generate(
**inputs,
max_new_tokens=512,
temperature=0.7,
do_sample=True,
top_p=0.9
)

response = tokenizer.decode(outputs[0], skip_special_tokens=True)
return response.split("### 回答:\n")[-1].strip()

常见问题

过拟合

  • 使用 dropout 和权重衰减
  • 增加数据集大小
  • 减少训练轮数

灾难性遗忘

  • 使用较小的学习率
  • 应用 LoRA 或适配器
  • 混合原始预训练数据

指令不一致

  • 标准化数据格式
  • 清理低质量样本
  • 平衡不同任务的数据

保存和加载

1
2
3
4
5
6
7
8
9
10
# 保存 LoRA 权重
model.save_pretrained("./sft_model")

# 加载
from peft import PeftModel
base_model = AutoModelForCausalLM.from_pretrained("base_model")
model = PeftModel.from_pretrained(base_model, "./sft_model")

# 可选:合并权重
model = model.merge_and_unload()

总结

指令微调(SFT)使预训练模型能够遵循人类指令。使用 LoRA 等参数高效方法可以在单个 GPU 上微调大型模型。关键是高质量的指令数据和正确的训练策略。

下一步

下一课将介绍 RLHF(基于人类反馈的强化学习),使模型更符合人类偏好。

📚 本文改编自 AI Engineering from Scratch(MIT License · 作者 Rohit Ghumare),中文内容来自官方中文镜像。原课程共 503 课 · 20 阶段 · 免费开源,教程网站见 aiengineeringfromscratch.com

📝 自我检查(课程配套测验)

Q1(学前) 基础语言模型与指令微调模型的根本区别是什么?

A. 架构不同
B. 基础模型续写文本;指令微调模型学会遵循指令并以对话方式回答
C. 指令微调模型更大
D. 基础模型不能生成文本

答案: B 解析: 基础模型(如 GPT-3 base)训练目标是下一 token 预测,会续写任何 prompt。SFT 在(指令,回答)对上微调,教模型理解并执行用户指令。

Q2(学前) SFT 训练数据通常是什么格式?

A. 原始网页文本
B. (指令,理想回答)对,常带 system prompt 和对话结构
C. 仅代码片段
D. 带标签的图像

答案: B 解析: SFT 数据是 curated 的指令-回答对,如「写一首关于 AI 的俳句」→「硅思流转…」。模型学习将指令映射到有用、格式正确的回答。

Q3(学后) SFT 中为什么要 masking prompt token 的损失?

A. 加快训练
B. 只对模型生成的回答 token 计算损失,不对 prompt token 计算,使模型学习生成而非重复指令
C. 减少内存
D. 防止过拟合

答案: B 解析: 损失只在回答 token 上计算。若对 prompt token 也算损失,模型会学习重复指令而非生成回答。这是 instruction tuning 的标准做法。

Q4(学后) SFT 通常需要多少训练数据?

A. 与预训练相同(数万亿 token)
B. 通常 1 万到 10 万条 curated 指令-回答对,远少于预训练
C. 至少 10 亿样本
D. 一条样本就够

答案: B 解析: SFT 从已具备语言能力的预训练模型出发,只需数千到数万高质量示例即可显著改变行为。质量比数量重要——1 万条 curated 对胜过 100 万条低质量对。

Q5(学后) SFT 之后模型仍可能有什么问题?

A. 无法生成任何文本
B. 可能产生有害、有偏见或不真实的回答;需要 RLHF/DPO 等对齐进一步改善
C. 上下文窗口缩小
D. 失去所有预训练知识

答案: B 解析: SFT 教模型遵循指令,但不保证安全、真实或无害。模型可能学会遵循有害指令或产生自信但错误的回答。对齐方法(RLHF、DPO)在此基础上进一步改善行为。

  • 标题: 【大模型】- 指令微调(SFT)
  • 作者:
  • 创建于 : 2026-08-19 09:06:00
  • 更新于 : 2026-08-21 16:20:11
  • 链接: https://sxl-space.tk/2026/08/19/010_LLM/010_LLM-06-InstructionTuningSFT/
  • 版权声明: 版权所有 © 宋,禁止转载。