如何在PyTorch中不转密集识别稀疏二进制矩阵的重复行?
识别稀疏二进制矩阵中的重复行(PyTorch实现)
问题背景
给定形状为n×m的二进制PyTorch稀疏张量A,需识别其中的重复行(即存在另一行与该行在所有位置的元素完全相同)。严格禁止将稀疏张量转换为密集表示(因矩阵规模极大,内存受限)。
原思路问题分析
此前尝试通过计算矩阵点积得到行与行的1的交集数量,结合列和判断重复,但逻辑错误:误将列和(A.sum(0))当作行和使用,导致条件判断失效。正确的判断逻辑需基于行的1的总数与行之间1的交集数量。
解决方案(纯PyTorch稀疏操作)
核心逻辑
两行完全等价的充要条件:
- 两行的1的总数(行和)相等;
- 两行共有的1的位置数量等于行和(说明1的位置完全一致);
- 两行不是同一行(排除自匹配)。
代码实现
import torch # 生成示例稀疏二进制矩阵 A = torch.randint(0, 2, (10, 100), dtype=torch.float32).to_sparse() # 1. 计算每行的1的数量(行和),转密集张量仅需存储n个元素,内存压力可忽略 row_sums = A.sum(dim=1).to_dense() # 2. 计算行与行的1的交集数量:稀疏矩阵点积A@A.T,结果为n×n稀疏张量 row_intersection = A @ A.T # 3. 提取点积结果的索引与值,判断重复条件 indices = row_intersection.indices() i, j = indices[0], indices[1] intersection_counts = row_intersection.values() # 组合三个判断条件 cond_same_sum = row_sums[i] == row_sums[j] cond_full_overlap = intersection_counts == row_sums[i] cond_not_self = i != j is_duplicate = torch.logical_and(torch.logical_and(cond_same_sum, cond_full_overlap), cond_not_self) # 筛选重复行对,并提取所有存在重复的行索引(去重) duplicate_pairs = indices[:, is_duplicate] duplicate_row_indices = torch.unique(torch.cat([duplicate_pairs[0], duplicate_pairs[1]])) print("存在重复的行索引:", duplicate_row_indices)
关键说明
- 行和计算:
A.sum(dim=1).to_dense()仅生成n个元素的张量,远小于m维度的内存占用,符合内存限制要求; - 稀疏点积:
A@A.T利用PyTorch稀疏矩阵乘法优化,仅存储非零的交集结果,避免密集矩阵的内存爆炸; - 条件判断:通过张量广播实现批量判断,无Python循环,保证性能。
内容的提问来源于stack exchange,提问作者daqh
相关产品推荐
相关产品推荐

