Mixed precision on a T4
Running the matrix math in 16-bit floats nearly doubled training speed with no measurable loss. The first attempt picked bfloat16, which the T4 only emulates, and ran slower than float32.
Every model up to Day 3 trained in float32, 4 bytes per number. GPUs have units built for 16-bit matrix math that run far faster. EXP-066 turned them on for the Day 4 model and checked that the loss did not suffer.
Three ways to spend the bits
A floating-point number stores a sign, an exponent and a mantissa. The exponent sets the range, how large and how small a number can be. The mantissa sets the precision, how many significant digits it keeps.
| Format | Bytes | Exponent bits | Mantissa bits | Largest value | Decimal digits |
|---|---|---|---|---|---|
| float32 | 4 | 8 | 23 | about 3.4e38 | about 7 |
| float16 | 2 | 5 | 10 | 65,504 | about 3 |
| bfloat16 | 2 | 8 | 7 | about 3.4e38 | about 2 |
float16 keeps more precision and loses range. bfloat16 keeps float32's range and loses precision. That difference decides how each one fails.
What happened?
- 0.1 cannot be stored exactly in any binary format. float16 is off by about 2e-4 relative, bfloat16 by about 1e-3.
- 70,000 overflows to infinity in float16 and stores fine in bfloat16, rounded to 70,144.
- 3e-8 rounds up to about 6e-8, the smallest value float16 can hold, a relative error near 100 percent. 1e-9 becomes 0. bfloat16 stores both with its usual 2 to 3 digits.
Where the 16-bit math happens
"Mixed" precision keeps the model's weights and the optimizer's state in float32 and runs selected operations in 16 bits. PyTorch's autocast decides per operation. Matrix multiplies and attention run in 16 bits, while softmax, norms and the loss run in float32, because sums over many small numbers lose too much in 16 bits. octlm's RMSNorm and RoPE already compute in float32 internally for the same reason.
Why float16 needs a GradScaler
Gradients are often tiny, around 1e-6 to 1e-8 for many weights. In float16, anything below about 3e-8 rounds to 0, and anything below 6e-5 keeps fewer significant bits the smaller it gets. A gradient that rounds to 0 stops its weight from learning. The fix is to multiply the loss by a large factor before the backward pass. Backpropagation is linear in the loss, so every gradient grows by the same factor and moves into float16's range. Before the optimizer step, the gradients are converted to float32 and divided by the factor again.
What happened?
- With no scale and gradients around 1e-6, a large share of them falls below float16's range and would be lost.
- Raising the scale shifts the whole distribution right. Somewhere around 2^10 to 2^16 almost nothing is lost and nothing overflows.
- Too large a scale pushes the biggest gradients past 65,504. PyTorch's
GradScalerdetects the resulting infinities, skips that step, and halves the scale. After 2,000 clean steps it doubles the scale again, so it settles near the largest safe value on its own.
bfloat16 has float32's exponent range, so it does not need loss scaling. That is why modern training prefers it, on hardware that has it.
The first run chose bfloat16, and the T4 does not have it
The first version of autocast_dtype picked bfloat16 whenever torch.cuda.is_bf16_supported() returned True. On the T4 it does. The T4 is compute capability 7.5, a Turing card, and it has tensor cores for float16 only. PyTorch supports bfloat16 on it by emulation. The "mixed precision" run was slower than float32.
The fix checks the compute capability instead. bfloat16 needs 8.0 or newer, which is an A100 or later. Anything older gets float16 with a GradScaler.
def autocast_dtype(device: torch.device) -> torch.dtype | None:
if device.type != "cuda":
return None
return torch.bfloat16 if torch.cuda.get_device_capability(device)[0] >= 8 else torch.float16Inside the training step, the scaler wraps the backward pass, unscales before clipping so the clip threshold means the same thing as in float32, and skips the step if it found an infinity:
with torch.autocast(target.type, dtype=dtype, enabled=dtype is not None):
loss = mtp_loss(model, inputs, targets, tokenizer.pad_id)
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
gradient_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), config.training.grad_clip)
scaler.step(optimizer)
scaler.update()
The checkpoint stores the scaler's state too, so a resumed run continues with the scale it had reached.
1.88 times faster, loss within 0.0013
Hypothesis: mixed precision at least doubles tokens per second, and validation loss at step 1,000 stays within 0.02 of float32. Stop condition: a gap above 0.02 means training the main run in float32.
The same 1,000 steps of the Day 4 model, on a Colab T4:
| dtype | Train seconds | Seconds per step | Validation loss | Bits per byte |
|---|---|---|---|---|
| float32 | 1,173 | 1.17 | 1.2500 | 0.8001 |
| bfloat16 (emulated) | 1,529 | 1.53 | 1.2590 | 0.8058 |
| float16 | 625 | 0.62 | 1.2513 | 0.8009 |
What happened?
- float16 is 1.88 times faster than float32. That misses the hypothesis's 2x, and the note says so.
- Its validation loss is 0.0013 above float32's, far inside the 0.02 limit.
- Emulated bfloat16 was 1.30 times slower than float32, with a slightly worse loss.
Decision: keep float16 for the main run.
The real throughput, and the new estimate
One step is 32 blocks of 512 tokens, 16,384 tokens. At 0.625 seconds per step that is about 26,200 tokens per second. Training costs about 6 FLOPs per parameter per token, so 26,200 tokens per second on a 26.3M-parameter model is about 4.1 TFLOPs sustained.
The plan had assumed 15 TFLOPs, which is closer to the T4's float16 peak than to what a whole training step achieves. At the measured rate the main run's 24,000 steps take about 4.2 hours, not 70 minutes. Why the plan changed uses that rate in the compute estimator.
The estimate held. The main run trained all 24,000 steps in float16 on a Kaggle T4 in 4.1 hours, 0.614 seconds per step. The loss never diverged, and the step-1,000 validation loss, 1.2513, matches the float16 row above.
What we did not build, and why
- The TPU. The Colab account also offers a TPU v5e. PyTorch reaches it through
torch_xla, and octlm's device choice, autocast, GradScaler, attention kernel and checkpointing are all written for CUDA or CPU. The port would cost more than it saves on a four-hour run, and no later day needs a TPU. - float8 and int8 training. The T4 has no float8, and int8 training is research.
- Mixed precision for Days 1 to 3. It is opt-in inside
train_model. Earlier days keep float32 so their numbers stay comparable.