Spec-Zone.ru › PyTorch 2.14

torch.diag_embed

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

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

Аргумент offset определяет, какую диагональ следует рассматривать:

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

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

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

Параметры:
  • input (Tensor) – входной тензор. Должен иметь как минимум одно измерение.
  • 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]]])

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

Spec-Zone.ru

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