Llama 3.1 8B Instruct with its decoder projections in FP8: RedHatAI's FP8-dynamic checkpoint, made with llm-compressor. Each query, key, value, output, gate, up and down projection holds E4M3 weights (one byte each) with one bf16 scale per output row; the embedding and the output head stay bf16.
The source is the llama-3.1-8b-instruct package with those projections as std.quant::Fp8Linear. bindings.json binds each scale to the checkpoint's weight_scale and each FP8 weight to its bytes.
Loading
from linnet import nest
model = nest.load("llama-3.1-8b-instruct-fp8", backend="torch", numerics="fast")On a GPU with compute capability 9 or later, numerics="fast" multiplies one sequence's decoding step by the FP8 weights directly and rounds longer inputs to FP8 a row at a time (torch._scaled_mm). Elsewhere the weights are decoded to bf16 for each product.
Against bf16
One H100, the same host; CUDA graphs for decoding, a 512-token prompt run eagerly:
| Perplexity, WikiText-2 | Time to first token | Decoding | Peak memory | |
|---|---|---|---|---|
bf16 (llama-3.1-8b-instruct) | 9.55 | 19.1 ms | 162 tokens/s | 16.3 GiB |
| FP8 | 9.62 | 20.3 ms | 197 tokens/s | 9.9 GiB |
Backends
The same source and checkpoint in each framework; numerics and compile are the loaders' options.
PyTorch
from linnet import nest
model = nest.load("llama-3.1-8b-instruct-fp8", backend="torch", numerics="fast", compile="inductor")
logits = model(tokens)JAX
from linnet import nest
f = nest.load("llama-3.1-8b-instruct-fp8", backend="jax") # StableHLO compiled by XLA
logits = f(tokens)Flax NNX
from linnet import nest
model = nest.load("llama-3.1-8b-instruct-fp8", backend="nnx")
logits = model(tokens)Command line
python -m linnet.nest pull llama-3.1-8b-instruct-fp8
linnet stablehlo --root Model --entry forward --bind Vocab=128256 --bind H=4096 --bind Heads=32 --bind KvHeads=8 --bind Inner=14336 --bind Layers=32 --bind Batch=1 --bind MaxSeq=8192 --bind T=bf16 --bind B=1 --bind S=8 <model dir>/src/lib.linnet