One email each morning — yesterday's AI, sortedGet it in your inbox

METAL LAB

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

  1. 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
  2. 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
  3. A strong type system checks tensor shapes ahead of time, preventing shape-mismatch errors during execution
  4. 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
  5. 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.
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).
BackendCheckpointsAxon ≤1×Axon >1×Median ratioMean ratio
<4B
PyTorch7648 (63%)28 (37%)0.9030.925
Triton7560 (80%)15 (20%)0.8430.878
JAX7464 (86%)10 (14%)0.4811.122
<4B total225172 (76%)53 (24%)0.8040.974
4–32B
PyTorch8755 (63%)32 (37%)0.9870.980
Triton8767 (77%)20 (23%)0.9240.931
JAX6855 (81%)13 (19%)0.5890.981
4–32B total242177 (73%)65 (27%)0.9080.963
Figure 2: Autoregressive generation performance with decoder-only models with up to 4​B 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.
Figure 2: Autoregressive generation performance with decoder-only models with up to 4​B 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).
BackendCheckpointsAxon ≤1×Axon >1×Median ratioMean ratio
≤4B
PyTorch2615 (58%)11 (42%)0.8831.786
Triton269 (35%)17 (65%)1.3322.169
JAX267 (27%)19 (73%)1.0791.260
≤4B total7831 (40%)47 (60%)1.0841.738
4–32B
PyTorch2013 (65%)7 (35%)0.9901.506
Triton208 (40%)12 (60%)1.1331.546
JAX140 (0%)14 (100%)3.7473.356
4–32B total5421 (39%)33 (61%)1.1752.001
Figure 3: Benchmarking Autoregressive models between 4​B and 32​B models. Log-log plot where the diagonal parity line indicates equal performance. Generally, Axon produces comparable or superior throughput on the bigger models.
Figure 3: Benchmarking Autoregressive models between 4​B and 32​B 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 sizeCheckpointsAxon ≤1×Axon >1×Median ratioMean ratio
Small (≤4B)5438 (70%)16 (30%)0.6561.430
Large (4B–32B)3427 (79%)7 (21%)0.6092.889
Total8865 (74%)23 (26%)0.6311.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.
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).
GroupPointsAxon ≤1×Axon >1×Median ratioMean ratio
BF166361 (97%)2 (3%)0.4230.500
FP326358 (92%)5 (8%)0.5030.588
causal_lm8482 (98%)2 (2%)0.3900.437
seq2seq_lm4237 (88%)5 (12%)0.5740.752
len=644242 (100%)0 (0%)0.4170.448
len=1284239 (93%)3 (7%)0.4830.525
len=2564238 (90%)4 (10%)0.4940.652
Total126119 (95%)7 (5%)0.4830.544
Figure 5: MLX. Benchmarking conventional HF models against Axon derived standalone model definitions. Axon yields some considerable speed-ups across the board.
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
PhaseRepresentationPrimary invariant
Parseone-file ASTSyntactic structure and explicit MAIN pragma insertion
Loadloaded AST setImports and builtins located without rewriting semantics
Materializeone-file ASTOptional checkpoint/config specialization for generic models
Resolve/validate-closedclosed ASTNo unresolved imports or names; unreachable definitions pruned from MAIN
Normalizenormalized ASTCall syntax, pipes, path sugar, and zero-arg call/name distinctions made explicit
Elaborate/validate-elaboratedelaborated ASTDefault arguments filled and call arguments positionalized
Flatten/validate-flatflat ASTExplicit evaluation order; flat calls and binds accepted by typecheck and Graph IR lowering
Typecheck/validate-typedTyped flat ASTexpression types, arities, dimensions, and primitive rules applied to a fixpoint
Optimize-asttyped flat ASTOptional conservative AST cleanup with retype/validation
Graph lowering/validationGraph IRTyped graph modules, multi-output nodes, structured paths, constraints, and metadata
Optimize-graph/validationGraph IROptional graph cleanup, specialization, backend-neutral rewrites, and opt-in backend intrinsics
Backendgenerated/runtime codeExecutable tensor program consuming the validated Graph IR contract
Figure 12: Benchmarking Encoder-only and Encoder–Decoder up to 4​B 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.
Figure 12: Benchmarking Encoder-only and Encoder–Decoder up to 4​B 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)
MetricTransformersAxon-TorchAxon-JAXTransformersAxon-TorchAxon-JAX
Wall-clock (ms)13801185541109110711160
Speedup vs Transformers1.0×0.86×0.39×1.0×0.98×1.07×
GPU active (ms)373370557561
GPU idle (ms)1007815534510
GPU idle %73.0%68.8%0.0%48.9%47.6%0.0%
CUDA kernels129,189120,987101,935104,123
Figure 13: Benchmarking Encoder-only and Encoder-Decoder between 4​B and 32​B models. Log-log plot where the diagonal parity line indicates equal performance.
Figure 13: Benchmarking Encoder-only and Encoder-Decoder between 4​B and 32​B 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)
MetricTransformersAxon-TorchAxon-JAXTransformersAxon-TorchAxon-JAX
Total Python calls455,769508,390137,805359,909330,92982,210
Calls vs Transformers1.0×1.12×0.30×1.0×0.92×0.23×
cProfile time (s)1.130.950.540.870.781.16
Per-op dispatch (Python calls per generate pass):
nn.Module.__call__35,6160028,03200
F.linear18,92818,92814,88014,880
SDPA2,6882,6882,1122,112
rope_apply05,37604,224
forward (jit dispatch)11296
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.
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

Read on arXiv

Latest papers

All papers →

Latest from METAL LAB

Figures: Jacob Nielsen et al., arXiv:2608.19889, arxiv-nonexclusive