PyTorch torch.amax Function


Pytorch torch 参考手册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))

The output result is:

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

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions