数据发散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:

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:

out = torch.zeros(4, 4)
index = torch.tensor([[2, 1],
[1, 3],
[0, 2],
[3, 0]])
src = torch.tensor([[

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

2044

被折叠的 条评论
为什么被折叠?



