Skip to content

训练与调参

模型会「前向」了,这一页教它「学习」:手写一个完整训练循环,包含取批次、优化器、学习率调度、周期性评估与保存。全部跑通后,你就会拥有一个训练好的 NanoGPT。

1. 取一个批次(对应「数据集处理」的切块思路)

python
# train.py
import os, pickle, math, time
import numpy as np
import torch
from model import GPT, Config

device = "cuda" if torch.cuda.is_available() else ("mps" if torch.backends.mps.is_available() else "cpu")
torch.manual_seed(1337)

# 载入数据与词表(先运行 prepare.py)
train_data = np.memmap("data/train.bin", dtype=np.uint16, mode="r")
val_data   = np.memmap("data/val.bin",   dtype=np.uint16, mode="r")
meta = pickle.load(open("data/meta.pkl", "rb"))

cfg = Config(vocab_size=meta["vocab_size"])
batch_size = 32
block = cfg.block_size

def get_batch(split):
    data = train_data if split == "train" else val_data
    # 随机起点, 取 batch_size 段长度为 block+1 的序列
    ix = np.random.randint(len(data) - block - 1, size=(batch_size,))
    x = torch.stack([torch.from_numpy(data[i:i+block].astype(np.int64)) for i in ix])
    y = torch.stack([torch.from_numpy(data[i+1:i+1+block].astype(np.int64)) for i in ix])
    return x.to(device), y.to(device)   # y = x 右移一位 → 预测下一字符

y 就是 x 右移一位,完美复刻 数据集处理 里「输入/目标错位一位」的自监督做法。

2. 模型、优化器与学习率调度

python
model = GPT(cfg).to(device)
print(f"参数量: {sum(p.numel() for p in model.parameters())/1e6:.2f}M")

# AdamW 是 Transformer 训练标配; 对不同权重用不同学习率是现代做法,这里统一简化
opt = torch.optim.AdamW(model.parameters(), lr=3e-4, betas=(0.9, 0.95), weight_decay=0.1)

@torch.no_grad()
def get_lr(step):
    warmup = 100
    if step < warmup:                     # ① 线性预热, 稳住早期
        return 3e-4 * step / warmup
    total = 5000
    if step > total:                      # ④ 到最低学习率后保持
        return 3e-5
    frac = (step - warmup) / (total - warmup)
    return 3e-5 + 0.5 * (3e-4 - 3e-5) * (1 + math.cos(math.pi * frac))  # ② 余弦衰减到 ③

@torch.no_grad()
def estimate_loss():
    model.eval()
    out = {}
    for split in ["train", "val"]:
        losses = []
        for _ in range(20):
            x, y = get_batch(split)
            _, loss = model(x, y)
            losses.append(loss.item())
        out[split] = sum(losses) / len(losses)
    model.train()
    return out

3. 训练循环本体(五步一气呵成)

python
out_dir = "out"; os.makedirs(out_dir, exist_ok=True)
iter_num, max_iters = 0, 5000
model.train()

while iter_num <= max_iters:
    x, y = get_batch("train")
    lr = get_lr(iter_num)
    for pg in opt.param_groups:
        pg["lr"] = lr                          # 每步按调度设置学习率

    logits, loss = model(x, y)                 # ① 前向
    opt.zero_grad(set_to_none=True)            # ② 清空梯度
    loss.backward()                            # ③ 反向传播
    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)  # 梯度裁剪, 防爆
    opt.step()                                 # ④ 更新参数

    if iter_num % 200 == 0:
        losses = estimate_loss()
        print(f"step {iter_num:4d} | lr {lr:.2e} | "
              f"train {losses['train']:.3f} | val {losses['val']:.3f}")
        torch.save({"model": model.state_dict(), "cfg": cfg.__dict__},
                   f"{out_dir}/ckpt_{iter_num}.pt")   # ⑤ 保存 checkpoint

    iter_num += 1

手写循环的核心永远是那 5 步:前向 → 清梯度 → 反向 → 裁剪 → 更新。把它和 训练流程Trainer 里的 Trainer 对照,你就明白 Trainer 帮你自动做了哪些事。

4. 看懂 loss:从 ln(V) 往下走

初始 loss ≈ ln(vocab_size)   ← 模型瞎猜, 交叉熵的理论起点
  例: vocab=65 → ln65≈4.17;  训练良好应一路下滑到 ~1.x
loss 快速下降  → 在学规律
loss 平台不动  → 该调学习率/加数据/查 bug
val 反而上升   → 过拟合了, 早停或加正则

5. 调参实战清单

症状优先尝试
loss 不降 / NaN调小学习率;确认因果掩码已加;检查数据是否读对
收敛太慢适度加大学习率或 batch;确认 warmup 正常
过拟合(val 升)减小模型/加 dropout/早停/加数据
显存不足(OOM)降 batch_size、降 block_size、开梯度检查点
CPU 跑太慢先用 n_layer=2、max_iters=1000 跑通流程再放大

关键超参经验:

  • 学习率:AdamW 下小模型常用 1e-4 ~ 6e-4;过大是 loss 不降的头号元凶。
  • block_size:决定能看多远,也影响显存(注意力 O(T²))。
  • n_layer / n_embd:容量旋钮,从小往大加,先看流程再谈规模。
  • batch_size:越大越稳但越吃显存;用梯度累积模拟大 batch。

6. 让它更专业的小增强(选做)

python
# ① 混合精度(GPU): 用 autocast 包裹前向, bf16/fp16 提速省显存
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
    logits, loss = model(x, y)

# ② 梯度累积: 等效放大 batch(显存不够时的救星)
accum = 4
loss = loss / accum
loss.backward()
if (iter_num + 1) % accum == 0:
    opt.step(); opt.zero_grad()

# ③ 断点续训: 启动时若有 ckpt 就 load_state_dict 继续

7. 训练产出

跑完你会得到 out/ckpt_*.pt,内含 model 权重与 cfg。下一页加载它,让模型开口说「人话」。

小结

  • 训练循环五步:前向 → 清梯度 → 反向 → 裁剪 → 更新
  • 自监督:yx 右移一位,学「预测下一字符」。
  • 学习率 + warmup + 余弦衰减 + 梯度裁剪 是 Transformer 稳定训练的四大护法。
  • 用 loss 起点 ln(V) 判断是否正常,用 train/val 关系判断是否过拟合。

下一步

模型练好了,来评估它、并动手用它生成文本 → 评估与文本生成