Chapter 10 β Pre-Training at Scale
βThe remarkable thing about LLMs isnβt the architecture β itβs that next-token prediction on internet-scale data produces intelligence.β
The Pre-Training Pipeline
Pre-training is where the bulk of compute and data investment goes. A modern LLM pre-training run involves months of GPU time, trillions of tokens, and careful orchestration.
Training Objective
The training objective is deceptively simple β minimize cross-entropy loss at every position:
\[\mathcal{L} = -\frac{1}{T}\sum_{t=1}^{T} \log P_\theta(x_t \mid x_{<t})\]where $P_\theta$ is the modelβs predicted probability for the correct next token.
This is computed efficiently: a single forward pass processes all T tokens in parallel (thanks to causal masking), and the loss is computed at every position simultaneously.
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
import torch
import torch.nn.functional as F
def compute_lm_loss(model, input_ids):
"""Standard causal LM loss computation."""
# input_ids: [B, T]
outputs = model(input_ids[:, :-1]) # predict from tokens 0..T-2
logits = outputs.logits # [B, T-1, vocab_size]
targets = input_ids[:, 1:] # targets are tokens 1..T-1
loss = F.cross_entropy(
logits.reshape(-1, logits.size(-1)), # [B*(T-1), vocab_size]
targets.reshape(-1), # [B*(T-1)]
)
return loss
Learning Rate Schedule
Almost all modern LLMs use a warmup + cosine decay schedule:
- Linear ramp from ~0 to peak LR
- Typically 0.1-2% of total steps (500β2000 steps)
- Stabilizes training at the start
- Smoothly decays from peak to min LR
- Min LR typically 0.1Γ or 0.01Γ of peak
- Gradual decay preserves learned representations
1
2
3
4
5
6
7
8
import math
def cosine_lr_schedule(step, warmup_steps, total_steps, max_lr, min_lr):
"""Warmup + cosine decay learning rate schedule."""
if step < warmup_steps:
return max_lr * step / warmup_steps
progress = (step - warmup_steps) / (total_steps - warmup_steps)
return min_lr + 0.5 * (max_lr - min_lr) * (1 + math.cos(math.pi * progress))
Typical Hyperparameters
| Hyperparameter | Small (1-3B) | Medium (7-13B) | Large (70B+) |
|---|---|---|---|
| Peak Learning Rate | 3e-4 | 3e-4 | 1.5e-4 |
| Min Learning Rate | 3e-5 | 3e-5 | 1.5e-5 |
| Warmup Steps | 2000 | 2000 | 2000 |
| Weight Decay | 0.1 | 0.1 | 0.1 |
| Batch Size (tokens) | 1-4M | 4-8M | 8-16M |
| Optimizer | AdamW | AdamW | AdamW |
| Ξ²β, Ξ²β | 0.9, 0.95 | 0.9, 0.95 | 0.9, 0.95 |
| Gradient Clipping | 1.0 | 1.0 | 1.0 |
| Total Tokens | 1-3T | 3-15T | 15T+ |
Pre-Training Datasets
| Dataset | Size (tokens) | Sources | Used By |
|---|---|---|---|
| FineWeb | 15T | CommonCrawl (filtered) | Community standard |
| FineWeb-Edu | 1.3T | FineWeb educational subset | Quality benchmark |
| RedPajama v2 | 30T | CommonCrawl, multi-source | Open replication |
| The Pile | 300B | 22 diverse sources | GPT-NeoX, early open models |
| Dolma | 3T | CommonCrawl, code, books, papers | OLMo |
| StarCoder data | 780B | GitHub (licensed code) | StarCoder, CodeLLaMA |
Data Mix
Models arenβt trained on just web text. The training mix typically includes:
Training Loop with Gradient Accumulation
Real LLM training uses large effective batch sizes (millions of tokens) achieved through gradient accumulation across micro-batches:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
def train_step(model, optimizer, dataloader, accumulation_steps, max_grad_norm=1.0):
"""Training loop with gradient accumulation and mixed precision."""
model.train()
optimizer.zero_grad()
total_loss = 0.0
for micro_step in range(accumulation_steps):
batch = next(dataloader) # [micro_batch_size, seq_len]
input_ids = batch["input_ids"].to(device)
with torch.amp.autocast("cuda", dtype=torch.bfloat16):
outputs = model(input_ids[:, :-1], labels=input_ids[:, 1:])
loss = outputs.loss / accumulation_steps
loss.backward()
total_loss += loss.item()
# Gradient clipping
torch.nn.utils.clip_grad_norm_(model.parameters(), max_grad_norm)
optimizer.step()
return total_loss
Example: micro_batch_size=4, seq_len=8192, accumulation_steps=32, 8 GPUs with DDP β effective batch = 4 Γ 8192 Γ 32 Γ 8 = ~8M tokens per step.
Training Infrastructure
Checkpointing and Recovery
Training runs crash. Hardware fails. Pre-training runs need robust checkpointing:
- Periodic checkpoints: Save every 500β1000 steps (model + optimizer state + RNG state)
- Async checkpointing: Save to storage without blocking training
- Loss spike detection: Monitor for loss spikes indicative of instability
- Automatic restart: Resume from last checkpoint on failure
- Deterministic recovery: Save RNG states for exact reproducibility
Loss Curves
A healthy pre-training loss curve shows rapid initial improvement followed by slow, steady decline:
| Phase | Steps | Behavior |
|---|---|---|
| Early | 0β1K | Rapid drop from ~11 (random) to ~4β5 |
| Middle | 1Kβ100K | Steady decline, ~log-linear |
| Late | 100K+ | Diminishing returns, approaching data limit |
| Anomalies | β | Loss spikes β learning rate issues, data quality, hardware faults |
Typical final pre-training loss: ~1.5β2.0 cross-entropy on general web text for a well-trained 7B model.
Compute Estimation
The compute needed for one forward+backward pass through the full dataset:
\[C \approx 6ND\]where $N$ = parameters, $D$ = tokens, and the factor 6 accounts for forward (2N FLOPs per token) + backward (4N FLOPs per token).
| Model | Params (N) | Tokens (D) | Compute (6ND) | GPU-hours (H100) |
|---|---|---|---|---|
| LLaMA-3 8B | 8B | 15T | 7.2Γ10Β²Β³ | ~180K |
| LLaMA-3 70B | 70B | 15T | 6.3Γ10Β²β΄ | ~1.6M |
| LLaMA-3.1 405B | 405B | 15T | 3.6Γ10Β²β΅ | ~9M |
Whatβs Next
Pre-training produces a raw βbase modelβ β capable but not aligned to human instructions or specialised domains. Before exploring the training stages further, the next chapter covers the optimisers and loss functions that drive all LLM training runs.
β Previous: Chapter 9 β Data for LLMs Β· Next: Chapter 11 β Optimizers & Loss Functions β
Last updated: April 2026