JAX中被jit装饰的函数使用迭代器得正确结果是否为偶然?
JAX中jit装饰含迭代器函数的结果是否偶然?
作为JAX新手,我查阅文档时了解到被jit装饰的函数不应包含迭代器(文档纯函数章节提及)。文档给出的示例如下:
import jax.numpy as jnp import jax.lax as lax from jax import jit # lax.fori_loop array = jnp.arange(10) print(lax.fori_loop(0, 10, lambda i,x: x+array[i], 0)) # expected result 45 iterator = iter(range(10)) print(lax.fori_loop(0, 10, lambda i,x: x+next(iterator), 0)) # unexpected result 0
为了触发错误而非未定义行为,我编写了以下测试代码:
@jit def f(x, arr): for i in range(10): x += arr[i] return x @jit def f1(x, arr): it = iter(arr) for i in range(10): x += next(it) return x print(f(0,array)) # 45 as expected print(f1(0,array)) # still 45
请问被jit装饰的函数f1得到正确结果是否属于偶然情况?
是的,这个正确结果完全是偶然情况,原因如下:
- JAX的
jit编译会将Python代码转换为XLA计算图,而迭代器属于Python运行时的有状态对象,XLA无法追踪其内部状态变化。编译f1时,JAX的追踪机制只是在编译阶段执行了一次迭代器遍历,把结果硬编码到了计算图中——这次刚好匹配预期,但换个JAX版本、运行环境,或者修改函数逻辑(比如迭代次数和数组长度不匹配),结果立刻会出现异常。 - 对比
f函数:range(10)会被JAX转换为静态循环次数,arr[i]是纯索引操作,XLA能正常编译为可靠的循环计算;但iter(arr)和next(it)是Python侧操作,不属于XLA可编译的范畴,JAX对这类操作的处理本身就是未定义行为。 - 本质上,JAX要求
jit装饰的函数必须是纯函数:输入相同则输出一致,且不能依赖外部或内部可变状态(迭代器就属于带内部状态的对象)。使用迭代器违反了这个核心原则,哪怕某次运行结果正确,也绝对不能依赖,后续随时可能出问题。
内容的提问来源于stack exchange,提问作者Niccolò Tiezzi
相关产品推荐
相关产品推荐

