A training loop that resumes exactly
Batches, AdamW, gradient clipping, a warmup and cosine schedule, and checkpoints that restore every random generator, so a run split in two ends with the same weights as a run done in one go.
A model and a loss are not enough to train. The loop around them decides what data each step sees, how far each step moves the weights, and what happens when the run stops halfway. Day 1 built that loop and held it to one strict test: stop a run, resume it, and end with exactly the weights an uninterrupted run produces.
From documents to blocks
The trainer encodes each document with <|bos|> at the start and <|eos|> at the end, joins them into one long stream, and cuts the stream into blocks of tokens. The first tokens of a block are the input and the last are the targets, the same block shifted by one. A short stream is padded with <|pad|>, and PAD targets are excluded from the loss.
def make_blocks(texts: str | list[str], tokenizer: Tokenizer, context_length: int) -> TokenBlocks:
documents = [texts] if isinstance(texts, str) else texts
token_ids: list[int] = []
for document in documents:
token_ids.extend(tokenizer.encode(document, add_special_tokens=True))
width = context_length + 1
if len(token_ids) < width:
token_ids.extend([tokenizer.pad_id] * (width - len(token_ids)))
block_count = len(token_ids) // width
tokens = torch.tensor(token_ids[: block_count * width]).view(block_count, width)
inputs = tokens[:, :-1]
targets = tokens[:, 1:]
byte_lengths = torch.tensor(
[[tokenizer.token_bytes(int(token)) for token in row] for row in targets]
)
return TokenBlocks(inputs, targets, byte_lengths)The third tensor, byte_lengths, records how many bytes each target token decodes to. Bits per byte needs it, and it is computed once here instead of at every evaluation.
Each step draws a batch of blocks with a seeded generator. Day 1 uses 8 blocks of 128 tokens, so one step sees 1,024 target tokens.
AdamW: a per-weight step size, and decay kept separate
Plain gradient descent moves every weight by the same multiple of its gradient. Adam keeps two running averages per weight, the gradient and the squared gradient , and divides one by the square root of the other:
The hats mark a bias correction for the early steps, when both averages start at 0. The ratio is about for a weight whose gradient keeps its sign, and smaller for one whose gradient flips around. So each weight gets its own effective step size, and the learning rate sets the scale of all of them.
The last term is weight decay, which pulls weights toward 0. The "W" in AdamW means the decay is applied directly to the weights, outside the gradient average. In the older form the decay went into and was then divided by , so weights with large gradients were barely decayed. Loshchilov and Hutter showed the decoupled form regularizes the way decay is supposed to.
octlm decays only matrices. Norm gains, biases and anything else with one dimension are excluded, because shrinking a norm's gain toward 0 would shrink the signal, not regularize it.
def optimizer_for(model: nn.Module, config: ProjectConfig) -> torch.optim.AdamW:
decay, no_decay = [], []
for name, parameter in model.named_parameters():
if not parameter.requires_grad:
continue
(no_decay if parameter.ndim == 1 or name.endswith("bias") else decay).append(parameter)
groups = [
{"params": decay, "weight_decay": config.training.weight_decay},
{"params": no_decay, "weight_decay": 0.0},
]
return torch.optim.AdamW(groups, lr=config.training.learning_rate, betas=(0.9, 0.95))and are the LLaMA values, and weight decay is 0.1.
Gradient clipping caps one bad step
Occasionally a batch produces a very large gradient, and one step at full size can undo many good ones. Clipping rescales the whole gradient when its norm across every parameter exceeds a limit :
The direction stays the same, only the length is capped. octlm uses and logs the norm before clipping at every evaluation, as gradient_norm in each run file. A norm that sits far above 1 for many steps means the learning rate is too high.
Warm up, then decay along a cosine
The learning rate changes every step. It rises linearly from near 0 to its peak over the warmup steps, then follows half a cosine down to a floor:
Warmup exists because Adam's second-moment estimate is noisy for the first steps, and a full-size step on a bad estimate can push the model somewhere it struggles to leave. The cosine spends most of the run near the peak, where learning is fast, then slows down smoothly so the final steps make small adjustments.
What happened?
- Day 1 warms up for 20 of 200 steps, 10 percent of the run. Day 4 warms up for 500 of 24,000, about 2 percent. Both reach the same shape.
- The floor is 10 percent of the peak in every config. The LLaMA schedule ends at the same ratio.
- Past the last step the rate stays at the floor. The function clamps the cosine's progress at 1.
def learning_rate(step: int, config: ProjectConfig) -> float:
settings = config.training
if step < settings.warmup_steps:
return settings.learning_rate * (step + 1) / max(1, settings.warmup_steps)
progress = (step - settings.warmup_steps) / max(1, settings.steps - settings.warmup_steps)
cosine = 0.5 * (1 + math.cos(math.pi * min(progress, 1.0)))
return settings.min_learning_rate + cosine * (
settings.learning_rate - settings.min_learning_rate
)What a checkpoint must hold to resume exactly
Resuming needs more than the weights. The Adam averages and are part of the state, and so is every random generator, because the next batch depends on the sampler's state. octlm saves all of it, plus three fingerprints that must match on load: the config, the dataset, and the tokenizer.
def checkpoint_state(
model: Decoder,
optimizer: torch.optim.Optimizer,
generator: torch.Generator,
step: int,
config: ProjectConfig,
tokenizer: Tokenizer,
dataset_hash: str,
scaler: torch.amp.GradScaler | None = None,
) -> dict[str, object]:
return {
"config_hash": config.fingerprint,
"data_hash": dataset_hash,
"format": "octlm-checkpoint-v1",
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
"python_rng": random.getstate(),
"sampler_rng": generator.get_state(),
"scaler": scaler.state_dict() if scaler else {},
"step": step,
"tokenizer_hash": tokenizer.fingerprint,
"tokens_processed": step * config.training.batch_size * config.model.context_length,
"torch_rng": torch.get_rng_state(),
}The file is written atomically. save_checkpoint writes to a temporary file and renames it over the old one, and a rename either happens completely or not at all. A crash during the write leaves the previous checkpoint intact instead of a half-written one.
def save_checkpoint(path: Path, state: dict[str, object]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
temporary = path.with_suffix(path.suffix + ".tmp")
torch.save(state, temporary)
os.replace(temporary, path)The fingerprints earned their place on Day 2. Adding config fields changed the Day 1 config fingerprint, and old checkpoints refused to load. That refusal is the check working. A checkpoint trained under one config must not quietly continue under another.
Two tests, exact resume and overfitting one block
Exact resume. A CPU run was split in two, saved and resumed, and compared with an uninterrupted run. The weights were identical and the next loss was the same. PyTorch promises this only on one release, one platform and one device, and the claim has the same limit.
Overfit one block. A model that can learn anything should be able to memorize a single block. On one fixed BPE block, training loss fell from 5.5891 at step 20 to 0.9121 at step 200, with perplexity 2.4830 and 0.6338 bits per byte on that block.
Greedy decoding from the start of the block then reproduced the memorized text:
# Miniature modern LLM stack from first principles
## Goal
Build one long runn
The note is careful about what this proves. The tokenizer, the causal model, the loss, the optimizer and generation fit together and can drive the loss down. It says nothing about generalization, since the model saw one block.
What we did not build, and why
- Learning-rate search. Day 1 took common values, 3e-4 with a floor of 3e-5. Nothing in the results pointed at the learning rate.
- Gradient accumulation and mixed precision. Day 1 batches fit in CPU memory in float32. Mixed precision arrived on Day 4.
- Resume across machines. Bit-exact resume holds on one device. A resumed GPU run is expected to drift slightly, and Day 4 does not claim otherwise.