JAX函数poly_func JIT编译报错:TypeError: unhashable type: 'DynamicJaxprTracer'
解决JAX JIT编译poly_func时的TypeError: unhashable type: 'DynamicJaxprTracer'问题
错误原因
JAX JIT编译时,若输入的order和index被当作动态参数,会被转换为DynamicJaxprTracer类型的追踪对象——这类对象不可哈希,无法直接作为字典的键进行查找,因此触发报错。
解决方案
根据order和index是否为编译时确定的静态参数,分两种处理方式:
情况1:order和index是静态参数(编译时固定)
如果调用poly_func时,order和index的值在编译阶段就已确定,可通过jax.jit的static_argnums参数将这两个参数标记为静态参数。编译时静态参数会被当作常量处理,不会被追踪,因此可以正常用字典索引。
代码示例:
import jax poly_dict = { (0, 0): lambda x, y, z: 1., (1, 0): lambda x, y, z: x, (1, 1): lambda x, y, z: y, (1, 2): lambda x, y, z: z, (2, 0): lambda x, y, z: x*x, (2, 1): lambda x, y, z: y*y, (2, 2): lambda x, y, z: z*z, (2, 3): lambda x, y, z: x*y, (2, 4): lambda x, y, z: y*z, (2, 5): lambda x, y, z: z*x } def poly_func(order: int, index: int): try: return poly_dict[(order, index)] except KeyError: print("(order, index) must be a key in poly_dict!") return # 将order(第0个参数)和index(第1个参数)标记为静态参数 jit_poly_func = jax.jit(poly_func, static_argnums=(0, 1)) # 使用示例 selected_func = jit_poly_func(1, 0) print(selected_func(2., 3., 4.)) # 输出: 2.0
情况2:order和index是动态参数(运行时可变)
如果order和index的值需要在运行时动态变化,不能用Python原生字典的哈希查找,需改用JAX提供的可追踪分支工具jax.lax.switch实现函数选择。
步骤:
- 将所有多项式函数整理为有序列表;
- 建立
(order, index)到列表索引的映射逻辑; - 用
jax.lax.switch根据动态参数选择对应函数执行。
代码示例:
import jax import jax.numpy as jnp # 将poly_dict中的函数按顺序整理为列表 poly_list = [ lambda x, y, z: 1., # (0, 0) lambda x, y, z: x, # (1, 0) lambda x, y, z: y, # (1, 1) lambda x, y, z: z, # (1, 2) lambda x, y, z: x*x, # (2, 0) lambda x, y, z: y*y, # (2, 1) lambda x, y, z: z*z, # (2, 2) lambda x, y, z: x*y, # (2, 3) lambda x, y, z: y*z, # (2, 4) lambda x, y, z: z*x # (2, 5) ] def poly_func_dynamic(order: int, index: int, x, y, z): # 根据order和index计算对应的列表索引 idx = jax.lax.switch( order, [ # order=0时,仅index=0有效,否则返回0(默认值) lambda idx_val: jnp.where(idx_val == 0, 0, 0), # order=1时,index对应列表索引为index+1 lambda idx_val: idx_val + 1, # order=2时,index对应列表索引为index+4 lambda idx_val: idx_val + 4 ], index ) # 通过switch选择对应函数并执行 return jax.lax.switch(idx, poly_list, x, y, z) # JIT编译动态版本 jit_poly_dynamic = jax.jit(poly_func_dynamic) # 使用示例 print(jit_poly_dynamic(1, 0, 2., 3., 4.)) # 输出: 2.0 print(jit_poly_dynamic(2, 3, 2., 3., 4.)) # 输出: 6.0
内容的提问来源于stack exchange,提问作者Jingyang Wang
相关产品推荐
相关产品推荐

