PyTorch torch.select_scatter Function
PyTorch torch Reference Manual
torch.select_scatteris a function in PyTorch used to scatter the source tensor to a specified index position. Itsrcscatters the values ofinputinto the specified dimensiondimand indexindexposition.
Function Definition
torch.select_scatter(input, src, dim, index)
Parameters:
input(Tensor): The input tensor, i.e., the tensor to be modified.src(Tensor): The source tensor, the values to be scattered into input.dim(int): The dimension of scattering.index(int): The position index of scattering.
Return Value:
torch.Tensor: Returns the modified tensor.
Usage Examples
Example
import torch
# Create the input tensor and source tensor
input = torch.zeros(4, 4)
src = torch.ones(4)
# Scatter src at the index 1 position in the first dimension (rows)
output = torch.select_scatter(input, src, dim=0, index=1)
print("Input tensor:")
print(input)
print("nSource tensor:")
print(src)
print("nScatter result at index 1 position:")
print(output)
# Create the input tensor and source tensor
input = torch.zeros(4, 4)
src = torch.ones(4)
# Scatter src at the index 1 position in the first dimension (rows)
output = torch.select_scatter(input, src, dim=0, index=1)
print("Input tensor:")
print(input)
print("nSource tensor:")
print(src)
print("nScatter result at index 1 position:")
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., 1., 1., 1.])
在索引1位置散布结果:
tensor([[0., 0., 0., 0.],
[1., 1., 1., 1.],
[0., 0., 0., 0.],
[0., 0., 0., 0.]])
Example
import torch
# Create a 3D tensor
input = torch.zeros(3, 4, 5)
src = torch.ones(5)
# Scatter in the second dimension (index 2)
output = torch.select_scatter(input, src, dim=1, index=2)
print("Input shape:", input.shape)
print("Source shape:", src.shape)
print("Result shape:", output.shape)
print("nResult (dim=1, index=2):")
print(output[0])
# Create a 3D tensor
input = torch.zeros(3, 4, 5)
src = torch.ones(5)
# Scatter in the second dimension (index 2)
output = torch.select_scatter(input, src, dim=1, index=2)
print("Input shape:", input.shape)
print("Source shape:", src.shape)
print("Result shape:", output.shape)
print("nResult (dim=1, index=2):")
print(output[0])
The output result is:
输入形状: torch.Size([3, 4, 5]) 源形状: torch.Size([5]) 结果形状: torch.Size([3, 4, 5]) 结果 (dim=1, index=2): tensor([1., 1., 1., 1., 1.])
Example
import torch
# Use different values
input = torch.arange(16).reshape(4, 4).float()
src = torch.tensor([100, 100, 100, 100])
# Scatter at the last row
output = torch.select_scatter(input, src, dim=0, index=-1)
print("Original tensor:")
print(input)
print("nScattered to index -1 position:")
print(output)
# Use different values
input = torch.arange(16).reshape(4, 4).float()
src = torch.tensor([100, 100, 100, 100])
# Scatter at the last row
output = torch.select_scatter(input, src, dim=0, index=-1)
print("Original tensor:")
print(input)
print("nScattered to index -1 position:")
print(output)
The output result is:
原始张量:
tensor([[ 0., 1., 2., 3.],
[ 4., 5., 6., 7.],
[ 8., 9., 10., 11.],
[12., 13., 14., 15.]])
散布到索引-1位置:
tensor([[ 0., 1., 2., 3.],
[ 4., 5., 6., 7.],
[ 8., 9., 10., 11.],
[100., 100., 100., 100.]])
Note:torch.select_scatterDoes not modify the original input tensor, but returns a new tensor.srcThe dimensions must match the dimensions of the slice to be replaced.
Other Extensions