JAX使用教程:3个变换(jit/grad/vmap)跑通你的NumPy代码
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
读完你能判断:手头的数值/模型项目是否值得切到 JAX,以及用哪三行变换代码起步。
概念对齐:Jaxpr 和"变换"到底在说什么
Jaxpr:把函数变成"菜谱"的中间表示
Jaxpr 是 JAX 把你的 Python 函数追踪(trace)后生成的一份"操作清单",只记录算子顺序,不含真实数据。类比:你写的是菜谱,Jaxpr 是厨房根据菜谱整理的执行工单,后面的编译、求导都在工单上操作。
变换(Transformation):为什么一行代码就能改行为
JAX 的jax.jit、jax.grad、jax.vmap不是"调用某个优化函数",而是对函数本身做变换:拿 Jaxpr 重新生成一份函数再还给你。所以它们可以任意叠在一起用,比如先jax.grad再jax.jit,互不干扰。
功能模块实战拆解
用 jax.jit 编译 NumPy 代码的 3 步
解决的是:NumPy 代码在 GPU/TPU 上慢,手工搬设备太繁琐。
- 把
numpy换成jax.numpy; - 在函数上挂
@jax.jit; - 第一次调用触发编译,之后直接复用编译产物。
import jax import jax.numpy as jnp @jax.jit # 挂上即编译,XLA 负责优化到 GPU/TPU def selu(x): return 1.05 * jnp.where(x > 0, x, 1.67 * jnp.exp(x) - 1.67)坑点提醒:第一次调用会明显变慢(在编译),测性能时要用block_until_ready(),且尽量把jax.jit放在最外层调用上。参考 docs/jit-compilation.md。
不写 for 循环做批处理:jax.vmap
解决的是:单样本逻辑想跑成 batch,不想手写循环和内存布局。
# 只按"单个样本"写函数 def forward_one(x, W): return jnp.dot(x, W) forward = jax.vmap(forward_one) # 自动向量化成批量版 # forward(x_batch, W) 等价于对 x_batch 每行做 dot坑点提醒:vmap映射的轴默认是每批量的第一维;如果函数里有显式循环读输入长度,长度不一致会直接报错。参考 docs/automatic-vectorization.md。
3 行代码换掉梯度计算:jax.grad
解决的是:不想搭计算图、不想手动推导反向公式。
def loss(params, x, y): return jnp.mean((jnp.dot(x, params) - y) ** 2) grad_loss = jax.grad(loss) # 对第 1 个参数求导 g = grad_loss(params, x, y) # g 是可直接进优化器的数组坑点提醒:JAX 要求函数是"纯"的——改全局变量、打印副作用会让追踪结果不可靠;另外默认是 32 位浮点,做高精度数值计算要先开 x64 配置。参考 docs/key-concepts.md。
数据佐证:性能问题的官方口径
仓库自带 benchmarks/ 目录(含 linalg、random、api 等基准套件),用 google_benchmark 框架在 CI 上跑。JAX 官方没有给出"对比某框架快 X%"的统一数字,但有 4 个高频性能疑问的官方解释,都记录在 docs/benchmarking.md:
| 高频疑问 | 官方解释(见 docs/benchmarking.md) |
|---|---|
| 第一次调用为什么慢 | 触发 JIT 编译,后续调用走编译产物 |
| 测出来"快"是不是假的 | 异步调度,需block_until_ready()再计时 |
| 和 NumPy 比不公平 | JAX 默认 32 位 dtype,对比前先对齐精度 |
| 小代码没提速 | 数据传输到加速器本身耗时,先device_put |
选型决策:JAX 还是 TensorFlow
场景 A 选 JAX:算法研究/快速原型;需要高阶导数或变换自由组合;目标硬件是 TPU 或多卡 GPU。
场景 B 选 TensorFlow:已有生产级模型服务(Serving)和移动端(TFLite)部署链路;团队依赖 Keras 生态与现成工具链。
TF 代码往 JAX 迁移,抓住 2 步即可:
tf.Tensor运算整体换jax.numpy(API 与 NumPy 对齐);tf.GradientTape换成jax.grad,with块直接删掉。
import jax.numpy as jnp def loss(params, x, y): # 原 tf 函数体基本原样搬 return jnp.mean((jnp.dot(x, params) - y) ** 2) g = jax.grad(loss)(params, x, y) # 替代 GradientTape 取梯度延伸资源
- 核心概念(变换/追踪/函数性):docs/key-concepts.md
- GPU 性能调优清单:docs/gpu_performance_tips.md
- 可跑的示例与交互 notebook:examples/、cloud_tpu_colabs/
你手上现在最想要哪个变换:更快的 jit、更省的 vmap,还是直接 grad 出梯度?说说你的场景,我们可以对着 docs/ 里的章节再拆一层。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考