PyTorch torch.slice_scatter Function


Pytorch torch 参考手册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:

ParameterTypeRequiredDefault ValueDescription
inputTensorRequired—Input tensor, the tensor to be modified.
srcTensorRequired—Source tensor, the values to be scattered into the input slice,its size must match the size of the slice.。
dimintOptional0The dimension for scattering.
startintOptionalNoneStarting index (inclusiveof that index), defaults to 0.
endintOptionalNoneEnding index (exclusiveof that index), defaults to the end.
stepintOptional1Step 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

import torch

# 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

import torch

# 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

import torch

# 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 istorch.slicethe inverse operation of.


Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions