Models / qwen3-8b-fp8

Qwen3 8B FP8

Qwen3 8B with every decoder projection in FP8 E4M3, one scale per block of 128 by 128 weights: Qwen's own FP8 checkpoint. 60% of bf16's memory, faster decoding.

8.2B parametersqwen3Apache-2.0text-generationdecoder-onlygrouped-query-attentionqk-normfp8

Qwen3 8B with its decoder projections in FP8: Qwen's own FP8 checkpoint. Each query, key, value, output, gate, up and down projection holds E4M3 weights (one byte each) with one bf16 scale per block of 128 by 128 weights, the scheme DeepSeek-V3 uses; the embedding, the norms and the output head stay bf16.

The source is the qwen3-8b package with those projections as std.quant::Fp8BlockLinear. Its blocks need every projection's widths to be multiples of 128, which the where clauses state. bindings.json binds each scale to the checkpoint's weight_scale_inv (a multiplier, despite its name) and each FP8 weight to its bytes.

Loading ​

python
from linnet import nest

model = nest.load("qwen3-8b-fp8", backend="torch", numerics="fast")

On a GPU with compute capability 9 or later, the projections multiply the FP8 weights as they are: one row in a kernel that reads each byte once, up to 32 rows on tensor cores, and longer inputs rounded to FP8 128 inputs at a time (torch._scaled_mm with block scales). 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 (qwen3-8b)12.7121.6 ms157 tokens/s16.9 GiB
FP812.8423.2 ms238 tokens/s10.4 GiB

Under numerics="equivalent", which never rounds the input, the perplexity is 12.77. The prompt runs slower than bf16's: cuBLAS's block-scaled FP8 product is no faster than bf16 at these sizes.

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("qwen3-8b-fp8", backend="torch", numerics="fast", compile="inductor")
logits = model(tokens)

JAX ​

python
from linnet import nest

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

Flax NNX ​

python
from linnet import nest

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

Command line ​

bash
python -m linnet.nest pull qwen3-8b-fp8
linnet stablehlo --root Model --entry forward --bind Vocab=151936 --bind H=4096 --bind Heads=32 --bind KvHeads=8 --bind HeadDim=128 --bind Inner=12288 --bind Layers=36 --bind Batch=1 --bind MaxSeq=40960 --bind T=bf16 --bind B=1 --bind S=8 <model dir>/src/lib.linnet