Numpy形状不匹配时用4D布尔矩阵对3D矩阵计算条件均值
实现步骤
核心是先对齐矩阵维度逻辑,过滤布尔矩阵的无效值,再按连通关系筛选有效元素计算平均值:
- 首先调整
A的维度顺序,和E的轴逻辑对齐:A原始形状为(4,3,3),E是对A.T做差生成的,因此先将A转置为形状(3,3,4)的A_t,轴顺序对应(列, 行, A的axis0索引)。 - 过滤
E的假阳性值:E中所有第三个轴索引小于第四个轴索引(即k<l)的位置,是np.tril将上三角强制置0产生的无效结果,需要构造掩码只保留k>=l的有效差判断。 - 对每个输出位置
(i,j),用并查集合并所有差小于等于阈值的索引对,得到axis0维度上的连通分量。 - 仅保留大小≥2的连通分量中的元素计算平均值,孤立的单个元素不参与计算;如果没有符合要求的连通分量,对应位置返回
np.nan。 - 最后将结果转置回和
A后两个维度一致的(3,3)顺序即可。
实现代码
import numpy as np # 并查集实现,用于合并连通的索引 class DSU: def __init__(self, size): self.parent = list(range(size)) def find(self, x): if self.parent[x] != x: self.parent[x] = self.find(self.parent[x]) return self.parent[x] def union(self, x, y): xr, yr = self.find(x), self.find(y) if xr != yr: self.parent[yr] = xr threshold = 2 A = np.array([ [ [90, 84, 88], [10, 30, 17], [7, 0, 4] ], [ [88, 83, 102], [12, 14, 15], [12, 17, 7]], [ [94, 14, 85], [8, 23, 20], [25, 5,27]], [ [150, 90, 103], [9, 16, 21], [17, 7, 12] ] ]) # 按提供的逻辑生成E B = A.T[..., None, :] - A.T[...,None] C = np.tril(B, k=0) C = np.absolute(C) D = np.delete(C, 3, axis=3) E = np.where(D <= threshold, True, False) A_t = A.T # 转置后形状为(3,3,4),和E前两个轴对齐 # 生成有效对掩码:仅保留k>=l位置的布尔值,过滤tril置0产生的假阳性 k_grid = np.arange(4)[:, None] l_grid = np.arange(3)[None, :] valid_pair_mask = E & (k_grid >= l_grid) result = np.full((3, 3), np.nan) for i in range(3): for j in range(3): dsu = DSU(4) # 合并所有差满足阈值的索引对(k>l,对角线无需合并) for k in range(4): for l in range(min(k, 3)): if valid_pair_mask[i, j, k, l]: dsu.union(k, l) # 统计每个连通分量的元素 components = {} for idx in range(4): root = dsu.find(idx) if root not in components: components[root] = [] components[root].append(A_t[i, j, idx]) # 收集大小≥2的分量元素计算平均 valid_elements = [] for elems in components.values(): if len(elems) >= 2: valid_elements.extend(elems) if valid_elements: result[i, j] = np.mean(valid_elements) # 转置回和A后两个维度匹配的顺序 result = result.T print(result)
内容的提问来源于stack exchange,提问作者equanimity
相关产品推荐
相关产品推荐

