PyTorch torch.lu Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.luIt is a function in PyTorch used to compute the LU decomposition of a matrix. LU decomposition decomposes matrix A into A = P * L * U, where L is a lower triangular matrix, U is an upper triangular matrix, and P is a permutation matrix.

Function Definition

torch.lu(A, pivot=True, get_infos=False, out=None)

Parameters:

  • A(Tensor): Input matrix.
  • pivot(bool, optional): Whether to perform LU decomposition with column pivoting. Defaults to True.
  • get_infos(bool, optional): If True, return information. Defaults to False.
  • out(tuple, optional): Output tuple.

Return Value:

  • tuple: Returns the tuple (pivot, L, U).

Usage Examples

Example

import torch

# Create matrix
A = torch.tensor([[1.0, 2.0, 3.0],
                  [4.0, 5.0, 6.0],
                  [7.0, 8.0, 9.0]])

# LU decomposition
LU, pivots = torch.lu(A, pivot=True)

print("Matrix A:")
print(A)
print("\nLU matrix:")
print(LU)
print("\nPivot indices:")
print(pivots)

The output result is:

矩阵 A:
tensor([[1., 2., 3.],
        [4., 5., 6.],
        [7., 8., 9.]])
LU 矩阵:
tensor([[7., 8., 9.],
        [0.2, 0.4, 0.6],
        [0.6, 0.8, 0.0]])
主元索引:
tensor([3, 3, 3], dtype=torch.int32)

Example - Using get_infos

import torch

A = torch.tensor([[1.0, 2.0, 3.0],
                  [4.0, 5.0, 6.0],
                  [7.0, 8.0, 9.0]])

# LU decomposition, including information
LU, pivots, info = torch.lu(A, pivot=True, get_infos=True)

print("Information:", info)

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions