nanochat 的分词与采样看得懂、训练循环和显存占用得自己动手

文章导读
看 nanochat 这类教学向项目时,分词和采样两条路径通常几十分钟就能读懂,真正卡住人的是训练循环和显存:代码通读一遍觉得都明白,一跑起来却不知道时间花在前向还是数据加载、显存是被 batch 还是被序列长度吃掉的。可行的切法是把项目分成两半——分词与采样按输入输出读懂即可,训练循环和显存必须自己改配置、看日志、做对比,才形成判断。
📋 目录
  1. 先读懂分词与采样两个环节的输入输出
  2. 把训练循环拆成数据加载、前向、反向、保存四段
  3. 用小规模训练观察显存随配置的变化
  4. 在采样结果里判断训练是否真的产生了变化
  5. 把看懂的部分与需要动手的部分列成学习清单
A A

看 nanochat 这类教学向项目时,分词和采样两条路径通常几十分钟就能读懂,真正卡住人的是训练循环和显存:代码通读一遍觉得都明白,一跑起来却不知道时间花在前向还是数据加载、显存是被 batch 还是被序列长度吃掉的。可行的切法是把项目分成两半——分词与采样按输入输出读懂即可,训练循环和显存必须自己改配置、看日志、做对比,才形成判断。

分词与采样先读懂、不用急着改:看清输入样本形态与输出形态,跑一次最小调用就够。训练循环与显存占用必须动手:把循环拆成数据加载、前向、反向、保存四段,分别计时和查显存,再按固定顺序小规模改配置。判断训练是否真的生效,靠对比不同检查点的采样输出,而不是只看 loss 数字。所有判断以本机日志和 nvidia-smi 观察为准,换硬件或换数据后需重新确认。

先读懂分词与采样两个环节的输入输出

分词环节的输入样本形态,通常是一段纯文本批,经编码后变成整数 id 的序列张量,形状类似 [batch, seq_len];训练时标签是同一序列右移一位,用来做下一个 token 预测。分词环节的输出形态则是 id 列表以及解码回来的字符串,方便肉眼核对是否可逆。采样环节的输入是提示词文本,输出是追加生成后的 id 序列和解码文本,中间可以打印概率或 top-k 候选。

可以直接运行观察的最小方式,是先只跑编码解码的往返,不碰模型:

# 入口名以项目实际实现为准,这里只示意形态
from tokenizer import Tokenizer

tok = Tokenizer.load("path/to/tokenizer.model")
ids = tok.encode("你好,nanochat")
print(type(ids), len(ids), ids[:16])
print(repr(tok.decode(ids)))

# 观察批内长度不齐的形态
batch = [tok.encode(t) for t in ["hello", "hello world"]]
print([len(x) for x in batch])

采样用命令行跑一次即可,重点是把参数显式写出来,后面才可复现:

python sample.py `--ckpt` out/ckpt.pt `--prompt` "Once upon" \
  `--max`_new_tokens 64 `--temperature` 0.8 `--top`_k 50 `--seed` 0

参数名和脚本入口需要结合项目实际确认;seed 固定后同一检查点应能重复出同一段文本,这是后续对比的前提。

把训练循环拆成数据加载、前向、反向、保存四段

读训练循环时不要整段读,按四段标注行号,每段都要求它单独产出可打印的中间结果,否则出错时无法定位。

环节代码位置应产出的中间结果
数据加载取 batch 的那几行,通常紧接在 for step in range(...) 之后输入张量形状、标签形状、batch 内序列长度分布、该步取数耗时
前向模型调用与 loss 计算处logits 形状、loss 标量值、forward 耗时
反向loss.backward() 到 optimizer.step() 之间梯度是否非 None、梯度范数、backward 耗时、是否触发梯度裁剪
保存checkpoint 保存与日志写入分支保存路径、保存前后的 step、保存耗时、文件大小

做法是在四段各自前后插入计时或打印,只跑几十步就停。若某一段明显偏慢,再单独压测该段:例如把模型调用换成空操作,只留数据加载循环,就能看出数据管线是不是瓶颈。

用小规模训练观察显存随配置的变化

把显存从黑盒变成可观察量,需要同时看进程外的总量和进程内的分配峰值:

# 终端 A:每秒刷新一次
nvidia-smi `--query-gpu`=memory.used,memory.total,utilization.gpu `--format`=csv -l 1

# 训练脚本内,在若干步后打印
import torch
if step % 20 == 0:
    print("peak MiB", torch.cuda.max_memory_allocated() // 1024**2)
    torch.cuda.reset_peak_memory_stats()

配置改动顺序建议每次只动一个变量:先固定 seq_len 和 batch_size 跑通基线;单独把 batch_size 翻倍;回到基线单独把 seq_len 翻倍;再依次试梯度累积、混合精度、激活重计算;模型宽度和层数放到最后动。改了模型结构又改 batch,出现 OOM 时无法归因。

观察记录至少包含这些字段:本地提交号、batch_size、seq_len、梯度累积步数、精度设置、是否开启激活重计算、参数量、峰值显存、每步耗时、是否 OOM。同一硬件上重复跑同一组配置,数值应当接近;明显波动就要怀疑是否有其他进程占用或数据长度分布突变。

在采样结果里判断训练是否真的产生了变化

训练是否生效,最直接的验证是用同一提示词、同一 seed、同一解码参数去采样不同检查点:

for ck in out/ckpt_step100.pt out/ckpt_step500.pt out/ckpt_step2000.pt; do
  echo "==== $ck"
  python sample.py `--ckpt` "$ck" `--prompt` "The capital of" \
    `--max`_new_tokens 48 `--temperature` 0.7 `--top`_k 50 `--seed` 0
done

解码参数一旦变动,输出差异就分不清来自训练还是来自采样设置,所以对比时必须锁死参数。观察点包括:是否从随机字符变成语料里的常见词、是否出现合理的空格与标点、有没有长段重复或复读提示词、换不同 seed 时输出的多样性是否合理。若早中期检查点输出几乎一样,先回查学习率、数据是否打乱、梯度是否真的在更新,而不必急着加大训练规模。

把看懂的部分与需要动手的部分列成学习清单

把上面的动作整理成表,用来决定下一步练哪里。反馈速度是经验估计,实际取决于硬件和数据规模。

环节是否需动手验证方式预计反馈速度
分词 encode/decode打印 id 与解码文本,检查往返是否一致秒级
采样参数部分固定提示词与 seed,只调 temperature/top_k分钟级
数据加载单独循环计时,看长度分布秒到分钟
前向与反向分段计时加峰值显存打印分钟级
检查点保存与恢复保存后重新加载并采样,确认 step 与输出一致分钟级
显存随配置变化单变量改动加 nvidia-smi 与峰值显存记录每次几分钟

建议的推进顺序是先跑通分词往返与一次采样,再做数据加载计时,然后补前向反向计时与显存记录,最后做多检查点采样对比。每完成一格就在本机留下日志,下一步换配置时才有对照,而不是凭感觉判断训练跑得好不好。