如何找出3D数组中垂直翻转的2D数组对的索引?
找出3D数组中垂直翻转的2D数组索引对
我们有一个shape为(4, 4, 4)的3D数组,沿axis=0维度的每个2D数组都存在一个垂直翻转的匹配项。例如示例中arr[0]与arr[3]、arr[1]与arr[2]分别构成翻转对:
import numpy as np arr = np.array([[[1, 2, 3, 4], [3, 3, 3, 3], [4, 6, 0, 9], [4, 3, 2, 1]], [[5, 6, 5, 6], [8, 8, 8, 8], [0, 9, 0, 9], [1, 2, 4, 8]], [[1, 2, 4, 8], [0, 9, 0, 9], [8, 8, 8, 8], [5, 6, 5, 6]], [[4, 3, 2, 1], [4, 6, 0, 9], [3, 3, 3, 3], [1, 2, 3, 4]]])
需要找出这些翻转2D数组对的索引,示例预期输出为:
matched_pairs = [[0, 3], [1, 2]]
实际场景中,3D数组的shape通常为(20000, 8, 8),此时输出的matched_pairs shape为(10000, 2)。
解决方案
核心思路
- 对每个2D数组执行垂直翻转操作(使用
np.flipud) - 为每个原数组和翻转后的数组生成唯一可哈希标识,建立翻转数组到原索引的映射
- 遍历原数组索引,跳过已配对的项,找到对应翻转数组的索引,组成配对结果
高效实现代码
针对大数组(如20000个8x8数组),使用字节数组哈希来提升效率:
import numpy as np def find_flip_pairs(arr): # 对所有2D数组执行垂直翻转 flipped_arr = np.flipud(arr) # 将每个2D数组展平为一维数组,方便生成哈希标识 arr_flat = arr.reshape(arr.shape[0], -1) flipped_flat = flipped_arr.reshape(flipped_arr.shape[0], -1) # 建立翻转数组哈希值到原索引的映射 flip_hash_map = {} for idx, subarr in enumerate(flipped_flat): # 转换为字节串后哈希,确保唯一性 hash_val = hash(subarr.tobytes()) flip_hash_map[hash_val] = idx matched_pairs = [] used_indices = set() for idx in range(arr.shape[0]): if idx not in used_indices: current_hash = hash(arr_flat[idx].tobytes()) pair_idx = flip_hash_map[current_hash] matched_pairs.append([idx, pair_idx]) # 标记两个索引为已使用,避免重复配对 used_indices.add(idx) used_indices.add(pair_idx) return np.array(matched_pairs) # 测试示例 arr = np.array([[[1, 2, 3, 4], [3, 3, 3, 3], [4, 6, 0, 9], [4, 3, 2, 1]], [[5, 6, 5, 6], [8, 8, 8, 8], [0, 9, 0, 9], [1, 2, 4, 8]], [[1, 2, 4, 8], [0, 9, 0, 9], [8, 8, 8, 8], [5, 6, 5, 6]], [[4, 3, 2, 1], [4, 6, 0, 9], [3, 3, 3, 3], [1, 2, 3, 4]]]) print(find_flip_pairs(arr))
代码说明
np.flipud:快速完成垂直翻转操作,时间复杂度为O(n*m)(n为2D数组行数,m为列数)- 字节数组哈希:相比将数组转为tuple,字节串哈希的处理速度更快,适合处理大规模数据
- 已使用索引集合:确保每个索引只被配对一次,避免重复结果
内容的提问来源于stack exchange,提问作者user109387
相关产品推荐
相关产品推荐

