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

如何计算维度0相同、维度1不同的两个2D Tensor逐行交集?

计算PyTorch中两个2D Tensor对应行的交集

现有两个2D Tensor,维度0的尺寸相同(均为8),维度1的尺寸分别为2和32,需要计算它们对应行的交集。示例如下:

输入Tensor:

import torch

t1 = torch.tensor([[1,2,3], [3,4,5], [4,5,6]])
t2 = torch.tensor([[1],[3],[9]])

期望得到的结果:

t3 = [[1],[3],[]]

解决方案

可以利用PyTorch的广播机制结合布尔索引实现,步骤如下:

  • 扩展两个Tensor的维度,使其能够进行逐元素的广播比较
  • 找出每行中存在匹配关系的元素
  • 收集每行的交集元素,自动处理空交集的情况

实现代码(返回列表格式)

import torch

def row_intersection(t1, t2):
    # 扩展维度:t1变为[N, C1, 1],t2变为[N, 1, C2],支持广播比较
    t1_exp = t1.unsqueeze(2)
    t2_exp = t2.unsqueeze(1)
    # 生成匹配矩阵,标记t1元素是否在t2的对应行中存在
    matches = (t1_exp == t2_exp)
    # 对每行,获取t1中存在匹配的元素掩码
    row_masks = matches.any(dim=2)
    # 逐行收集交集元素
    result = []
    for i in range(t1.size(0)):
        inter_elements = t1[i][row_masks[i]].tolist()
        result.append(inter_elements)
    return result

# 测试示例
t1 = torch.tensor([[1,2,3], [3,4,5], [4,5,6]])
t2 = torch.tensor([[1],[3],[9]])
t3 = row_intersection(t1, t2)
print(t3)  # 输出: [[1], [3], []]

实现代码(返回嵌套Tensor格式)

如果需要保持Tensor格式而非列表,可以使用PyTorch的嵌套Tensor存储变长结果:

import torch

def row_intersection_tensor(t1, t2):
    t1_exp = t1.unsqueeze(2)
    t2_exp = t2.unsqueeze(1)
    matches = (t1_exp == t2_exp)
    row_masks = matches.any(dim=2)
    # 构建嵌套Tensor存储每行的交集
    nested_result = torch.nested.nested_tensor([t1[i][row_masks[i]] for i in range(t1.size(0))])
    return nested_result

# 测试
nested_t3 = row_intersection_tensor(t1, t2)
print(nested_t3)
# 输出: nested_tensor([[1], [3], []])

内容的提问来源于stack exchange,提问作者user21339477

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 03:52:40