PyTorch torch.lu_unpack Function
PyTorch torch Reference Manual
torch.lu_unpackIt is a function in PyTorch used to unpack the L, U, and P matrices from the result of LU decomposition.
Function Definition
torch.lu_unpack(LU_data, LU_pivots, unpack_data=True, out=None)
Parameters:
LU_data(Tensor): The matrix obtained from LU decomposition.LU_pivots(Tensor): The pivot indices from LU decomposition.unpack_data(bool, optional): Whether to unpack the data. Defaults to True.out(tuple, optional): Output tuple.
Return Value:
tuple: Returns the tuple (pivot, L, U).
Usage Example
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)
# Unpack L, U, P
P, L, U = torch.lu_unpack(LU, pivots)
print("Matrix A:")
print(A)
print("nPermutation matrix P:")
print(P)
print("nLower triangular matrix L:")
print(L)
print("nUpper triangular matrix U:")
print(U)
print("nVerification: P @ L @ U =")
print(P @ L @ U)
# 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)
# Unpack L, U, P
P, L, U = torch.lu_unpack(LU, pivots)
print("Matrix A:")
print(A)
print("nPermutation matrix P:")
print(P)
print("nLower triangular matrix L:")
print(L)
print("nUpper triangular matrix U:")
print(U)
print("nVerification: P @ L @ U =")
print(P @ L @ U)
The output result is:
矩阵 A:
tensor([[1., 2., 3.],
[4., 5., 6.],
[7., 8., 9.]])
置换矩阵 P:
tensor([[0., 0., 1.],
[1., 0., 0.],
[0., 1., 0.]])
下三角矩阵 L:
tensor([[1.0000, 0.0000, 0.0000],
[0.1429, 1.0000, 0.0000],
[0.5714, 0.5000, 1.0000]])
上三角矩阵 U:
tensor([[7.0000, 8.0000, 9.0000],
[0.0000, 0.8571, 1.7143],
[0.0000, 0.0000, 0.0000]])
验证: P @ L @ U =
tensor([[1., 2., 3.],
[4., 5., 6.],
[7., 8., 9.]])
Other Extensions