PyTorch torch.amax Function
PyTorch torch Reference Manual
torch.amaxIt is a function in PyTorch used to return the maximum value of a tensor along the specified dimension.
Function Definition
torch.amax(input, dim, keepdim=False)
Usage Examples
Example
import torch
x = torch.tensor([[1, 3, 2], [4, 1, 3]])
# Return the maximum value of all elements
print("Global maximum:", torch.amax(x))
# Maximum along dim=0
print("dim=0 maximum:", torch.amax(x, dim=0))
# Maximum along dim=1
print("dim=1 maximum:", torch.amax(x, dim=1))
x = torch.tensor([[1, 3, 2], [4, 1, 3]])
# Return the maximum value of all elements
print("Global maximum:", torch.amax(x))
# Maximum along dim=0
print("dim=0 maximum:", torch.amax(x, dim=0))
# Maximum along dim=1
print("dim=1 maximum:", torch.amax(x, dim=1))
The output result is:
全局最大: tensor(4) dim=0 最大: tensor([4, 3, 3]) dim=1 最大: tensor([3, 4])
Other Extensions