PyTorch torch.lu_solve Function


Pytorch torch 参考手册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)

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)

Pytorch torch 参考手册PyTorch torch Reference Manual

Other Extensions