Write Once, Run Everywhere: The Axon DSL for Shape-Safe and Framework-Agnostic LLM Architectures
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.
METAL MEDIA explanatory visual
A new language, Axon, lets you write an LLM once and run it on PyTorch, JAX, MLX or vLLM
- 01Problem: 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
- 02Solution: 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
- 03A strong type system checks tensor shapes ahead of time, preventing shape-mismatch errors during execution
- 04Results: 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
- 05Training experiments showed an Axon-derived PyTorch model matched the loss curve of the standard Transformers implementation while running about 9.6% faster per step
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

| 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 |
| 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 |
| 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 |
| 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 |
| 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 |
| 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 | — |
| 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 |
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.
Read on arXivLatest papers
- SWE-bench Science: Can Coding Agents Resolve Engineering Tasks in Science?AI coding agents were tested on fixing real scientific software, and even the best one failed more than half the time
- FlashPrefill V2: Block-Sparse Prefill Attention for Long-Context LLM ServingMaking sparse attention fast enough and accurate enough for real LLM serving, not just papers
- PolicyGuide: From Guarding One Action to Guiding the Whole Workflow for Policy-Compliant LLM AgentsMaking customer-service AI agents follow the whole procedure, not just avoid one bad action
- EXIMO: VLM Guided Exploration of VLA PoliciesTeaching a robot new chores without human teleoperation, by letting a chatty AI supervise it
- EnvHarness: Awakening Static Worlds for Agent LearningInstead of building new training worlds from scratch, this work adds a plug-in layer that reshapes existing ones around each agent's actual weaknesses
- Bounded Sovereignty and the Control Tax: Pricing AI Oversight When the Deployer Does Not Own the ModelCompanies that rent AI instead of owning it can only do half of AI safety oversight
- PersonalBench: Measuring the Authorship Gap in LLM PersonalizationAI can be prompted to write 'like someone,' but its own voice never fully disappears
- Automated Summarization of Financial News Using Large Language Models and Retrieval-Augmented Generation: An Early Empirical Study (Fall 2023)Testing AI summaries of stock news, the simple approach beat the trendy retrieval-based one
Latest from METAL MEDIA
Figures: Jacob Nielsen et al., arXiv:2608.19889, arxiv-nonexclusive