如何用纯Numpy实现二维数组索引过滤(适配JAX JIT)
问题描述
现有二维MxN数组A,每行是一组索引,末尾用-1填充,示例如下:
import numpy as np A = np.array([ [2, 1, -1, -1, -1], [1, 4, 3, -1, -1], [3, 1, 0, -1, -1] ])
另有同维度的浮点数组B:
B = np.array([ [0.7, 0.4, 1.5, 2.0, 4.4], [0.8, 4.0, 0.3, 0.11, 0.53], [0.6, 7.4, 0.22, 0.71, 0.06] ])
需用A中的索引过滤B:每行仅保留A中有效索引(非-1)对应的B值,其余位置设为0.0,期望结果如下:
[[0.0, 0.4, 1.5, 0.0, 0.0], [0.0, 4.0, 0.0, 0.11, 0.53], [0.6, 7.4, 0.0, 0.71, 0.0]]
要求纯Numpy实现,且适配JAX的JIT编译。
纯Numpy实现方案
import numpy as np def filter_B(A, B): # 生成列索引的广播矩阵 col_indices = np.arange(B.shape[1])[np.newaxis, :] # 构建掩码:判断列索引是否在该行的有效索引列表中 mask = (A[..., np.newaxis] == col_indices) & (A != -1)[..., np.newaxis] # 压缩掩码维度,得到每行每个位置的有效性标记 valid_mask = mask.any(axis=1) # 按掩码保留B的对应值,其余置0 result = np.where(valid_mask, B, 0.0) return result # 测试示例 A = np.array([[2,1,-1,-1,-1],[1,4,3,-1,-1],[3,1,0,-1,-1]]) B = np.array([[0.7,0.4,1.5,2.0,4.4],[0.8,4.0,0.3,0.11,0.53],[0.6,7.4,0.22,0.71,0.06]]) output = filter_B(A, B) print(output)
JAX JIT适配说明
该实现完全基于向量化张量操作,没有循环、动态条件分支等JAX JIT不支持的语法。切换到JAX只需替换numpy为jax.numpy,并添加jax.jit装饰器:
import jax import jax.numpy as jnp @jax.jit def jax_filter_B(A, B): col_indices = jnp.arange(B.shape[1])[jnp.newaxis, :] mask = (A[..., jnp.newaxis] == col_indices) & (A != -1)[..., jnp.newaxis] valid_mask = mask.any(axis=1) result = jnp.where(valid_mask, B, 0.0) return result # JAX测试 jax_A = jnp.array(A) jax_B = jnp.array(B) jax_output = jax_filter_B(jax_A, jax_B) print(jax_output)
内容的提问来源于stack exchange,提问作者thesilverbail
相关产品推荐
相关产品推荐

