现在位置: 首页 > PyTorch 教程 > 正文

PyTorch torch.slice_scatter 函数


Pytorch torch 参考手册 Pytorch torch 参考手册

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 的各个参数:

参数类型是否必填默认值说明
inputTensor必填输入张量,即要修改的张量。
srcTensor必填源张量,要散布到 input 切片中的值,大小必须与切片后的大小一致
dimint可选0散布的维度。
startint可选None起始索引(包含该索引),默认从 0 开始。
endint可选None结束索引(不包含该索引),默认到末尾结束。
stepint可选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)

输出结果为:

输入张量:
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)

输出结果为:

输入: 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)

输出结果为:

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