TT Lab
开始
学习 学习路径 课程

MiniMind — 亲手从头到尾训练一个小型语言模型

续训文件里除了权重还有三样东西

在 TT Lab 中继续学习

一句话总结

训练会被中断——会话结束、节点重启、因成本而停下。要从停下的位置接着训练,并做到 与一次性跑到底完全相同,光有权重是不够的。MiniMind 的 lm_checkpoint 在权重之外另设恢复文件,把优化器状态、步数,以及(分布式训练时)进程数一起放进去。这个模块要让一次性训练 60 步的结果,与在 30 步停下再接着训练的结果逐比特一致,并测量每漏掉一样会相差多少。

为什么需要它

只保存权重、再接着训练,结果会悄悄地变。不会报错,损失曲线也看着合理地衔接上了。可是有两样东西对不上。

学习率安排不必另外保存——它是由步数计算出来的函数(get_lr),有步数就够了。分布式训练中如果 GPU 数量变了,一步看到的数据量就不同,所以 MiniMind 会用保存的 world_size 和现在的数量来换算步数。

工作原理

MiniMind 的保存会写两个文件。

state_dict = {k: v.half().cpu() for k, v in raw_model.state_dict().items()}
torch.save(state_dict, ckp_path + ".tmp");  os.replace(ckp_path + ".tmp", ckp_path)
resume = {"model": state_dict, "optimizer": optimizer.state_dict(),
          "epoch": epoch, "step": step, "world_size": ..., "wandb_id": ...}
torch.save(resume, resume_path + ".tmp");  os.replace(resume_path + ".tmp", resume_path)

临时文件 → 改名。保存到一半进程就挂了,会留下写了一半的文件。如果这个文件正覆盖着原来的名字,连好好的上一个检查点也丢了。先全部写进 .tmp,再用 os.replace 一次性改名,在同一个文件系统里是原子性的,无论何时挂掉,都会留下一个完好的文件。

半精度。权重用 .half() 保存,文件缩小一半。加载后展开成 fp32,logits 会略有不同(在这门课程的模型上,最大约 0.003)。推理没有问题,但如果恢复文件里的权重也是半精度,那么接着训练的结果就不可能与一次性训练的结果逐比特相同。这个实验为了便于比较,在恢复文件里放入 fp32 权重。

安全地读取。torch.load 使用的 pickle 会还原 Python 对象。如果直接读取别人给的检查点,里面的代码可能被执行。用 weights_only=True 读取,就只还原张量、数、字符串和 dict——这个实验的恢复文件(包括随机数生成器状态)也能照常读出来。

在现场相遇的样子

用抢占式(spot)实例训练,每隔几个小时节点就会被收回。如果恢复得不精确,每次都会变成略有不同的训练,最终模型无法重新做出来,出了问题也缩小不了原因范围。所以要先用短训练确认“接着训练 = 一次性训练”。如果保存不是原子性的,就会发生相反方向的事故——保存途中节点挂掉,最新的检查点损坏,而之前那个已经被覆盖掉了。

本课程与 MiniMind 原版的不同之处

MiniMind 不是每一步重新抽取,而是采用 以轮次为单位的顺序。每个轮次用 setup_seed(seed + epoch) 生成 torch.randperm 的顺序,恢复时重新生成同样的顺序,再用 SkipBatchSampler 跳过已经看过的批次数。这门课程的循环是每一步抽取随机位置,所以用保存、恢复批次随机数生成器的状态(get_state、set_state)来做同样的事。做法不同,守住的东西相同——“接下来要看的数据”从停下的位置接续上去。另外,MiniMind 以半精度保存权重,所以接着训练的结果不会逐比特相同。这门课程为了把这个差别也消除后再确认,在恢复文件里放入 fp32 权重。

下一项实验要做什么

以一次性训练 60 步的结果为基准,在 30 步处停下,用原子方式保存恢复文件,再接着训练到 60 步,看权重是否相同。测量丢掉优化器状态后接着训练、把批次顺序倒回去,各自相差多少,并数出半精度保存的大小和 logits 差异,以及优化器状态在恢复文件中所占的份额。