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
关键说明
- 避免递归调用JIT函数:JAX的JIT编译对递归支持有限,容易触发递归深度限制或追踪变量问题,迭代式实现更可靠。
- 静态参数的正确使用:标记为
static_argnums的参数必须在编译时是确定的Python整数/数值,不能是JAX追踪的动态变量。 - 选择合适的控制流:
jlax.switch适合编译时确定分支的场景,动态条件判断用jnp.where或jlax.cond,循环用while_loop或scan。
内容的提问来源于stack exchange,提问作者linkhyrule5
相关产品推荐
相关产品推荐

