Python中统计列表内大型稀疏数组重复数量的技术问询
解决方案
问题分析
你的核心需求是统计60个大型稀疏2D数组中完全相同的数组总数量,当前代码存在两个关键问题:
- 每次循环重置
count,最终仅输出最后一个数组的重复数,无法得到全局统计结果 - 直接使用
(arr1 == arr2).all()对比大型数组,内存占用极高且效率极低
从你的测试数据来看,数组属于稀疏结构(大部分元素为0),稀疏矩阵是最优处理方案,以下是具体实现:
方案1:基于稀疏矩阵的高效统计
利用稀疏矩阵仅存储非零元素的特性,我们可以通过提取非零元素的位置+值生成唯一哈希标识,再统计每个标识的出现次数,避免冗余的两两对比。
实现代码
import numpy as np import scipy.sparse as sp from collections import defaultdict # 生成测试数据(模拟真实场景) a = np.zeros((6,6)) a[1,2] = 1 a[2,5] = 1 a[3,2] = 1 a[4,1] = 1 b = np.zeros((6,6)) b[1,2] = 1 b[2,5] = 1 b[3,2] = 1 b[4,1] = 1 c = np.zeros((6,6)) c[1,3] = 1 c[2,5] = 1 d = np.zeros((6,6)) d[1,3] = 1 d[2,4] = 1 # 转换为CSR格式稀疏矩阵(适合快速访问非零元素) list_sparse = [sp.csr_matrix(arr) for arr in [a,b,c,d]] # 将稀疏矩阵转换为可哈希的唯一标识 def sparse_to_hash(sp_mat): non_zero = [] # 遍历所有非零元素的行、列索引和值 for row in range(sp_mat.shape[0]): start = sp_mat.indptr[row] end = sp_mat.indptr[row+1] cols = sp_mat.indices[start:end] vals = sp_mat.data[start:end] for col, val in zip(cols, vals): non_zero.append((row, col, val)) # 排序后转为元组,确保相同矩阵的哈希一致 return tuple(sorted(non_zero)) # 统计每个哈希对应的数组索引 hash_map = defaultdict(list) for idx, mat in enumerate(list_sparse): hash_key = sparse_to_hash(mat) hash_map[hash_key].append(idx) # 计算重复数组总数量并输出结果 total_duplicates = 0 print("重复数组组:") for indices in hash_map.values(): if len(indices) > 1: print(f"数组索引{indices}完全相同") total_duplicates += len(indices) # 统计所有重复的数组个数 # 若需统计重复对数(如a和b算1对),则用 len(indices)*(len(indices)-1)//2 print(f"\n完全相同的数组总数量:{total_duplicates}")
关键说明
- 哈希生成逻辑:提取所有非零元素的
(行,列,值)三元组并排序,确保相同稀疏矩阵生成一致哈希,不受存储顺序影响 - 效率优化:时间复杂度为O(n*m)(n为数组数量,m为非零元素数),远优于两两对比的O(n²)
- 内存优化:稀疏矩阵仅存储非零元素,对于(30000,30000)的大型数组,内存占用仅为原numpy数组的几十分之一
方案2:普通numpy数组的优化方案(非稀疏场景)
如果数组并非稀疏结构,可通过分块哈希降低内存压力:
import numpy as np from collections import defaultdict def array_to_hash(arr, block_size=1000): # 将大数组分块计算哈希,拼接成总标识 hashes = [] for i in range(0, arr.shape[0], block_size): for j in range(0, arr.shape[1], block_size): block = arr[i:i+block_size, j:j+block_size] hashes.append(hash(block.tobytes())) return tuple(hashes) # 统计逻辑同方案1 hash_map = defaultdict(list) for idx, arr in enumerate(list2d): hash_key = array_to_hash(arr) hash_map[hash_key].append(idx) total_duplicates = 0 print("重复数组组:") for indices in hash_map.values(): if len(indices) > 1: print(f"数组索引{indices}完全相同") total_duplicates += len(indices) print(f"\n完全相同的数组总数量:{total_duplicates}")
原代码的修正(仅作参考,不适合大型数组)
若仅需修正原代码的统计逻辑,而非优化效率,可调整如下:
import numpy as np list2d = [a,b,c,d] n = len(list2d) counted_pairs = set() total_count = 0 # 仅对比i<j的组合,避免重复统计 for i in range(n): for j in range(i+1, n): if (list2d[i] == list2d[j]).all(): counted_pairs.add((i,j)) total_count += 2 # 每个重复对包含2个相同数组 print(f'list2d[{i}] 与 list2d[{j}] 完全相同') print(f'完全相同的数组总数量:{total_count}')
内容的提问来源于stack exchange,提问作者Y. Pat
相关产品推荐
相关产品推荐

