高效检查PyTorch张量中每列是否存在对应反转列
高效检查(2,ncol)张量列反转匹配的实现
核心思路
- 快速过滤:若
ncol为奇数,直接返回False——每对反转列需要占用两个位置,奇数列数无法完全配对。 - 列归一化:将每一列
[a,b]转换为有序元组(min(a,b), max(a,b)),让[a,b]和[b,a]映射为同一个键,消除反转差异。 - 频次校验:统计每个有序元组的出现次数,所有频次必须为偶数(包括
a==b的自反列,这类列也需要成对出现)。
高效实现代码
import torch def has_all_reversed_columns(tensor): # 第一步:快速排除奇数列情况 ncol = tensor.size(1) if ncol % 2 != 0: return False # 第二步:将每列转换为(min, max)的有序形式,向量化操作避免循环 col_mins = torch.min(tensor, dim=0).values col_maxes = torch.max(tensor, dim=0).values ordered_cols = torch.stack([col_mins, col_maxes], dim=1) # 第三步:统计唯一有序列的出现次数,用PyTorch底层优化的unique函数 _, counts = torch.unique(ordered_cols, dim=0, return_counts=True) # 第四步:检查所有计数是否为偶数 return torch.all(counts % 2 == 0)
代码优势说明
- 时间复杂度:整体为O(ncol),所有操作均为PyTorch向量化实现,避免Python循环,适配1e5量级的列数场景。
- 内存效率:无需额外存储所有列的元组,直接通过张量操作完成归一化和统计。
- 鲁棒性:自动处理
a==b的自反列,这类列自身就是反转列,必须成对出现才能满足条件。
测试示例
# 测试1:奇数列数,直接返回False t1 = torch.tensor([[1, 2, 3, 7, 8], [3, 3, 1, 8, 7]], dtype=torch.long) print(has_all_reversed_columns(t1)) # 输出: False # 测试2:所有列都有对应反转,返回True t2 = torch.tensor([[1, 2, 3, 7, 8, 4], [3, 3, 1, 8, 7, 2]], dtype=torch.long) print(has_all_reversed_columns(t2)) # 输出: True # 测试3:存在无法配对的列,返回False t3 = torch.tensor([[1,2],[3,4]], dtype=torch.long) print(has_all_reversed_columns(t3)) # 输出: False
内容的提问来源于stack exchange,提问作者DeltaIV
相关产品推荐
相关产品推荐

