PyTorch torch.slice_scatter Function
Pytorch torch Reference Manual
torch.slice_scatterIt is a function in PyTorch used to scatter the values of a source tensor into the slice positions of the input tensor.
Its effect is equivalent toinput[dim][start:end:step] = src, but it does not modify the original tensor; instead, it returns a new tensor.
Function Definition
torch.slice_scatterThe complete function definition is as follows:
torch.slice_scatter(input, src, dim=0, start=None, end=None, step=1)
Parameter Description
The following table liststorch.slice_scatterits various parameters:
| Parameter | Type | Required | Default Value | Description |
|---|---|---|---|---|
input | Tensor | Required | — | Input tensor, the tensor to be modified. |
src | Tensor | Required | — | Source tensor, the values to be scattered into the input slice,its size must match the size of the slice.。 |
dim | int | Optional | 0 | The dimension for scattering. |
start | int | Optional | None | Starting index (inclusiveof that index), defaults to 0. |
end | int | Optional | None | Ending index (exclusiveof that index), defaults to the end. |
step | int | Optional | 1 | Step size, with the same meaning as step in Python slicing. |
Return Value
Returns the scattered new tensor (torch.Tensor), the originalinputwill not be modified.
Usage Examples
The following examples demonstratetorch.slice_scatterits common usage.
Example
# Create input tensor and source tensor
input = torch.zeros(8, 4)
src = torch.ones(2, 4)
# Scatter src into the first two rows of input
output = torch.slice_scatter(input, src, dim=0, end=2)
print("Input tensor:")
print(input)
print("\nSource tensor:")
print(src)
print("\nScatter result:")
print(output)
The output is:
输入张量:
tensor([[0., 0., 0., 0.],
[0., 0., 0., 0.],
[0., 0., 0., 0.],
[0., 0., 0., 0.],
[0., 0., 0., 0.],
[0., 0., 0., 0.],
[0., 0., 0., 0.],
[0., 0., 0., 0.]])
源张量:
tensor([[1., 1., 1., 1.],
[1., 1., 1., 1.]])
散布结果:
tensor([[1., 1., 1., 1.],
[1., 1., 1., 1.],
[0., 0., 0., 0.],
[0., 0., 0., 0.],
[0., 0., 0., 0.],
[0., 0., 0., 0.],
[0., 0., 0., 0.],
[0., 0., 0., 0.]])
Example
# Use start and end to specify the range
input = torch.zeros(10)
src = torch.tensor([1, 2, 3])
# Scatter src to positions at indices 2-4 of input
output = torch.slice_scatter(input, src, dim=0, start=2, end=5)
print("Input:", input)
print("Source:", src)
print("Result:", output)
The output is:
输入: tensor([0., 0., 0., 0., 0., 0., 0., 0., 0., 0.]) 源: tensor([1., 2., 3.]) 结果: tensor([0., 0., 1., 2., 3., 0., 0., 0., 0., 0.])
Example
# Use the step parameter
input = torch.zeros(10)
src = torch.tensor([1, 2])
# Step size is 2, from index 0 to 4 (exclusive), actually writes indices 0 and 2
output = torch.slice_scatter(input, src, dim=0, start=0, end=4, step=2)
print("Input:", input)
print("Source:", src)
print("Result with step 2:", output)
The output is:
输入: tensor([0., 0., 0., 0., 0., 0., 0., 0., 0., 0.]) 源: tensor([1., 2.]) 步长为2的结果: tensor([1., 0., 2., 0., 0., 0., 0., 0., 0., 0.])
Note:
torch.slice_scatterdoes not modify the original input tensor, but returns a new tensor.This function is
torch.slicethe inverse operation of.
Other Extensions