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

Jax中如何对BCOO稀疏数组按指定索引取值并返回稀疏数组?

JAX稀疏BCOO数组按高级索引切片的解决方法

JAX的BCOO稀疏数组目前不支持直接使用sparse_array[jnp.arange(3), indices]这类高级索引语法,手动构造新的BCOO数组是最直接的解决方案——核心思路是筛选原稀疏数组中符合切片条件的非零元素,再重新组装成目标稀疏数组。

实现代码

import jax.numpy as jnp
from jax.experimental import sparse

indices = jnp.array([1,1,0])

full_array = jnp.array(
        [
            [
                [0,0,0],
                [2,2,2],
                [0,0,0],
                [0,0,0]
            ],
            [
                [1,1,1],
                [0,0,0],
                [0,0,0],
                [0,0,0]
            ],
            [
                [1,1,1],
                [0,0,0],
                [0,0,0],
                [0,0,0]
            ]
        ]
)
sparse_array = sparse.bcoo_fromdense(full_array)

# 提取原稀疏数组的非零元素索引与值
orig_indices = sparse_array.indices
orig_values = sparse_array.values

# 筛选符合切片条件的元素:原索引的第二维度等于对应第一维度的目标索引
target_j = indices[orig_indices[:, 0]]
mask = orig_indices[:, 1] == target_j

# 构造新稀疏数组的索引(保留原索引的第一、第三维度)和值
new_indices = orig_indices[mask][:, [0, 2]]
new_values = orig_values[mask]

# 创建目标BCOO稀疏数组
result = sparse.BCOO((new_values, new_indices), shape=(3, 3))

# 验证结果
print(result.todense())
# 输出:
# [[2 2 2]
#  [0 0 0]
#  [1 1 1]]

原理说明

  1. 提取原数据:BCOO数组通过indices和values存储所有非零元素的位置与数值,先取出这两部分核心数据。
  2. 筛选匹配元素:针对每个非零元素的第一维度索引i,检查其第二维度索引是否等于indices[i](即切片目标位置),生成匹配掩码。
  3. 重构稀疏数组:用筛选后的非零元素构造新的indices和values,指定目标形状后生成新的BCOO数组。

这种方法避免了将稀疏数组转为稠密数组的性能损耗,适合处理大规模稀疏数据。

内容的提问来源于stack exchange,提问作者Simon P.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 04:01:13