PyTorch torch.solve Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.solveIt is a function in PyTorch for solving linear equations. It solves AX = B, where A is the coefficient matrix and B is the right-hand side matrix or vector.

Function Definition

torch.solve(B, A, out=None)

Parameters:

  • B(Tensor): Right-hand side matrix or vector.
  • A(Tensor): Coefficient matrix.
  • out(tuple, optional): Output tuple.

Return Value:

  • tuple: Returns a tuple of (X, LU), where X is the solution and LU is the decomposition.

Usage Example

Example

import torch

# Create coefficient matrix and right-hand side vector
A = torch.tensor([[1.0, 2.0], [3.0, 4.0]])
B = torch.tensor([5.0, 11.0])

# Solve AX = B
X, _ = torch.solve(B, A)

print("Coefficient 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.]])
右侧向量 B:
tensor([ 5., 11.])
解 X:
tensor([1., 2.])
验证: A @ X =
tensor([ 5., 11.])

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions