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]]
原理说明
- 提取原数据:BCOO数组通过
indices和values存储所有非零元素的位置与数值,先取出这两部分核心数据。 - 筛选匹配元素:针对每个非零元素的第一维度索引
i,检查其第二维度索引是否等于indices[i](即切片目标位置),生成匹配掩码。 - 重构稀疏数组:用筛选后的非零元素构造新的
indices和values,指定目标形状后生成新的BCOO数组。
这种方法避免了将稀疏数组转为稠密数组的性能损耗,适合处理大规模稀疏数据。
内容的提问来源于stack exchange,提问作者Simon P.
相关产品推荐
相关产品推荐

