PyTorch中如何比较二维张量与一维张量并获取整行匹配结果?
实现二维张量与一维张量的逐行完全匹配比较
要得到预期的逐行匹配结果,核心是在逐元素比较后,对每行的结果进行聚合,判断整行是否与目标一维张量完全匹配。
方法一:使用torch.all()(推荐,高效向量化操作)
通过广播完成逐元素比较后,调用all(dim=1)沿着列维度(dim=1)检查每行所有元素是否均为True,直接生成每行的匹配结果:
import torch tensor_2d = torch.tensor([[1,2],[1,2],[1,3]]) tensor_1d = torch.tensor([1,2]) # 逐元素比较后聚合每行结果 match_result = (tensor_2d == tensor_1d).all(dim=1) print(match_result)
输出结果:
tensor([True, True, False])
方法二:使用torch.equal()遍历行(适合小张量场景)
如果张量规模较小,也可以遍历二维张量的每一行,用torch.equal()直接比较单行与目标一维张量是否完全相等:
import torch tensor_2d = torch.tensor([[1,2],[1,2],[1,3]]) tensor_1d = torch.tensor([1,2]) match_result = torch.tensor([torch.equal(row, tensor_1d) for row in tensor_2d]) print(match_result)
输出结果同样符合预期。
内容的提问来源于stack exchange,提问作者JobHunter69
相关产品推荐
相关产品推荐

