PyTorch torch.lu_solve Function
PyTorch torch Reference Manual
torch.lu_solveThis function uses LU decomposition in PyTorch to solve linear equations. It efficiently solves AX = B by utilizing the already computed LU decomposition.
Function Definition
torch.lu_solve(B, LU, pivots, out=None)
Parameters:
B(Tensor): The right-hand side matrix or vector.LU(Tensor): The matrix obtained from LU decomposition.pivots(Tensor): The pivot indices of the LU decomposition.out(Tensor, optional): Output tensor.
Return Value:
torch.Tensor: Returns the solution of the linear equation system.
Usage Example
Example
import torch
# Create matrix and right-hand side vector
A = torch.tensor([[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0],
[7.0, 8.0, 9.0]], dtype=torch.float64)
B = torch.tensor([14.0, 32.0, 50.0], dtype=torch.float64)
# LU decomposition
LU, pivots = torch.lu(A)
# Solve using LU decomposition
X = torch.lu_solve(B, LU, pivots)
print("Matrix A:")
print(A)
print("\nRight-hand side vector B:")
print(B)
print("\nSolution X:")
print(X)
print("\nVerification: A @ X =")
print(A @ X)
# Create matrix and right-hand side vector
A = torch.tensor([[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0],
[7.0, 8.0, 9.0]], dtype=torch.float64)
B = torch.tensor([14.0, 32.0, 50.0], dtype=torch.float64)
# LU decomposition
LU, pivots = torch.lu(A)
# Solve using LU decomposition
X = torch.lu_solve(B, LU, pivots)
print("Matrix A:")
print(A)
print("\nRight-hand side vector B:")
print(B)
print("\nSolution X:")
print(X)
print("\nVerification: A @ X =")
print(A @ X)
The output result is:
矩阵 A:
tensor([[1., 2., 3.],
[4., 5., 6.],
[7., 8., 9.]], dtype=torch.float64)
右侧向量 B:
tensor([14., 32., 50.], dtype=torch.float64)
解 X:
tensor([1., 2., 3.], dtype=torch.float64)
验证: A @ X =
tensor([14., 32., 50.], dtype=torch.float64)
Other Extensions