每天早上一封邮件,把昨天的 AI 梳理好订阅邮件

METAL LAB

Write Once, Run Everywhere: The Axon DSL for Shape-Safe and Framework-Agnostic LLM Architectures

arXiv:2608.198892026-08-21

新语言Axon让语言模型代码一次编写,可在PyTorch、JAX、MLX、vLLM上运行

开源语言模型生态系统实际上依赖单一平台Hugging Face,把模型移植到其他框架往往需要人工重写,还容易丢失原有的优化。研究团队开发了强类型、语法类似Haskell的专用语言Axon,只需描述一次模型结构,编译器就能自动生成PyTorch、Triton、JAX、MLX、vLLM五种后端的独立代码。在参数量从1.35亿到320亿的模型上进行的467次推理基准测试中,Axon生成的模型普遍比参考的Transformers实现更快。

他们做了什么

  1. 问题:语言模型代码被锁定在特定平台的实现方式里,移植到其他框架需要人工重写,过程中容易出错并丢失针对性优化
  2. 方案:用Axon描述模型的层结构、张量形状、参数位置等,编译器将其转为共享的中间表示(Graph IR),再自动生成五种后端的独立可运行代码
  3. 强类型系统在编译阶段就检查张量形状,避免运行时出现形状不匹配的错误
  4. 结果:相较Transformers参考实现,PyTorch后端中位数快7%,结合Triton的PyTorch快12%,JAX快91%,MLX快107%;以vLLM原生方式部署并使用PagedAttention和KV缓存时,中位数比Transformers快58%
  5. 训练实验显示,Axon生成的PyTorch模型与原版Transformers实现的损失曲线几乎完全重合,且每步训练速度快约9.6%
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.

为什么重要

当前语言模型工具链高度集中在少数平台上,一旦该平台出问题整个生态都会受影响;Axon提出用共享的语言规范取代共享框架来减少这种依赖。这对希望摆脱单一厂商优化基础设施、又想获得快速可移植模型的研究者和小团队具有实际意义。

本文术语

  • DSL(领域专用语言) · 为特定用途设计的编程语言,这里用于描述神经网络结构
  • 强类型 · 编译器在运行前严格检查数值的种类和形状
  • Graph IR(图中间表示) · 各后端代码生成器共用的编译中间结构
  • PagedAttention/KV缓存 · vLLM用来高效管理生成过程中缓存令牌的内存管理技术,可提升服务速度
  • top-1 token parity(首选词一致性) · 检验不同实现在每一步是否预测出相同下一个词的正确性标准

论文原文摘要(英文)

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.

作者 · Jacob Nielsen, Danial Namazifard, Lukas Galke Poech, Peter Schneider-Kamp

在 arXiv 阅读

最新论文

全部论文 →

METAL LAB 最新报道

图片来源: Jacob Nielsen et al., arXiv:2608.19889, arxiv-nonexclusive