PyTorch torch.diagonal_scatter Function
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)
# 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)
# 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])
# 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.
Other Extensions