如何用PyTorch高效生成掩码张量:判断元素是否存在于其他张量
高效实现PyTorch元素存在性掩码
解决方案
可以直接使用PyTorch内置的torch.isin()函数高效完成需求,该函数是向量化实现,性能远优于手动循环判断,适合处理大规模张量。具体步骤如下:
- 将输入列表转换为PyTorch张量;
- 合并张量
b和c,得到包含所有目标元素的张量; - 调用
torch.isin()判断a中每个元素是否存在于合并后的张量中,直接生成掩码张量。
代码示例
import torch # 定义输入张量 a = torch.tensor([1, 234, 54, 6543, 55, 776]) b = torch.tensor([234, 54]) c = torch.tensor([55, 776]) # 合并b和c target_elements = torch.cat([b, c]) # 生成掩码张量 a_masked = torch.isin(a, target_elements) print(a_masked) # 输出:tensor([False, True, True, False, True, True])
补充说明
torch.isin(input, other)会逐个检查input中的元素是否在other中存在,返回与input同形状的布尔张量;- 也可以通过
torch.isin(a, b) | torch.isin(a, c)实现相同效果,两种方式性能差异可忽略,按需选择即可; - 该方法支持任意维度的张量,内部基于高效向量化运算实现,无需手动编写循环逻辑。
内容的提问来源于stack exchange,提问作者Ofek Glick
相关产品推荐
相关产品推荐

