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

如何在PyTorch中不转密集识别稀疏二进制矩阵的重复行?

识别稀疏二进制矩阵中的重复行(PyTorch实现)

问题背景

给定形状为n×m的二进制PyTorch稀疏张量A,需识别其中的重复行(即存在另一行与该行在所有位置的元素完全相同)。严格禁止将稀疏张量转换为密集表示(因矩阵规模极大,内存受限)。

原思路问题分析

此前尝试通过计算矩阵点积得到行与行的1的交集数量,结合列和判断重复,但逻辑错误:误将列和(A.sum(0))当作行和使用,导致条件判断失效。正确的判断逻辑需基于行的1的总数与行之间1的交集数量。

解决方案(纯PyTorch稀疏操作)

核心逻辑

两行完全等价的充要条件:

  1. 两行的1的总数(行和)相等;
  2. 两行共有的1的位置数量等于行和(说明1的位置完全一致);
  3. 两行不是同一行(排除自匹配)。

代码实现

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 03:21:13