PyTorch如何判断第一个二维张量的每行是否存在于第二个张量中?
PyTorch实现张量行级存在性判断
给定两个张量:
t1 = torch.tensor([[1,2],[3,4],[5,6]]) t2 = torch.tensor([[1,2],[5,6]])
需要判断t1的每一行是否存在于t2中,返回布尔结果[True, False, True]。以下是两种可行实现方案:
方法一:广播+全匹配判断
这是通用程度最高的方案,不依赖元素类型和数值范围:
import torch t1 = torch.tensor([[1,2],[3,4],[5,6]]) t2 = torch.tensor([[1,2],[5,6]]) # 扩展维度实现逐行元素比较 row_element_matches = (t1.unsqueeze(1) == t2.unsqueeze(0)).all(dim=-1) # 判断t1每行是否在t2中有完全匹配的行 result = row_element_matches.any(dim=-1) print(result) # 输出: tensor([ True, False, True])
原理说明:
t1.unsqueeze(1)将t1形状转为(3,1,2),t2.unsqueeze(0)将t2形状转为(1,2,2),通过广播实现t1每行与t2每行的逐元素比较all(dim=-1)检查每行的所有元素是否完全匹配,得到t1每行与t2每行的匹配情况张量any(dim=-1)确认t1每行是否在t2中存在至少一个匹配行,输出最终结果
方法二:行编码+元素级存在判断
如果张量元素为整数且数值范围较小,可将每行编码为单个整数,再用torch.isin判断:
import torch t1 = torch.tensor([[1,2],[3,4],[5,6]]) t2 = torch.tensor([[1,2],[5,6]]) def encode_rows(tensor): # 计算编码权重,避免不同行出现编码冲突 base = tensor.max() + 1 weights = base ** torch.arange(tensor.size(1), device=tensor.device) return tensor @ weights # 对两行张量进行编码 encoded_t1 = encode_rows(t1) encoded_t2 = encode_rows(t2) # 用torch.isin判断编码后的行是否存在 result = torch.isin(encoded_t1, encoded_t2) print(result) # 输出: tensor([ True, False, True])
注意:该方法需确保编码后不会出现整数溢出,仅适合小规模整数张量场景。
内容的提问来源于stack exchange,提问作者xc-2021
相关产品推荐
相关产品推荐

