ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

JAX 可组合变换:grad、jit、vmap 全好使,坑也全在 tracing

JAX 可组合变换:grad、jit、vmap 全好使,坑也全在 tracing JAX 可组合变换grad、jit、vmap 全好使坑也全在 tracing【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax昨晚把训练从单卡迁到多卡jax.jit一贴上去第一个 epoch 反而更慢了把 Python 的 if/while 写进 jitted 函数直接报错更邪门的是 print 出来的不再是数组而是什么 Tracer。JAX 是对纯 Python 和 NumPy 函数做可组合变换的库jax.grad、jax.jit、jax.vmap任意叠加这套便利全部建立在 tracing追踪机制之上而绝大多数坑也出自这里。把 tracing 这条主线吃透你就能解释 JAX 里大多数反直觉行为。核心机制JAX 变换靠 tracing 把函数变成 jaxpr把函数传给jax.jit或jax.grad时JAX 并不是真的执行它而是把每个参数换成一个 tracer——一个只携带 shape 和 dtype、不带具体数值的影子值。函数照着影子值走一遍所有运算被记录在案JAX 再据此重建出函数的完整操作序列即 jaxprJAX exPRession。打个比方这不是让演员上场演出而是给他一份台词本让他念一遍——你要的是台词顺序不是演出本身。import jax import jax.numpy as jnp def log2(x): return jnp.log(x) / jnp.log(2.0) # make_jaxpr 只追踪不执行打印出函数的中间表示 print(jax.make_jaxpr(log2)(3.0)) # 追踪得到的 jaxpr 交给 XLA 编译再叠加 grad 就是编译版梯度 selu_jit jax.jit(jax.grad(selu))注意一个容易被忽略的细节jaxpr 只记录 JAX 操作。往全局 list 里 append、用 Pythonprint打日志这类副作用不会出现在 jaxpr 里——JAX 的变换只理解无副作用的纯函数。所以当你看到输出里冒出Tracer对象时背后实际发生的是你在用运行时数值做了追踪时看不到的事情取值判断、写日志而不是数组本身出了问题。实战视角三个递进用法预热 JIT 编译消除首步卡顿jit的首次调用承担全部编译开销之后靠缓存直接复用编译产物。jax.jit def step(x): return x * x x * 2.0 # 正式计时/训练前单独预热一次把编译时间隔离出去 step(x).block_until_ready()训练第一步的耗时和其余步骤是两个数量级预热之后步与步之间的时间就稳定了。配置 static 参数处理输入分支函数里必须对参数取值做 Python 分支时可以把它标成静态# n 标为静态每个不同的 n 值各编译一份 g_jit jax.jit(g, static_argnames[n]) g_jit(10, 20)代价是 static 值一变就重新编译一份所以只适合取值有限、变化不频繁的参数比如层数、模式开关。用 lax.cond 改写运行时值判断运行时数值要决定分支时把判断搬进编译区jax.jit def f(x): # 用 lax.cond 取代 Python if分支在编译区内展开 return jax.lax.cond(x 0, lambda x: x, lambda x: 2 * x, x)整体没法搬进去的循环比如迭代次数本身是运行时值就只把循环体jit、循环留在 Python 侧前提是把被编译的函数定义在循环外否则每次迭代都是一个新函数、缓存永远不命中。数据说话JAX 分布式扩展的三种模式模式数据视图显式分片显式集合通信Auto全局无无Explicit全局有无Manual逐设备有有这张表来自 README 对编译器自动并行 / 显式分片 / 手动逐设备编程三种扩展模式的划分。差距意味着想让编译器替你切数据就交出控制反过来想要通信与计算重叠这类精细调度Manual 模式是必须的——shard_map文档里对 all_gather 矩阵乘法有没有 overlap 的 profile 对比就是给这个决策用的。团队规模小、通信需求简单时我倾向于从 Auto 起步等 profile 说话了再往 Manual 走。选型判断适合上 JAX 的场景算法原型和科研迭代快 →grad、vmap、jit任意组合改完一个函数立刻有编译版梯度手上已有 NumPy/科学计算代码 →jax.numpy沿用 NumPy 接口迁移成本接近换 import要往多卡、TPU 扩 → 同一套纯函数代码只改分片声明不建议硬上的场景生产级模型服务和完整数据管道 → JAX 只给计算与变换层Serving、数据管道要自己在生态里补齐代码里可变状态、类属性满天飞 → 与纯函数约束正面冲突要么先重构要么留在 TensorFlow/PyTorch只有单卡 CPU、模型很小 → JIT 编译开销可能比逐 op 执行还贵直接用 NumPy 反而省心踩坑与排障现象训练第一步特别慢之后恢复正常。根因首次调用要做追踪、XLA 编译并写缓存。解法正式流程前对jit函数做一次预热调用别把编译时间算进步数里。现象jitted 函数里if x 0或运行时值控制的 while 直接抛错。根因被追踪的数值不能控制编译期流程只有 shape、dtype 这类静态属性可以参与分支。解法改写为jax.lax.cond或只编译循环体、循环留在 Python 侧。现象把自定义对象标进static_argnums后改了它的属性结果却不刷新。根因编译缓存拿对象的 id 当 keyJAX 无从知道它变过。解法给对象重写__hash__和__eq__或干脆别标 static动态部分走 PyTree 的叶子。生态与收尾仓库自带的 examples/ 覆盖了从 单卡 MNIST 分类器 到 VAE 再到 SPMD 多设备实现 的完整路径cloud_tpu_colabs/ 提供 ODE 求解、pmap cookbook、并行思维等交互式教程benchmarks/ 则放着官方基准脚本想看jit前后的差距可以直接跑。JAX 不是功能最全的深度学习框架它是把函数变换组合这件事做到最彻底的那一个——上层叠模型库和训练器下层交给 XLA中间这层薄得刚好。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表