使用JAX cond时出现不可哈希静态参数错误的原因咨询
JAX中
jax.lax.cond调用带静态参数的jit函数报错问题 问题描述
我编写了一段简化版JAX程序:
import jax.numpy as jnp import jax from functools import partial from jax import jit def main(): N = 3 x = jnp.ones((4, 5)) a = True answer = jax.lax.cond(a, run_true, run_false, N, x) print(answer) @partial(jit, static_argnames=["N"]) def run_true(N, x): return x.reshape((1, 4, 5)) * jnp.ones((N, 4, 5)) @partial(jit, static_argnames=["N"]) def run_false(N, x): return jnp.zeros((N, 4, 5)) if __name__ == "__main__": main()
运行时抛出如下错误:
ValueError: Non-hashable static arguments are not supported, as this can lead to unexpected cache-misses. Static argument (index 0) of type <class 'jax.interpreters.partial_eval.DynamicJaxprTracer'> for function run_true is non-hashable.
我对此感到困惑,run_true和run_false都是已用jit编译的函数,为何会触发报错?
改用Python原生条件分支可以正常运行,但我需要在jit编译的函数内部使用jax.lax.cond,因此该方案不可行:
if a: answer = run_true(N, x) else: answer = run_false(N, x) print(answer)
问题原因与解决方案
原因分析
问题出在jax.lax.cond的调用逻辑上:jax.lax.cond会将传入的分支函数当作未编译的纯函数处理,不会保留外部的jit编译配置。当你把标记了N为静态参数的jit函数直接传给cond时,cond会在追踪过程中将N作为动态参数传递给分支函数,但静态参数要求编译时必须是可哈希的确定值,而动态追踪状态下的N(类型为DynamicJaxprTracer)不具备哈希性,因此触发报错。
解决方案
有两种可行的解决方式:
将jit编译逻辑整合到
cond所在的jit函数中
去掉分支函数的外层jit,把整个cond调用包裹在带静态参数的jit函数内:import jax.numpy as jnp import jax from functools import partial from jax import jit def run_true(N, x): return x.reshape((1, 4, 5)) * jnp.ones((N, 4, 5)) def run_false(N, x): return jnp.zeros((N, 4, 5)) @partial(jit, static_argnames=["N"]) def cond_logic(N, x, a): return jax.lax.cond(a, run_true, run_false, N, x) def main(): N = 3 x = jnp.ones((4, 5)) a = True answer = cond_logic(N, x, a) print(answer) if __name__ == "__main__": main()提前用
functools.partial绑定静态参数
在传给jax.lax.cond之前,先把静态参数N绑定到分支函数上,让cond仅传递动态参数x:import jax.numpy as jnp import jax from functools import partial from jax import jit def main(): N = 3 x = jnp.ones((4, 5)) a = True # 提前绑定静态参数N run_true_bound = partial(run_true, N=N) run_false_bound = partial(run_false, N=N) answer = jax.lax.cond(a, run_true_bound, run_false_bound, x) print(answer) @partial(jit, static_argnames=["N"]) def run_true(N, x): return x.reshape((1, 4, 5)) * jnp.ones((N, 4, 5)) @partial(jit, static_argnames=["N"]) def run_false(N, x): return jnp.zeros((N, 4, 5)) if __name__ == "__main__": main()
补充说明
jax.lax.cond是JAX的控制流原语,它会将整个分支逻辑编译为单一计算图,而非像Pythonif/else那样在运行时选择分支。因此传入的分支函数必须能被JAX正常追踪,且参数要符合动态/静态的规则。
内容的提问来源于stack exchange,提问作者Scott Swarthout
相关产品推荐
相关产品推荐

