PyTorch torch.diagonal_scatter Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.diagonal_scatteris a function in PyTorch used to scatter values into the diagonal positions of a tensor. Itsrcscatters the values ofinputonto the specified diagonal of.

Function Definition

torch.diagonal_scatter(input, src, offset=0, dim1=0, dim2=1)

Parameters:

  • input(Tensor): The input tensor, i.e., the tensor to be modified.
  • src(Tensor): The source tensor, the values to be scattered into the diagonal positions.
  • offset(int, optional): The diagonal offset. A positive value indicates a superdiagonal, a negative value indicates a subdiagonal, and 0 indicates the main diagonal.
  • dim1(int, optional): The first dimension, defaults to 0.
  • dim2(int, optional): The second dimension, defaults to 1.

Return Value:

  • torch.Tensor: Returns the modified tensor.

Usage Examples

Example

import torch

# Create the input tensor
input = torch.zeros(4, 4)
src = torch.tensor([1, 2, 3, 4])

# Scatter the values into the main diagonal
output = torch.diagonal_scatter(input, src)

print("Input tensor:")
print(input)
print("nSource tensor:")
print(src)
print("nScatter to the main diagonal:")
print(output)

The output result is:

输入张量:
tensor([[0., 0., 0., 0.],
        [0., 0., 0., 0.],
        [0., 0., 0., 0.],
        [0., 0., 0., 0.]])

源张量:
tensor([1, 2, 3, 4])

散布到主对角线:
tensor([[1., 0., 0., 0.],
        [0., 2., 0., 0.],
        [0., 0., 3., 0.],
        [0., 0., 0., 4.]])

Example

import torch

# Create the input tensor
input = torch.zeros(4, 4)
src = torch.tensor([1, 2, 3])

# Scatter the values into the superdiagonal (offset=1)
output = torch.diagonal_scatter(input, src, offset=1)

print("Scatter to the superdiagonal (offset=1):")
print(output)

# Scatter the values into the subdiagonal (offset=-1)
output2 = torch.diagonal_scatter(input, src, offset=-1)

print("nScatter to the subdiagonal (offset=-1):")
print(output2)

The output result is:

散布到上对角线 (offset=1):
tensor([[0., 1., 0., 0.],
        [0., 0., 2., 0.],
        [0., 0., 0., 3.],
        [0., 0., 0., 0.]])

散布到下对角线 (offset=-1):
tensor([[0., 0., 0., 0.],
        [1., 0., 0., 0.],
        [0., 2., 0., 0.],
        [0., 0., 3., 0.]])

Example

import torch

# Using diagonals in a 3D tensor
input = torch.zeros(3, 4, 4)
src = torch.tensor([10, 20, 30])

# Scatter on the specified two dimensions
output = torch.diagonal_scatter(input, src, dim1=1, dim2=2)

print("Input shape:", input.shape)
print("Source shape:", src.shape)
print("Result shape:", output.shape)

# View the first batch
print("nResult of the first batch:")
print(output[0])

The output result is:

输入形状: torch.Size([3, 4, 4])
源形状: torch.Size([3])
结果形状: torch.Size([3, 4, 4])

第一个batch的结果:
tensor([[10., 0., 0., 0.],
        [0., 20., 0., 0.],
        [0., 0., 30., 0.],
        [0., 0., 0., 0.]])

Note:torch.diagonal_scatterIt does not modify the original input tensor, but returns a new tensor.srcThe size of must match the number of diagonal elements.


Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions