gather筛选规则:

import torch
data = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
indices = torch.tensor([0, 2]) # 在轴上筛选坐标
out = torch.gather(data, dim=0, index=torch.tensor([[0,1],[1,2]]))
print(out)
结果:
tensor([[1, 2, 3],
[7, 8, 9]])
tensor([[1, 5],
[4, 8]])
筛选规则:dim为0,表示筛选是行,列是自动的,按照当前位置来映射
0,0 1,1
1,0 2,1
按列筛选:
列参数是给定的,行是自动按照当前位置来同步映射。
import torch
data = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
indices = torch.tensor([0, 2]
本文介绍了PyTorch中的torch.gather函数,通过示例解释了该函数如何进行数据筛选。当dim为0时,筛选按行进行;当dim为1时,筛选按列进行。详细解析了索引的使用方式,并配以图表帮助理解。
订阅专栏 解锁全文
7万+

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



