Spec-Zone.ru › PyTorch 1

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() разные значения по умолчанию для измерений, поэтому их нужно явно указать.

Parameters:
  • input (Тензор) – входной тензор. Должен быть по крайней мере одномерным.
  • 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/1.13/generated/torch.diag_embed.html

Spec-Zone.ru

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