PyTorch torch.max Function
Pytorch torch Reference Manual
torch.maxIt is a function in PyTorch used to compute the maximum value of tensors.
Function Definition
torch.max(input, dim, keepdim, out)
Usage Examples
Example
import torch
x = torch.tensor([[1, 2, 3], [4, 5, 6]])
# Global maximum
print("Global maximum:", torch.max(x))
# Maximum along dimension
print("dim=0 maximum:", torch.max(x, dim=0))
x = torch.tensor([[1, 2, 3], [4, 5, 6]])
# Global maximum
print("Global maximum:", torch.max(x))
# Maximum along dimension
print("dim=0 maximum:", torch.max(x, dim=0))
The output result is:
全局最大: tensor(6) dim=0 最大: torch.return_types.max(values=tensor([4, 5, 6]), indices=tensor([1, 1, 1]))
Other Extensions