Spec-Zone.ru › PyTorch 2.14

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.

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

Примечание

Поддержка градиентов:

  • COO @ Dense: обратный проход поддерживается для обоих входных аргументов. Градиент для разреженного входа возвращается в виде разреженного тензора COO.
  • CSR @ Dense: обратный проход поддерживается для обоих входных аргументов. Градиент для разреженного входа возвращается в виде разреженного тензора CSR.
  • CSC/BSR/BSC @ Dense: не поддерживается.
  • Sparse @ Sparse (COO @ COO, CSR @ CSR): прямой проход выполняется, но обратный проход не поддерживается.
  • Смешанные форматы (COO @ CSR, CSR @ COO): не поддерживаются.

Эта функция также принимает необязательный аргумент reduce, который позволяет задать необязательную операцию редукции. Математически выполняется следующая операция:

zij=⨁k=0K−1xikykjz_{ij} = \bigoplus_{k = 0}^{K - 1} x_{ik} y_{kj}

где ⨁\bigoplus обозначает оператор редукции. reduce реализован только для формата хранения CSR на устройстве CPU.

Параметры:
  • mat1 (Tensor) – первая разреженная матрица для умножения
  • mat2 (Tensor) – вторая матрица для умножения, которая может быть разреженной или плотной
  • reduce (str, необязательно) – операция редукции для неуникальных индексов ("sum", "mean", "amax", "amin"). По умолчанию "sum".
Форма:

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

Пример:

>>> a = torch.tensor([[1., 0, 2], [0, 3, 0]]).to_sparse().requires_grad_()
>>> a
tensor(indices=tensor([[0, 0, 1],
                       [0, 2, 1]]),
       values=tensor([1., 2., 3.]),
       size=(2, 3), nnz=3, layout=torch.sparse_coo, requires_grad=True)
>>> b = torch.tensor([[0, 1.], [2, 0], [0, 0]], requires_grad=True)
>>> b
tensor([[0., 1.],
        [2., 0.],
        [0., 0.]], requires_grad=True)
>>> y = torch.sparse.mm(a, b)
>>> y
tensor([[0., 1.],
        [6., 0.]], grad_fn=<SparseAddmmBackward0>)
>>> y.sum().backward()
>>> a.grad
tensor(indices=tensor([[0, 0, 1],
                       [0, 2, 1]]),
       values=tensor([1., 0., 2.]),
       size=(2, 3), nnz=3, layout=torch.sparse_coo)
>>> c = a.detach().to_sparse_csr()
>>> c
tensor(crow_indices=tensor([0, 2, 3]),
       col_indices=tensor([0, 2, 1]),
       values=tensor([1., 2., 3.]), size=(2, 3), nnz=3,
       layout=torch.sparse_csr)
>>> y1 = torch.sparse.mm(c, b, 'sum')
>>> y1
tensor([[0., 1.],
        [6., 0.]], grad_fn=<SparseMmReduceImplBackward0>)
>>> y2 = torch.sparse.mm(c, b, 'max')
>>> y2
tensor([[0., 1.],
        [6., 0.]], grad_fn=<SparseMmReduceImplBackward0>)

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

Spec-Zone.ru

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