在torch包内有两个方法值得关注,它们分别是:torch.max和torch.argmax。
在PyTorch中,max方法有两种形式:
(1)返回向量里面最大值的元素。
torch.max(input)
(2)返回一个命名元组 (values, indices),包含最大值以及对应的索引值
torch.max(input, dim, keepdim=False, *, out=None)
上面函数返回一个命名元组 (values, indices),其中:values 是输入张量在给定维度 dim 上每一行的最大值。indices 是每个最大值所在的位置索引(即 argmax 的结果)。简单来说,这个函数会返回每一行在指定维度上的最大值及其对应的索引位置。
torch.argmax()方法返回的是torch.max()方法的第二个值,即每个最大值所在的位置索引。所以,不要单独学习argmax方法,应该重点抓住max方法,然后顺手学习argmax方法,这个前后次序要明白。
不要单独学习argmax方法,应该重点抓住max方法,然后顺手学习argmax方法,这个前后次序要明白。