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

JAX调用Scipy方法求导遇TracerArrayConversionError问题求助

JAX调用Scipy求导报错的解决方案

问题原因

scipy.spatial.distance.cdist是纯NumPy实现的函数,不支持JAX的自动微分追踪机制。当JAX的grad对函数求导时,会将输入包装为Tracer对象以记录计算路径,但Scipy函数会尝试将Tracer转换为NumPy数组,直接触发TracerArrayConversionError。

解决方案

方案1:用JAX原生操作实现欧氏距离矩阵

直接用JAX数组运算替代Scipy的cdist,完全兼容自动微分:

import jax.numpy as jnp
from jax import grad, random

key = random.PRNGKey(1)
size = 3
x = random.uniform(key, (size, size), dtype=jnp.float32)

def error_func(x):
    # 计算所有点对的欧氏距离
    diff = x[:, None, :] - x[None, :, :]
    dists = jnp.sqrt(jnp.sum(jnp.square(diff), axis=-1))
    return jnp.sum(dists)

error_diff = grad(error_func)

print(error_func(x))
print(error_diff(x))

方案2:使用JAX兼容的Scipy函数

JAX提供了部分Scipy函数的适配版本,直接导入jax.scipy下的cdist即可:

import jax.numpy as jnp
from jax import grad, random
from jax.scipy.spatial.distance import cdist  # 导入JAX版cdist

key = random.PRNGKey(1)
size = 3
x = random.uniform(key, (size, size), dtype=jnp.float32)

def error_func(x):
    dists = cdist(x, x, metric='euclidean')
    return jnp.sum(dists)

error_diff = grad(error_func)

print(error_func(x))
print(error_diff(x))

注意事项

  • 涉及自动微分的计算流程,优先使用JAX生态内的函数(jax.numpy、jax.scipy等),避免混用原版Scipy/NumPy函数
  • JAX的Tracer对象是自动微分的核心载体,不能直接转换为NumPy数组,否则会丢失梯度追踪能力

内容的提问来源于stack exchange,提问作者siamak attarian

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 18:54:58