PyTorch如何构建按列标记张量最大值位置的布尔索引掩码
PyTorch 构建逐列最大值位置布尔掩码方案
不需要依赖冗余的自定义循环或复杂索引操作,直接用PyTorch原生向量化算子即可实现,根据需求可以选择两种实现逻辑:
方案1:标记所有并列最大值位置
如果需要把列内所有等于最大值的位置都标记为True,直接逐列计算最大值后做相等比较即可,代码最简洁:
import torch as T x = T.tensor([[0, 3, 0, 5, 9, 8, 2, 0], [0, 4, 9, 6, 7, 9, 1, 0]]) # 计算逐列最大值,相等判断直接生成同形状布尔掩码 col_max = x.max(dim=0).values mask = x == col_max
该方案输出会把最后一列两个相等的0都标记为最大值位置:
tensor([[ True, False, False, False, True, False, True, True], [False, True, True, True, False, True, False, True]])
方案2:仅标记首个最大值位置(匹配示例输出)
如果需要和torch.argmax行为一致,每列仅标记第一个出现的最大值位置(并列值不重复标记,和给出的期望输出完全匹配),可以用索引赋值的方式构建掩码:
import torch as T x = T.tensor([[0, 3, 0, 5, 9, 8, 2, 0], [0, 4, 9, 6, 7, 9, 1, 0]]) # 生成列坐标、取每列首个最大值的行坐标 col_indices = T.arange(x.shape[1]) row_indices = x.argmax(dim=0) # 初始化全False布尔掩码,对应位置赋值为True mask = T.zeros_like(x, dtype=bool) mask[row_indices, col_indices] = True
运行后输出和期望结果完全一致:
tensor([[ True, False, False, False, True, False, True, True], [False, True, True, True, False, True, False, False]])
两种方案均为纯张量向量化实现,没有Python层循环开销,执行效率和PyTorch内置算子持平,生成的布尔掩码可以直接用于同形状张量的索引操作。
内容的提问来源于stack exchange,提问作者JVGD
相关产品推荐
相关产品推荐

