使用jax.vmap结合广播向量化时遇Tracer对象转换错误问题
解决JAX vmap中Tracer对象转换错误的问题
错误根源是你在vmap包裹的函数里混用了原生numpy(np)函数。JAX的vmap会将输入转为Tracer对象以追踪计算流程,但numpy无法直接处理这类对象,因此触发了__array__()转换错误。
修正方案很简单:将所有numpy相关调用替换为JAX的numpy接口(jax.numpy,常用别名jnp),让计算全程在JAX的追踪体系内进行。
修正后的代码
import jax import jax.numpy as jnp import numpy as np # 定义处理单行的函数,全部使用jax.numpy接口 cfun = lambda x: jnp.sum(jnp.sin(x - x[:, jnp.newaxis]), axis=1) # 用vmap将函数向量化,处理二维数组的每一行 cfuns = jax.vmap(cfun) # 测试二维输入(可以是JAX数组或numpy数组,JAX会自动转换) x = jnp.arange(6).reshape(3, 2) print(cfuns(x)) # 如果输入是numpy数组,也可以直接传入 x_np = np.arange(6).reshape(3, 2) print(cfuns(x_np))
说明
jax.numpy的函数是为JAX的自动微分和向量化机制设计的,能正确识别并处理Tracer对象,避免转换错误。- vmap会自动将
cfun应用到二维数组x的每一行,最终输出形状为(3, 2)的数组,对应每一行的计算结果。
内容的提问来源于stack exchange,提问作者Abolfazl
相关产品推荐
相关产品推荐

