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

NumPy:根据掩码获取每行前N个符合条件的列

向量化实现提取NumPy数组每行mask为True的前N列

问题描述

给定形状为(m, n)的NumPy数组arr,以及同形状的布尔数组mask,编写向量化函数提取arr中每行对应mask为True的前N列,输出形状为(m, N)的数组。

示例输入输出

arr = np.array([[1,2,3,4,5],
                [6,7,8,9,10],
                [11,12,13,14,15]])

mask = np.array([[False, True, True, True, True],
                [True, False, False, True, False],
                [True, True, False, False, False]]) 

N = 2

期望输出:

output = np.array([[2,3],[6,9],[11,12]])

向量化解决方案

import numpy as np

def maskify_n_columns(arr, mask, N):
    # 按行计算每个True元素是该行的第几个True
    true_position_counts = np.cumsum(mask, axis=1)
    # 筛选出每行前N个True的位置
    valid_positions = (true_position_counts <= N) & mask
    # 提取元素并重塑为目标形状
    return arr[valid_positions].reshape(-1, N)

代码解释

  • np.cumsum(mask, axis=1):对mask按行累加,将每行中True的位置标记为它是该行的第几个True元素(True被视为1,False视为0)。比如示例第一行的累加结果为[0,1,2,3,4]。
  • (true_position_counts <= N) & mask:双重过滤——既保留累加计数≤N的位置(确保是前N个True),又通过& mask排除原本为False的位置(避免累加产生的0被误判)。
  • arr[valid_positions].reshape(-1, N):提取所有符合条件的元素,由于每行恰好有N个有效元素,直接重塑为(m, N)的二维数组。

测试验证

运行示例代码后,输出结果与预期完全一致:

>>> output = maskify_n_columns(arr, mask, N)
>>> print(output)
[[ 2  3]
 [ 6  9]
 [11 12]]

内容的提问来源于stack exchange,提问作者Vanillihoot

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 10:10:20