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 # XLAthroughput.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:
| Model | Linnet | vLLM |
|---|---|---|
| Llama 3.1 8B, 256 requests | 5600 tokens/s | 5449 |
| gpt-oss 20B, 256 requests | 6031 tokens/s | 4313 |
| Llama 3.1 8B, one request | 167 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 (
Batchrows ofMaxSeqpositions), 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")