如何统计一维数组与三维数组指定[y][x][所有元素]的公共元素数量
统计一维数组与三维数组指定维度的公共元素数量
需求背景
已实现统计一维数组与整个三维数组公共元素数量的代码,现需统计一维数组arr1=[1,2,3]与三维数组arr2指定[y][x][所有元素]的公共元素数量(预期输出3),要求不使用循环,通过NumPy函数完成一维数组与三维数组最后维度元素的对比统计。
原实现代码(统计整个三维数组)
import numpy as np arr1 = np.array([1, 2, 3]) arr2 = np.array([[[1, 2, 3, 4, 5, 6]], [[7, 8, 9, 10, 11, 12]], [[3, 4, 5, 6, 7, 8]]]) flatarray = arr2.flatten() common_elements = np.intersect1d(arr1, flatarray) a = len(common_elements) print(a)
针对指定位置的解决方案
直接提取三维数组指定[y][x]对应的最后一维子数组,再计算与arr1的交集元素数量:
import numpy as np arr1 = np.array([1, 2, 3]) arr2 = np.array([[[1, 2, 3, 4, 5, 6]], [[7, 8, 9, 10, 11, 12]], [[3, 4, 5, 6, 7, 8]]]) # 提取指定[y=0, x=0]位置的最后一维所有元素 target_subarr = arr2[0, 0, :] # 计算交集并统计数量 common_count = len(np.intersect1d(arr1, target_subarr)) print(common_count) # 输出:3
批量统计所有[y][x]位置的方案
若需一次性统计所有[y][x]对应的子数组与arr1的公共元素数量,可使用np.apply_along_axis避免显式循环:
# 沿最后一维(axis=2)对每个子数组执行统计 counts = np.apply_along_axis(lambda subarr: len(np.intersect1d(arr1, subarr)), axis=2, arr=arr2) print(counts) # 输出:[[3], [0], [1]]
关键说明
arr2[y, x, :]:快速定位三维数组中指定[y][x]的最后一维元素,得到目标一维子数组。np.intersect1d:高效计算两个一维数组的交集,返回唯一的公共元素数组,通过len()获取数量。np.apply_along_axis:沿指定轴批量处理子数组,替代手动循环,保持代码简洁且利用NumPy的向量化优势。
内容的提问来源于stack exchange,提问作者Razor
相关产品推荐
相关产品推荐

