导读:学过《数据结构》的我们都知道,从数组中取数靠的数组索引,而在PyTorch中也不例外,但是对于“数组索引”这个东西,PyTorch进行了深加工,从而增加了理解难度,大家要多加注意。gather函数是PyTorch常用的取数工具,它可以帮助我们从tensor中取出指定乱序索引下的数据。
数组所传递的信息有两个:数组的值是什么,以及数组值所对应的索引是什么。如下所示:
index = torch.tensor([[2, 1, 0]])
其表达的意思是:
$$index = [[ \frac{ 2 }{(0,0)}, \frac{ 1 }{(0,1)}, \frac{ 0 }{(0,2)} ]] $$
当dim=0的时候,用上半部分的数替换下半部分左侧的数,如下所示:
$$index = [[ \frac{ \color{#F00}{2} }{( \color{#F00}{2} , \color{#000}{0} )}, \frac{ \color{#F00}{1} }{(\color{#F00}{1},\color{#000}{1})}, \frac{ \color{#F00}{0} }{(\color{#F00}{0},\color{#000}{2})} ]] $$
当dim=1的时候,用上半部分的数替换下半部分右侧的数,如下所示:
$$index = [[ \frac{ \color{#F00}{2} }{( \color{#000}{0} , \color{#F00}{2} )}, \frac{ \color{#F00}{1} }{(\color{#000}{0},\color{#F00}{1})}, \frac{ \color{#F00}{0} }{(\color{#000}{0},\color{#F00}{0})} ]] $$
如上文所示,PyTorch可以获得新的索引值,根据这个索引值,利用gather函数,就可以任意取数了。下面请看一下gather函数的语法定义:
torch.gather(input, dim, index, *, sparse_grad=False, out=None)
其中,dim和index参数对应上文提到的:PyTorch对数组的索引的深加工。
最后需要注意的一点是:输出out要与索引index具有相同的形状。
下面请看一下gather函数的例子介绍:
import torch
data = torch.arange(1, 10).view(3, 3)
print("原始数据:\n", data)
index = torch.tensor([[2, 1, 0]])
result = data.gather(dim=0, index=index)
print("dim=0取数结果:", result)
result = data.gather(dim=1, index=index)
print("dim=1取数结果:", result)
结果为:
原始数据:
tensor([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])
dim=0取数结果: tensor([[7, 5, 3]])
dim=1取数结果: tensor([[3, 2, 1]])