Jax Pallas内核中静态参数被追踪及对角矩阵乘法逻辑失效问题求助
Jax Pallas内核中静态参数被追踪及对角矩阵乘法逻辑失效问题求助
看起来你遇到了两个JAX/Pallas新手常见的问题:静态参数没正确传入内核导致被追踪,还有Python循环在设备内核里的适配问题,我来帮你一步步拆解解决:
一、为什么offsets还是被追踪了?
你在外层jax.jit里标记了static_argnums=(1,),但Pallas的pallas_call默认会把所有输入参数都作为动态张量传递给内核,哪怕外层已经标记为静态。要让offsets在内核里保持静态,需要在pallas_call中也明确指定静态参数,或者把offsets提前绑定到内核函数上。
解决方法1:在pallas_call中指定静态参数
pallas_call同样支持static_argnums参数,对应传入的参数位置(这里offsets是第二个参数,所以位置是1):
@functools.partial(jax.jit, static_argnums=(1, )) def dia_matmul(diags: Array, offsets: tuple[int], other: Array) -> Array: return pl.pallas_call( dia_matmul_kernel, out_shape=jax.ShapeDtypeStruct(other.shape, other.dtype), static_argnums=(1,) # 这里也要指定静态参数 )(diags, offsets, other)
解决方法2:用functools.partial绑定静态参数到内核
把offsets直接绑定到内核函数上,这样内核就不需要接收这个参数,自然就是静态的:
def dia_matmul_kernel(offsets, diags_ref, other_ref, o_ref): # 这里offsets已经是静态值,不会被追踪 diags, other = diags_ref[...], other_ref[...] # 后续逻辑不变... @functools.partial(jax.jit, static_argnums=(1, )) def dia_matmul(diags: Array, offsets: tuple[int], other: Array) -> Array: # 绑定offsets到内核 bound_kernel = functools.partial(dia_matmul_kernel, offsets) return pl.pallas_call( bound_kernel, out_shape=jax.ShapeDtypeStruct(other.shape, other.dtype) )(diags, other)
二、对角矩阵乘法逻辑失效的问题
你的内核里还有几个适配Pallas设备运行的问题:
- Python循环无法被设备正确编译:JAX/Pallas的设备内核里,Python的
for循环会在追踪阶段被展开,但对于设备端循环,更适合用JAX原生的jax.lax.fori_loop来实现。 - 用了Python的
min()而非JAX原生函数:Python的min()会在追踪阶段就被求值为常量,无法处理设备上的动态计算,需要替换为jax.lax.min()。 - 数组更新的效率问题:多次用
out.at[...]创建新数组,在循环里会累积不必要的开销,建议用JAX的索引更新来累积结果。
修改后的内核代码
def dia_matmul_kernel(diags_ref, offsets, other_ref, o_ref): diags, other = diags_ref[...], other_ref[...] N = other.shape[0] out = jnp.zeros((N, N)) # 用jax.lax.fori_loop实现设备端循环 def loop_body(i, val): out = val offset = offsets[i] diag = diags[i] start = jax.lax.max(0, offset) end = jax.lax.min(N, N + offset) top = jax.lax.max(0, -offset) bottom = top + end - start # 用索引更新累积结果 update = diag[start:end, None] * other[start:end, :] return out.at[top:bottom, :].add(update) out = jax.lax.fori_loop(0, len(offsets), loop_body, out) o_ref[...] = out
完整可运行代码
把上面的修改整合后,完整代码如下:
import functools import jax from jax.experimental import pallas as pl import jax.numpy as jnp import numpy as np from jaxtyping import Array key = jax.random.PRNGKey(52) other = jax.random.normal(key, (10, 10)) diags = jax.random.normal(key, (3, 10)) offsets = (-2, 1, 2) def dia_matmul_kernel(diags_ref, offsets, other_ref, o_ref): diags, other = diags_ref[...], other_ref[...] N = other.shape[0] out = jnp.zeros((N, N)) def loop_body(i, val): out = val offset = offsets[i] diag = diags[i] start = jax.lax.max(0, offset) end = jax.lax.min(N, N + offset) top = jax.lax.max(0, -offset) bottom = top + end - start update = diag[start:end, None] * other[start:end, :] return out.at[top:bottom, :].add(update) out = jax.lax.fori_loop(0, len(offsets), loop_body, out) o_ref[...] = out @functools.partial(jax.jit, static_argnums=(1, )) def dia_matmul(diags: Array, offsets: tuple[int], other: Array) -> Array: return pl.pallas_call( dia_matmul_kernel, out_shape=jax.ShapeDtypeStruct(other.shape, other.dtype), static_argnums=(1,) )(diags, offsets, other) # 测试运行 result = dia_matmul(diags, offsets, other) print(result.shape) # 应该输出(10,10)
这样修改后,offsets会保持静态不会被追踪,同时对角矩阵乘法的逻辑也能正确在Pallas内核中运行啦!
备注:内容来源于stack exchange,提问作者bsaoptima
相关产品推荐
相关产品推荐

