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отличных от , порядок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