You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.02 01:25:14