Функция linalg.svd
Функция torch.linalg.svd выполняет сингулярное разложение (SVD) матрицы или пакета матриц. Она раскладывает входную матрицу на три матрицы: U, S и V^T, где U и V - ортогональные матрицы, а S - диагональная матрица с сингулярными значениями. Первым параметром функция принимает тензор с входными данными. Вторым параметром можно указать full_matrices для выбора полного или экономного разложения. Третьим параметром driver можно выбрать алгоритм вычисления. Четвертым параметром out можно передать кортеж из трех тензоров для сохранения результата.
Синтаксис
torch.linalg.svd(A, full_matrices=True, driver=None, out=None)
Пример
Давайте выполним сингулярное разложение для матрицы размером 3x2:
import torch
t = torch.tensor([
[1.0, 2.0],
[3.0, 4.0],
[5.0, 6.0],
])
U, S, Vh = torch.linalg.svd(t)
print("U:", U)
print("S:", S)
print("Vh:", Vh)
Результат выполнения кода:
U: tensor([
[-0.2298, 0.8835, 0.4082],
[-0.5247, 0.2408, -0.8165],
[-0.8196, -0.4019, 0.4082],
])
S: tensor([9.5255, 0.5143])
Vh: tensor([
[-0.6196, -0.7849],
[-0.7849, 0.6196],
])
Пример
Используем параметр full_matrices=False для получения экономного разложения:
import torch
t = torch.tensor([
[1.0, 2.0],
[3.0, 4.0],
[5.0, 6.0],
])
U, S, Vh = torch.linalg.svd(t, full_matrices=False)
print("U:", U)
print("S:", S)
print("Vh:", Vh)
Результат выполнения кода:
U: tensor([
[-0.2298, 0.8835],
[-0.5247, 0.2408],
[-0.8196, -0.4019],
])
S: tensor([9.5255, 0.5143])
Vh: tensor([
[-0.6196, -0.7849],
[-0.7849, 0.6196],
])
Пример
Выполним SVD для пакета матриц (двух матриц 2x2):
import torch
t = torch.tensor([
[[1.0, 2.0], [3.0, 4.0]],
[[5.0, 6.0], [7.0, 8.0]],
])
U, S, Vh = torch.linalg.svd(t)
print("U shape:", U.shape)
print("S:", S)
print("Vh shape:", Vh.shape)
Результат выполнения кода:
U shape: torch.Size([2, 2, 2])
S: tensor([
[5.4650, 0.3660],
[13.0277, 0.7599],
])
Vh shape: torch.Size([2, 2, 2])
Смотрите также
-
функцию
inv,
которая вычисляет обратную матрицу -
функцию
qr,
которая выполняет QR-разложение матрицы -
функцию
eig,
которая вычисляет собственные значения и векторы -
функцию
matrix_norm,
которая вычисляет норму матрицы