JAX 易踩坑指南(Common Gotchas):纯函数、原地更新、jit 与动态形状的完整避坑手册
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
JAX 是一套对 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 functions)
JAX 的变换(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(even=0, odd=0) 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: x+array[i], 0)) # 期望 45 iterator = iter(range(10)) print(lax.fori_loop(0, 10, lambda i,x: x+next(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: x+1, 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 updates)
NumPy 中常见的原地索引更新:
numpy_array = np.zeros((3,3), dtype=np.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), dtype=jnp.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), dtype=jnp.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] = y |
x = x.at[idx].add(y) | x[idx] += y |
x = x.at[idx].subtract(y) | x[idx] -= y |
x = x.at[idx].multiply(y) | x[idx] *= y |
x = x.at[idx].divide(y) | x[idx] /= y |
x = x.at[idx].power(y) | x[idx] **= y |
x = 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(类型CustomClass),JAX 不知道如何处理这种类型。文档给出三种基本策略。
策略 1:JIT 编译的辅助函数(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_argnums=0) 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_argnums=0) 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_argnums=0) 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)、哪些作为动态(children):
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 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 indexing)
NumPy 越界索引会抛错:
np.arange(10)[11]但 JAX 无法(或极难)在加速器上抛出运行时错误,因此必须为越界索引选择某种"非报错"行为(类似于无效浮点运算产生NaN):
- 索引更新(如
index_add、scatter 类原语):越界索引处的更新被跳过; - 索引读取(如 NumPy 索引、gather 类原语):索引被钳制(clamp)到数组边界,因为必须返回"某个东西"。例如下面的索引操作会返回数组最后一个值:
jnp.arange(10)[11]这与底层GatherScatterMode的默认行为一致。查看 slicing.py 中GatherScatterMode的定义:
CLIP:把索引钳到最近的界内值,保证要 gather 的整个窗口都在界内;FILL_OR_DROP:gather 时若窗口任何部分越界则整个窗口用常量填充;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(mode='fill', fill_value=jnp.nan)mode取值:"promise_in_bounds"(默认,get 钳制 / set、add 等丢弃)、"clip"(钳制)、"drop"(丢弃)、"fill"("drop"的别名,对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. JAX
NumPy 通常乐意接受 Python list / tuple 作为 API 输入:
np.sum([1, 2, 3])JAX 则相反,通常会给出有用的报错:
jnp.sum([1, 2, 3])这是刻意的设计决策:把 list/tuple 传给被追踪(traced)的函数可能导致难以察觉的静默性能退化。例如下面这个"宽容版"jnp.sum:
def 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 numbers)
JAX 的伪随机数生成与 NumPy 有本质区别。NumPy 使用隐式、全局、有状态的随机状态;JAX 采用显式、无状态的 PRNG key体系:每个随机操作都接收一个key参数并消费它,新 key 通过random.split派生。快速上手可参考 docs/101/random.md 与 docs/101/index.rst 对应的教程;底层实现在 jax/_src/random/core.py(key、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.md:jit对 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.py:
jax_debug_nans:默认False。给每个操作添加 NaN 检查;当在 jit 编译计算输出中检测到 NaN 时,回退到未编译版本以更精确地定位产生 NaN 的操作;jax_debug_infs:默认False。与上同理,针对 Inf 检查。
详细使用方式见 docs/debugging/flags.md。
🔪 双精度(64-bit precision)
JAX 默认强制单精度,以缓解 NumPy API 将操作数激进提升到double的倾向。这对许多机器学习应用是期望行为,但可能出乎你的意料:
x = random.uniform(random.key(0), (1000,), dtype=jnp.float64) x.dtype输出仍是float32。要使用双精度,必须在**启动时(startup)**设置jax_enable_x64配置变量,有以下几种方式:
- 设置环境变量
JAX_ENABLE_X64=True; - 启动时手动设置配置 flag:
# 注意:这只能在启动时生效! import jax jax.config.update("jax_enable_x64", True)- 用
absl.app.run(main)解析命令行 flags:
import 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()注意方式 2~4 对 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,), dtype=jnp.float64) x.dtype # --> dtype('float64')从源码看,jax_enable_x64是 config.py 中定义的布尔配置项(默认False),且被标记为include_in_jit_key=True、include_in_trace_context=True,即它会参与 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], dtype=uint8) >>> jnp.arange(254.0, 258.0).astype('uint8') Array([254, 255, 255, 255], dtype=uint8)这类不一致典型出现在浮点与整数类型之间转换极端值时。
- 次正规数(subnormal)刷新为零:在部分后端上,JAX 对次正规浮点数采用 flush-to-zero 语义:
>>> import jax.numpy as jnp >>> subnormal = jnp.float32(1E-45) >>> subnormal # 次正规数本身可表示 Array(1.e-45, dtype=float32) >>> subnormal + 0 # 但在运算内被刷新为零 Array(0., dtype=float32)次正规数的详细运算语义通常随后端而异。
教程中覆盖的其他坑
- docs/201/control-flow.md:讲解如何在
jit对 Python 控制流与逻辑运算符的约束下工作; - docs/stateful-computations.md:鉴于 JAX 变换只能作用于纯函数,该文给出在 JAX 程序中正确管理状态的建议。
Fin.
如果本文没有覆盖到让你抓狂的问题,欢迎反馈以便扩充这份入门避坑指南。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考