如何计算维度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
相关产品推荐
相关产品推荐

