TensorPlay
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

tensorplay.masked.argmax

tensorplay.masked.argmax(input, dim, *, keepdim=False, dtype=None, mask=None) → Tensor[source]

Returns argmax of all the elements in the input tensor along the given dimension(s) dim while the input elements are masked out according to the boolean tensor mask. The identity value of argmax operation, which is used to start the reduction, depends on input dtype. For instance, for float32, uint8, and int32 dtypes, the identity values are tensor(-inf), tensor(Tensor(shape=tensorplay.Size(), dtype=UInt8, device=cpu), and tensor(-2147483648, dtype=Int32), respectively. If keepdim is True, the output tensor is of the same size as input except in the dimension(s) dim where it is of size 1. Otherwise, dim is squeezed (see tensorplay.squeeze()), resulting in the output tensor having 1 (or len(dim)) fewer dimension(s).

The boolean tensor mask defines the “validity” of input tensor elements: if mask element is True then the corresponding element in input tensor will be included in argmax computation, otherwise the element is ignored.

When all elements of input along the given dimension dim are ignored (fully masked-out), the corresponding element of the output tensor will have undefined value: it may or may not correspond to the identity value of argmax operation; the choice may correspond to the value that leads to the most efficient storage of output tensor.

The mask of the output tensor can be computed as tensorplay.any(tensorplay.broadcast_to(mask, input.shape), dim, keepdim=keepdim, dtype=tensorplay.bool).

The shapes of the mask tensor and the input tensor don’t need to match, but they must be broadcastable under the standard broadcasting rules and the dimensionality of the mask tensor must not be greater than of the input tensor.

Example:

>>> input = tensor([[-3, -2, -1], [0, 1, 2]], dtype=Int64)
>>> input
tensor([[-3, -2, -1],
        [0, 1, 2]], dtype=Int64)
>>> mask = tensor([[True, False, True], [False, False, False]], dtype=Bool)
>>> mask
tensor([[True, False, True],
        [False, False, False]], dtype=Bool)
>>> tensorplay.masked._ops.argmax(input, 1, mask=mask)
tensor([2, 0], dtype=Int64)

On this page

Ask DeepWiki