pytorch torch.gather函數(shù)的使用
pytorch torch.gather函數(shù)
torch.gather 是 PyTorch 中的一個(gè)用于從給定維度上按索引取值的函數(shù)。
它根據(jù)一個(gè)索引張量 index,從源張量 input 中收集值,并返回一個(gè)新的張量。
torch.gather 常用于需要從張量的特定位置抽取元素的操作。
1. 函數(shù)簽名
torch.gather(input, dim, index, *, sparse_grad=False, out=None)
input:輸入張量,表示要從中收集元素的源張量。dim:要收集的維度索引。例如,對(duì)于一個(gè)二維張量,0 表示沿著行的維度,1 表示沿著列的維度。index:索引張量,其形狀應(yīng)與input張量在除了dim維度之外的其他維度上保持一致。索引張量中的值表示在input張量對(duì)應(yīng)維度上要收集的元素的索引。out(可選):輸出張量,如果提供,結(jié)果將存儲(chǔ)在這個(gè)張量中。
2. 工作原理
torch.gather 在 dim 維度上,通過(guò) index 指定的索引,從 input 中選取元素。
返回的張量的形狀與 index 的形狀相同。
3. 示例代碼
以下是一個(gè)簡(jiǎn)單的示例代碼,演示如何使用 torch.gather 函數(shù):
import torch
# 創(chuàng)建一個(gè)源張量
input = torch.tensor([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])
# 創(chuàng)建一個(gè)索引張量
index = torch.tensor([[0, 2, 1],
[2, 0, 1],
[1, 2, 0]])
# 在 dim=1 維度上使用 gather 函數(shù)
result = torch.gather(input, dim=1, index=index)
print("Input Tensor:")
print(input)
print("\nIndex Tensor:")
print(index)
print("\nResult Tensor:")
print(result)4. 輸出結(jié)果
Input Tensor:
tensor([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])Index Tensor:
tensor([[0, 2, 1],
[2, 0, 1],
[1, 2, 0]])Result Tensor:
tensor([[1, 3, 2],
[6, 4, 5],
[8, 9, 7]])
5. 解釋
- 輸入張量 (
input) 是一個(gè)3x3的矩陣,每個(gè)元素代表一個(gè)值。 - 索引張量 (
index) 指定了要從input中提取的元素的索引。 - 結(jié)果張量 (
result) 是根據(jù)index從input中提取的元素形成的張量。
在這個(gè)例子中:
- 對(duì)于
input的第一行,index提取了索引0, 2, 1對(duì)應(yīng)的元素1, 3, 2。 - 對(duì)于
input的第二行,index提取了索引2, 0, 1對(duì)應(yīng)的元素6, 4, 5。 - 對(duì)于
input的第三行,index提取了索引1, 2, 0對(duì)應(yīng)的元素8, 9, 7。
總結(jié)
torch.gather 通過(guò)索引在指定維度上提取張量中的元素,是用于基于索引選擇數(shù)據(jù)的有用工具。
函數(shù)對(duì)批處理數(shù)據(jù)特別有用,例如在分類(lèi)任務(wù)中提取對(duì)應(yīng)類(lèi)別的概率或得分。
索引張量的形狀必須與源張量在指定維度的形狀相匹配,以確保正確的取值操作。
以上為個(gè)人經(jīng)驗(yàn),希望能給大家一個(gè)參考,也希望大家多多支持腳本之家。
相關(guān)文章
pandas篩選某列出現(xiàn)編碼錯(cuò)誤的解決方法
今天小編就為大家分享一篇pandas篩選某列出現(xiàn)編碼錯(cuò)誤的解決方法,具有很好的參考價(jià)值,希望對(duì)大家有所幫助。一起跟隨小編過(guò)來(lái)看看吧2018-11-11
使用Python Tkinter創(chuàng)建一個(gè)動(dòng)態(tài)祝福彈窗的詳細(xì)教程
本文手把手教你用Python的Tkinter庫(kù)創(chuàng)建一個(gè)浪漫的彈窗程序,包含淡入淡出動(dòng)畫(huà)、多線程管理、隊(duì)列控制等高級(jí)特性,通過(guò)完整的代碼解析和配置指南,帶你掌握GUI編程的核心技巧,需要的朋友可以參考下2025-11-11
Python中schedule模塊關(guān)于定時(shí)任務(wù)使用方法
這篇文章主要介紹了Python中schedule模塊關(guān)于定時(shí)任務(wù)使用方法,文章圍繞主題展開(kāi)詳細(xì)的內(nèi)容介紹,具有一定的參考價(jià)值,需要的小伙伴可以參考一下2022-05-05
Python 對(duì)象序列化與反序列化之pickle json詳細(xì)解析
我們知道在Python中,一切皆為對(duì)象,實(shí)例是對(duì)象,類(lèi)是對(duì)象,元類(lèi)也是對(duì)象。本文正是要聊聊如何將這些對(duì)象有效地保存起來(lái),以供后續(xù)使用2021-09-09
Python實(shí)現(xiàn)簡(jiǎn)單石頭剪刀布小游戲的示例代碼
石頭剪刀布是一種簡(jiǎn)單而又經(jīng)典的游戲,常常用于決定勝負(fù)或者娛樂(lè)消遣,本文將使用Python實(shí)現(xiàn)一個(gè)簡(jiǎn)單的石頭剪刀布游戲,需要的可以參考一下2023-06-06
在Python中合并字典模塊ChainMap的隱藏坑【推薦】
在Python中,當(dāng)我們有兩個(gè)字典需要合并的時(shí)候,可以使用字典的 update 方法,接下來(lái)通過(guò)本文給大家介紹在Python中合并字典模塊ChainMap的隱藏坑,感興趣的朋友一起看看吧2019-06-06
Python使用scrapy爬取陽(yáng)光熱線問(wèn)政平臺(tái)過(guò)程解析
這篇文章主要介紹了Python使用scrapy爬取陽(yáng)光熱線問(wèn)政平臺(tái)過(guò)程解析,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友可以參考下2019-08-08

