如何高效统计3D NumPy数组中各2D子数组的唯一行数量
高性能无循环实现方案
针对形状为(N, M, 2)的3D数组,以下方案完全基于NumPy矢量化操作,无Python级循环,可高效处理百万级规模输入:
实现代码
import numpy as np def count_unique_rows_per_block(a): # 1. 将所有行转为唯一整数编码 rows = a.reshape(-1, 2) _, row_inverse = np.unique(rows, axis=0, return_inverse=True) row_inverse = row_inverse.reshape(a.shape[0], a.shape[1]) # 2. 生成带块标识的唯一ID,避免不同块的相同行被误判为同一值 max_row_id = row_inverse.max() + 1 block_uids = row_inverse + np.arange(a.shape[0])[:, None] * max_row_id # 3. 排序后统计每个块的唯一值数量 block_ids = np.repeat(np.arange(a.shape[0]), a.shape[1]) sorted_idx = np.lexsort((block_uids.ravel(), block_ids)) sorted_uids = block_uids.ravel()[sorted_idx] sorted_block_ids = block_ids[sorted_idx] # 识别唯一值的变化点 change_points = np.r_[True, sorted_uids[1:] != sorted_uids[:-1]] # 统计每个块的唯一值数量 return np.bincount(sorted_block_ids[change_points], minlength=a.shape[0])
效果验证
用提供的示例数组测试:
a = np.array([[[1, 2], [1, 2], [2, 3]], [[2, 3], [2, 3], [3, 4]], [[1, 2], [1, 2], [1, 2]]]) print(count_unique_rows_per_block(a))
输出结果:
array([2, 2, 1])
完全符合预期。
性能说明
对于规模为(1000000, 256, 2)的数组,该方案仅需一次全局np.unique调用、一次排序和少量简单数组操作,时间复杂度为O(NM log(NM)),远快于循环遍历每个块的实现。
注:如果数组元素为浮点数,需注意
np.unique的浮点数精度匹配问题,可根据业务场景先做精度截断再处理。
内容的提问来源于stack exchange,提问作者Eric Truett
相关产品推荐
相关产品推荐

