如何使用Numpy获取数组中多个可能被替换的不同元素的索引
纯Numpy实现指定行列索引提取
问题背景
现有结构如下的Numpy二维数组,每行长度固定为10:
0, 50, 50, 2, 50, 1, 50, 99, 50, 50 50, 2, 1, 50, 50, 50, 98, 50, 50, 50 0, 50, 50, 98, 50, 1, 50, 50, 50, 50 0, 50, 50, 50, 50, 99, 50, 50, 2, 50 2, 50, 50, 0, 98, 1, 50, 50, 50, 50
数组元素遵循以下规则:
- 包含0到n的所有整数(n<50),最多缺失1个,示例中n=2
- 若存在缺失值,缺失位置优先用98填充,无98时用99填充
- 其余位置均为50
期望输出数组规则:第一行是原数组每行中0的索引,第二行是原数组每行中1的索引,第三行是原数组每行中2的索引,以此类推。示例期望输出如下:
0, 6, 0, 0, 3 5, 2, 5, 5, 5 3, 1, 3, 8, 0
纯Numpy实现代码
import numpy as np def get_target_indices(arr, n): # 构造列索引矩阵 col_idx = np.arange(arr.shape[1])[None, :].repeat(arr.shape[0], axis=0) # 过滤掉值为50的无效位置 mask = arr != 50 valid_vals = arr[mask].reshape(arr.shape[0], n+1) valid_indices = col_idx[mask].reshape(arr.shape[0], n+1) # 计算每行缺失的目标值(如果有) total_target_sum = np.sum(np.arange(n+1)) exist_sum = np.sum(np.where(valid_vals < 50, valid_vals, 0), axis=1) missing_val = total_target_sum - exist_sum # 将98/99替换为对应的缺失值 replace_mask = (valid_vals == 98) | (valid_vals == 99) valid_vals[replace_mask] = np.repeat(missing_val[:, None], n+1, axis=1)[replace_mask] # 向量化匹配获取最终结果 target_arr = np.arange(n+1)[:, None, None] match_pos = (valid_vals[None, :, :] == target_arr).argmax(axis=2) res = valid_indices[np.arange(arr.shape[0])[None, :], match_pos] return res # 测试示例 arr = np.array([ [0, 50, 50, 2, 50, 1, 50, 99, 50, 50], [50, 2, 1, 50, 50, 50, 98, 50, 50, 50], [0, 50, 50, 98, 50, 1, 50, 50, 50, 50], [0, 50, 50, 50, 50, 99, 50, 50, 2, 50], [2, 50, 50, 0, 98, 1, 50, 50, 50, 50] ]) n = 2 print(get_target_indices(arr, n))
代码说明
- 首先构造和输入数组同形状的列索引矩阵,用于后续提取对应位置的索引
- 过滤掉所有值为50的无效位置,每行剩余的n+1个元素为目标值+可能的98/99
- 通过0~n的总和减去每行现有目标值的总和,计算出缺失的目标值,将98/99替换为对应缺失值
- 利用广播机制一次性匹配所有目标值对应的位置,提取索引生成最终结果,全程无Python级别的行遍历,性能远高于for循环实现
内容的提问来源于stack exchange,提问作者416E64726577
相关产品推荐
相关产品推荐

