意外构建:JAX 的 LLVM 编译器
We accidentally built an LLVM compiler for Jax

我们在为 PennyLane 开发量子编译器 Catalyst 时,意外发现了一个有趣的副产品。原本为了优化混合量子 - 经典工作流,我们利用 MLIR 将 JAX 代码直接编译到 LLVM,却意外绕过了 XLA 运行时。这意味着,即使不使用任何量子指令,Catalyst 也能将纯 JAX 代码编译成独立的 AOT 二进制文件。这不仅支持动态形状数组和原生 Python 控制流,还让我们能利用 Enzyme 进行反向传播。虽然它无法在标准深度学习任务上超越 XLA,但为边缘设备部署和自定义硬件提供了全新的可能性。
有时候当你构建软件时,你会为了修一个小问题而深入挖掘,结果意外得到了一件令人惊喜的毛衣(而那头牦牛大概会疑惑为什么自己突然变得冷飕飕的)。
- ndesaulniers
挺有意思。XLA 比 MLIR 更早出现,背后有不少故事。想听的话,得去湾区的 LLVM 月度聚会才行。:-X
> 那么,既然 XLA 已经在使用 LLVM,我们的方法有什么不同?
我们用了 MLIR,而 XLA 没有。
> 所以……这有什么意义?
> 老实说?我们自己也还不完全确定。
> 让我说清楚一点:对于标准的深度学习工作负载,这并不会比 XLA 更强。XLA 在 GPU 和 TPU 上的线性代数运算方面,已经积累了多年的高度特定优化。如果你要训练一个巨大的 Transformer,还是用标准的 JAX 吧。
> 但我们觉得挺酷的是,当你把 JAX 直接连接到更广泛的 LLVM 生态系统,并去掉沉重的 XLA 运行时时会发生什么。(另外,也不需要再用 Bazel 构建 XLA 了!谢天谢地。)