本节摘要:本节把前两节的原理装进一份完整的训练脚本——HuggingFace Trainer 风格(标注:写法示意,教学用;nanochat 原生是自带极简训练循环,接口以官方仓库 README 为准)。脚本串起全链条:读 JSONL → 3.1 的掩码构造 → 整理 batch → Trainer 配 3.2 的超参 → 存 checkpoint;约 60 行,124M 模型在 24GB 卡上单卡数小时可跑完数万条数据。下半场是三故障排查表:显存不足(CUDA OOM 的六连降级)、loss 不降(从掩码到数据的五嫌疑)、灾难性遗忘(英文能力滑坡的防与治)。跑通的那一刻,回到 0.1 节做行为复验。
阅读完本节,你应当能够:
# train_sft_hf.py —— HuggingFace Trainer 风格 SFT 脚本(写法示意,以官方仓库 README 为准) import json import torch from torch.utils.data import Dataset from transformers import (AutoModelForCausalLM, AutoTokenizer, Trainer, TrainingArguments) from build_sft_example import build_example # 3.1 节的掩码构造 TOK_PATH, BASE_PATH = "base/tokenizer", "base/model" TRAIN_JSONL, VAL_JSONL = "data_sft/train.jsonl", "data_sft/val.jsonl" OUT_DIR = "runs/sft_124m" def load_jsonl(path: str) -> list[dict]: return [json.loads(l) for l in open(path, encoding="utf-8")] class SftDataset(Dataset): def __init__(self, rows: list[dict]): self.items = [build_example(r["messages"]) for r in rows] def __len__(self): return len(self.items) def __getitem__(self, i): return self.items[i] def collate(batch: list[dict]) -> dict: """右侧 pad;labels 的 pad 值用 -100(不计 loss)。""" pad_id = tok.pad_token_id or tok.eos_token_id input_ids = torch.nn.utils.rnn.pad_sequence( [b["input_ids"] for b in batch], batch_first=True, padding_value=pad_id) labels = torch.nn.utils.rnn.pad_sequence( [b["labels"] for b in batch], batch_first=True, padding_value=-100) return {"input_ids": input_ids, "labels": labels, "attention_mask": (input_ids != pad_id).long()} tok = AutoTokenizer.from_pretrained(TOK_PATH) model = AutoModelForCausalLM.from_pretrained(BASE_PATH, dtype=torch.bfloat16) model.config.use_cache = False # 训练时关 KV cache,省显存 args = TrainingArguments( output_dir=OUT_DIR, per_device_train_batch_size=8, # 0.2 节显存估算的落点 gradient_accumulation_steps=4, # 有效 batch = 32(3.2 量级表) learning_rate=2e-5, # 3.2:预训练的 1/5~1/10 num_train_epochs=3, # 3.2:2~3 封顶,盯 val 拐点提前收 lr_scheduler_type="cosine", warmup_ratio=0.03, bf16=True, # 0.2 三开关之一 gradient_checkpointing=True, # 0.2 三开关之二 logging_steps=20, eval_strategy="steps", eval_steps=200, save_strategy="steps", save_steps=200, save_total_limit=5, # 留足回退空间(3.2 过拟合处置) report_to="none", ) trainer = Trainer(model=model, args=args, train_dataset=SftDataset(load_jsonl(TRAIN_JSONL)), eval_dataset=SftDataset(load_jsonl(VAL_JSONL)), data_collator=collate) trainer.train() trainer.save_model(f"{OUT_DIR}/final") tok.save_pretrained(f"{OUT_DIR}/final")
四处与前面章节的钩子(改脚本前先认它们):build_example 是 3.1 的掩码;gradient_accumulation_steps × batch = 32 是 3.2 的有效 batch;learning_rate=2e-5 是 3.2 的轻力度;save_steps=200 是 3.2 的"回退用存档"。验证集从训练数据里切 1%~2% 即可(2.2 产出时顺手分好)。
💡 nanochat 原生实现走的是"自带极简循环"路线(数据提前离线 token化、填充对齐 2048 边界、单文件训练脚本),没有 Trainer 依赖——教学取舍不同,原理与本脚本一一对应。要对照原味实现,读官方仓库的 SFT 脚本(当前 master 里对应
chat_sft.py一类命名,以 README 为准)。
六连降级,按序尝试,从上到下代价递增:
| 序 | 手段 | 代价 | 备注 |
|---|---|---|---|
| 1 | 减 per_device_train_batch_size,等量加 gradient_accumulation_steps |
速度略降 | 有效 batch 不变,首选 |
| 2 | 开 gradient_checkpointing |
慢约 30% | 0.2 三开关之一 |
| 3 | 确认 bf16=True、use_cache=False |
几乎无 | 两个最常见的"忘了关/忘开" |
| 4 | 减 max_seq_len(截到 p90 长度) |
长样本受损 | 长度预算见 2.1 末 |
| 5 | 换 8-bit 优化器(如 bitsandbytes 的 AdamW8bit) |
精度略降 | Adam 状态省 4 份参数量 |
| 6 | 换更大显存的卡 / 多卡 | 钱 | 124M 很少走到这步 |
| 症状 | 首要嫌疑 | 处置 |
|---|---|---|
| 起点 7~10(瞎猜量级) | 掩码错位:labels 与 input_ids 没对齐 | 重跑 3.1 的前向验收;打印一条样本逐 token 对照 |
| 起点 <0.1 | 模型看见了答案(泄漏) | 查 collate 是否把 labels 当输入拼了 |
| 起点正常但几百步不降 | lr 过小 / 数据全是重复模板 | lr 升一档试跑;回 2.2 查去重 |
| train 降、生成无行为变化 | 掩码把回答也 mask 掉了(全 -100) | 打印 labels 非零比例,应约 20%~60% |
| loss 剧烈震荡 | warmup 太短 / batch 太小 | warmup 加长;累积步数加大 |
典型表现:SFT 后英文问答能力明显变差、代码能力消失、只会说中文聊天腔。
| 处置 | 做法 | 适用 |
|---|---|---|
| 数据回掺 | SFT 数据里掺 10%~30% 预训练语料(或英文开源集) | 首选,2.3 的 70/30 配比正是此理 |
| 减力度 | lr 减半、epoch 减一 | 忘记不严重时 |
| 回退 | 用行为最均衡的中间 checkpoint | 应急 |
| 上手段 | LoRA 等参数高效微调(冻结主干) | 数据大、反复迭代时(超出本书,详见《Happy-LLM:从零训练大语言模型》微调章节思路) |
训练完成的验收仪式——回到 0.1 节的那个问题:
# chat_smoke_test.py —— SFT 完成的行为复验(写法示意) import torch from transformers import AutoModelForCausalLM, AutoTokenizer tok = AutoTokenizer.from_pretrained("runs/sft_124m/final") model = AutoModelForCausalLM.from_pretrained("runs/sft_124m/final", dtype=torch.bfloat16) ids = tok.apply_chat_template( [{"role": "user", "content": "中国最长的河流是哪一条?"}], tokenize=True, add_generation_prompt=True, return_tensors="pt") out = model.generate(ids, max_new_tokens=100, do_sample=True, temperature=0.8, top_p=0.9) print(tok.decode(out[0][ids.shape[1]:], skip_special_tokens=False)) # 期望:一句回答 + <|endoftext|>,然后停 —— 与 0.1 节 base 模型的"续写不止"对照
三问三答都正常(停下、身份、听懂指令),SFT 主干即告完成。生成参数 temperature=0.8, top_p=0.9 是什么、怎么调——正是第 4 章下半场的主题。
模型会回答了,但它听到的每个字都经过一道翻译——第 4 章讲这道翻译本身:chat template,全流水线最容易错位的一环。