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