pytorch之torch_scatter.scatter_max()用法
torch_scatter.scatter_max()
torch_scatter.scatter_max(src, index, dim=-1, out=None, dim_size=None, fill_value=None)

- 根據(jù)index將src分組,求每一組中的最大值輸出到out
- dim是維度

from torch_scatter import scatter_max src = torch.Tensor([[2, 0, 1, 4, 3], [0, 2, 1, 3, 4]]) index = torch.tensor([[4, 5, 4, 2, 3], [0, 0, 2, 2, 1]]) out = src.new_zeros((2, 6)) '''src根據(jù)index進行分組''' out, argmax = scatter_max(src, index, out=out) print(out) print(argmax)
輸出
tensor([[0., 0., 4., 3., 2., 0.],
[2., 4., 3., 0., 0., 0.]])
tensor([[-1, -1, 3, 4, 0, 1],
[ 1, 4, 3, -1, -1, -1]])
解釋

torch_scatter.scatter()使用
1. 參數(shù)

具體來講,scatter函數(shù)的作用就是將index中相同索引對應(yīng)位置的src元素進行某種方式的操作,例如 sum 、 mean 等,然后將這些操作結(jié)果按照索引順序進行拼接。
下面我用具體的例子來進行講解。
2. 示例
2.1 簡單示例
首先初始化src和index:
src = torch.Tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]]) # (3, 3) index = torch.tensor([0, 0, 1], dtype=torch.int64)
接著使用scatter函數(shù):
out = scatter(src, index, dim=0, reduce='mean')
我們觀察 index=[0, 0, 1] ,第0個位置和第1個位置都為0,第2個位置為1。也就是說,我們需要將src中第0個元素和第1個元素求平均變成一個元素,然后第2個元素求mean也就是本身為一個元素。如果 index=[1, 0, 0] ,則意味著我們需要將src中第1個元素和第2個元素求平均變成一個元素,而第0個元素保持不變。
那么src中第幾個元素到底是如何定義的呢?這就需要用到 dim 參數(shù)了。
dim=0 意味著我們需要對src的維度0進行操作:
tensor([[1., 2., 3.],
[4., 5., 6.],
[7., 8., 9.]])即src中第0個元素為 [1, 2, 3] ,第1個元素為 [4, 5, 6] ,第2個元素為 [7, 8, 9] 。
而如果 dim=1 ,則第0個元素為 [1, 4, 7] ,第1個元素為 [2, 5, 8] ,第2個元素為 [3, 6, 9] 。
因此,如果有以下代碼:
src = torch.Tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]]) # (3, 3) index = torch.tensor([0, 0, 1], dtype=torch.int64) out = scatter(src, index, dim=0, reduce='mean')
那么我們就應(yīng)該將src中的第0個元素為 [1, 2, 3] 和第1個元素為 [4, 5, 6] 求平均為 [2.5, 3.5, 4.5] ,然后第2個元素 [7, 8, 9] 保持不變,即:
tensor([[2.5000, 3.5000, 4.5000],
[7.0000, 8.0000, 9.0000]])2.2 順序問題
上面的例子中 index=[0, 0, 1] ,最后結(jié)果是將src中第0個元素和第1個元素求平均放到了位置0,然后src中第2個元素保持不變放到了位置1。
如果 index=[1, 1, 0] ,結(jié)果為:
tensor([[7.0000, 8.0000, 9.0000],
[2.5000, 3.5000, 4.5000]])可以發(fā)現(xiàn),上述結(jié)果是將src中第2個元素 [7, 8, 9] 保持不變放到了位置0,然后將src中第0個元素 [1, 2, 3] 和第1個元素 [4, 5, 6] 求平均保持不變放到了位置1。
也就是說,無論index怎么變化,都是優(yōu)先將index中0對應(yīng)位置的操作結(jié)果進行放置。
2.3 維度問題
如果src的維度為(4, 3),而我們需要對 dim=0 操作,也就是一共有四個元素,那么index的長度應(yīng)該為4,即以下操作是不合法的:
src = torch.Tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9], [10, 11, 12]]) # (4, 3) index = torch.tensor([1, 1, 0], dtype=torch.int64) out = scatter(src, index, dim=0, reduce='mean') print(out)
報錯為:
RuntimeError: The expanded size of the tensor (4) must match the existing size (3) at non-singleton dimension 0. Target sizes: [4, 3]. Tensor sizes: [3, 1]
正確做法應(yīng)該是:
src = torch.Tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9], [10, 11, 12]]) # (4, 3) index = torch.tensor([1, 1, 0, 2], dtype=torch.int64) out = scatter(src, index, dim=0, reduce='mean') print(out)
輸出為:
tensor([[ 7.0000, 8.0000, 9.0000],
[ 2.5000, 3.5000, 4.5000],
[10.0000, 11.0000, 12.0000]])
總結(jié)
以上為個人經(jīng)驗,希望能給大家一個參考,也希望大家多多支持腳本之家。
相關(guān)文章
使用Python設(shè)置PDF中圖片的透明度的實現(xiàn)方法
在PDF文檔的設(shè)計與內(nèi)容創(chuàng)作過程中,圖像的透明度設(shè)置是一個重要的操作,尤其是在處理圖文密集型PDF文檔時,本文將介紹如何使用Python添加指定透明度的圖片到PDF文檔或調(diào)整PDF文檔中現(xiàn)有圖片的透明度,需要的朋友可以參考下2024-09-09
python如何實現(xiàn)excel數(shù)據(jù)添加到mongodb
本文介紹了python是如何實現(xiàn)excel數(shù)據(jù)添加到mongodb,為了將數(shù)據(jù)導入mongodb,引入了pymongo,xlrd包,需要的朋友可以參考下2015-07-07
torch.utils.data.DataLoader與迭代器轉(zhuǎn)換操作
這篇文章主要介紹了torch.utils.data.DataLoader與迭代器轉(zhuǎn)換操作,文章內(nèi)容接受非常詳細,對正在學習或工作的你有一定的幫助,需要的朋友可以參考一下2022-02-02
python調(diào)用有道智云API實現(xiàn)文件批量翻譯
這篇文章主要介紹了python如何調(diào)用有道智云API實現(xiàn)文件批量翻譯,幫助大家更好得理解和使用python,感興趣的朋友可以了解下2020-10-10

