Spec-Zone.ru › PyTorch 2

torch.diag_embed

torch.diag_embed(input, offset=0, dim1=-2, dim2=-1) → Tensor

Создает тензор, диагонали определённых 2D плоскостей (указанных dim1 и dim2) заполняются input. Для облегчения создания пакетных диагональных матриц по умолчанию выбираются 2D плоскости, образованные двумя последними измерениями возвращаемого тензора.

Аргумент offset управляет выбором диагонали:

  • Если offset = 0, это главная диагональ.
  • Если offset > 0, это диагональ выше главной.
  • Если offset < 0, это диагональ ниже главной.

Размер новой матрицы вычисляется для того, чтобы заданная диагональ имела размер последнего измерения входных данных. Обратите внимание, что для offset отличного от 00, порядок dim1 и dim2 имеет значение. Их перемена эквивалентна изменению знака offset.

Применение torch.diagonal() к результату этой функции с теми же аргументами даёт матрицу, идентичную входной. Однако, у torch.diagonal() разные значения по умолчанию для измерений, поэтому их нужно явно указать.

Параметры
  • input (Tensor) – входной тензор. Должен быть не менее 1-мерным.
  • offset (int, необязательно) – какая диагональ должна рассматриваться. По умолчанию: 0 (главная диагональ).
  • dim1 (int, необязательно) – первое измерение относительно которого брать диагональ. По умолчанию: -2.
  • dim2 (int, необязательно) – второе измерение относительно которого брать диагональ. По умолчанию: -1.

Пример:

>>> a = torch.randn(2, 3)
>>> torch.diag_embed(a)
tensor([[[ 1.5410,  0.0000,  0.0000],
         [ 0.0000, -0.2934,  0.0000],
         [ 0.0000,  0.0000, -2.1788]],

        [[ 0.5684,  0.0000,  0.0000],
         [ 0.0000, -1.0845,  0.0000],
         [ 0.0000,  0.0000, -1.3986]]])

>>> torch.diag_embed(a, offset=1, dim1=0, dim2=2)
tensor([[[ 0.0000,  1.5410,  0.0000,  0.0000],
         [ 0.0000,  0.5684,  0.0000,  0.0000]],

        [[ 0.0000,  0.0000, -0.2934,  0.0000],
         [ 0.0000,  0.0000, -1.0845,  0.0000]],

        [[ 0.0000,  0.0000,  0.0000, -2.1788],
         [ 0.0000,  0.0000,  0.0000, -1.3986]],

        [[ 0.0000,  0.0000,  0.0000,  0.0000],
         [ 0.0000,  0.0000,  0.0000,  0.0000]]])

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

Spec-Zone.ru

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