PyTorch torch.slice_scatter 函数
torch.slice_scatter 是 PyTorch 中用于将源张量的值散布到输入张量切片位置的函数。
它的效果相当于 input[dim][start:end:step] = src,但不会修改原始张量,而是返回一个新的张量。
函数定义
torch.slice_scatter 的完整函数定义如下:
torch.slice_scatter(input, src, dim=0, start=None, end=None, step=1)
参数说明
下表列出了 torch.slice_scatter 的各个参数:
| 参数 | 类型 | 是否必填 | 默认值 | 说明 |
|---|---|---|---|---|
input | Tensor | 必填 | — | 输入张量,即要修改的张量。 |
src | Tensor | 必填 | — | 源张量,要散布到 input 切片中的值,大小必须与切片后的大小一致。 |
dim | int | 可选 | 0 | 散布的维度。 |
start | int | 可选 | None | 起始索引(包含该索引),默认从 0 开始。 |
end | int | 可选 | None | 结束索引(不包含该索引),默认到末尾结束。 |
step | int | 可选 | 1 | 步长,含义与 Python 切片中的 step 相同。 |
返回值
返回散布后的新张量(torch.Tensor),原始 input 不会被修改。
使用示例
下面的实例演示了 torch.slice_scatter 的常见用法。
实例
import torch
# 创建输入张量和源张量
input = torch.zeros(8, 4)
src = torch.ones(2, 4)
# 将 src 散布到 input 的前两行
output = torch.slice_scatter(input, src, dim=0, end=2)
print("输入张量:")
print(input)
print("\n源张量:")
print(src)
print("\n散布结果:")
print(output)
# 创建输入张量和源张量
input = torch.zeros(8, 4)
src = torch.ones(2, 4)
# 将 src 散布到 input 的前两行
output = torch.slice_scatter(input, src, dim=0, end=2)
print("输入张量:")
print(input)
print("\n源张量:")
print(src)
print("\n散布结果:")
print(output)
输出结果为:
输入张量:
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.]])
实例
import torch
# 使用 start 和 end 指定范围
input = torch.zeros(10)
src = torch.tensor([1, 2, 3])
# 将 src 散布到 input 的索引 2-4 位置
output = torch.slice_scatter(input, src, dim=0, start=2, end=5)
print("输入:", input)
print("源:", src)
print("结果:", output)
# 使用 start 和 end 指定范围
input = torch.zeros(10)
src = torch.tensor([1, 2, 3])
# 将 src 散布到 input 的索引 2-4 位置
output = torch.slice_scatter(input, src, dim=0, start=2, end=5)
print("输入:", input)
print("源:", src)
print("结果:", output)
输出结果为:
输入: 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.])
实例
import torch
# 使用 step 参数
input = torch.zeros(10)
src = torch.tensor([1, 2])
# 步长为 2,从索引 0 开始到 4 结束(不包含 4),实际写入索引 0、2
output = torch.slice_scatter(input, src, dim=0, start=0, end=4, step=2)
print("输入:", input)
print("源:", src)
print("步长为2的结果:", output)
# 使用 step 参数
input = torch.zeros(10)
src = torch.tensor([1, 2])
# 步长为 2,从索引 0 开始到 4 结束(不包含 4),实际写入索引 0、2
output = torch.slice_scatter(input, src, dim=0, start=0, end=4, step=2)
print("输入:", input)
print("源:", src)
print("步长为2的结果:", output)
输出结果为:
输入: 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.])
注意:
torch.slice_scatter不会修改原始输入张量,而是返回一个新的张量。该函数是
torch.slice的逆操作。

Pytorch torch 参考手册