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

为何jax.grad(lambda v: jnp.linalg.norm(v-v))(jnp.ones(2))返回NaN?

问题解释与分析

这不是bug,是JAX自动微分的特性导致的,具体原因如下:

  • 第一个案例:grad(lambda v: jnp.linalg.norm(v-v))(x)
    虽然v-v的结果是全零数组,但JAX的自动微分会完整追踪计算路径:先计算v与自身的差,再对这个结果求L2范数。L2范数||u||对u的导数是u / ||u||,当u是全零数组时,就会出现0除以0的情况,最终得到NaN。而由于u是由v计算而来,JAX会继续链式求导,将这个NaN传递到对v的梯度中,所以最终结果是[nan, nan]。

  • 第二个案例:grad(lambda v: jnp.linalg.norm(0))(x)
    这里直接传入jnp.linalg.norm的是常数0,JAX会识别出这个值与输入变量v完全无关,因此对v的梯度为0,最终返回[0., 0.]。

补充:JAX常见陷阱相关内容(翻译自官方文档)

JAX的自动微分基于计算图追踪,而非符号式的代数简化。也就是说,即使某个表达式在数学上等价于常数,但只要它是通过输入变量计算得到的,JAX就会保留完整的计算路径来计算导数,不会提前将其简化为常数。这就是为什么第一个案例中v-v没有被当成常数0处理,而是保留了与v的关联,导致导数计算出现NaN。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 18:05:17