Models / llama-3.1-8b-instruct-fp8

Llama 3.1 8B Instruct FP8

Llama 3.1 8B Instruct with every decoder projection in FP8 E4M3, one scale per output row: RedHatAI's FP8-dynamic checkpoint. Half the weight memory of bf16, faster single-sequence decoding.

8B parametersllamallama3.1text-generationdecoder-onlygrouped-query-attentionchatfp8

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 ​

python
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-2Time to first tokenDecodingPeak memory
bf16 (llama-3.1-8b-instruct)9.5519.1 ms162 tokens/s16.3 GiB
FP89.6220.3 ms197 tokens/s9.9 GiB

Backends ​

The same source and checkpoint in each framework; numerics and compile are the loaders' options.

PyTorch ​

python
from linnet import nest

model = nest.load("llama-3.1-8b-instruct-fp8", backend="torch", numerics="fast", compile="inductor")
logits = model(tokens)

JAX ​

python
from linnet import nest

f = nest.load("llama-3.1-8b-instruct-fp8", backend="jax")          # StableHLO compiled by XLA
logits = f(tokens)

Flax NNX ​

python
from linnet import nest

model = nest.load("llama-3.1-8b-instruct-fp8", backend="nnx")
logits = model(tokens)

Command line ​

bash
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