JAX使用教程3个变换jit/grad/vmap跑通你的NumPy代码【免费下载链接】jaxComposable transformations of PythonNumPy 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.numpyAPI 与 NumPy 对齐tf.GradientTape换成jax.gradwith块直接删掉。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.mdGPU 性能调优清单docs/gpu_performance_tips.md可跑的示例与交互 notebookexamples/、cloud_tpu_colabs/你手上现在最想要哪个变换更快的 jit、更省的 vmap还是直接 grad 出梯度说说你的场景我们可以对着 docs/ 里的章节再拆一层。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考