Benchmarks
Linnet compiles one checked source to each framework's own fast path. Measured in bf16 on H100s: 24 real checkpoints from the model zoo against the stacks people already run them with, Llama 3.1 8B fine-tuned against TRL, and a synthetic Llama that isolates the generated code.
Faster than vLLM at batch 1
Faster on 10 of 10 models: 1.08× to 2.72×, median 1.21×.
the stack Linnet replacesLinnetthe faster of the twodecode speed in tok/s, higher is better. Each card is one model on its own scale, the stack's bar above Linnet's; the faster of the two is red. Each side is its fastest configuration.
One request at a time, the speed a chat turn or agent step sees, Linnet's fastest path (generated PyTorch as CUDA graphs, or XLA) decodes faster than vLLM on the same checkpoint. Portability costs nothing here: the generated code calls the kernels hand-written code would, and each decoding step runs as one CUDA graph or XLA program, with the host out of the loop.
Faster in PyTorch and JAX
Faster on 24 of 24 models: 1.00× to 4.81×, median 2.09×.
the stack Linnet replacesLinnetthe faster of the twoEach model by its own measure: decode speed for a decoder, transcription for Whisper, batch-1 latency for the rest. Each card is one model on its own scale, the stack's bar above Linnet's; the faster of the two is red. Each side is its fastest configuration.
Against the implementation people run on each framework today, Linnet's code is faster on every model in PyTorch and JAX. The PyTorch gap is widest where the host sets the pace (small models, single decoding steps, batch-1 encoders), because one CUDA graph removes that overhead. It closes on the large convolutional models: the SD VAE decoder and SDXL UNet come out even with torch.compile, though Linnet's XLA program decodes the VAE fastest (7.4 ms). On ONNX Runtime and behind Triton, Linnet is ahead on most models, but the reference export is faster on MiniLM by 16% and on ResNet-50 by 13%, and 21% ahead on ResNet-50 behind Triton.
Exports run as the originals do
Within 5% of the original in 16 of 19 engine and model pairs, 2.9% apart at the median; the same first token in 14 of 14.
the stack Linnet replacesLinnetdecode speed in tok/s, higher is better. Each card is one model on its own scale, the original's bar above the export's. Each side is its fastest configuration.
A model exported with linnet.hf or linnet.gguf runs in the engines people deploy like the original checkpoint, one request at a time and serving: the engine cannot tell the difference. The smallest models' serving pairs vary the most, since their runs last a second or two.
Serving many requests
Faster on 10 of 10 models: 1.02× to 1.84×, median 1.09×.
the stack Linnet replacesLinnetthe faster of the twoserving throughput in tok/s, higher is better. Each card is one model on its own scale, the stack's bar above Linnet's; the faster of the two is red. Each side is its fastest configuration.
linnet.serve serves faster than vLLM on every decoder. The lead is widest on small models, where a step is bound by kernel launches that one CUDA graph per step removes. From 4 B up it is 2-4%, and on gpt-oss 40%. Behind Triton Inference Server, linnet.serve outpaces Triton's own vLLM backend on every decoder.
Training faster than TRL
The same runs in Linnet and in TRL, step for step. On one GPU, Linnet's PyTorch step is 1.76 times TRL's for LoRA SFT, 2.4 times for DPO and 2.7 times for GRPO. Fully fine-tuned on four GPUs, the three stacks are within 5% of each other, with JAX the fastest. Each pair of runs reaches the same held-out loss, accuracy or reward. Training lists the conditions that differ.
| Stack | Step | Tokens/s | Peak a GPU | held-out loss after 30 steps, one evaluator |
|---|---|---|---|---|
| Linnet PyTorch | 1.04 s | 15.6K | 36.4 GiB | 1.369 |
| Linnet JAX | 1.32 s | 12.2K | 38.0 GiB | 1.374 |
| TRL + PEFT | 1.82 s | 8.9K | 42.8 GiB | 1.373 |
| Stack | Step | Tokens/s | Peak a GPU | held-out loss after 30 steps |
|---|---|---|---|---|
| Linnet PyTorch | 0.55 s | 29.4K | 47.5 GiB | 1.342 |
| Linnet JAX | 0.54 s | 30.0K | 65.5 GiB | 1.345 |
| Linnet JAX, remat | 0.67 s | 24.0K | 42.7 GiB | 1.344 |
| TRL, FSDP2 | 0.57 s | 28.6K | 51.5 GiB | last-5 training loss 1.26 (Linnet 1.24) |
| Stack | Step | Peak a GPU | accuracy over the last five steps |
|---|---|---|---|
| Linnet PyTorch | 1.91 s | 35.6 GiB | 0.68 |
| Linnet JAX | 2.41 s | 38.3 GiB | 0.65 |
| TRL + PEFT (2 pairs at a time; 4 and 8 ran out of memory) | 4.63 s | 47.9 GiB | 0.70 |
| Stack | Step | mean reward, first three steps to last three |
|---|---|---|
| Linnet PyTorch | 1.18 s | 0.44 → 0.39 |
| Linnet JAX | 1.44 s | 0.44 → 0.40 |
| TRL + vLLM (8 completions at a time; vLLM held 35% of the GPU outside this peak) | 3.14 s | 0.43 → 0.41 |
Two GPUs and offloading
Faster on 1 of 2 models: 0.96× to 1.02×, median 0.99×.
the stack Linnet replacesLinnetthe faster of the twodecode speed in tok/s, higher is better. Each card is one model on its own scale, the stack's bar above Linnet's; the faster of the two is red. Each side is its fastest configuration.
Split across two GPUs at batch 1, the two are within 4%: Linnet is ahead on Llama 3.1 8B, vLLM on Qwen3 8B. The offloaded rows run Llama 3.1 8B and Qwen3 8B on a GPU capped at 8 GiB, streaming half the layers from host memory: Llama peaks at 8.9 GiB and decodes 5.7 tokens per second, where otherwise it would not load at all. These rows show what a small GPU can do, so they are never marked best.
One source, every runtime
| Model | PyTorch | XLA | ONNX Runtime | TensorRT | Triton | linnet.serve | vLLM | SGLang | TGI | llama.cpp |
|---|---|---|---|---|---|---|---|---|---|---|
| all-MiniLM-L6-v2 | ✓ | ✓ | ✓ | ✓ | ✓ | · | · | · | · | · |
| BERT Base (uncased) | ✓ | ✓ | ✓ | ✓ | ✓ | · | · | · | · | · |
| RoBERTa Base | ✓ | ✓ | ✓ | ✓ | ✓ | · | · | · | · | · |
| ModernBERT-base | ✓ | ✓ | ✓ | ✓ | ✓ | · | · | · | · | · |
| ResNet-18 | ✓ | ✓ | ✓ | ✓ | ✓ | · | · | · | · | · |
| ResNet-50 | ✓ | ✓ | ✓ | ✓ | ✓ | · | · | · | · | · |
| ViT-Base/16 224 | ✓ | ✓ | ✓ | ✓ | ✓ | · | · | · | · | · |
| DINOv2-Base | ✓ | ✓ | ✓ | ✓ | ✓ | · | · | · | · | · |
| SAM ViT-Base | ✓ | ✓ | ✓ | ✓ | · | · | · | · | · | · |
| SigLIP Base/16 224 | ✓ | ✓ | ✓ | ✓ | ✓ | · | · | · | · | · |
| Whisper tiny | ✓ | ✓ | ✓ | ✓ | · | · | · | · | · | · |
| Whisper large-v3 | ✓ | ✓ | ✓ | ✓ | · | · | · | · | · | · |
| SD VAE ft-MSE (decoder) | ✓ | ✓ | ✓ | ✓ | · | · | · | · | · | · |
| SDXL Base UNet | ✓ | ✓ | ✓ | ✓ | · | · | · | · | · | · |
| GPT-2 (124M) | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ |
| Qwen2.5 0.5B Instruct | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | · | · | ✓ |
| TinyLlama 1.1B Chat v1.0 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ |
| SmolLM2 1.7B Instruct | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ |
| Phi-3 Mini 4K Instruct | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | · | · | ✓ |
| Qwen3 4B | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | · | · | ✓ |
| Mistral 7B Instruct v0.3 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ |
| Llama 3.1 8B Instruct | ✓ | ✓ | ✓ | ✓ 2 of 3 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ |
| Qwen3 8B | ✓ | ✓ | ✓ | ✓ 2 of 3 | ✓ | ✓ | ✓ | · | · | ✓ |
| GPT-OSS 20B | ✓ | ✓ | ✓ 2 of 3 | ✕ failed | ✓ | ✓ 2 of 3 | · | · | · | · |
✓ ran · ✓ n of m some of its configurations ran · ✕ failed · · not a runtime this model takes part in. Hover a cell for each configuration.
Every model in the zoo is one .linnet source. The runs that failed:
- TensorRT cannot build an engine for gpt-oss (its Myelin compiler fails inside NVRTC) or for the 8 B decoders in
f32. - gpt-oss in
f32is 84 GB of weights, more than the GPU holds, and its ONNX serving row runs out of memory. - transformers'
generate_batchreturns no tokens for gpt-oss, and KerasHub's batched gpt-oss fails in XLA's autotuner.
Every measurement
Hover a bar for its notes and how it compares with the stack it replaces. Each model's page in the zoo has its samples and every number.
Decoders, one request at a time
Serving
Encoders, vision, audio, and diffusion
Models with no KerasHub implementation have no JAX baseline.
Every row
| Method | TTFT ms | tok/s | peak GiB | load s | max |diff| | Notes |
|---|---|---|---|---|---|---|
| PyTorch | ||||||
| transformers | 22.2 | 51.2 | 16.0 | 5.29 | 0 | |
| transformers, compiled | 18.6 | 110 | 16.1 | 4.29 | 0.063 | |
| Linnet, generated PyTorch | 17.4 | 79.0 | 15.9 | 3.59 | 0.094 | KV cache compiled for 768 positions |
| Linnet, inductor | 13.1 | 141 | 22.8 | 3.35 | 0.063 | KV cache compiled for 768 positions |
| Linnet, CUDA graphs | 13.0 | 161 | 22.8 | 3.10 | 0.063 | KV cache compiled for 768 positions |
| JAX | ||||||
| KerasHub | 27.5 | 160 | 17.3 | 31.0 | 0.094 | peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions |
| Linnet, XLA (StableHLO) | 16.6 | 167 | 16.3 | 10.7 | 0.094 | KV cache compiled for 768 positions; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions |
| Linnet, XLA (generated JAX) | 15.7 | 166 | 16.3 | 10.6 | 0.096 | KV cache compiled for 768 positions; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions |
| ONNX Runtime and TensorRT | ||||||
| Linnet, ONNX Runtime f32 | 42.0 | 53.7 | 35.2 | 0.50 | 0.065 | f32; ONNX Runtime's CUDA execution provider; first calls build the sessions; logits copied to the host each step for the argmax |
| Linnet, ONNX Runtime f16 | 30.5 | 71.0 | 17.5 | 0.51 | 0.070 | f16; ONNX Runtime's CUDA execution provider; first calls build the sessions; logits copied to the host each step for the argmax |
| Linnet, ONNX Runtime bf16 | 31.1 | 71.8 | 16.5 | 0.48 | 0.094 | bf16; ONNX Runtime's CUDA execution provider; first calls build the sessions; logits copied to the host each step for the argmax |
| Linnet, TensorRT f32 | 65.5 | failed: Fail: [ONNXRuntimeError] : 1 : FAIL : TensorRT EP failed to create engine from network for fused node: TensorrtExecutionProvider_TRTKernel_graph_main_3492707504 | ||||
| Linnet, TensorRT f16 | 28.6 | 47.5 | 42.0 | 0.49 | 9.8 | f16; ONNX Runtime's TensorRT execution provider; first calls build the sessions; logits copied to the host each step for the argmax |
| Linnet, TensorRT bf16 | 28.5 | 49.5 | 39.9 | 0.49 | 0.13 | bf16; ONNX Runtime's TensorRT execution provider; first calls build the sessions; logits copied to the host each step for the argmax |
| LLM engines | ||||||
| vLLM | 16.5 | 152 | 67.4 | 74.6 | reserves a KV-cache pool up front (gpu_memory_utilization 0.85), so its memory is a setting, not a need | |
| vLLM on Linnet's export | 12.2 | 157 | 67.4 | 39.3 | linnet.hf.export (16 s), then vLLM on the exported checkpoint; reserves a KV-cache pool up front (gpu_memory_utilization 0.85), so its memory is a setting, not a need | |
| SGLang | 21.0 | 158 | 68.0 | 55.7 | reserves a KV-cache pool up front (mem_fraction_static 0.85), so its memory is a setting, not a need | |
| SGLang on Linnet's export | 19.9 | 158 | 67.9 | 38.4 | linnet.hf.export (18 s), then SGLang on the exported checkpoint; reserves a KV-cache pool up front (mem_fraction_static 0.85), so its memory is a setting, not a need | |
| TGI | 27.9 | 112 | 59.2 | 44.1 | over HTTP, streamed; the prompt is the token ids decoded and tokenized again; reserves a KV-cache pool up front (cuda-memory-fraction 0.85), so its memory is a setting, not a need | |
| TGI on Linnet's export | 25.7 | 115 | 59.2 | 40.1 | linnet.hf.export (19 s), then Text Generation Inference on the exported checkpoint; over HTTP, streamed; the prompt is the token ids decoded and tokenized again; reserves a KV-cache pool up front (cuda-memory-fraction 0.85), so its memory is a setting, not a need | |
| llama.cpp on Linnet's GGUF | 22.0 | 171 | 15.0 | 75.5 | linnet.gguf.export, then llama-bench with every layer on the GPU; the first token is the prompt at llama-bench's prompt rate | |
| Beyond one GPU | ||||||
| vLLM, tensor parallel | 11.9 | 233 | 139.4 | 286 | reserves a KV-cache pool up front (gpu_memory_utilization 0.85), so its memory is a setting, not a need | |
| Linnet, tensor parallel (PyTorch) | 10.8 | 238 | 27.1 | 4.77 | one process per GPU under torchrun, NCCL, CUDA graphs; KV cache compiled for 768 positions | |
| Linnet, tensor parallel (XLA) | 16.2 | 181 | 18.4 | 16.1 | 0.13 | KV cache compiled for 768 positions; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions |
| Linnet, layers on two GPUs | 22.5 | 87.8 | 23.3 | 7.00 | 0.094 | KV cache compiled for 768 positions |
| Linnet, offloaded to host | 180 | 5.73 | 8.9 | 19.6 | 0.094 | KV cache compiled for 768 positions; cuda:0: embedding, layers.0-13; host, streamed in: layers.14-31, norm, lm_head |
| Serving many requests | ||||||
| vLLM | 68.7 | 135 | 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at gpu_memory_utilization 0.85 | |||
| linnet.serve, CUDA graphs | 27.8 | 6.99 | 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; 64 fixed cache rows of 640 positions | |||
| linnet.serve, XLA | 22.1 | 12.0 | 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; 64 fixed cache rows of 640 positions; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions | |||
| linnet.serve, ONNX Runtime | 37.6 | 0.26 | 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; 64 fixed cache rows of 640 positions | |||
| vLLM on Linnet's export | 68.7 | 34.7 | linnet.hf.export (0 s), then vLLM on the exported checkpoint; 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at gpu_memory_utilization 0.85 | |||
| SGLang | 70.0 | 55.8 | 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at mem_fraction_static 0.85 | |||
| SGLang on Linnet's export | 69.9 | 44.9 | linnet.hf.export (0 s), then SGLang (offline, continuous batching) on the exported checkpoint; 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; KV-cache pool at mem_fraction_static 0.85 | |||
| TGI | 60.7 | 42.1 | 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; over HTTP, streamed; stops at end-of-sequence (TGI cannot ignore it), so throughput counts the tokens produced; KV-cache pool at cuda-memory-fraction 0.85 | |||
| TGI on Linnet's export | 61.2 | 40.1 | linnet.hf.export (0 s), then Text Generation Inference (continuous batching) on the exported checkpoint; 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; over HTTP, streamed; stops at end-of-sequence (TGI cannot ignore it), so throughput counts the tokens produced; KV-cache pool at cuda-memory-fraction 0.85 | |||
| Triton, vLLM backend | 67.5 | 61.1 | 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; over HTTP, streamed, first tokens timed at the client | |||
| Triton, linnet.serve backend | 28.3 | 645 | 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; over HTTP, streamed, first tokens timed at the client | |||
| transformers, batched | 79.1 | 3.98 | 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; paged|sdpa attention; reserves a paged KV-cache pool up front, so its memory is a setting, not a need; no per-request timestamps | |||
| KerasHub, static batches | 25.8 | 30.1 | 256 requests of 128-512 prompt tokens and 128 new tokens, at most 64 in flight; every batch runs to the longest possible prompt plus the new tokens; peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions | |||
The generated code
Measured 2026-09-24 on device NVIDIA H100 80GB HBM3, torch 2.9.1+cu128, python 3.12.3, os Linux 6.8.0-90-generic, jax 0.11.2 (gpu), linnet linnet 0.1.0.
Llama-style decoder · H=512 L=8 heads=8/8 B=1 S=512 bf16 · forward
| Variant | Latency (ms) | tokens/s | vs reference | Kernels | max |Δ| |
|---|---|---|---|---|---|
| PyTorch reference (eager) | 2.71 | 188,954 | 358 | 0 | |
| PyTorch reference (torch.compile) | 1.33 | 384,208 | 2.03× | 117 | 4.9e-4 |
| Linnet → PyTorch (numerics=equivalent) Core IR interpreted per call; library ops on native kernels | 3.02 | 169,341 | 0.90× | 246 | 2.4e-4 |
| Linnet → PyTorch (generated source) straight-line PyTorch from `linnet torch`, native kernels | 3.33 | 153,569 | 0.81× | 246 | 2.4e-4 |
| Linnet → PyTorch (generated source, numerics=fast) straight-line PyTorch from `linnet torch`, native kernels | 3.24 | 158,084 | 0.84× | 221 | 0 |
| Linnet → PyTorch (generated source + torch.compile) straight-line PyTorch from `linnet torch`, native kernels | 1.77 | 288,779 | 1.53× | 122 | 4.9e-4 |
| Linnet → PyTorch (generated source + torch.compile, numerics=fast) straight-line PyTorch from `linnet torch`, native kernels | 1.54 | 333,531 | 1.77× | 106 | 4.9e-4 |
| Linnet → PyTorch (generated source + CUDA graphs, numerics=fast) straight-line PyTorch from `linnet torch`, native kernels | 0.81 | 635,823 | 3.36× | 109 | 4.9e-4 |
| Linnet → XLA (jax, gpu) whole-program compile of the StableHLO export | 0.91 | 565,059 | 2.99× | — | 2.4e-4 |
| Linnet → XLA (jax, gpu, numerics=fast) whole-program compile of the StableHLO export | 1.34 | 381,419 | 2.02× | — | 2.4e-4 |
Llama-style decoder · H=512 L=8 heads=8/8 B=1 S=512 bf16 · decode
| Variant | Latency (ms) | steps/s | vs reference | Kernels | max |Δ| |
|---|---|---|---|---|---|
| PyTorch reference (eager, recomputes prefix) no KV cache: full forward over 257 tokens | 3.49 | 286.5 | 352 | 0 | |
| Linnet → PyTorch (numerics=equivalent) KV cache at position 256 of 512 | 3.66 | 273.3 | 0.95× | 317 | 2.4e-4 |
| Linnet → PyTorch (generated source) KV cache at position 256 of 512 | 3.64 | 274.5 | 0.96× | 317 | 2.4e-4 |
| Linnet → PyTorch (generated source, numerics=fast) KV cache at position 256 of 512 | 3.42 | 292.8 | 1.02× | 285 | 2.4e-4 |
| Linnet → PyTorch (generated source + torch.compile) KV cache at position 256 of 512 | 2.18 | 459.4 | 1.60× | 144 | 2.9e-4 |
| Linnet → PyTorch (generated source + torch.compile, numerics=fast) KV cache at position 256 of 512 | 2.02 | 495.4 | 1.73× | 119 | 3.1e-4 |
| Linnet → PyTorch (generated source + CUDA graphs, numerics=fast) KV cache at position 256 of 512 | 1.09 | 915.6 | 3.20× | 155 | 3.1e-4 |
| Linnet → XLA (jax, gpu) KV cache threaded as inputs/outputs, position 256 | 1.94 | 516.3 | 1.80× | — | 2.4e-4 |
| Linnet → XLA (jax, gpu, numerics=fast) KV cache threaded as inputs/outputs, position 256 | 1.70 | 586.5 | 2.05× | — | 2.4e-4 |
Llama-style decoder · H=2048 L=22 heads=32/8 B=1 S=512 bf16 · forward
| Variant | Latency (ms) | tokens/s | vs reference | Kernels | max |Δ| |
|---|---|---|---|---|---|
| PyTorch reference (eager) | 8.63 | 59,328 | 900 | 0 | |
| PyTorch reference (torch.compile) | 3.68 | 139,245 | 2.35× | 325 | 1.1e-3 |
| Linnet → PyTorch (numerics=equivalent) Core IR interpreted per call; library ops on native kernels | 11.42 | 44,822 | 0.76× | 1030 | 9.8e-4 |
| Linnet → PyTorch (generated source) straight-line PyTorch from `linnet torch`, native kernels | 11.24 | 45,571 | 0.77× | 1030 | 9.8e-4 |
| Linnet → PyTorch (generated source, numerics=fast) straight-line PyTorch from `linnet torch`, native kernels | 7.15 | 71,587 | 1.21× | 568 | 4.9e-4 |
| Linnet → PyTorch (generated source + torch.compile) straight-line PyTorch from `linnet torch`, native kernels | 6.18 | 82,843 | 1.40× | 388 | 1.1e-3 |
| Linnet → PyTorch (generated source + torch.compile, numerics=fast) straight-line PyTorch from `linnet torch`, native kernels | 4.15 | 123,289 | 2.08× | 278 | 1.1e-3 |
| Linnet → PyTorch (generated source + CUDA graphs, numerics=fast) straight-line PyTorch from `linnet torch`, native kernels | 3.79 | 134,947 | 2.27× | 282 | 1.1e-3 |
| Linnet → XLA (jax, gpu) whole-program compile of the StableHLO export | 6.39 | 80,106 | 1.35× | — | 9.8e-4 |
| Linnet → XLA (jax, gpu, numerics=fast) whole-program compile of the StableHLO export | 5.85 | 87,575 | 1.48× | — | 9.8e-4 |
Llama-style decoder · H=2048 L=22 heads=32/8 B=1 S=512 bf16 · decode
| Variant | Latency (ms) | steps/s | vs reference | Kernels | max |Δ| |
|---|---|---|---|---|---|
| PyTorch reference (eager, recomputes prefix) no KV cache: full forward over 257 tokens | 9.06 | 110.4 | 895 | 0 | |
| Linnet → PyTorch (numerics=equivalent) KV cache at position 256 of 512 | 12.71 | 78.6 | 0.71× | 1119 | 6.1e-4 |
| Linnet → PyTorch (generated source) KV cache at position 256 of 512 | 12.57 | 79.6 | 0.72× | 1118 | 6.1e-4 |
| Linnet → PyTorch (generated source, numerics=fast) KV cache at position 256 of 512 | 12.74 | 78.5 | 0.71× | 1140 | 6.1e-4 |
| Linnet → PyTorch (generated source + torch.compile) KV cache at position 256 of 512 | 6.01 | 166.3 | 1.51× | 514 | 7.3e-4 |
| Linnet → PyTorch (generated source + torch.compile, numerics=fast) KV cache at position 256 of 512 | 6.07 | 164.7 | 1.49× | 513 | 7.3e-4 |
| Linnet → PyTorch (generated source + CUDA graphs, numerics=fast) KV cache at position 256 of 512 | 3.52 | 283.9 | 2.57× | 598 | 7.3e-4 |
| Linnet → XLA (jax, gpu) KV cache threaded as inputs/outputs, position 256 | 3.51 | 285.0 | 2.58× | — | 7.3e-4 |
| Linnet → XLA (jax, gpu, numerics=fast) KV cache threaded as inputs/outputs, position 256 | 3.78 | 264.8 | 2.40× | — | 8.1e-4 |
Compile and load
| Step | Seconds |
|---|---|
| linnet check examples/05-llama | 1.85 |
| linnet.torch.load (small) | 2.31 |
| torch.compile (small) | 1.17 |
| linnet torch + first call (small, generated source) | 4.82 |
| linnet torch + first call (small, generated source, numerics=fast) | 4.62 |
| linnet torch + first call (small, generated source + torch.compile) | 6.85 |
| linnet torch + first call (small, generated source + torch.compile, numerics=fast) | 4.94 |
| linnet torch + first call (small, generated source + CUDA graphs, numerics=fast) | 4.77 |
| linnet stablehlo + XLA compile (small, forward, numerics=equivalent) | 6.53 |
| linnet stablehlo + XLA compile (small, forward, numerics=fast) | 4.24 |
| linnet stablehlo + XLA compile (small, decode, numerics=equivalent) incl. 256 steps | 5.63 |
| linnet stablehlo + XLA compile (small, decode, numerics=fast) incl. 256 steps | 5.46 |
| linnet.torch.load (medium) | 2.22 |
| torch.compile (medium) | 2.75 |
| linnet torch + first call (medium, generated source) | 2.92 |
| linnet torch + first call (medium, generated source, numerics=fast) | 4.82 |
| linnet torch + first call (medium, generated source + torch.compile) | 8.24 |
| linnet torch + first call (medium, generated source + torch.compile, numerics=fast) | 4.86 |
| linnet torch + first call (medium, generated source + CUDA graphs, numerics=fast) | 5.58 |
| linnet stablehlo + XLA compile (medium, forward, numerics=equivalent) | 8.62 |
| linnet stablehlo + XLA compile (medium, forward, numerics=fast) | 5.79 |
| linnet stablehlo + XLA compile (medium, decode, numerics=equivalent) incl. 256 steps | 7.76 |
| linnet stablehlo + XLA compile (medium, decode, numerics=fast) incl. 256 steps | 6.89 |
A Llama-shaped model with random weights, against a hand-written PyTorch implementation, shows what the generated code costs on its own. On the medium model, plain generated PyTorch runs a forward pass in 7.2 ms against 8.6 ms for the eager reference, and CUDA graphs bring it to 3.8 ms against 3.7 ms for the compiled reference. The small model's plain generated code is slower than eager (3.2 ms against 2.7 ms) until compiled (1.5 ms) or replayed as CUDA graphs (0.8 ms). A decode step is launch-bound: the medium model's spends most of its 12.7 ms in Python dispatch, so use CUDA graphs in a decoding loop, which bring it to 3.5 ms.
How it was measured
- Hardware: one H100 80GB (SXM), each method in its own process.
- Dtype:
bf16. - Workloads: one request at a time is a 512-token prompt, then 128 greedy tokens. Serving is 256 requests of 128 to 512 prompt tokens, each wanting 128 new tokens, all waiting from the start, at most 64 in flight. Encoders, vision, audio and diffusion run one forward pass at batch 1 (latency) and at a large batch (throughput).
- Baselines: transformers or diffusers (eager or
torch.compile, whichever is faster), sentence-transformers, KerasHub in JAX, the model's owntorch.onnxexport on ONNX Runtime and behind Triton Inference Server, vLLM (prefix cache off, since every timed call sends the same prompt), SGLang, Text Generation Inference, and llama.cpp. - Linnet's row: its fastest configuration across generated PyTorch (as is, under
torch.compile, or as CUDA graphs), XLA, ONNX Runtime (CUDA or TensorRT), and Triton. Every output is checked against the reference stack's. - Memory: peak GPU memory from the driver, or from JAX's allocator for JAX rows. The memory views leave out the offloaded rows and every row that starts vLLM, which reserves 85% of the GPU for its cache pool.
- Generated code:
examples/01-llamawith random weights, small (about 60 M parameters) and medium (TinyLlama shape, about 1.1 B), timingforwardoverB=1, S=512and onedecodestep at position 256. The variants are thelinnet.torch.loadmodes andlinnet.jax.loadunderjax.jit. - Training: Llama 3.1 8B Instruct on H100s, packed into 4096-token rows. SFT uses Alpaca, DPO UltraFeedback pairs, and GRPO two-number multiplications with an exact-answer reward. The reference is TRL with PEFT, FSDP2 for the four-GPU run and vLLM colocated for GRPO. Every stack uses the same hyperparameters and step count.
- Data: the charts and tables come from
bench/results/zoo.json,bench/results/latest.jsonandbench/results/training.json; the prose quotes them. The zoo'sbench/compare.jsonpairs the rows, and itsbench/setup-pod.shandbench/run-all.shreproduce them on NVIDIA's Triton Inference Server image.python bench/zoo.py <zoo checkout>refreshes this page's copy.
To rerun the generated-code benchmark (bench/setup-pod.sh prepares a fresh GPU machine; bench/README.md lists the published environment):
cd python/linnet && uv sync --all-extras
cd ../..
LINNET_BIN=build/release/linnet python bench/run.py --device cuda --out bench/results/latest.json