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

JAX中实现矩阵快速幂递归遇静态参数不可哈希问题求助

JAX中矩阵快速幂(快速平方)实现问题及解决方法

我尝试在JAX中编写矩阵快速幂(exponentiate-by-squaring)算法,但对追踪变量(traced variables)了解不足,导致实现遇到困难。

我的代码如下:

import numpy as np
import jax
import jax.numpy as jnp
import jax.lax as jlax
from functools import partial

@partial(jax.jit, static_argnums=(1,))
def matpow(A, n):
    dim = A.shape[0]
    return jlax.switch(
        n,
        [lambda: jnp.identity(dim),
         lambda: A,
         lambda: jlax.cond(
             jnp.floor_divide(n, 2) == jnp.true_divide(n, 2),
             lambda: matpow(jnp.dot(A, A), jnp.floor_divide(n, 2)),
             lambda: jnp.dot(A, matpow(jnp.dot(A, A), jnp.floor_divide(n, 2)))
         )])

运行matpow(2 * jnp.eye(4), 5)时,编译抛出错误:

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

我完全不明白这个错误的含义,困惑的是n明明是整数,哈希应该很简单。还尝试过其他方法:使用jnp.binary_repr(尚未实现),使用np.binary_repr(出现TracerIntegerConversionError,尽管n已标记为静态参数),以及在matpow内部定义独立函数recpow的递归封装版本(触发递归限制)。


问题根源

错误的核心是:你在jlax.switch中使用静态参数n作为分支索引,但递归调用时传递的jnp.floor_divide(n, 2)是动态追踪变量,而非静态整数。即使n被标记为静态,在JAX的JIT编译流程中,递归调用时的参数会被追踪,无法直接作为静态参数传递,导致哈希失败。

另外,jlax.switch的分支索引需要在编译时确定为静态值,而你的第三个分支里包含动态条件判断,这和switch的静态分支设计冲突。

正确实现方式

JAX中实现递归快速幂,需要用jax.lax.while_loop代替递归,或者使用jax.lax.scan遍历指数的二进制位,同时确保静态参数在编译时是确定的。这里提供两种可行方案:

方案1:基于while_loop的迭代式快速幂

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

@partial(jax.jit, static_argnums=(1,))
def matpow(A, n):
    # 初始化结果为单位矩阵
    result = jnp.identity(A.shape[0])
    base = A

    # 定义循环条件:n > 0
    def cond(carry):
        _, n_remaining = carry
        return n_remaining > 0

    # 定义循环体:处理当前二进制位
    def body(carry):
        res, n_remaining = carry
        # 如果当前位是1,结果乘base
        res = jnp.where(n_remaining % 2 == 1, jnp.dot(res, base), res)
        # base平方,n右移一位
        base_new = jnp.dot(base, base)
        n_new = n_remaining // 2
        return (res, n_new)

    # 执行循环
    final_result, _ = jax.lax.while_loop(cond, body, (result, n))
    return final_result

方案2:基于scan的二进制位遍历法

先将指数n的二进制位转为静态列表(因为n是静态参数,编译时可确定),然后用scan遍历每一位:

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

@partial(jax.jit, static_argnums=(1,))
def matpow(A, n):
    # 静态生成n的二进制位(从最低位到最高位)
    bits = [int(bit) for bit in bin(n)[2:][::-1]]
    dim = A.shape[0]

    # 定义scan的步函数
    def step(carry, bit):
        res, base = carry
        # 如果当前位是1,结果乘base
        res = jnp.where(bit == 1, jnp.dot(res, base), res)
        # base平方
        base_new = jnp.dot(base, base)
        return (res, base_new), None

    # 初始状态:结果为单位矩阵,base为输入矩阵
    initial_carry = (jnp.identity(dim), A)
    (final_result, _), _ = jax.lax.scan(step, initial_carry, bits)
    return final_result

关键说明

  1. 避免递归调用JIT函数:JAX的JIT编译对递归支持有限,容易触发递归深度限制或追踪变量问题,迭代式实现更可靠。
  2. 静态参数的正确使用:标记为static_argnums的参数必须在编译时是确定的Python整数/数值,不能是JAX追踪的动态变量。
  3. 选择合适的控制流:jlax.switch适合编译时确定分支的场景,动态条件判断用jnp.where或jlax.cond,循环用while_loop或scan。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 22:54:32