尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

JAX 易踩坑指南(Common Gotchas):纯函数、原地更新、jit 与动态形状的完整避坑手册

JAX 易踩坑指南(Common Gotchas):纯函数、原地更新、jit 与动态形状的完整避坑手册 JAX 易踩坑指南Common Gotchas纯函数、原地更新、jit 与动态形状的完整避坑手册【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jaxJAX 是一套对 Python NumPy 数值程序进行可组合变换微分、向量化、JIT 编译到 GPU/TPU的框架但它的变换与编译机制只对满足特定约束的程序生效。本文以 docs/notebooks/Common_Gotchas_in_JAX.md 为骨架结合当前仓库源码如 array_methods.py、slicing.py、config.py逐一拆解 JAX 中最常踩的坑纯函数约束、原地更新、类方法 jit、越界索引、非数组输入、随机数、控制流、动态形状、NaN/Inf 调试与 64 位精度并给出可复制、可运行的解决方案。读完本文你将能识别并绕开这些陷阱写出既正确又高性能的 JAX 代码。 纯函数Pure functionsJAX 的变换jit、grad、vmap等与编译只设计用于函数式纯函数所有输入数据通过函数参数传入所有结果通过函数返回值输出。纯函数对相同的输入总是返回相同的结果。副作用Side-effects第一次运行与缓存命中的差异考虑带print副作用的函数import numpy as np from jax import jit from jax import lax from jax import random import jax import jax.numpy as jnp def impure_print_side_effect(x): print(Executing function) # 这是副作用 return x # 副作用在第一次运行时出现 print (First call: , jit(impure_print_side_effect)(4.)) # 相同类型/形状参数的后续调用可能不再显示副作用命中编译缓存 print (Second call: , jit(impure_print_side_effect)(5.)) # 当参数类型或形状改变时JAX 会重新执行 Python 函数 print (Third call, different type: , jit(impure_print_side_effect)(jnp.array([5.])))由于 JAX 在参数类型与形状不变时直接复用缓存的编译结果print这类副作用只在首次或参数签名变化时出现。需要注意这些行为并不被 JAX 系统保证正确用法是只对纯函数使用 JAX 变换。全局变量Globals值被捕获还是实时读取g 0. def impure_uses_globals(x): return x g # JAX 在第一次运行时捕获全局变量的值 print (First call: , jit(impure_uses_globals)(4.)) g 10. # 更新全局变量 # 后续相同签名的调用可能静默使用缓存值 print (Second call: , jit(impure_uses_globals)(5.)) # 当参数类型/形状变化导致重新执行 Python 时才会读到最新的全局值 print (Third call, different type: , jit(impure_uses_globals)(jnp.array([4.])))修改全局变量小心内部 Traced 值泄漏g 0. def impure_saves_global(x): global g g x return x # JAX 用特殊的 Traced 值执行一次变换后的函数 print (First call: , jit(impure_saves_global)(4.)) print (Saved global: , g) # 全局变量 g 被写入了 JAX 内部 Traced 值从源码结构看这正是 JAX 追踪tracing模型的必然结果变换期间参数被替换为Traced抽象值写入外部状态的代码会把内部表示泄漏到 Python 对象中。正确的做法是把所有状态当作函数参数显式传入。内部状态是允许的只要不读写外部状态一个 Python 函数只要不读写外部状态即使内部使用了有状态对象也仍然是纯函数def pure_uses_internal_state(x): state dict(even0, odd0) for i in range(10): state[even if i % 2 0 else odd] x return state[even] state[odd] print(jit(pure_uses_internal_state)(5.))不要在 jit 或控制流中使用迭代器迭代器是带状态的 Python 对象靠内部状态取下一个元素与 JAX 的函数式模型不兼容。在jit或任何控制流原语中使用迭代器大部分会直接报错有些会静默产生意外结果import jax.numpy as jnp from jax import make_jaxpr # lax.fori_loop直接用数组索引没问题 array jnp.arange(10) print(lax.fori_loop(0, 10, lambda i,x: xarray[i], 0)) # 期望 45 iterator iter(range(10)) print(lax.fori_loop(0, 10, lambda i,x: xnext(iterator), 0)) # 意外结果 0 # lax.scan迭代器作为 elems 会抛错 def func11(arr, extra): ones jnp.ones(arr.shape) def body(carry, aelems): ae1, ae2 aelems return (carry ae1 * ae2 extra, carry) return lax.scan(body, 0., (arr, ones)) make_jaxpr(func11)(jnp.arange(16), 5.) # make_jaxpr(func11)(iter(range(16)), 5.) # 抛错 # lax.cond迭代器作为 operand 会抛错 array_operand jnp.array([0.]) lax.cond(True, lambda x: x1, lambda x: x-1, array_operand) iter_operand iter(range(10)) # lax.cond(True, lambda x: next(x)1, lambda x: next(x)-1, iter_operand) # 抛错lax.fori_loop、lax.scan、lax.cond的签名与语义可参见仓库源码 lax/control_flow/loops.py 与 lax/control_flow/conditionals.py 相关定义。迭代器场景下应改用数组或jnp.arange之类的无状态数据结构。 原地更新In-place updatesNumPy 中常见的原地索引更新numpy_array np.zeros((3,3), dtypenp.float32) print(original array:) print(numpy_array) # 原地、可变更新 numpy_array[1, :] 1.0 print(updated array:) print(numpy_array)而jax.Array禁止下标赋值%xmode Minimal jax_array jnp.zeros((3,3), dtypejnp.float32) # 对 JAX 数组做原地更新会直接报错 jax_array[1, :] 1.0__iadd__的差异重绑定而非原地修改jax_array jnp.array([10, 20]) jax_array_new jax_array jax_array_new 10 print(jax_array_new) # jax_array_new 被重绑定到新值 [20, 30]但... print(jax_array) # 原始数组保持 [10, 20] 不变 numpy_array np.array([10, 20]) numpy_array_new numpy_array numpy_array_new 10 print(numpy_array_new) # numpy_array_new is numpy_array被原地更新 print(numpy_array) # 两者都是 [20, 30]原因在于NumPy 的__iadd__执行原地修改而jax.Array不定义__iadd__Python 把jax_array_new 10当作jax_array_new jax_array_new 10的语法糖只发生变量重绑定不修改任何数组。允许原地修改变量会让程序分析与变换变得困难JAX 要求程序是纯函数因此改用函数式数组更新。函数式数组更新x.at[idx].set(y)JAX 通过数组的.at属性提供函数式纯函数的索引更新。上面的更新可改写为jax_array jnp.zeros((3,3), dtypejnp.float32) updated_array jax_array.at[1, :].set(1.0) print(updated array:\n, updated_array)与 NumPy 版本不同JAX 的数组更新函数**就地外out-of-place**操作返回新数组原始数组不被修改。print(original array unchanged:\n, jax_array)不过在jit 编译代码内部如果x.at[idx].set(y)的输入x之后不再被复用编译器会自动把该数组更新优化为原地执行——这正是函数式写法兼顾正确性与性能的关键。更多索引更新操作索引更新不只覆盖数值还可以做索引加法等操作print(original array:) jax_array jnp.ones((5, 6)) print(jax_array) new_jax_array jax_array.at[::2, 3:].add(7.) print(new array post-addition:) print(new_jax_array)从源码看.at由 array_methods.py 中的_IndexUpdateHelper实现其 docstring 给出了完整的等价对照表x.at写法等价的原地表达式x x.at[idx].set(y)x[idx] yx x.at[idx].add(y)x[idx] yx x.at[idx].subtract(y)x[idx] - yx x.at[idx].multiply(y)x[idx] * yx x.at[idx].divide(y)x[idx] / yx x.at[idx].power(y)x[idx] ** yx x.at[idx].min(y)x[idx] minimum(x[idx], y)x x.at[idx].max(y)x[idx] maximum(x[idx], y)x x.at[idx].apply(ufunc)ufunc.at(x, idx)x x.at[idx].get()x x[idx]这些方法分别映射到底层 scatter 原语lax_slicing.scatter、scatter_add、scatter_sub、scatter_mul等见 array_methods.py。源码 docstring 还提示与 NumPy 原地操作不同若多个索引指向同一位置所有更新都会被应用NumPy 只保留最后一次且冲突更新的应用顺序是实现定义的、可能在部分硬件上不确定。越界切片尺寸限制在jit代码以及lax.while_loop/lax.fori_loop内部切片的大小不能是参数值的函数只能依赖参数形状切片起始索引不受此限制详见下文控制流一节。 使用jax.jit装饰类方法大多数jax.jit示例针对独立函数装饰类方法会引入额外问题。看一个朴素写法import jax.numpy as jnp from jax import jit class CustomClass: def __init__(self, x: jnp.ndarray, mul: bool): self.x x self.mul mul jit # ---- 如何正确做到这一点 def calc(self, y): if self.mul: return self.x * y return y调用c CustomClass(2, True); c.calc(3)会报错因为函数第一个参数是self类型CustomClassJAX 不知道如何处理这种类型。文档给出三种基本策略。策略 1JIT 编译的辅助函数helper function最直接的方式是在类外定义一个可正常 JIT 的辅助函数class CustomClass: def __init__(self, x: jnp.ndarray, mul: bool): self.x x self.mul mul def calc(self, y): return _calc(self.mul, self.x, y) jit(static_argnums0) def _calc(mul, x, y): if mul: return x * y return yc CustomClass(2, True) print(c.calc(3))优点简单、显式且无需教 JAX 如何处理CustomClass类型代价是方法逻辑被拆到了类外。策略 2把self标记为静态static用static_argnums标记self为静态参数但要小心意外结果class CustomClass: def __init__(self, x: jnp.ndarray, mul: bool): self.x x self.mul mul # 警告下面的例子是坏的别直接复制粘贴 jit(static_argnums0) def calc(self, y): if self.mul: return self.x * y return y首次调用c CustomClass(2, True); print(c.calc(3))不再报错但陷阱在于首次调用后修改对象属性后续调用可能返回错误结果c.mul False print(c.calc(3)) # 本应打印 3原因对象被标记为静态后会作为字典键进入 JIT 的内部编译缓存因此 JAX 假定其哈希hash(obj)、相等性obj1 obj2与对象同一性obj1 is obj2行为一致。自定义对象的默认__hash__是其对象 ID所以 JAX 无从得知对象被修改应触发重新编译。部分解决方法是定义合适的__hash__与__eq__class CustomClass: def __init__(self, x: jnp.ndarray, mul: bool): self.x x self.mul mul jit(static_argnums0) def calc(self, y): if self.mul: return self.x * y return y def __hash__(self): return hash((self.x, self.mul)) def __eq__(self, other): return (isinstance(other, CustomClass) and (self.x, self.mul) (other.x, other.mul))只要永不修改对象这种方式能配合 JIT 与其他变换正常工作。对象一旦被修改作为哈希键使用会引发多种微妙问题——这正是可变容器dict、list不定义__hash__、而不可变容器tuple定义的原因。若你的类依赖原地修改如方法内self.attr ...对象就不是真正静态的标记为静态会出问题——这时应该用策略 3。策略 3把CustomClass注册为 PyTree最灵活的做法是把类型注册为自定义 PyTree 节点精确指定哪些组件作为静态aux data、哪些作为动态childrenclass CustomClass: def __init__(self, x: jnp.ndarray, mul: bool): self.x x self.mul mul jit def calc(self, y): if self.mul: return self.x * y return y def _tree_flatten(self): children (self.x,) # 数组 / 动态值 aux_data {mul: self.mul} # 静态值 return (children, aux_data) classmethod def _tree_unflatten(cls, aux_data, children): return cls(*children, **aux_data) from jax import tree_util tree_util.register_pytree_node(CustomClass, CustomClass._tree_flatten, CustomClass._tree_unflatten)该方案解决了前述所有问题c CustomClass(2, True) print(c.calc(3)) c.mul False # 修改会被检测到 print(c.calc(3)) c CustomClass(jnp.array(2), True) # 不可哈希的 x 也能支持 print(c.calc(3))只要tree_flatten/tree_unflatten正确处理类的所有相关属性就能不加任何特殊注解直接将该类型的对象作为 JIT 函数的参数。PyTree 注册机制的底层实现可参见 tree_util.py 中的register_pytree_node系列函数。 越界索引Out-of-bounds indexingNumPy 越界索引会抛错np.arange(10)[11]但 JAX 无法或极难在加速器上抛出运行时错误因此必须为越界索引选择某种非报错行为类似于无效浮点运算产生NaN索引更新如index_add、scatter 类原语越界索引处的更新被跳过索引读取如 NumPy 索引、gather 类原语索引被钳制clamp到数组边界因为必须返回某个东西。例如下面的索引操作会返回数组最后一个值jnp.arange(10)[11]这与底层GatherScatterMode的默认行为一致。查看 slicing.py 中GatherScatterMode的定义CLIP把索引钳到最近的界内值保证要 gather 的整个窗口都在界内FILL_OR_DROPgather 时若窗口任何部分越界则整个窗口用常量填充scatter 时若窗口任何部分越界则整个窗口丢弃PROMISE_IN_BOUNDS用户承诺索引在界内不做额外检查。当前 XLA 实现下越界 gather 会被钳制、越界 scatter 会被丢弃索引越界时梯度不正确。.at默认使用promise_in_bounds语义mode参数缺省值映射关系见 array_methods.py。用.at[...].get()精细控制越界行为如果需要对越界索引进行更精细的控制可用ndarray.at的可选参数jnp.arange(10.0).at[11].get()jnp.arange(10.0).at[11].get(modefill, fill_valuejnp.nan)mode取值promise_in_bounds默认get 钳制 / set、add 等丢弃、clip钳制、drop丢弃、filldrop的别名对get()可用fill_value指定返回值。另有wrap_negative_indices默认 True负索引表示从数组末尾数、indices_are_sorted与unique_indices提示实现索引已排序/唯一部分后端可优化执行若声明与实际不符则输出未定义。fill_value默认对非精确类型为NaN、有符号类型为最大负值、无符号类型为最大正值、布尔为True。两个连带影响由于索引读取的钳制行为jnp.nanargmin、jnp.nanargmax对全 NaN 切片返回 -1而 NumPy 会抛错。上述两种行为更新跳过 vs 读取钳制互不为逆运算因此反向模式自动微分把索引更新转成索引读取、反之亦然不会保持越界索引语义。最好把 JAX 中的越界索引视为未定义行为undefined behavior。 非数组输入NumPy vs. JAXNumPy 通常乐意接受 Python list / tuple 作为 API 输入np.sum([1, 2, 3])JAX 则相反通常会给出有用的报错jnp.sum([1, 2, 3])这是刻意的设计决策把 list/tuple 传给被追踪traced的函数可能导致难以察觉的静默性能退化。例如下面这个宽容版jnp.sumdef permissive_sum(x): return jnp.sum(jnp.array(x)) x list(range(10)) permissive_sum(x)输出符合预期但掩盖了性能问题在 JAX 的追踪与 JIT 编译模型里Python list/tuple 的每个元素都被当作独立的 JAX 变量被单独处理并推送到设备。用make_jaxpr可以直观看到make_jaxpr(permissive_sum)(x)每个 list 元素都被当作独立输入处理追踪与编译开销随 list 长度线性增长。为避免此类意外JAX 不隐式转换 list/tuple 为数组。若确实要向 JAX 函数传 tuple/list请先显式转成数组jnp.sum(jnp.array(x)) 随机数Random numbersJAX 的伪随机数生成与 NumPy 有本质区别。NumPy 使用隐式、全局、有状态的随机状态JAX 采用显式、无状态的 PRNG key体系每个随机操作都接收一个key参数并消费它新 key 通过random.split派生。快速上手可参考 docs/101/random.md 与 docs/101/index.rst 对应的教程底层实现在 jax/_src/random/core.pykey、split、fold_in等核心函数默认使用 Threefry 计数器模式算法jax_default_prng_impl配置项默认threefry2x32见 config.py。key random.key(0) # 生成一个 PRNG key key, subkey random.split(key) # 分裂出新 key x random.uniform(subkey, (1000,))务必遵守每条随机路径使用独立 key、用后即 split的规则避免可复现性被破坏。 控制流Control flow控制流细节已从本文移入专门的指南 docs/201/control-flow.mdjit对 Python 控制流与逻辑运算符的使用施加了约束需要用lax.cond、lax.while_loop、lax.fori_loop、lax.scan等结构化控制流原语来表达依赖数据值的分支与循环。 动态形状Dynamic shapes用于jax.jit、jax.vmap、jax.grad等变换的 JAX 代码要求所有输出数组与中间数组具有静态形状即形状不能依赖其他数组中的值。例如自己实现jnp.nansum时可能这样写def nansum(x): mask ~jnp.isnan(x) # 选择非 NaN 值的布尔掩码 x_without_nans x[mask] return x_without_nans.sum()在 JIT 之外它可以正常工作x jnp.array([1, 2, jnp.nan, 3, 4]) print(nansum(x))但对其应用jax.jit或其他变换就会报错jax.jit(nansum)(x)问题在于x_without_nans的尺寸依赖x中的值即它是动态的。JAX 中通常可以用其他手段绕开动态形状数组例如用三参数形式jnp.where把 NaN 替换为 0得到相同结果同时避免动态形状jax.jit def nansum_2(x): mask ~jnp.isnan(x) # 选择非 NaN 值的布尔掩码 return jnp.where(mask, x, 0).sum() print(nansum_2(x))其他出现动态形状数组的场景也可采用类似技巧。 调试 NaN 与 Inf使用jax_debug_nans与jax_debug_infs两个 flag 定位函数与梯度中 NaN/Inf 的来源。它们定义于 config.pyjax_debug_nans默认False。给每个操作添加 NaN 检查当在 jit 编译计算输出中检测到 NaN 时回退到未编译版本以更精确地定位产生 NaN 的操作jax_debug_infs默认False。与上同理针对 Inf 检查。详细使用方式见 docs/debugging/flags.md。 双精度64-bit precisionJAX 默认强制单精度以缓解 NumPy API 将操作数激进提升到double的倾向。这对许多机器学习应用是期望行为但可能出乎你的意料x random.uniform(random.key(0), (1000,), dtypejnp.float64) x.dtype输出仍是float32。要使用双精度必须在**启动时startup**设置jax_enable_x64配置变量有以下几种方式设置环境变量JAX_ENABLE_X64True启动时手动设置配置 flag# 注意这只能在启动时生效 import jax jax.config.update(jax_enable_x64, True)用absl.app.run(main)解析命令行 flagsimport jax jax.config.config_with_absl()让 JAX 替你运行 absl 解析import jax if __name__ __main__: # 调用 jax.config.config_with_absl() 并执行 absl 解析 jax.config.parse_flags_with_absl()注意方式 24 对 JAX 的任意配置选项都适用。确认 x64 已启用import jax import jax.numpy as jnp from jax import random jax.config.update(jax_enable_x64, True) x random.uniform(random.key(0), (1000,), dtypejnp.float64) x.dtype # -- dtype(float64)从源码看jax_enable_x64是 config.py 中定义的布尔配置项默认False且被标记为include_in_jit_keyTrue、include_in_trace_contextTrue即它会参与 JIT 编译键与追踪上下文因此必须在启动时尽早设置否则已缓存的编译产物不会随开关切换而更新。注意事项⚠️ XLA 并非在所有后端都支持 64 位卷积 与 NumPy 的其他已知分歧jax.numpy尽力复刻 numpy API但存在一些 API 分歧的边界情况除上文各节外已知分歧还包括类型提升规则二元运算中JAX 的类型提升规则与 NumPy 略有不同详见 docs/101/type_promotion.rst。不安全类型转换unsafe cast当目标 dtype 无法表示输入值时JAX 的行为可能依赖后端总体上可能与 NumPy 不同。NumPy 通过astype的casting参数控制结果JAX 不提供此类配置直接继承 XLAConvertElementType的语义。例如 np.arange(254.0, 258.0).astype(uint8) array([254, 255, 0, 1], dtypeuint8) jnp.arange(254.0, 258.0).astype(uint8) Array([254, 255, 255, 255], dtypeuint8)这类不一致典型出现在浮点与整数类型之间转换极端值时。次正规数subnormal刷新为零在部分后端上JAX 对次正规浮点数采用 flush-to-zero 语义 import jax.numpy as jnp subnormal jnp.float32(1E-45) subnormal # 次正规数本身可表示 Array(1.e-45, dtypefloat32) subnormal 0 # 但在运算内被刷新为零 Array(0., dtypefloat32)次正规数的详细运算语义通常随后端而异。教程中覆盖的其他坑docs/201/control-flow.md讲解如何在jit对 Python 控制流与逻辑运算符的约束下工作docs/stateful-computations.md鉴于 JAX 变换只能作用于纯函数该文给出在 JAX 程序中正确管理状态的建议。Fin.如果本文没有覆盖到让你抓狂的问题欢迎反馈以便扩充这份入门避坑指南。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表