pytorch torch.gather函數(shù)的使用
pytorch torch.gather函數(shù)
torch.gather 是 PyTorch 中的一個用于從給定維度上按索引取值的函數(shù)。
它根據(jù)一個索引張量 index,從源張量 input 中收集值,并返回一個新的張量。
torch.gather 常用于需要從張量的特定位置抽取元素的操作。
1. 函數(shù)簽名
torch.gather(input, dim, index, *, sparse_grad=False, out=None)
input:輸入張量,表示要從中收集元素的源張量。dim:要收集的維度索引。例如,對于一個二維張量,0 表示沿著行的維度,1 表示沿著列的維度。index:索引張量,其形狀應(yīng)與input張量在除了dim維度之外的其他維度上保持一致。索引張量中的值表示在input張量對應(yīng)維度上要收集的元素的索引。out(可選):輸出張量,如果提供,結(jié)果將存儲在這個張量中。
2. 工作原理
torch.gather 在 dim 維度上,通過 index 指定的索引,從 input 中選取元素。
返回的張量的形狀與 index 的形狀相同。
3. 示例代碼
以下是一個簡單的示例代碼,演示如何使用 torch.gather 函數(shù):
import torch
# 創(chuàng)建一個源張量
input = torch.tensor([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]])
# 創(chuàng)建一個索引張量
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) 是一個3x3的矩陣,每個元素代表一個值。 - 索引張量 (
index) 指定了要從input中提取的元素的索引。 - 結(jié)果張量 (
result) 是根據(jù)index從input中提取的元素形成的張量。
在這個例子中:
- 對于
input的第一行,index提取了索引0, 2, 1對應(yīng)的元素1, 3, 2。 - 對于
input的第二行,index提取了索引2, 0, 1對應(yīng)的元素6, 4, 5。 - 對于
input的第三行,index提取了索引1, 2, 0對應(yīng)的元素8, 9, 7。
總結(jié)
torch.gather 通過索引在指定維度上提取張量中的元素,是用于基于索引選擇數(shù)據(jù)的有用工具。
函數(shù)對批處理數(shù)據(jù)特別有用,例如在分類任務(wù)中提取對應(yīng)類別的概率或得分。
索引張量的形狀必須與源張量在指定維度的形狀相匹配,以確保正確的取值操作。
以上為個人經(jīng)驗,希望能給大家一個參考,也希望大家多多支持腳本之家。
相關(guān)文章
使用Python Tkinter創(chuàng)建一個動態(tài)祝福彈窗的詳細教程
本文手把手教你用Python的Tkinter庫創(chuàng)建一個浪漫的彈窗程序,包含淡入淡出動畫、多線程管理、隊列控制等高級特性,通過完整的代碼解析和配置指南,帶你掌握GUI編程的核心技巧,需要的朋友可以參考下2025-11-11
Python中schedule模塊關(guān)于定時任務(wù)使用方法
這篇文章主要介紹了Python中schedule模塊關(guān)于定時任務(wù)使用方法,文章圍繞主題展開詳細的內(nèi)容介紹,具有一定的參考價值,需要的小伙伴可以參考一下2022-05-05
Python 對象序列化與反序列化之pickle json詳細解析
我們知道在Python中,一切皆為對象,實例是對象,類是對象,元類也是對象。本文正是要聊聊如何將這些對象有效地保存起來,以供后續(xù)使用2021-09-09
在Python中合并字典模塊ChainMap的隱藏坑【推薦】
在Python中,當(dāng)我們有兩個字典需要合并的時候,可以使用字典的 update 方法,接下來通過本文給大家介紹在Python中合并字典模塊ChainMap的隱藏坑,感興趣的朋友一起看看吧2019-06-06

