PyTorch torch.triangular_solve Function
Pytorch torch Reference Manual
torch.triangular_solveIt is a function in PyTorch used to solve triangular linear equations. This function solves AX = B, where A is a triangular matrix.
Function Definition
torch.triangular_solve(A, B, upper, transpose, unitriangular)
Parameter Description
A: coefficient matrix (must be a square matrix)B: right-hand side matrix or vectorupper: whether A is an upper triangular matrix (default True)transpose: whether to transpose A (default False)unitriangular: whether to use unit triangular (default False)
Usage Example
Example
import torch
# Create an upper triangular coefficient matrix
A = torch.tensor([[3.0, 1.0, 2.0],
[0.0, 2.0, 1.0],
[0.0, 0.0, 1.0]])
# Right-hand side vector
B = torch.tensor([9.0, 5.0, 2.0])
# Solve AX = B
X = torch.triangular_solve(B.unsqueeze(1), A, upper=True)
print("Solution X:")
print(X.solution)
# Create an upper triangular coefficient matrix
A = torch.tensor([[3.0, 1.0, 2.0],
[0.0, 2.0, 1.0],
[0.0, 0.0, 1.0]])
# Right-hand side vector
B = torch.tensor([9.0, 5.0, 2.0])
# Solve AX = B
X = torch.triangular_solve(B.unsqueeze(1), A, upper=True)
print("Solution X:")
print(X.solution)
The output is:
解 X:
tensor([[1.],
[2.],
[2.]])
Other Extensions