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не равно , порядок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