当前位置:首页> AI教程> Pytorch中torch.gather的用法和代码示例

Pytorch中torch.gather的用法和代码示例

释放双眼,带上耳机,听听看~!
了解Pytorch中torch.gather算子的用法和代码示例,帮助初学者快速掌握如何使用该功能。本文详细解释了torch.gather的用法,并给出了实际的代码示例,帮助读者更好地理解。

本文已参与「新人创作礼」活动,一起开启掘金创作之路。

一、用法:

torch.gather 算子用于返回给定索引/下标的 Tensor 元素,在 pytorch 官网文档中的定义如下:

torch.gather( input, dim, index, *, sparse_grad=False, out=None) → Tensor

其用法等价于:

input.gather( dim, index, *, sparse_grad=False, out=None) → Tensor

其中,input 是目标 Tensor ,即被搜索的 Tensor ;dim 是搜索维度(也是 Tensor ),index 是索引。

返回值类型:Tensor

二、代码示例:

概念看不懂没关系,一看代码便知用法。

a = torch.tensor([1, 5, 3, 6, 8])
b = torch.tensor([3])    # 索引为3
c = a.gather(0, b)    # 输出a中第0维索引是3的元素:6
# 等价于 c=torch.gather(a,0,b)
print(c)     # tensor([6])
a = torch.tensor([[1.3, 2, 3, 4.5, 5],
                  [2.0, 3, 0.3, 4.1, 2],
                  [6, 7, 8, 9, 2],
                  [10, 5, 0, 6, 8]])
b = torch.tensor([[1],
                  [2],
                  [3],
                  [4]])
c = torch.gather(a, 1, b)    # 输出a中第1维索引分别是1,2,3,4的元素:2,0.3,9,8
print(c)      # tensor([[2.0000],[0.3000],[9.0000],[8.0000]])
本网站的内容主要来自互联网上的各种资源,仅供参考和信息分享之用,不代表本网站拥有相关版权或知识产权。如您认为内容侵犯您的权益,请联系我们,我们将尽快采取行动,包括删除或更正。
AI教程

分布式文件系统下的回收站设计及性能优化

2023-12-20 23:08:14

AI教程

MegEngine开源4位量化模型:实现精度与速度的双赢

2023-12-21 0:58:14

个人中心
购物车
优惠劵
今日签到
有新私信 私信列表
搜索