JAX:兼容JIT的布尔稀疏矩阵高效切片方案问询
问题描述
我有一个布尔稀疏矩阵,通过True值的行索引和列索引来表示,数据初始化代码如下:
import numpy as np import jax from jax import numpy as jnp N = 10000 M = 1000 X = np.random.randint(0, 100, size=(N, M)) == 0 # data setup rows, cols = np.where(X == True) rows = jax.device_put(rows) cols = jax.device_put(cols)
希望仅通过行索引和列索引,实现类似X[:, 3]的列切片操作。
尝试用jnp.isin实现了该功能,但由于rows[cols == m]是数据依赖形状的数组,无法兼容JIT:
def not_jit_compatible_slice(rows, cols, m): return jnp.isin(jnp.arange(N), rows[cols == m])
改用三参数形式的jnp.where实现了兼容JIT的版本,但运行速度远慢于前者:
def jit_compatible_but_slow_slice(rows, cols, m): return jnp.isin(jnp.arange(N), jnp.where(cols == m, rows, -1))
请问是否存在既快速又兼容JIT的解决方案?
高效兼容JIT的解决方案
可以利用JAX内置的稀疏索引优化算子实现,以下两种方法既满足JIT要求,又能达到接近非JIT版本的性能:
方法一:使用jax.ops.segment_sum
@jax.jit def fast_jit_slice(rows, cols, m): # 生成目标列的掩码,将布尔值转为整数权重 col_mask = (cols == m).astype(jnp.int32) # 按行索引聚合权重,统计每行在目标列中True值的数量 row_counts = jax.ops.segment_sum(col_mask, rows, num_segments=N) # 数量大于0即表示该行在目标列有True值 return row_counts > 0
方法二:使用jnp.bincount
@jax.jit def fast_jit_slice_v2(rows, cols, m): # 生成目标列的权重掩码 weights = (cols == m).astype(jnp.int32) # 统计每行的权重总和,未出现的行自动填充0 counts = jnp.bincount(rows, weights=weights, minlength=N) return counts > 0
优势说明
- 这两个算子都是JAX针对稀疏索引场景优化的原生实现,内部利用硬件加速,避免了
jnp.isin配合where时的冗余遍历,性能接近甚至超过非JIT版本。 - 输出形状固定为
(N,),完全符合JAX JIT对静态形状的要求,不会触发数据依赖形状的错误。
内容的提问来源于stack exchange,提问作者Shuhei Iitsuka
相关产品推荐
相关产品推荐

