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

如何用纯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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 21:15:37