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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 19:30:34