
| 基本信息 | 详情 |
|---|---|
| 开发公司 | 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是标配
做新架构实验,需要灵活的自动微分和编译优化
物理模拟、微分方程求解,函数式范式匹配数学结构
可微分物理引擎(Brax)、策略梯度方法,JAX的函数变换很合适
用TPU、Colab、Google Cloud,JAX原生集成最好
如果你做的是传统深度学习工程(CV分类、NLP微调),PyTorch的生态和易用性仍然更好。但如果你要训练Gemini级别的大模型,或者在TPU上跑实验,JAX是更好的选择。DeepMind在2020年就全面转向JAX,说它"契合我们的工程哲学"。
同类竞品对比
| 维度 | JAX | PyTorch | TensorFlow | MLX |
|---|---|---|---|---|
| 范式 | 函数式 | 面向对象 | 声明式 | 函数式 |
| 编译 | XLA(JIT) | TorchDynamo/Inductor | Graph mode | Metal JIT |
| 自动微分 | 源到源转换 | 动态图追踪 | 静态图 | 源到源转换 |
| TPU支持 | 原生最佳 | 通过XLA | 原生 | 不支持 |
| 大模型训练 | Gemini/Claude/Grok | Llama/Mistral | BERT时代 | 本地推理 |
| 生态 | Flax/Optax/Orbax | HuggingFace/TorchVision | Keras | mlx-lm |
| 上手难度 | 高 | 低 | 中 | 低 |
FAQ
JAX和PyTorch学哪个? 看用途。做研究和工程落地,PyTorch生态更成熟、教程更多、HuggingFace集成更好。做大规模模型训练或TPU开发,JAX的编译和并行能力更强。两个都会最好。
JAX能在普通GPU上跑吗? 可以。NVIDIA GPU通过XLA编译器支持,CUDA后端性能不错。不需要TPU。但JAX在TPU上的优化是最深的,Google自家硬件配合最好。
Flax是什么? Flax是JAX上的神经网络库,类似PyTorch的nn.Module。提供Layer、Module抽象、训练循环工具。Google和社区共同维护,是JAX生态里最主流的NN框架。Keras 3也支持JAX后端。
JAX的函数式编程难学吗? 如果你习惯了PyTorch的面向对象风格,切换到JAX的函数式范式确实需要适应。纯函数、不可变状态、函数变换这些概念,数学背景好的人反而更容易上手。
最新版本是什么? 0.11.0,2026年7月发布。0.10.0在2026年4月发布。JAX采用"基于努力程度的版本控制",不是传统语义化版本。0.11.0放弃了Python 3.11和NumPy 2.0支持。
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。函数式范式是门槛也是优势,看你站在哪一边。
参考来源
- ▸JAX官方文档 | JAX: High Performance Array Computing,官方文档和API参考
- ▸GitHub | jax-ml/jax,开源代码仓库和变更日志
- ▸AIWiki | JAX,JAX百科条目,含历史和采用情况
- ▸密歇根大学课件 | JAX: A Domain-Specific Tracing JIT Compiler,JAX技术原理学术讲解
- ▸JAX中文文档 | 更新日志,JAX中文社区维护的版本更新记录
