PyTorch torch.triangular_solve Function


Pytorch torch 参考手册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 vector
  • upper: 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)

The output is:

解 X:
tensor([[1.],
        [2.],
        [2.]])

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions