PyTorch torch.cholesky function
PyTorch torch Reference Manual
torch.choleskyThis is a function in PyTorch used to compute the Cholesky decomposition of a symmetric positive definite matrix. Cholesky decomposition decomposes a positive definite matrix A into A = L * L^T, where L is a lower triangular matrix.
Function Definition
torch.cholesky(A, upper=False, out=None)
Parameters:
A(Tensor): The input symmetric positive definite matrix.upper(bool, optional): If True, returns the upper triangular matrix; otherwise, returns the lower triangular matrix. Default is False.out(Tensor, optional): The output tensor.
Return Value:
torch.Tensor: Returns the lower (or upper) triangular matrix of the Cholesky decomposition.
Usage Example
Example
import torch
# Create a symmetric positive definite matrix
A = torch.tensor([[4.0, 2.0, 2.0],
[2.0, 5.0, 3.0],
[2.0, 3.0, 6.0]])
# Cholesky decomposition
L = torch.cholesky(A)
print("Original matrix A:")
print(A)
print("nCholesky decomposition (lower triangular matrix L):")
print(L)
print("nVerification: L @ L.T = ")
print(L @ L.T)
# Create a symmetric positive definite matrix
A = torch.tensor([[4.0, 2.0, 2.0],
[2.0, 5.0, 3.0],
[2.0, 3.0, 6.0]])
# Cholesky decomposition
L = torch.cholesky(A)
print("Original matrix A:")
print(A)
print("nCholesky decomposition (lower triangular matrix L):")
print(L)
print("nVerification: L @ L.T = ")
print(L @ L.T)
The output result is:
原矩阵 A:
tensor([[4., 2., 2.],
[2., 5., 3.],
[2., 3., 6.]])
Cholesky 分解 (下三角矩阵 L):
tensor([[2.0000, 0.0000, 0.0000],
[1.0000, 2.0000, 0.0000],
[1.0000, 1.0000, 2.0000]])
验证: L @ L.T =
tensor([[4., 2., 2.],
[2., 5., 3.],
[2., 3., 6.]])
Example - Return Upper Triangular Matrix
import torch
A = torch.tensor([[4.0, 2.0, 2.0],
[2.0, 5.0, 3.0],
[2.0, 3.0, 6.0]])
# Return the upper triangular matrix
U = torch.cholesky(A, upper=True)
print("Upper triangular matrix U:")
print(U)
A = torch.tensor([[4.0, 2.0, 2.0],
[2.0, 5.0, 3.0],
[2.0, 3.0, 6.0]])
# Return the upper triangular matrix
U = torch.cholesky(A, upper=True)
print("Upper triangular matrix U:")
print(U)
Other Extensions