Models / gpt-oss-20b

GPT-OSS 20B

A 21B-parameter mixture-of-experts decoder, 32 experts per layer with 4 active per token, alternating 128-position sliding-window and full causal attention with learned attention sinks, YaRN rope scaling for a 131072-token context, and expert weights read straight from the checkpoint's MXFP4 blocks and scales.

20.9B parametersgpt_ossApache-2.0text-generationdecoder-onlymixture-of-expertsgrouped-query-attentionattention-sinkssliding-window-attentionmxfp4

NVIDIA H100 80GB HBM3 · 512 prompt tokens, 128 new · batch 1 · median of 10 runs after 3 warm-ups · measured 2026-09-28T20:56:55+00:00

Linnet 0.1.0, PyTorch 2.14.0, JAX 0.11.2, transformers 5.17.0, diffusers 0.40.0, ONNX Runtime 1.30.0, Python 3.12.3, driver 580.126.09

Linnet against the stack it replaces

Each side in its fastest configuration, on the same GPU and checkpoint. A speed-up is how many times the other's speed; an export is measured against the original checkpoint in the same engine.

vs vLLM, one request1.23× fasterdecode speed
  • Linnet, CUDA graphs368 tok/svLLM299 tok/s
vs the reference, in PyTorch3.67× fasterdecode speed
  • Linnet, CUDA graphs368 tok/stransformers, compiled100 tok/s
vs KerasHub, in JAX4.43× fasterdecode speed
  • Linnet, XLA (generated JAX)298 tok/sKerasHub67.4 tok/s
vs vLLM, serving1.40× fasterserving throughput
  • linnet.serve, CUDA graphs6,031 tok/svLLM4,313 tok/s
vs vLLM, behind Triton3.43× fasterserving throughput
  • Triton, linnet.serve backend5,839 tok/sTriton, vLLM backend1,701 tok/s

Each row also carries its distance from reference: the largest absolute difference between its output and transformers (eager)'s on the same input. bf16 outputs differ by rounding (about 0.1 on logits near 16), so a small number is expected; it is there so that a fast wrong answer cannot look like a win.

Time to first token

ms, lower is better; the best bar is red.

PyTorchJAXONNX Runtime and TensorRTLLM engines050100mstransformers41.9 mstransformers, compiled25.6 msLinnet, generated PyTorch90.0 msLinnet, inductor34.8 msLinnet, CUDA graphs14.1 msKerasHub72.8 msLinnet, XLA (StableHLO)122 msLinnet, XLA (generated JAX)25.2 msLinnet, ONNX Runtime f32failedLinnet, ONNX Runtime f16107 msLinnet, ONNX Runtime bf16113 msLinnet, TensorRT f32failedLinnet, TensorRT f16failedLinnet, TensorRT bf16failedvLLM17.8 ms

Decode speed

tok/s, higher is better; the best bar is red.

PyTorchJAXONNX Runtime and TensorRTLLM engines0100200300tok/stransformers44.9 tok/stransformers, compiled100 tok/sLinnet, generated PyTorch25.5 tok/sLinnet, inductor200 tok/sLinnet, CUDA graphs368 tok/sKerasHub67.4 tok/sLinnet, XLA (StableHLO)141 tok/sLinnet, XLA (generated JAX)298 tok/sLinnet, ONNX Runtime f32failedLinnet, ONNX Runtime f1621.4 tok/sLinnet, ONNX Runtime bf1621.2 tok/sLinnet, TensorRT f32failedLinnet, TensorRT f16failedLinnet, TensorRT bf16failedvLLM299 tok/s

Serving throughput

tok/s, higher is better; the best bar is red.

Serving many requests0200040006000tok/svLLM4,313 tok/slinnet.serve, CUDA graphs6,031 tok/slinnet.serve, XLA3,117 tok/slinnet.serve, ONNX RuntimefailedTriton, vLLM backend1,701 tok/sTriton, linnet.serve backend5,839 tok/stransformers, batchedfailedKerasHub, static batchesfailed

Serving time to first token

ms, lower is better; the best bar is red.

Serving many requests0250050007500mslinnet.serve, CUDA graphs2,302 mslinnet.serve, XLA4,654 mslinnet.serve, ONNX RuntimefailedTriton, vLLM backend8,389 msTriton, linnet.serve backend2,378 mstransformers, batchedfailedKerasHub, static batchesfailed

Load time

s, lower is better; the best bar is red.

PyTorchJAXONNX Runtime and TensorRTLLM enginesServing many requests010002000stransformers115 stransformers, compiled126 sLinnet, generated PyTorch2.54 sLinnet, inductor3.11 sLinnet, CUDA graphs3.21 sKerasHub2,105 sLinnet, XLA (StableHLO)8.61 sLinnet, XLA (generated JAX)9.20 sLinnet, ONNX Runtime f160.24 sLinnet, ONNX Runtime bf160.21 svLLM50.2 svLLM80.0 slinnet.serve, CUDA graphs7.28 slinnet.serve, XLA11.5 sTriton, vLLM backend106 sTriton, linnet.serve backend705 s

Peak GPU memory

What the driver reports the process holding at its peak, in GiB; lower is better; the best bar is red. Not drawn, since theirs is a setting rather than a need: vLLM, vLLM, Triton, vLLM backend (in the table).

PyTorchJAXONNX Runtime and TensorRTServing many requests020406080GiBtransformers40.4 GiBtransformers, compiled40.3 GiBLinnet, generated PyTorch53.3 GiBLinnet, inductor28.7 GiBLinnet, CUDA graphs28.8 GiBKerasHub44.1 GiBLinnet, XLA (StableHLO)14.9 GiBLinnet, XLA (generated JAX)50.1 GiBLinnet, ONNX Runtime f3277.7 GiBLinnet, ONNX Runtime f1668.4 GiBLinnet, ONNX Runtime bf1667.9 GiBLinnet, TensorRT f1668.4 GiBLinnet, TensorRT bf1667.9 GiBlinnet.serve, CUDA graphs30.6 GiBlinnet.serve, XLA52.1 GiBlinnet.serve, ONNX Runtime79.1 GiBTriton, linnet.serve backend31.2 GiBtransformers, batched77.6 GiBKerasHub, static batches60.0 GiB

Every number

MethodTime to first tokenDecode speedServing throughputServing time to first tokenLoad timePeak GPU memoryAgainst the stack it replacesDistance from referenceNotes
PyTorch
transformers41.9 ms44.9 tok/s––115 s40.4 GiB0 (the reference)eager attention: no SDPA for this architecture
transformers, compiled25.6 ms100 tok/s––126 s40.3 GiB0.188eager attention: no SDPA for this architecture
Linnet, generated PyTorch90.0 ms25.5 tok/s––2.54 s53.3 GiB0.25× transformers, compiled0.195KV cache compiled for 768 positions
Linnet, inductor34.8 ms200 tok/s––3.11 s28.7 GiB1.99× transformers, compiled0.188KV cache compiled for 768 positions
Linnet, CUDA graphs14.1 ms368 tok/s––3.21 s28.8 GiB3.67× transformers, compiled0.188KV cache compiled for 768 positions
JAX
KerasHub72.8 ms67.4 tok/s––2,105 s44.1 GiB8.91peak memory is JAX's allocator peak; the driver shows its pool, which grows in whole regions
Linnet, XLA (StableHLO)122 ms141 tok/s––8.61 s14.9 GiB2.09× KerasHub0.203KV 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)25.2 ms298 tok/s––9.20 s50.1 GiB4.43× KerasHub0.109KV 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–––––77.7 GiB––Failed: RuntimeError: Error in execution: Non-zero status code returned while running Mul node. Name:'' Status Message: /onnxruntime_src/onnxruntime/core/framework/bfc_arena.cc:360 void* onnxruntime::BFCArena::AllocateRawInternal(size_t, bool, onnxruntime::Stream*) Failed to allocate memory for requested buffer of size 2123366400
Linnet, ONNX Runtime f16107 ms21.4 tok/s––0.24 s68.4 GiB0.192f16; ONNX Runtime's CUDA execution provider; first calls build the sessions; the argmax taken in the graph, each step replayed as a CUDA graph on the CUDA provider
Linnet, ONNX Runtime bf16113 ms21.2 tok/s––0.21 s67.9 GiB0.188bf16; ONNX Runtime's CUDA execution provider; first calls build the sessions; the argmax taken in the graph, each step replayed as a CUDA graph on the CUDA provider
Linnet, TensorRT f32––––––––Failed: exit -9: [5] Failed to import initializer: In node -1 with name: and operator: (parseGraph): UNSUPPORTED_NODE: Assertion failed: ctx->getWeightsContext().convertOnnxWeights(initializer, &weights): Failed to import initializer: v582
Linnet, TensorRT f16–––––68.4 GiB––Failed: Fail: [ONNXRuntimeError] : 1 : FAIL : TensorRT EP failed to create engine from network for fused node: TensorrtExecutionProvider_TRTKernel_graph_main_8376132635522534149_0_0
Linnet, TensorRT bf16–––––67.9 GiB––Failed: Fail: [ONNXRuntimeError] : 1 : FAIL : TensorRT EP failed to create engine from network for fused node: TensorrtExecutionProvider_TRTKernel_graph_main_11065519248658182765_0_0
LLM engines
vLLM17.8 ms299 tok/s––50.2 s69.2 GiBnot comparedreserves a KV-cache pool up front (gpu_memory_utilization 0.85), so its memory is a setting, not a need
Serving many requests
vLLM––4,313 tok/s–80.0 s68.5 GiBnot compared256 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––6,031 tok/s2,302 ms7.28 s30.6 GiB1.40× vLLMnot compared256 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––3,117 tok/s4,654 ms11.5 s52.1 GiB0.72× vLLMnot compared256 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–––––79.1 GiB––Failed: RuntimeError: Error in execution: Non-zero status code returned while running Einsum node. Name:'' Status Message: /onnxruntime_src/onnxruntime/core/framework/bfc_arena.cc:360 void* onnxruntime::BFCArena::AllocateRawInternal(size_t, bool, onnxruntime::Stream*) Failed to allocate memory for requested buffer of size 301989888
Triton, vLLM backend––1,701 tok/s8,389 ms106 s68.2 GiBnot compared256 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––5,839 tok/s2,378 ms705 s31.2 GiB3.43× Triton, vLLM backendnot compared256 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–––––77.6 GiB––Failed: RuntimeError: generate_batch returned no tokens for any request
KerasHub, static batches–––––60.0 GiB––Failed: JaxRuntimeError: NOT_FOUND: Failed to get configs for: 2 out of 100 instructions. See logs for all failures. Example failure: All configs failed during profiling or were excluded from selection. Failures (19): EXECUTION FAILED: RESOURCE_EXHAUSTED: Out of memory while trying to allocate 14.08GiB with allocator GPU_0_bfc on device 0. [tf-allocator-allocation-error=''] EXECUTION FAILED: RESOURCE_EXH