Skip to content

Serving ​

Llama 3.1 8B from Nest, served by linnet.serve: continuous batching over a fixed set of cache rows, prompts packed end to end into one pass, and each step replayed as a CUDA graph. The model is the card's .linnet source; nothing of the original repository runs.

Run ​

bash
python examples/04-serve/throughput.py                  # generated PyTorch, CUDA graphs
python examples/04-serve/throughput.py --backend jax    # XLA

throughput.py gives the engine the benchmarks' load: 256 requests of 128 to 512 prompt tokens, at most 64 in flight, each generating 128 tokens. It prints the generated tokens per second: on one H100, about 5600 with PyTorch and 4700 with JAX.

On one H100 in bf16:

ModelLinnetvLLM
Llama 3.1 8B, 256 requests5600 tokens/s5449
gpt-oss 20B, 256 requests6031 tokens/s4313
Llama 3.1 8B, one request167 tokens/s (XLA)152

The benchmarks have every row and how it was measured.

An OpenAI-compatible server ​

bash
linnet serve llama-3.1-8b-instruct --batch 32 --max-seq 4096 --port 8000
curl http://127.0.0.1:8000/v1/chat/completions -H 'Content-Type: application/json' \
     -d '{"model": "llama-3.1-8b-instruct", "messages": [{"role": "user", "content": "Hi"}]}'

The server takes OpenAI's completions and chat completions, streamed or not. Serving and tools lists the options, and the exports to vLLM, Triton and llama.cpp.

How the engine is fast ​

  • Fixed rows. Each request in flight owns one row of the KV caches (Batch rows of MaxSeq positions), so a decoding step has one shape and replays as one captured CUDA graph.
  • Packed prompts. Prompts that arrive together are packed end to end into one pass with no padding, each attending only to itself (FlexAttention over the blocks it reaches). Prompts that arrive while other rows decode join that step's pass.
  • Shared prompts. Requests with the same prompt prefill it once and copy its cache rows.

Source ​

throughput.py ​

python
"""Serves Llama 3.1 8B from Nest with `linnet.serve` under the benchmarks'
load: 256 requests of 128 to 512 prompt tokens, at most 64 in flight, each
generating 128 tokens. Prints the generated tokens per second."""

import argparse
import random
import time

from linnet import nest
from linnet.serve import Engine, Request

parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--card", default="llama-3.1-8b-instruct")
parser.add_argument("--backend", choices=["torch", "jax"], default="torch")
parser.add_argument("--requests", type=int, default=256)
parser.add_argument("--rows", type=int, default=64, help="requests in flight")
parser.add_argument("--new", type=int, default=128, help="tokens each request generates")
args = parser.parse_args()

rng = random.Random(0)
prompts = [
    [rng.randrange(100, 20000) for _ in range(rng.randint(128, 512))] for _ in range(args.requests)
]
# One cache row per request in flight, long enough for the longest prompt
# and its completion.
generics = {"Batch": args.rows, "MaxSeq": 512 + args.new, "T": "bf16"}
if args.backend == "torch":
    # Generated PyTorch, each step replayed as one CUDA graph.
    model = nest.load(
        args.card,
        backend="torch",
        device="cuda",
        numerics="fast",
        compile=True,
        generics=generics,
        cast_dtype=True,
    )
else:
    # Every entry over one copy of the weights, compiled by XLA.
    model = nest.load(
        args.card, backend="jax_model", numerics="fast", generics=generics, cast_dtype=True
    )

engine = Engine(model)
start = time.perf_counter()
engine.warmup(len(prompt) for prompt in prompts)
print(f"compiled in {time.perf_counter() - start:.0f} s")

# No end token: every request generates all its tokens, as in the benchmarks.
done, stats = engine.run([Request(prompt=prompt, max_new_tokens=args.new) for prompt in prompts])
rate = stats.generated_tokens / stats.seconds
print(f"{stats.generated_tokens} tokens in {stats.seconds:.2f} s: {rate:.0f} tokens/s")

Browse on GitHub

Released under the MIT License.