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

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设备运行的问题:

  1. Python循环无法被设备正确编译:JAX/Pallas的设备内核里,Python的for循环会在追踪阶段被展开,但对于设备端循环,更适合用JAX原生的jax.lax.fori_loop来实现。
  2. 用了Python的min()而非JAX原生函数:Python的min()会在追踪阶段就被求值为常量,无法处理设备上的动态计算,需要替换为jax.lax.min()。
  3. 数组更新的效率问题:多次用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.20 08:13:00