如何将Numpy数组中的配对重组成相邻元素匹配的子数组?
Numpy数组连通配对合并问题
输入与期望输出
输入数组:
import numpy as np a = np.array([[1,2],[2,4],[5,6],[3,4],[7,8],[3,5]])
期望输出:
[np.array([1,2,2,4,4,3,3,5,5,6]), np.array([7,8])]
规则:
- 子数组及内部元素顺序无强制要求
- 二元组可翻转后纳入结果
- 已知每个二元组的两个元素互不相同
尝试的代码及错误输出
以下代码无法得到正确结果:
import numpy as np def concatenate_neighbouring_pairs(input): output = [] for i in range(len(input)): subarray = input[i] for j in range(1,len(input)): intersect1 = np.in1d(subarray, input[j]) intersect2 = np.in1d(input[j] ,subarray) if intersect1[0] == True and intersect1[-1] == False and intersect2[0] == True and intersect2[-1] == False: subarray = np.concatenate((np.flip(input[j]),subarray)) elif intersect1[0] == True and intersect1[-1] == False and intersect2[0] == False and intersect2[-1] == True: subarray = np.concatenate((input[j], subarray)) elif intersect1[0] == False and intersect1[-1] == True and intersect2[0] == True and intersect2[-1] == False: subarray = np.concatenate((subarray, input[j])) elif intersect1[0] == False and intersect1[-1] == True and intersect2[0] == False and intersect2[-1] == True: subarray = np.concatenate((subarray, np.flip(input[j]))) output.append(subarray) return output
调用后错误输出:
[array([1, 2, 2, 4, 4, 3, 3, 5]), array([2, 4, 4, 3, 3, 5]), array([3, 5, 5, 6]), array([5, 3, 3, 4, 4, 2]), array([7, 8]), array([4, 3, 3, 5, 5, 6])]
解决方案
这个问题本质是找连通分量:把每个二元组看作节点间的边,合并所有连通的边(允许翻转),再将连通分量内的边按链式拼接。
代码实现
import numpy as np def merge_connected_pairs(arr): # 并查集:找出所有连通的节点分组 parent = {} def find(u): if parent[u] != u: parent[u] = find(parent[u]) return parent[u] def union(u, v): u_root = find(u) v_root = find(v) if u_root != v_root: parent[v_root] = u_root # 初始化并查集 for pair in arr: u, v = pair parent.setdefault(u, u) parent.setdefault(v, v) union(u, v) # 按连通分量分组边 components = {} for pair in arr: root = find(pair[0]) components.setdefault(root, []).append(pair.tolist()) # 对每个分量拼接成链式数组 result = [] for comp in components.values(): merged = [] used = set() # 从第一条边开始 current = comp[0] merged.extend(current) used.add(tuple(current)) used.add(tuple(reversed(current))) # 依次拼接后续边 while len(used) < len(comp)*2: last_node = merged[-1] # 找未使用且包含last_node的边 for pair in comp: tp = tuple(pair) if tp not in used: if pair[0] == last_node: merged.append(pair[1]) used.add(tp) used.add(tuple(reversed(pair))) break elif pair[1] == last_node: merged.append(pair[0]) used.add(tp) used.add(tuple(reversed(pair))) break result.append(np.array(merged)) return result # 测试 a = np.array([[1,2],[2,4],[5,6],[3,4],[7,8],[3,5]]) print(merge_connected_pairs(a))
输出结果
[array([1, 2, 2, 4, 4, 3, 3, 5, 5, 6]), array([7, 8])]
说明
- 并查集分组:通过并查集将所有连通的节点归为一组,确定哪些边属于同一分量;
- 链式拼接:从分量内任意边开始,找到与当前末尾节点相连的边(可翻转),依次拼接,直到所有边都被使用;
- 孤立边直接作为单独子数组保留。
内容的提问来源于stack exchange,提问作者jercai
相关产品推荐
相关产品推荐

