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

使用JAX cond时出现不可哈希静态参数错误的原因咨询

JAX中jax.lax.cond调用带静态参数的jit函数报错问题

问题描述

我编写了一段简化版JAX程序:

import jax.numpy as jnp
import jax
from functools import partial
from jax import jit

def main():
    N = 3
    x = jnp.ones((4, 5))
    a = True

    answer = jax.lax.cond(a, run_true, run_false, N, x)
    print(answer)


@partial(jit, static_argnames=["N"])
def run_true(N, x):
    return x.reshape((1, 4, 5)) * jnp.ones((N, 4, 5))

@partial(jit, static_argnames=["N"])
def run_false(N, x):
    return jnp.zeros((N, 4, 5))


if __name__ == "__main__":
    main()

运行时抛出如下错误:

ValueError: Non-hashable static arguments are not supported, as this can lead to unexpected cache-misses. Static argument (index 0) of type <class 'jax.interpreters.partial_eval.DynamicJaxprTracer'> for function run_true is non-hashable.

我对此感到困惑,run_true和run_false都是已用jit编译的函数,为何会触发报错?

改用Python原生条件分支可以正常运行,但我需要在jit编译的函数内部使用jax.lax.cond,因此该方案不可行:

if a:
    answer = run_true(N, x)
else:
    answer = run_false(N, x)
print(answer)

问题原因与解决方案

原因分析

问题出在jax.lax.cond的调用逻辑上:jax.lax.cond会将传入的分支函数当作未编译的纯函数处理,不会保留外部的jit编译配置。当你把标记了N为静态参数的jit函数直接传给cond时,cond会在追踪过程中将N作为动态参数传递给分支函数,但静态参数要求编译时必须是可哈希的确定值,而动态追踪状态下的N(类型为DynamicJaxprTracer)不具备哈希性,因此触发报错。

解决方案

有两种可行的解决方式:

  1. 将jit编译逻辑整合到cond所在的jit函数中
    去掉分支函数的外层jit,把整个cond调用包裹在带静态参数的jit函数内:

    import jax.numpy as jnp
    import jax
    from functools import partial
    from jax import jit
    
    def run_true(N, x):
        return x.reshape((1, 4, 5)) * jnp.ones((N, 4, 5))
    
    def run_false(N, x):
        return jnp.zeros((N, 4, 5))
    
    @partial(jit, static_argnames=["N"])
    def cond_logic(N, x, a):
        return jax.lax.cond(a, run_true, run_false, N, x)
    
    def main():
        N = 3
        x = jnp.ones((4, 5))
        a = True
        answer = cond_logic(N, x, a)
        print(answer)
    
    if __name__ == "__main__":
        main()
    
  2. 提前用functools.partial绑定静态参数
    在传给jax.lax.cond之前,先把静态参数N绑定到分支函数上,让cond仅传递动态参数x:

    import jax.numpy as jnp
    import jax
    from functools import partial
    from jax import jit
    
    def main():
        N = 3
        x = jnp.ones((4, 5))
        a = True
    
        # 提前绑定静态参数N
        run_true_bound = partial(run_true, N=N)
        run_false_bound = partial(run_false, N=N)
        answer = jax.lax.cond(a, run_true_bound, run_false_bound, x)
        print(answer)
    
    @partial(jit, static_argnames=["N"])
    def run_true(N, x):
        return x.reshape((1, 4, 5)) * jnp.ones((N, 4, 5))
    
    @partial(jit, static_argnames=["N"])
    def run_false(N, x):
        return jnp.zeros((N, 4, 5))
    
    if __name__ == "__main__":
        main()
    

补充说明

jax.lax.cond是JAX的控制流原语,它会将整个分支逻辑编译为单一计算图,而非像Pythonif/else那样在运行时选择分支。因此传入的分支函数必须能被JAX正常追踪,且参数要符合动态/静态的规则。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 19:05:22