PyTorch torch.lu Function
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)
# 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)
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)
Other Extensions