Pytorch中scatter与gather操作

本文详细解析PyTorch中的scatter和gather操作,包括scatter_、scatter_add_及ONNX中的scatterND,阐述它们如何用于张量数据的发散与聚集,并通过实例说明不同约束条件的应用场景。

数据发散scatter

函数原型pytorch官方文档scatter_

scatter_(dim, index, src) → Tensor

注: scatter_是scatter的就地操作。

对于一个三维的张量来说,张量self(即调用scatter_的张量)的更新公式如下所示:

self[index[i][j][k]][j][k] = src[i][j][k]  # if dim == 0
self[i][index[i][j][k]][k] = src[i][j][k]  # if dim == 1
self[i][j][index[i][j][k]] = src[i][j][k]  # if dim == 2

其中需要注意的是,scatter对张量self,张量index和张量src之间的维度关系有三个约束:

(1)张量self,张量index和张量src的维度数量必须相同(即三者的.dim()必须相等,注意不是维度大小);

(2)对于每一个维度d,有index.size(d)<=src.size(d)

(3)对于每一个维度d,如果d!=dim,有index.size(d)<=self.size(d)

同时,张量index中的数值大小也有2个约束:

(4)张量index中的每一个值大小必须在[0, self.size(dim)-1]之间;

(5)张量index沿dim维的那一行中所有值都必须是唯一的(弱约束,违反不会报错,但是会造成没有意义的操作)。

其实只要记住scatter的目的是将张量src中的值根据index放入到self中,这几个约束就很好理解,为了进一步方便理解,请看下面的例子:

例1:
例1.jpg

out = torch.zeros(4, 4)
index = torch.tensor([[2, 1],
                      [1, 3],
                      [0, 2],
                      [2, 1]])
src = torch.tensor([[1, 2],
                    [3, 4],
                    [5, 6],
                    [7, 8]]).float()
res = out.scatter_(1, index, src)
# tensor([[0., 2., 1., 0.],
#         [0., 3., 0., 4.],
#         [5., 0., 6., 0.],
#         [0., 8., 7., 0.]])

例2:
例2

out = torch.zeros(4, 4)
index = torch.tensor([[2, 1],
                      [1, 3],
                      [0, 2],
                      [3, 0]])
src = torch.tensor([[
评论 2
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值