PyTorch torch.argmax Function


Pytorch torch 参考手册PyTorch torch Reference Manual

torch.argmaxIt is a function in PyTorch used to return the index of the maximum value along a dimension.

Function Definition

torch.argmax(input, dim, keepdim)

Usage Example

Example

import torch

x = torch.tensor([[1, 3, 2], [4, 1, 3]])

# Index of the global maximum value
print("Global max index:", torch.argmax(x))

# Index of the maximum value along dim=0
print("dim=0 max index:", torch.argmax(x, dim=0))

# Index of the maximum value along dim=1
print("dim=1 max index:", torch.argmax(x, dim=1))

The output result is:

全局最大索引: tensor(4)
dim=0 最大索引: tensor([1, 0, 1])
dim=1 最大索引: tensor([1, 0])

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions