PyTorch:torch.argmax返回值解析与多维张量索引获取方法
问题解析与解决方法
为什么argmax()返回10701?
torch.argmax()默认会把你的4维张量展平成一维后,返回最大值所在的一维索引。这个数值是按照**行优先(C-style)**的展平规则计算出来的,我们可以手动验证:
你的张量形状是(1,8,4,576),展平顺序是从最后一个维度往第一个维度遍历,索引计算公式为:
一维索引 = dim0索引 × (dim1大小 × dim2大小 × dim3大小) + dim1索引 × (dim2大小 × dim3大小) + dim2索引 × dim3大小 + dim3索引
代入你的目标索引(0,4,2,333):
0*(8×4×576) + 4*(4×576) + 2×576 + 333 = 0 + 4×2304 + 1152 + 333 = 10701
这就是你得到这个数值的原因。
如何获取多维索引(0,4,2,333)?
有两种实用方法可以拿到完整的多维索引:
方法1:用torch.unravel_index()直接转换
这是最直接的方式,它能把一维索引还原成对应原张量形状的多维索引:
import torch x = torch.zeros(1,8,4,576) x[0,4,2,333] = 1.0 flat_idx = x.argmax() multi_idx = torch.unravel_index(flat_idx, x.shape) print(multi_idx) # 输出:(tensor(0), tensor(4), tensor(2), tensor(333))
方法2:指定argmax()的dim参数(按需使用)
如果你只需要在某个特定维度上找最大值索引,可以给argmax()传入dim参数,但它没法直接返回全维度索引。比如要在最后一个维度找:
idx = x.argmax(dim=-1) print(idx) # 输出每个前三维位置对应的最后一维最大值索引
如果要完整的4维索引,优先推荐第一种方法。
内容的提问来源于stack exchange,提问作者relot
相关产品推荐
相关产品推荐

