torch.sparse.sampled_addmm
-
torch.sparse.sampled_addmm(input, mat1, mat2, *, beta=1., alpha=1., out=None) → Tensor -
Выполняет матричное произведение плотных матриц
mat1иmat2в позициях, заданных структурой разреженностиinput. Матрицаinputдобавляется к конечному результату.Математически это выполняет следующую операцию:
где — матрица структуры разреженности
input,alphaиbeta— множители масштабирования. имеет значение 1 в позициях, гдеinputимеет ненулевые значения, и 0 в противном случае.Примечание
inputдолжна быть разреженной тензором CSR.mat1иmat2должны быть плотными тензорами.- Параметры
-
-
input (Tensor) – разреженная CSR-матрица размера
(m, n), которая должна быть добавлена и использована для вычисления выборочного матричного произведения -
mat1 (Tensor) – плотная матрица размера
(m, k), которая должна быть умножена -
mat2 (Tensor) – плотная матрица размера
(k, n), которая должна быть умножена
-
input (Tensor) – разреженная CSR-матрица размера
- Ключевые аргументы
-
-
beta (Число, необязательно) – множитель для
input() - alpha (Число, необязательно) – множитель для ()
-
out (Tensor, необязательно) – выходной тензор. Игнорируется, если
None. По умолчанию:None.
-
beta (Число, необязательно) – множитель для
Примеры:
>>> input = torch.eye(3, device='cuda').to_sparse_csr() >>> mat1 = torch.randn(3, 5, device='cuda') >>> mat2 = torch.randn(5, 3, device='cuda') >>> torch.sparse.sampled_addmm(input, mat1, mat2) tensor(crow_indices=tensor([0, 1, 2, 3]), col_indices=tensor([0, 1, 2]), values=tensor([ 0.2847, -0.7805, -0.1900]), device='cuda:0', size=(3, 3), nnz=3, layout=torch.sparse_csr) >>> torch.sparse.sampled_addmm(input, mat1, mat2).to_dense() tensor([[ 0.2847, 0.0000, 0.0000], [ 0.0000, -0.7805, 0.0000], [ 0.0000, 0.0000, -0.1900]], device='cuda:0') >>> torch.sparse.sampled_addmm(input, mat1, mat2, beta=0.5, alpha=0.5) tensor(crow_indices=tensor([0, 1, 2, 3]), col_indices=tensor([0, 1, 2]), values=tensor([ 0.1423, -0.3903, -0.0950]), device='cuda:0', size=(3, 3), nnz=3, layout=torch.sparse_csr)
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.sparse.sampled_addmm.html