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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 12:36:24