Spec-Zone.ru › PyTorch 1

torch.sparse.mm

torch.sparse.mm()

Выполняет матричное умножение разреженной матрицы mat1 и (разреженной или сплошной) матрицы mat2. Аналогично torch.mm(), если mat1 является тензором (n×m)(n \times m), то mat2 является тензором (m×p)(m \times p), out будет тензором (n×p)(n \times p). Когда mat1 является тензором COO, он должен иметь sparse_dim = 2. При использовании тензоров COO эта функция также поддерживает обратное распространение для обоих входных данных.

Поддерживает форматы хранения CSR и COO.

Примечание

Эта функция не поддерживает вычисление производных относительно матриц CSR.

Аргументы:

mat1 (Tensor): первая разреженная матрица, которая должна быть умножена mat2 (Tensor): вторая матрица, которая должна быть умножена, которая может быть разреженной или плотной

Форма:

Формат выходного тензора этой функции следующий: - разреженная х разреженная -> разреженная - разреженная х плотная -> плотная

Пример:

>>> a = torch.randn(2, 3).to_sparse().requires_grad_(True)
>>> a
tensor(indices=tensor([[0, 0, 0, 1, 1, 1],
                       [0, 1, 2, 0, 1, 2]]),
       values=tensor([ 1.5901,  0.0183, -0.6146,  1.8061, -0.0112,  0.6302]),
       size=(2, 3), nnz=6, layout=torch.sparse_coo, requires_grad=True)

>>> b = torch.randn(3, 2, requires_grad=True)
>>> b
tensor([[-0.6479,  0.7874],
        [-1.2056,  0.5641],
        [-1.1716, -0.9923]], requires_grad=True)

>>> y = torch.sparse.mm(a, b)
>>> y
tensor([[-0.3323,  1.8723],
        [-1.8951,  0.7904]], grad_fn=<SparseAddmmBackward>)
>>> y.sum().backward()
>>> a.grad
tensor(indices=tensor([[0, 0, 0, 1, 1, 1],
                       [0, 1, 2, 0, 1, 2]]),
       values=tensor([ 0.1394, -0.6415, -2.1639,  0.1394, -0.6415, -2.1639]),
       size=(2, 3), nnz=6, layout=torch.sparse_coo)

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.sparse.mm.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API