Write Once, Run Everywhere: The Axon DSL for Shape-Safe and Framework-Agnostic LLM Architectures
arXiv:2608.198892026-08-21
A new language, Axon, lets you write an LLM once and run it on PyTorch, JAX, MLX or vLLM
The open-source language model ecosystem effectively depends on a single platform, Hugging Face, and porting models to other frameworks has meant manually rewriting them, often losing optimizations along the way. The authors built Axon, a strongly typed, Haskell-inspired language for describing model architectures once, from which a compiler automatically generates standalone code for PyTorch, Triton, JAX, MLX and vLLM. Across 467 inference benchmarks on models from 135M to 32B parameters, Axon-derived models were generally faster than reference Transformers implementations.
What they did
Problem: LLM code is locked into a specific platform's conventions, so porting to other frameworks requires manual, error-prone rewrites that often drop backend-specific optimizations
Solution: Axon lets researchers describe a model's structure (layers, tensor shapes, parameter locations) once; the compiler lowers this into a shared Graph IR and then generates standalone code for five backends
A strong type system checks tensor shapes ahead of time, preventing shape-mismatch errors during execution
Results: median speedups of 7% on PyTorch, 12% on PyTorch+Triton, 91% on JAX, and 107% on MLX versus Transformers reference implementations; 58% median speedup when deployed natively on vLLM with PagedAttention and KV-cache
Training experiments showed an Axon-derived PyTorch model matched the loss curve of the standard Transformers implementation while running about 9.6% faster per step
Figure 1: Write-once, run everywhere. Axon DSL compiles axon definitions (.axon) to standalone model definitions for PyTorch, JAX, MLX and vLLM. All that is needed is a standard safetensors checkpoint.
Table 1: Per-backend comparison of Axon’s autoregressive generation performance with decoder-only models against PyTorch, Triton, and JAX. “Axon ≤1×” = count (and %) of checkpoints where Axon is faster than Transformers. “Axon >1×” = count (%) of checkpoints where Axon is slower. Median and mean report runtime ratio (lower means Axon is faster).
Backend
Checkpoints
Axon ≤1×
Axon >1×
Median ratio
Mean ratio
<4B
PyTorch
76
48 (63%)
28 (37%)
0.903
0.925
Triton
75
60 (80%)
15 (20%)
0.843
0.878
JAX
74
64 (86%)
10 (14%)
0.481
1.122
<4B total
225
172 (76%)
53 (24%)
0.804
0.974
4–32B
PyTorch
87
55 (63%)
32 (37%)
0.987
0.980
Triton
87
67 (77%)
20 (23%)
0.924
0.931
JAX
68
55 (81%)
13 (19%)
0.589
0.981
4–32B total
242
177 (73%)
65 (27%)
0.908
0.963
Figure 2: Autoregressive generation performance with decoder-only models with up to 4B parameters. Log-log plot where the diagonal line indicates equal performance. Points above the parity line indicate Axon is faster, color-fill denotes backends and border-line denotes dtype.
Table 2: Per-backend comparison of Axon’s Forward performance with Encoder-Only and Encoder-Decoder Models against PyTorch, Triton, and JAX backends. “Axon ≤1×” = count (and %) of checkpoints where Axon is faster than Transformers. “Axon >1×” = count (%) of checkpoints where Axon is slower. Median and mean report runtime ratio (lower means Axon is faster).
Backend
Checkpoints
Axon ≤1×
Axon >1×
Median ratio
Mean ratio
≤4B
PyTorch
26
15 (58%)
11 (42%)
0.883
1.786
Triton
26
9 (35%)
17 (65%)
1.332
2.169
JAX
26
7 (27%)
19 (73%)
1.079
1.260
≤4B total
78
31 (40%)
47 (60%)
1.084
1.738
4–32B
PyTorch
20
13 (65%)
7 (35%)
0.990
1.506
Triton
20
8 (40%)
12 (60%)
1.133
1.546
JAX
14
0 (0%)
14 (100%)
3.747
3.356
4–32B total
54
21 (39%)
33 (61%)
1.175
2.001
Figure 3: Benchmarking Autoregressive models between 4B and 32B models. Log-log plot where the diagonal parity line indicates equal performance. Generally, Axon produces comparable or superior throughput on the bigger models.
Table 3: Autoregressive generation comparison of Axon (vLLM native) against Transformers. “Axon ≤1×” = number (and %) of checkpoints where Axon is faster than Transformers. “Axon >1×” = number (and %) of checkpoints where Axon is slower than Transformers. Median and mean report Axon-to-Transformers runtime ratio (lower is faster).
Model size
Checkpoints
Axon ≤1×
Axon >1×
Median ratio
Mean ratio
Small (≤4B)
54
38 (70%)
16 (30%)
0.656
1.430
Large (4B–32B)
34
27 (79%)
7 (21%)
0.609
2.889
Total
88
65 (74%)
23 (26%)
0.631
1.994
Figure 4: vLLM: Axon (vLLM native) vs. Transformers generation throughput on 88 checkpoints. Points above the parity line indicate Axon is faster. 74% fall above.
Table 4: Breakdown of Axon performance by precision, model type, and sequence length. “Axon ≤1×” = number (and %) of checkpoints where Axon is faster than Transformers. “Axon >1×” = number (and %) of checkpoints where Axon is slower than Transformers. Median and mean report Axon-to-Transformers runtime ratio (lower = Axon is faster).
Group
Points
Axon ≤1×
Axon >1×
Median ratio
Mean ratio
BF16
63
61 (97%)
2 (3%)
0.423
0.500
FP32
63
58 (92%)
5 (8%)
0.503
0.588
causal_lm
84
82 (98%)
2 (2%)
0.390
0.437
seq2seq_lm
42
37 (88%)
5 (12%)
0.574
0.752
len=64
42
42 (100%)
0 (0%)
0.417
0.448
len=128
42
39 (93%)
3 (7%)
0.483
0.525
len=256
42
38 (90%)
4 (10%)
0.494
0.652
Total
126
119 (95%)
7 (5%)
0.483
0.544
Figure 5: MLX. Benchmarking conventional HF models against Axon derived standalone model definitions. Axon yields some considerable speed-ups across the board.
Table 5: Axon Compiler Phase and Stage with Invariants
Phase
Representation
Primary invariant
Parse
one-file AST
Syntactic structure and explicit MAIN pragma insertion
Load
loaded AST set
Imports and builtins located without rewriting semantics
Materialize
one-file AST
Optional checkpoint/config specialization for generic models
Resolve/validate-closed
closed AST
No unresolved imports or names; unreachable definitions pruned from MAIN
Normalize
normalized AST
Call syntax, pipes, path sugar, and zero-arg call/name distinctions made explicit
Elaborate/validate-elaborated
elaborated AST
Default arguments filled and call arguments positionalized
Flatten/validate-flat
flat AST
Explicit evaluation order; flat calls and binds accepted by typecheck and Graph IR lowering
Typecheck/validate-typed
Typed flat AST
expression types, arities, dimensions, and primitive rules applied to a fixpoint
Optimize-ast
typed flat AST
Optional conservative AST cleanup with retype/validation
Graph lowering/validation
Graph IR
Typed graph modules, multi-output nodes, structured paths, constraints, and metadata
Optimize-graph/validation
Graph IR
Optional graph cleanup, specialization, backend-neutral rewrites, and opt-in backend intrinsics
Backend
generated/runtime code
Executable tensor program consuming the validated Graph IR contract
Figure 12: Benchmarking Encoder-only and Encoder–Decoder up to 4B models. Log-log plot where the diagonal parity line indicates equal performance. Here, we see the majority of the points in the vicinity below the parity line.
Table 6: Wall-clock time and GPU idle for autoregressive generation (112 tokens for Qwen 0.5B, 96 for Pleias 3.8B). “—” indicates the metric is not directly measurable: JAX has no torch.profiler equivalent for per-kernel timing, and GPU active/idle is a synchronous-execution concept that does not apply to JAX’s async dispatch model.
Qwen2.5-0.5B (0.67B)
Pleias-3b-Preview (3.8B)
Metric
Transformers
Axon-Torch
Axon-JAX
Transformers
Axon-Torch
Axon-JAX
Wall-clock (ms)
1380
1185
541
1091
1071
1160
Speedup vs Transformers
1.0×
0.86×
0.39×
1.0×
0.98×
1.07×
GPU active (ms)
373
370
—
557
561
—
GPU idle (ms)
1007
815
—
534
510
—
GPU idle %
73.0%
68.8%
0.0%
48.9%
47.6%
0.0%
CUDA kernels
129,189
120,987
—
101,935
104,123
—
Figure 13: Benchmarking Encoder-only and Encoder-Decoder between 4B and 32B models. Log-log plot where the diagonal parity line indicates equal performance.
Table 7: Python function call counts and cProfile cumulative time. “—” indicates the category does not apply: JAX fuses all per-op dispatches into a single XLA compilation per step, so individual F.linear and SDPA calls do not exist as Python-level dispatches.
Qwen2.5-0.5B (0.67B)
Pleias-3b-Preview (3.8B)
Metric
Transformers
Axon-Torch
Axon-JAX
Transformers
Axon-Torch
Axon-JAX
Total Python calls
455,769
508,390
137,805
359,909
330,929
82,210
Calls vs Transformers
1.0×
1.12×
0.30×
1.0×
0.92×
0.23×
cProfile time (s)
1.13
0.95
0.54
0.87
0.78
1.16
Per-op dispatch (Python calls per generate pass):
nn.Module.__call__
35,616
0
0
28,032
0
0
F.linear
18,928
18,928
—
14,880
14,880
—
SDPA
2,688
2,688
—
2,112
2,112
—
rope_apply
0
5,376
—
0
4,224
—
forward (jit dispatch)
—
—
112
—
—
96
Figure 14: Training of Gemma 3 270M on summarization task for 2000 steps, demonstrating identical training behaviour between conventional Transformers-definition and Axon derived PyTorch model. Note the Axon and Transformers loss curves are on top of each other together.
Why it matters
Today's LLM tooling concentrates around a handful of platforms, creating a single point of failure for the whole ecosystem; Axon proposes a shared language specification instead of a shared framework to reduce that dependency. This matters for anyone who wants fast, portable models without being locked into one company's optimization stack.
Terms in this paper
DSL (domain-specific language) · a programming language built for one narrow purpose, here describing neural network architectures
strongly typed · the compiler strictly checks the kind and shape of values before running the program
Graph IR · a shared intermediate representation that all backend code generators consume
PagedAttention / KV-cache · vLLM's memory management technique that speeds up serving by efficiently handling cached tokens during generation
top-1 token parity · a correctness check confirming different implementations predict the same next token at every step
Original abstract (English)
The entire ecosystem of open-source language models effectively relies on a single platform. What if this platform was forced to shut down tomorrow? Implementing and maintaining efficient model definitions and translating them between different training and inference regimes is a resource-heavy task that severely limits model efficiency and portability, hindering both scaling and deployment. Here, we present Axon, a strongly typed domain-specific language with Haskell-like syntax, that enables a write-once, run everywhere paradigm for LLM architectures. By basing collaboration on a language specification rather than a specific framework's vision, Axon fosters open cooperation and empowers researchers to implement highly specialized architectures without giving up optimization infrastructure or accepting deployment lock-in. Axon allows for concise, auditable specifications that can be automatically compiled to standalone implementations for leading frameworks: PyTorch, PyTorch with Triton, JAX, MLX and vLLM. In 467 inference benchmarking experiments on models ranging from 135M to 32B parameters, we demonstrate median speedups of 7% on PyTorch, 12% on PyTorch with Triton, 91% on JAX, and 107% on MLX, compared to the reference implementations from Transformers. When deployed as native vLLM architectures with PagedAttention and KV-cache, Axon models achieve a 58% median speedup over Transformers implementations.
Authors · Jacob Nielsen, Danial Namazifard, Lukas Galke Poech, Peter Schneider-Kamp