MiniMind — Train a Small Language Model Yourself, End to End
A resume file holds three things besides the weights
In one line
Training gets interrupted — the session ends, a node restarts, or it stops because of cost. To resume from where it stopped exactly as if it had run straight to the end, weights alone are not enough. MiniMind's lm_checkpoint keeps a separate resume file apart from the weights, and puts into it the optimizer state, the step count, and (for distributed training) the number of processes. In this module, you make the result of 60 steps done at once and the result of stopping at step 30 and resuming not differ by a single bit, and measure how far they diverge when you leave out one thing at a time.
Why this was needed
If you save only the weights and resume training, the result changes silently. No error is raised. The loss curve even continues plausibly. But two things are out of line.
- Optimizer state. AdamW carries, for each parameter, a moving average of the gradient (
exp_avg) and a moving average of its square (exp_avg_sq), and decides the step size with them. If you resume with a new optimizer, this memory starts again from 0, and the first few steps right after resuming move at completely different sizes. This state is twice the size of the parameters, so the resume file is about three times larger than the weights file. - Position in the data. If you stopped in the middle of an epoch, you must skip the batches already seen in that epoch. If you draw again from the start, you see the first part twice and never see the latter part. MiniMind skips as many batches as were already seen with
SkipBatchSampler, and rebuilds the batch order withseed + epochto reproduce the same order.
The learning rate schedule does not need to be saved separately — it is a function (get_lr) computed from the step number, so only the step count is needed. In distributed training, if the number of GPUs changes, the amount of data one step sees changes, so MiniMind converts the step count using the saved world_size and the current number.
How it works
MiniMind's save writes two files.
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)
Temporary file → rename. If the process dies while saving, a half-written file is left. If that file was overwriting the original name, you lose even the previous checkpoint that was fine. If you write everything to .tmp and then swap with os.replace in one step, within the same filesystem the swap is atomic, so a whole file is left no matter when it dies.
Half precision. The weights are saved with .half() to cut the file in half. If you load them and expand to fp32, the logits differ slightly (up to about 0.003 for this course's model). This is no problem for inference, but if the weights in the resume file are also half precision, training that was resumed cannot be bit-for-bit the same as training done at once. For comparison, this lab puts fp32 weights in the resume file.
Reading safely. torch.load uses pickle, which brings Python objects back to life. If you just read a checkpoint someone else gave you, the code inside it may run. If you read with weights_only=True, it restores only tensors, numbers, strings, and dicts — the resume file of this lab (including the random number generator state) is read as it is, too.
What it looks like in the field
If you train on spot (preemptible) instances, the node is taken away every few hours. If resuming is not exact, each time becomes a slightly different training run, so you cannot rebuild the final model and cannot narrow down the cause when a problem occurs. That is why you first check "resumed = done at once" with a short training. If the save is not atomic, the opposite accident happens — the node dies mid-save so the most recent checkpoint is broken, and the one before it has already been overwritten and is gone.
How this course differs from the original MiniMind
Instead of drawing anew at each step, MiniMind uses an order per epoch. In each epoch it builds a torch.randperm order with setup_seed(seed + epoch), and on resume it rebuilds the same order and then skips as many batches as were already seen with SkipBatchSampler. The loop in this course draws a random position at each step, so it does the same thing by saving and restoring the state of the batch random number generator (get_state and set_state). The methods differ but what they keep is the same — "the next data to see" continues from where it stopped. Also, MiniMind saves weights in half precision, so training that was resumed is not bit-for-bit the same. To eliminate even that difference for verification, this course puts fp32 weights in the resume file.
What you will do in the next lab
Taking the result of training 60 steps at once as the reference, you stop at step 30, save the resume file atomically, resume up to step 60, and see whether the weights are the same. You measure how far they diverge when you resume with the optimizer state thrown away and when you rewind the batch order, and count the size and logit difference of half-precision saving and the share of the resume file that the optimizer state takes.