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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 03:18:14