FP8 inference is often summarised as "cast the weights to 8 bits and enjoy twice the tensor-core throughput." The work that decides whether accuracy holds is elsewhere: which of the two FP8 formats each tensor uses, how its scale factor is computed, and at what granularity. Get those right and an FP8 model is usually hard to tell apart from BF16 on standard evaluations. Get them wrong and you see silent drift, NaNs from one outlier token, or a model that quietly degrades on long contexts.
This article covers those decisions, not the serving stack. The serving side (scaled GEMMs, FP8 KV cache, shipping a checkpoint with vLLM) is in FP8 inference explained. Here you will decode both formats bit by bit, meet the FNUZ variants that break ports between GPU vendors, compute scales by hand, compare per-tensor, per-channel, per-token and block scaling, and build a small harness that measures per-layer error so you can choose a recipe from evidence instead of folklore.
Two formats in one byte
An FP8 number is a sign, an exponent field and a mantissa field. For normal values, value = (-1)s × 2e - bias × (1 + m / 2M), where M is the number of mantissa bits. When the exponent field is zero, the number is subnormal: there is no implicit leading 1, which extends the range downward with fewer significant bits.
| Property | E4M3 (OCP, torch.float8_e4m3fn) | E5M2 (OCP, torch.float8_e5m2) |
|---|---|---|
| Exponent / mantissa bits | 4 / 3 | 5 / 2 |
| Exponent bias | 7 | 15 |
| Largest finite | 448 | 57,344 |
| Smallest normal | 2-6 = 0.015625 | 2-14 ≈ 6.1e-5 |
| Smallest subnormal | 2-9 ≈ 0.00195 | 2-16 ≈ 1.5e-5 |
| Infinity | none (top exponent reused for finite values) | yes |
| Worst-case relative rounding error (normals) | 2-4 = 6.25% | 2-3 = 12.5% |
The last row is the key one. Each mantissa bit halves the rounding error, which is worth about 6 dB of signal-to-quantization-noise ratio, so E5M2 gives up roughly 6 dB against E4M3 in exchange for about 14 more binades of range. Inference tensors (weights and activations) have a fairly stable range once scaled, so precision wins: current inference recipes use E4M3 for both GEMM operands. E5M2 earns its place where range is unpredictable and precision matters less. Gradients in training are the classic case, which is why Transformer Engine's HYBRID format uses E4M3 forward and E5M2 backward. Some engines also offer E5M2 as a KV-cache option.
OCP versus FNUZ: the porting trap
There are two families of FP8 in use. The OCP 8-bit floating point specification, implemented by NVIDIA Hopper, Ada and Blackwell, defines the formats above. AMD's MI300 series implemented a different pair, called FNUZ in PyTorch (float8_e4m3fnuz, float8_e5m2fnuz), whose name stands for finite, no negative zero. These formats have no negative zero, a single NaN encoded where negative zero would be, and an exponent bias one higher. For E4M3 that makes the bias 8 and the largest finite value 240, not 448. AMD's later CDNA 4 parts (MI350 series) are documented as supporting the OCP formats, so the variant now depends on the exact GPU. Check your accelerator's documentation rather than assuming one.
The resulting bug is easy to miss. A checkpoint quantized on NVIDIA hardware stores E4M3 bytes and per-tensor scales computed as amax / 448. Loaded on an FNUZ device and reinterpreted byte for byte, every value is off by a factor of two because of the bias, and anything above 240 has no valid encoding. Correct loaders convert, either by adjusting the scale or by requantizing, and record the format in the checkpoint config. If you write your own kernels or loaders, never derive the maximum from a constant. Read it from torch.finfo(dtype).max and carry the dtype with the scale.
What a scale factor does
Because FP8 has so few binades, almost no real tensor fits without help. A scale factor maps the tensor's range onto the format's: s = amax / fmax, store q = round_to_fp8(clamp(x / s)), and recover x ≈ q × s. A GEMM between scaled operands multiplies the FP8 values with higher-precision accumulation, then multiplies the output by sa × sb. The code below emulates this with public PyTorch dtypes and is enough to measure errors on CPU or GPU.
import torch
def quantize_fp8(x, dtype=torch.float8_e4m3fn, dim=None, eps=1e-12):
# Scaled FP8 quantization. dim=None: per-tensor; dim=-1: one scale per row (token/channel).
fmax = torch.finfo(dtype).max # 448 for e4m3fn, 57344 for e5m2, 240 for e4m3fnuz
amax = x.abs().amax() if dim is None else x.abs().amax(dim=dim, keepdim=True)
scale = (amax.float() / fmax).clamp(min=eps)
q = (x.float() / scale).clamp(-fmax, fmax).to(dtype) # clamp first: do not rely on the cast to saturate
return q, scale
def dequantize(q, scale):
return q.float() * scale
def sqnr_db(ref, approx):
noise = (ref.float() - approx.float()).pow(2).mean()
return 10 * torch.log10(ref.float().pow(2).mean() / noise.clamp(min=1e-30))The explicit clamp is important. PyTorch does not document the plain cast as saturating, and users have reported NaN from out-of-range E4M3 casts. E4M3 has no infinity, so overflow cannot become inf; in E5M2 it can. A single NaN in a hidden state spreads through the next layer's matmul to the whole sequence.
Scale granularity
One scale per tensor is cheapest but lets the largest element set the step size for every element. LLM activations make that costly. A few hidden channels carry outliers 100 times larger than the median, so a per-tensor scale pushes ordinary values down into the subnormal range. Finer scales fix this at the cost of more metadata and more complex kernels.
| Granularity | One scale per | Typical use | Cost |
|---|---|---|---|
| Per-tensor | whole tensor | Weights of well-behaved layers; static activation scales | Cheapest; one multiply in the epilogue |
| Per-channel | output channel of a weight | Weights with uneven row norms | Vector in the epilogue, still cheap |
| Per-token (dynamic) | row of the activation | Activations with outlier tokens | Needs an amax reduction per row at runtime |
| Block 1x128 / 128x128 | tile of activation / weight | DeepSeek-V3 style training and inference | Scales applied inside the K loop; custom kernels |
| MXFP8 | 32 consecutive elements, E8M0 power-of-two scale | OCP Microscaling; native on Blackwell tensor cores | Hardware-handled; 1 extra byte per 32 values |
DeepSeek-V3's report is a useful reference point. It used E4M3 for all GEMM inputs, 1x128 tiles for activations and 128x128 blocks for weights, and periodically promoted partial sums to FP32 because it found tensor-core FP8 accumulation precision insufficient over long reductions. The fine granularity is what allowed E4M3 throughout. With finer blocks, the outlier problem shrinks to the block that contains the outlier.
Static, dynamic and delayed scaling
Granularity is one axis. The other is when the scale is computed.
- Static: activation scales are calibrated offline on a few hundred representative samples and frozen. This is the fastest at runtime, but inputs unlike the calibration set can exceed the range and saturate. Weights are always static.
- Dynamic (current): amax is computed from the tensor that is about to be quantized, per tensor or per token. It always fits, at the cost of a reduction before the GEMM. Tools such as llm-compressor call the weight-static, activation-dynamic-per-token recipe FP8_DYNAMIC, and it is the safe default.
- Delayed: the scale comes from a history of amax values from earlier steps. Transformer Engine's
DelayedScalingrecipe does this for training, where it overlaps the reduction with compute. It assumes ranges change slowly, which suits training. For inference it offers little, because each request is independent.
# Training-side reference: Transformer Engine delayed scaling (E4M3 forward, E5M2 backward).
# TE has renamed entry points across releases; check names against your installed version.
import transformer_engine.pytorch as te
from transformer_engine.common.recipe import DelayedScaling, Format
recipe = DelayedScaling(fp8_format=Format.HYBRID, amax_history_len=16, amax_compute_algo="max")
with te.fp8_autocast(enabled=True, fp8_recipe=recipe):
y = model(x)A margin argument lowers the scale by powers of two to leave headroom. Use the same idea for static inference scales: computing them from a high percentile, or from max times a small factor, trades a little precision for protection against unseen outliers.
Measure error per layer before you trust a recipe
Benchmarks such as MMLU are too coarse to choose a recipe with, because they average away the one layer that broke. Measure per layer first. The harness below runs a BF16 model on calibration prompts, quantizes each linear layer's real input and weight with several candidate recipes, and records the SQNR of the layer output against the BF16 reference.
RECIPES = {
"e4m3/tensor": dict(dtype=torch.float8_e4m3fn, a_dim=None, w_dim=None),
"e4m3/token+ch": dict(dtype=torch.float8_e4m3fn, a_dim=-1, w_dim=-1),
"e5m2/token+ch": dict(dtype=torch.float8_e5m2, a_dim=-1, w_dim=-1),
}
results = {}
def probe(name, mod):
def hook(m, inputs, out):
x = inputs[0].reshape(-1, inputs[0].shape[-1])
ref = x.float() @ m.weight.float().T
for rname, r in RECIPES.items():
qa, sa = quantize_fp8(x, r["dtype"], r["a_dim"])
qw, sw = quantize_fp8(m.weight, r["dtype"], r["w_dim"])
y = dequantize(qa, sa) @ dequantize(qw, sw).T
results.setdefault((name, rname), []).append(sqnr_db(ref, y).item())
return mod.register_forward_hook(hook)
handles = [probe(n, m) for n, m in model.named_modules() if isinstance(m, torch.nn.Linear)]
with torch.no_grad():
for batch in calib_batches: # 64-256 real prompts, including long ones
model(**batch)
for h in handles: h.remove()Sort layers by their worst SQNR under the recipe you plan to ship. Most layers sit in a narrow band. In typical decoder models, the outliers are often the MLP down-projections and the first and last blocks. Those are the candidates to keep in BF16, or to move to finer scaling, before you run a full evaluation. Then confirm end to end: perplexity on held-out text, the task benchmarks you care about, and a long-context test, because errors accumulate over thousands of positions.
Worked example: one outlier, three recipes
Consider an activation row of 4,096 values. Most have magnitude around 0.5 to 2, a typical small value is 0.01, and one outlier channel reaches 600.
- Per-tensor E4M3. s = 600 / 448 ≈ 1.34. The value 0.01 maps to 0.0075, which is below E4M3's smallest normal (0.0156), so it lands in the subnormal range, where the step is 2-9 ≈ 0.00195. It rounds to 4 × 0.00195 = 0.0078, about a 5% error, and smaller values lose more. A value of 0.003 maps to 0.0022 and rounds to 0.00195, a 13% error. Ordinary values near 1 map to about 0.75 and keep the normal 6.25% worst-case bound.
- Per-tensor E5M2. s = 600 / 57,344 ≈ 0.0105. Small values now stay normal, but every value has only two mantissa bits, so the bulk of the tensor near 1 suffers up to 12.5% error. Overall SQNR is worse, even though the tiny values improved.
- Per-token E4M3 with the outlier channel handled. If the outlier is systematic per channel, smoothing (migrating it into the weights, as SmoothQuant does) shrinks amax to around 10. Then s ≈ 0.022, 0.01 maps to 0.45, a normal value, and the error drops back to the 6.25% bound everywhere.
So the format choice and the scale choice are linked. E5M2's range does not make up for its lost precision in inference. The fix for outliers is finer scaling or outlier migration, not a wider format.
Failure modes
| Failure | Cause | Mitigation |
|---|---|---|
| NaN after quantization | Unclamped cast overflowed E4M3, which has no infinity | Clamp before cast; assert finite in tests |
| Accuracy off by a factor of two after porting | OCP E4M3 bytes read as FNUZ, or the reverse | Store the dtype with the scale; convert on load |
| Fine on benchmarks, bad on long prompts | Static activation scales calibrated on short text | Calibrate on long inputs; use dynamic per-token scales |
| One layer dominates the loss | Outlier channels with a per-tensor scale | Per-layer SQNR sweep; BF16 fallback or finer blocks |
| Wrong output with KV cache in FP8 | Default scale of 1.0 left in place | Supply calibrated KV scales; see the serving article |
What to do next
- Default to E4M3 for weights and activations, with per-channel weight scales and dynamic per-token activation scales.
- Keep
torch.finfo(dtype).maxand the dtype next to every scale, and test checkpoint loading on each GPU family you deploy to, OCP and FNUZ. - Run the per-layer SQNR harness on 64 to 256 real prompts and keep the worst layers in BF16 until finer scaling recovers them.
- If outliers dominate, try SmoothQuant-style migration before changing format, and review calibration.
- Confirm end to end with perplexity, task benchmarks and a long-context test, then follow FP8 inference explained to serve it.
- For training-side FP8, read FP8 training math and the overview in FP8 for training and inference.