4Day 4

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.

The formats

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.

FormatBytesExponent bitsMantissa bitsLargest valueDecimal digits
float324823about 3.4e38about 7
float16251065,504about 3
bfloat16287about 3.4e38about 2

float16 keeps more precision and loses range. bfloat16 keeps float32's range and loses precision. That difference decides how each one fails.

Playground
FormatSign · exponent · mantissa bitsStored valueRelative error
float32
8 exp, 23 mantissa
0 01111011 100110011001100110011010.1000000011.5e-8
bfloat16
8 exp, 7 mantissa
0 01111011 10011010.1000976569.8e-4
float16
5 exp, 10 mantissa
0 01011 10011001100.09997558592.4e-4

float16 spends 5 bits on the exponent, so it tops out at 65,504 and rounds values below about 3e-8 to zero. bfloat16 keeps float32's 8 exponent bits and so its range, but has only 7 mantissa bits, about 2 to 3 decimal digits. Try 70000, 3e-8 and 1e-9.

Type a number and see how each format stores it. The rounding runs in this page and is tested against PyTorch's own float16 and bfloat16 casts.

What happened?

Playground
float32bfloat16float162^-1402^-1002^-602^-202^02^202^602^100
float320.0000037% error1.240e-6
bfloat160.24% error1.237e-6
float160.96% error1.252e-6

Solid bars are the normal range, faded bars the subnormals, where precision thins out toward zero. float16 covers a narrow window around 1, from about 2^−24 to 65,504. Gradients often sit near 2^−20 to 2^−30, below it, which is why float16 training scales the loss up first. bfloat16's bar is as long as float32's.

The magnitudes each format can store, on a log scale from 2^−150 to 2^130. Slide a value across and see which formats keep it.
Mixed, not half

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.

Loss scaling

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.

Playground
1e−14scaled gradient magnitude, one bar per power of ten1e9
below 1e−7: float16 rounds these to 0 or keeps one or two significant bitsrepresentableabove 1e5: past float16's 65,504 maximum
Lost to zero in float1611.3%
Overflow to inf0.0%step is applied

Multiply the loss by the scale, and every gradient grows by the same factor, because the backward pass is linear in the loss. The scaler divides the gradients back down in float32 before the optimizer step. Slide the scale until the zero share falls without anything overflowing. GradScaler searches for that scale on its own: it starts at 2^16, halves on any inf, and doubles after 2,000 clean steps.

A spread of gradient magnitudes, rounded to float16 after scaling. Gray bars are small enough that float16 turns most of them into 0. Orange bars overflow. The rounding is octlm's tested float16 port.

What happened?

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 trap

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.

octlm/train.pyline 303
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.float16

Inside 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.

EXP-066

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:

dtypeTrain secondsSeconds per stepValidation lossBits per byte
float321,1731.171.25000.8001
bfloat16 (emulated)1,5291.531.25900.8058
float166250.621.25130.8009

What happened?

Decision: keep float16 for the main run.

A second result

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.

Skipped

What we did not build, and why