如何将PyTorch布尔张量转换为指定整数映射的张量?
布尔张量转自定义整数张量的PyTorch实现
给定如下PyTorch布尔张量:
>>> t = torch.tensor([0, 1, 0, 1], dtype=torch.bool) tensor([False, True, False, True])
要将其转换为True对应整数a、False对应整数b的张量,这里提供两种简单高效的实现方式:
方法一:算术运算转换
利用布尔张量的数值特性(True对应1,False对应0),通过线性变换直接得到目标结果:
import torch t = torch.tensor([0, 1, 0, 1], dtype=torch.bool) a = 5 # 替换成你需要的整数 b = 3 # 替换成你需要的整数 result = t.int() * (a - b) + b print(result) # 输出:tensor([3, 5, 3, 5])
原理:先把布尔张量转为整数类型,True变为1、False变为0;用a - b作为系数,True位置会得到a - b,False位置为0;最后加上b,True位置就得到a,False位置保留b,完全符合需求。
方法二:使用torch.where(更直观)
torch.where可以根据布尔条件直接选择对应值,代码可读性更强:
result = torch.where(t, a, b) print(result) # 输出:tensor([3, 5, 3, 5])
这个方法逻辑直白:遍历布尔张量,遇到True就取a,遇到False就取b,非常容易理解。
补充说明
如果a或b不是标量,而是和t同形状的张量,两种方法依然适用——算术运算会自动广播,torch.where也能对应匹配位置取值。
内容的提问来源于stack exchange,提问作者Gijs Groote
相关产品推荐
相关产品推荐

