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

Llama 3.1 8B Instruct W8A8

Llama 3.1 8B Instruct with every decoder projection in int8, one scale per output row, and prompts multiplied in int8 too: RedHatAI's W8A8 checkpoint. 60% of bf16's memory, faster decoding.

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

Llama 3.1 8B Instruct with its decoder projections in int8: RedHatAI's W8A8 checkpoint, made with llm-compressor (GPTQ). Each query, key, value, output, gate, up and down projection holds int8 weights with one bf16 scale per output row; the embedding and the output head stay bf16. The checkpoint was calibrated for inputs rounded to int8 a token at a time, as the products below run them.

The source is the llama-3.1-8b-instruct package with those projections as std.quant::Int8Linear. bindings.json binds each scale to the checkpoint's weight_scale.

Loading ​

python
from linnet import nest

model = nest.load("llama-3.1-8b-instruct-w8a8", backend="torch", numerics="fast")

On a CUDA GPU, the projections multiply the int8 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 int8 a row at a time and multiplied in int8 (Linnet's kernel, summing in int32). Elsewhere the weights are dequantized 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.5518.9 ms163 tokens/s16.3 GiB
W8A89.7021.0 ms271 tokens/s9.9 GiB

Under numerics="equivalent", which never rounds the input, the perplexity is 9.55. Run eagerly, the prompt's two kernels a projection cost the host more than its int8 products save on the device.

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

JAX ​

python
from linnet import nest

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

Flax NNX ​

python
from linnet import nest

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

Command line ​

bash
python -m linnet.nest pull llama-3.1-8b-instruct-w8a8
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