作用
scatter是“散开”的意思,顾名思义,是将一个Tensor按照index做分散。
形式
在pytorch中,scatter可以通过torch.scatter和torch.scatter_(修改自身数据),或者Tensor自生就有的方法scatter
Tensor.scatter_(dim, index, src, reduce=None) → Tensor
参数
-
input
输入参数,如果是通过Tensor直接调用的,没有该参数(就是自身嘛),仅仅在torch.scatter/torch.scatter_中需要指定 -
dim
维度,需要用于在哪一个维度 -
index
为Tensor做scatter时需要的索引,注意,这个index必须是int64型。另外它的-1维必须与input的-1维一致。 -
src
用于“分散”的源数据,它的shape必须与input的shape一致
它将用于具体的计算,公式如下

-
value
一个0-1.0之间的float数,它与src只要一个即可。如果指定了value,则按照value这个值去填充index指定的位置。
使用样例
一维数据
input = torch.zeros(3)
index = torch.LongTensor([0,2,1])
input =

本文详细介绍了PyTorch中scatter函数的功能与用法,包括如何使用scatter进行张量的分散操作,提供了多个示例帮助理解不同场景下scatter的使用方法。

3387

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



