3.3 训练实战:脚本与常见故障


3.3 训练实战:脚本与常见故障

本节摘要:本节把前两节的原理装进一份完整的训练脚本——HuggingFace Trainer 风格(标注:写法示意,教学用;nanochat 原生是自带极简训练循环,接口以官方仓库 README 为准)。脚本串起全链条:读 JSONL → 3.1 的掩码构造 → 整理 batch → Trainer 配 3.2 的超参 → 存 checkpoint;约 60 行,124M 模型在 24GB 卡上单卡数小时可跑完数万条数据。下半场是三故障排查表:显存不足(CUDA OOM 的六连降级)、loss 不降(从掩码到数据的五嫌疑)、灾难性遗忘(英文能力滑坡的防与治)。跑通的那一刻,回到 0.1 节做行为复验。

学习目标

阅读完本节,你应当能够:

  1. 跑通完整 SFT 脚本,知道每个配置项对应前面哪节的结论。
  2. 遇 OOM 时按六连降级顺序自救,而不是瞎删 batch。
  3. 用排查表定位"loss 不降"与"灾难性遗忘"并处置。

一、完整训练脚本

# 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 为准)。

二、故障排查表

故障一:显存不足(CUDA out of memory)

六连降级,按序尝试,从上到下代价递增:

手段 代价 备注
1 per_device_train_batch_size,等量加 gradient_accumulation_steps 速度略降 有效 batch 不变,首选
2 gradient_checkpointing 慢约 30% 0.2 三开关之一
3 确认 bf16=Trueuse_cache=False 几乎无 两个最常见的"忘了关/忘开"
4 max_seq_len(截到 p90 长度) 长样本受损 长度预算见 2.1 末
5 换 8-bit 优化器(如 bitsandbytes 的 AdamW8bit 精度略降 Adam 状态省 4 份参数量
6 换更大显存的卡 / 多卡 124M 很少走到这步

故障二:loss 不降(或起点异常)

症状 首要嫌疑 处置
起点 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 章下半场的主题。

本节要点回顾

  1. 完整脚本四钩子:3.1 掩码、3.2 的有效 batch 与 lr、200 步存档;验证集切 1%~2%。
  2. OOM 六连降级(先累积后检查点后截长度);loss 不降先跑掩码前向验收;灾难性遗忘首选数据回掺。
  3. 验收看行为复验:停下、身份、指令——训练闭环完成,剩下的是把模型"接上嘴"。

模型会回答了,但它听到的每个字都经过一道翻译——第 4 章讲这道翻译本身:chat template,全流水线最容易错位的一环。


作者与出处
原作者: 灏天文库
整理: 灏天文库整理
本站整理收录,版权归原作者/开源协议所有;欢迎通过原文链接访问源仓库。
发布者: 作者: 灏天文库 转发
评论区 (0)
U