PyTorch torch.trace Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.traceis a function in PyTorch used to compute the trace of a matrix. The trace is the sum of the elements on the main diagonal of a matrix.

Function Definition

torch.trace(input)

Parameter Description:

  • input: input tensor (at least two-dimensional)

Usage Example

Example

import torch

# Create a matrix
a = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])

# Compute the trace (sum of main diagonal elements)
y = torch.trace(a)
print(y)

The output result is:

tensor(15)

Example

import torch

# Create a non-square matrix
a = torch.tensor([[1, 2, 3], [4, 5, 6]])

# Compute the trace
y = torch.trace(a)
print(y)

The output result is:

tensor(6)

Example

import torch

# Create a multi-dimensional tensor, only consider the first two dimensions
a = torch.randn(3, 4, 4, 5)
y = torch.trace(a)
print(y.shape)

The output result is:

torch.Size([3, 5])

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions