在JAX中对子数组应用自定义分段函数

编程语言 2026-07-09

假设我有一个将向量映射到向量的函数 vec_fun。我想基于某个条件将这个函数应用到数组的一个子向量。比如:

mask = jnp.asarray([False, True, False, True, True, True, False]) # Some mask
vec = jnp.arange(mask.size, dtype=float) # Some vector

def vec_fun(vec): # Some function that maps a vector into a vector of the same size
    return (vec + jnp.flip(vec)**2)

@jax.jit
def func_segmented(vec, mask):
    return vec.at[mask].set(vec_fun(vec[mask])) # Try to replace a subvector by a function of it

如果没有 jit,这段代码可以运行,但在 jit 下会失败,因为 vec[mask] 的大小只有在运行时才会被知道:

NonConcreteBooleanIndexError: Array boolean indices must be concrete; got bool[7]

但我知道 jax 确实支持一些特殊的 分段函数,它们在 jit 下对子数组应用函数时没有问题。

那么难道不应该也能够对自定义函数做同样的事情,前提是分段的数量是已知且静态的吗?

如果是这样,我应该如何在子数组上定义一个自定义的分段函数?

如果这不可能,为什么呢?

解决方案

对于像用户问题中描述的非逐元素函数,我还没有发现一个很好的方法来实现,除了在JIT之外对数组进行掩蔽,使得 vec_masked = vec[mask] 可以在不出错的情况下被构造出来。多年来JAX团队一直在探索在JIT内实现动态形状或ragged(非规则)数组的可能实现,但目前还没有成熟到可以在一般情况下使用。

对于更简单的逐元素函数,推荐的做法是使用三项式 jnp.where

@jax.jit
def func_segmented(vec, mask):
    return jnp.where(mask, vec_fun(vec), vec)

其缺点是 vec_fun 会在整个向量上进行计算,因此包含一些无谓的计算。其好处是它很容易向量化(在使用加速器时也可并行),并且数组的大小不依赖于运行时的值,因此在 mask 的值改变时不需要重新编译。

如果你认为避免这种无谓计算很重要,你的最佳选项要么是

  1. 在JIT之外静态地拆分你的向量,只把所需的那一部分传给 vec_fun。这里的缺点是必须在JIT之外完成,并且每次大小改变时 vec_fun 会重新编译,这可能抵消避免无谓计算带来的收益。
  2. 使用 lax.map 对你的数组进行逐步遍历(或按 batch_size 参数分块遍历),仅在需要时调用 vec_fun。这里的好处是减少无谓计算,但缺点是顺序计算的成本可能超过任何改进。

如果运行时成本很重要,你可能会对这三种选项进行基准测试。但通常,jnp.where 的方法(以在不重新编译的前提下实现并行运行的能力换取无谓计算的代价)往往是性能最优的。

站内所有文章版权归属LeftHeroAI导航站,无授权禁止任何主体转载、抄袭、复制内容,亦不得私自架设镜像站点。一经侵权,本站将通过法律途径追责。

相关文章