对含奇偶数值的2D Tensor执行不同算术运算的实现疑问
对PyTorch张量按奇偶性执行不同运算的实现方案
你可以通过两种简洁的方式实现需求:掩码索引赋值或使用torch.where()函数,以下是具体代码和说明:
方法一:掩码索引直接赋值
这种方式直观,通过布尔掩码定位奇偶元素后分别计算:
import torch list1 = [ [10, 25, 75, 10, 50], [25, 30, 35, 40, 30], [45, 50, 55, 60, 20], [50, 20, 15, 20, 10], [10, 25, 40, 50, 35]] tensor2 = torch.tensor(list1) # 生成奇偶掩码:even_mask为True的位置是偶数,odd_mask取反得到奇数位置 even_mask = (tensor2 % 2) == 0 odd_mask = ~even_mask # 创建结果张量,若需保留整数类型可去掉dtype参数,用整数除法// result = torch.empty_like(tensor2, dtype=torch.float32) # 对偶数执行x/2,奇数执行3x+1 result[even_mask] = tensor2[even_mask] / 2 result[odd_mask] = 3 * tensor2[odd_mask] + 1 print(result)
方法二:使用torch.where()函数(更简洁)
torch.where()可以直接根据布尔条件选择对应运算,一行代码完成分支逻辑:
import torch list1 = [ [10, 25, 75, 10, 50], [25, 30, 35, 40, 30], [45, 50, 55, 60, 20], [50, 20, 15, 20, 10], [10, 25, 40, 50, 35]] tensor2 = torch.tensor(list1) even_mask = (tensor2 % 2) == 0 # 语法:torch.where(条件, 满足条件的运算, 不满足条件的运算) result = torch.where(even_mask, tensor2 / 2, 3 * tensor2 + 1) # 若需要整数结果,改用整数除法// # result = torch.where(even_mask, tensor2 // 2, 3 * tensor2 + 1) print(result)
补充说明
- 如果运算结果都是整数(比如偶数除以2、奇数3x+1后均为整数),可以使用整数除法
//并保持张量的整数类型,避免浮点转换。 - 两种方法都能高效处理张量运算,
torch.where()更适合简洁的分支逻辑,掩码索引则更灵活,适合复杂的多分支场景。
内容的提问来源于stack exchange,提问作者Jamie M
相关产品推荐
相关产品推荐

