AI万址
JAX

JAX

JAX是Google开源的高性能数值计算库,结合NumPy API、自动微分和XLA编译,用于训练Gemini、Claude等大模型,支持CPU/GPU/TPU。

JAX
访问官网
基本信息详情
开发公司Google(原Google Brain团队,现Google DeepMind)
上线时间2018年12月开源
官网https://jax.readthedocs.io/
支持平台CPU、GPU、TPU,跨平台
价格免费开源(Apache 2.0许可)
核心定位高性能可组合函数变换的数值计算框架

JAX是Google开发的高性能数值计算库。一句话定义它:NumPy的API + 自动微分 + XLA编译器加速。你写的Python代码可以自动编译到CPU、GPU和TPU上运行,还能任意组合微分、向量化、并行化这些变换。

2018年开源的时候,JAX还是Google Brain的实验项目。到了2026年,它已经是训练世界最大AI模型的首选框架。Google的Gemini、Gemma、Imagen、Veo,Anthropic的Claude,xAI的Grok,这些你叫得出名字的大模型,预训练阶段都在用JAX。GitHub上35,800颗Star,PyPI月下载量约1780万次。

JAX的核心设计理念是"可组合的函数变换"。jax.grad求导,jax.jit编译加速,jax.vmap自动向量化,jax.pmap跨设备并行。这些变换可以嵌套使用--你可以对梯度再求梯度算Hessian,也可以先vmap再jit编译。这种数学上的优雅是PyTorch做不到的,因为PyTorch的动态图和JAX的函数式范式根本是两条路线。

核心功能

自动微分(autodiff) 支持任意阶前向和反向模式微分。jax.grad返回一个新函数计算梯度。可以嵌套:jax.grad(jax.grad(f))直接算二阶导。autodiff基于Autograd项目,JAX的核心开发者里好几位就是Autograd的作者。

JIT编译 jax.jit把Python函数编译成XLA HLO,再生成CPU/GPU/TPU的目标代码。XLA做算子融合和内存优化,通常能加速好几倍。JIT和autodiff可以组合,编译后的梯度函数也能再被微分。

自动向量化 jax.vmap把对单个样本的函数自动扩展成批量处理。不用手写batch维度循环,JAX自动处理。和jax.jit组合后性能接近手写批处理。

并行化 jax.pmap做跨设备数据并行,jax.shard_map(shmap)做更细粒度的分片控制。配合GSPMD或新的Shardy分片系统,可以在TPU pod上跑超大规模模型训练。

Pallas 自定义内核 JAX的低级内核语言。可以在TPU和GPU上写自定义kernel,比XLA生成的代码更精确控制性能。支持TPU pipelining和Blackwell GPU矩阵乘法优化。

NumPy兼容API jax.numpy跟NumPy API几乎一致。很多NumPy代码改个import就能跑。也有jax.scipy提供SciPy兼容接口。

适合谁

大模型训练团队

需要在TPU/GPU集群上训练超大模型,JAX+Flax+Optax是标配

AI研究员

做新架构实验,需要灵活的自动微分和编译优化

科学计算

物理模拟、微分方程求解,函数式范式匹配数学结构

强化学习

可微分物理引擎(Brax)、策略梯度方法,JAX的函数变换很合适

Google生态用户

用TPU、Colab、Google Cloud,JAX原生集成最好

如果你做的是传统深度学习工程(CV分类、NLP微调),PyTorch的生态和易用性仍然更好。但如果你要训练Gemini级别的大模型,或者在TPU上跑实验,JAX是更好的选择。DeepMind在2020年就全面转向JAX,说它"契合我们的工程哲学"。

同类竞品对比

维度JAXPyTorchTensorFlowMLX
范式函数式面向对象声明式函数式
编译XLA(JIT)TorchDynamo/InductorGraph modeMetal JIT
自动微分源到源转换动态图追踪静态图源到源转换
TPU支持原生最佳通过XLA原生不支持
大模型训练Gemini/Claude/GrokLlama/MistralBERT时代本地推理
生态Flax/Optax/OrbaxHuggingFace/TorchVisionKerasmlx-lm
上手难度

FAQ

A

JAX和PyTorch学哪个? 看用途。做研究和工程落地,PyTorch生态更成熟、教程更多、HuggingFace集成更好。做大规模模型训练或TPU开发,JAX的编译和并行能力更强。两个都会最好。

A

JAX能在普通GPU上跑吗? 可以。NVIDIA GPU通过XLA编译器支持,CUDA后端性能不错。不需要TPU。但JAX在TPU上的优化是最深的,Google自家硬件配合最好。

A

Flax是什么? Flax是JAX上的神经网络库,类似PyTorch的nn.Module。提供Layer、Module抽象、训练循环工具。Google和社区共同维护,是JAX生态里最主流的NN框架。Keras 3也支持JAX后端。

A

JAX的函数式编程难学吗? 如果你习惯了PyTorch的面向对象风格,切换到JAX的函数式范式确实需要适应。纯函数、不可变状态、函数变换这些概念,数学背景好的人反而更容易上手。

A

最新版本是什么? 0.11.0,2026年7月发布。0.10.0在2026年4月发布。JAX采用"基于努力程度的版本控制",不是传统语义化版本。0.11.0放弃了Python 3.11和NumPy 2.0支持。

A

Pallas是什么? JAX的低级kernel编程接口。可以在TPU和GPU上写自定义计算kernel,直接控制内存访问和流水线。适合需要极致性能优化的场景,比如Blackwell GPU上的矩阵乘法。

更新动态

2026年的大变化是分片系统从GSPMD迁移到Shardy,autodiff切换到直接线性化。Pallas API持续扩展,支持Blackwell GPU和TPU SparseCore。社区生态在快速增长,MaxText、AXLearn等大模型训练框架都基于JAX构建。

客户评论

综合评分:4.5/5(基于GitHub社区和行业采用情况)

"用JAX训练了70B参数模型,shard_map让分片策略写起来比PyTorch FSDP清晰得多。但调试体验不如PyTorch直观。" - 大模型训练工程师

"函数式范式一开始很痛苦,习惯之后发现数学公式到代码的映射太自然了。grad嵌套grad直接出Hessian。" - ML研究员

"TPU上跑JAX是唯一选择。但NVIDIA GPU上和PyTorch的差距在缩小,TorchDynamo编译器进步很快。" - 基础设施工程师

JAX在顶级AI实验室的采用率很高,但在工程化团队中的普及度仍不如PyTorch。函数式范式是门槛也是优势,看你站在哪一边。

参考来源