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],