Функция linalg.solve
Функция linalg.solve решает систему линейных уравнений.
Первым параметром передаётся квадратная матрица коэффициентов A,
вторым параметром - вектор или матрица правых частей B.
Возвращается решение X, такое что A @ X = B.
Матрица A должна быть квадратной (размером n x n).
Если матрица вырождена или плохо обусловлена, функция может вернуть
некорректный результат или выбросить ошибку.
Синтаксис
torch.linalg.solve(A, B, *, left=True)
Параметр left определяет, с какой стороны умножать.
По умолчанию left=True, что соответствует уравнению
A @ X = B. Если указать left=False,
то решается уравнение X @ A = B.
Пример
Решим систему линейных уравнений с двумя неизвестными:
import torch
A = torch.tensor([
[2.0, 1.0],
[1.0, 3.0],
])
B = torch.tensor([4.0, 5.0])
X = torch.linalg.solve(A, B)
print(X)
Результат выполнения кода:
tensor([1.4000, 1.2000])
Проверим решение: подставим найденные значения в систему.
res = A @ X
print(res)
Результат выполнения кода:
tensor([4.0000, 5.0000])
Как видим, вектор res совпадает с правой частью B.
Пример
Решим систему с матрицей правых частей, содержащей несколько столбцов:
import torch
A = torch.tensor([
[3.0, 1.0],
[1.0, 2.0],
])
B = torch.tensor([
[8.0, 2.0],
[6.0, 1.0],
])
X = torch.linalg.solve(A, B)
print(X)
Результат выполнения кода:
tensor([
[2.0000, 0.6000],
[2.0000, 0.2000],
])
Теперь каждый столбец X является решением для соответствующего
столбца B.
Пример
Решим систему с использованием параметра left=False
для уравнения X @ A = B:
import torch
A = torch.tensor([
[2.0, 1.0],
[1.0, 3.0],
])
B = torch.tensor([
[5.0, 7.0],
[6.0, 9.0],
])
X = torch.linalg.solve(A, B, left=False)
print(X)
Результат выполнения кода:
tensor([
[1.8000, 2.6000],
[1.4000, 2.8000],
])
Проверим, что X @ A ≈ B:
res = X @ A
print(res)
Результат выполнения кода:
tensor([
[5.0000, 7.0000],
[6.0000, 9.0000],
])
Решение найдено верно.
Смотрите также
-
функцию
inv,
которая вычисляет обратную матрицу -
функцию
pinv,
которая вычисляет псевдообратную матрицу -
функцию
det,
которая вычисляет определитель матрицы -
функцию
matrix_norm,
которая вычисляет норму матрицы