You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.29 23:12:35