最新国产好看的视频,伊人天堂AV在线,国产Aaaaaa视频,蜜臀视频在线观看一区,人妻av色图,密臀久久久精品影片,青青视频免费观看毛片,久草在线观看视,国产三级精品色情在线

Pytorch中的Broadcasting問(wèn)題

 更新時(shí)間:2023年01月03日 09:53:28   作者:luputo  
這篇文章主要介紹了Pytorch中的Broadcasting問(wèn)題,具有很好的參考價(jià)值,希望對(duì)大家有所幫助。如有錯(cuò)誤或未考慮完全的地方,望不吝賜教

Numpy、Pytorch中的broadcasting

寫在前面

自己一直都不清楚numpy、pytorch里面不同維數(shù)的向量之間的element wise的計(jì)算究竟是按照什么規(guī)則來(lái)確認(rèn)維數(shù)匹配和不匹配的情況的,比如

>>> b = np.ones((4,5))
>>> a = np.arange(5)
>>> c = a + b
>>> c.shape
(4, 5)
>>> c
array([[1., 2., 3., 4., 5.],
? ? ? ?[1., 2., 3., 4., 5.],
? ? ? ?[1., 2., 3., 4., 5.],
? ? ? ?[1., 2., 3., 4., 5.]])

上面這種情況就會(huì)自動(dòng)讓a和b的維數(shù)匹配,a加到了b的每一行上

>>> b = np.ones((5,4))
>>> a = np.arange(5)
>>> c = a + b
Traceback (most recent call last):
? File "<stdin>", line 1, in <module>
ValueError: operands could not be broadcast together with shapes (5,) (5,4)

這種情況就無(wú)法匹配,此時(shí)我們希望的是a能自動(dòng)加到b的每一列上,但結(jié)果看來(lái)好像不行

雖然一直存在這種疑惑,但因?yàn)槠綍r(shí)遇到的各種運(yùn)算都比較簡(jiǎn)單,遇到這種不是直接匹配的array的加法第一直覺(jué)就是去console里面試一試,報(bào)錯(cuò)就換個(gè)姿勢(shì)再試一試,總歸問(wèn)題可以快速地解決,但是最近在寫模型的時(shí)候,遇到了繞不過(guò)去的問(wèn)題,所以去查了文檔,本文就以解決那個(gè)問(wèn)題為目標(biāo),來(lái)解釋清楚pytorch(numpy也是一樣)中的broadcasting semantics的問(wèn)題

問(wèn)題描述

我有一個(gè)數(shù)據(jù)Tensor,維數(shù)是64 × 2048 64\times204864×2048,現(xiàn)在我想通過(guò)對(duì)這64 6464個(gè)2048 20482048維的向量做attention(也就是做一個(gè)加權(quán)和)來(lái)得到一個(gè)2048 20482048維的向量,因?yàn)槟P偷男枰?,我需要用五組不同的權(quán)值向量來(lái)計(jì)算出五個(gè)不同的加權(quán)結(jié)果,也就是我的計(jì)算結(jié)果應(yīng)該是一個(gè)5 × 2048 5\times 20485×2048維的向量,因?yàn)樵?4 6464個(gè)向量上加權(quán),所以一組權(quán)值向量是64 6464維,五組就是5 × 64 5\times 645×64維

嘗試解決

現(xiàn)在我手頭上有兩個(gè)Tensor,一個(gè)是數(shù)據(jù)Tensor(64 × 2048 64\times 204864×2048)另一個(gè)是權(quán)值Tensor(5 × 64 5\times 645×64),我GAN!直到我寫到了這里,我才發(fā)現(xiàn)這不是一個(gè)矩陣乘法就能解決的問(wèn)題嘛+_+,當(dāng)然,我想給自己正名,這里我簡(jiǎn)化了一下問(wèn)題所以才發(fā)現(xiàn)原來(lái)這么容易就解決了,而原來(lái)我在寫代碼的時(shí)候因?yàn)檫€要考慮batch_size等問(wèn)題才云里霧里不知道咋辦,還好當(dāng)時(shí)沒(méi)想出來(lái),所以去查了文檔發(fā)現(xiàn)了新的東西,然后寫文章的時(shí)候想到也算是完滿了(不然也不會(huì)發(fā)現(xiàn)自己好澇)

以上都是題外話,現(xiàn)在,我們還是考慮用愚蠢的element wise的方法來(lái)解決,好在現(xiàn)在有兩種方法可以解決問(wèn)題,所以我們可以用來(lái)相互檢驗(yàn)一下,element wise的解決方法就是,我希望這5個(gè)64維的權(quán)值向量分別和這64個(gè)2048維的向量進(jìn)行element wise的乘法,也就是第一個(gè)64維權(quán)值向量先對(duì)64個(gè)2048維向量加權(quán)得到一個(gè)2048維的向量,然后第二個(gè)64維權(quán)值向量先對(duì)64個(gè)2048維向量加權(quán)得到一個(gè)2048維的向量…,以此類推總共五個(gè),最終得到五個(gè)64 × 2048 64×204864×2048維的向量,然后求和得到最后的5 × 2048 5×20485×2048維的向量

那么按照平常的習(xí)慣,我就去先試試pytorch能不能直接地理解我的想法

>>> import torch
>>> bs = 10 # batch_size
>>> x = torch.randn(bs,64,2048)
>>> att = torch.randn(5,64)
>>> out = att * x
Traceback (most recent call last):
? File "<stdin>", line 1, in <module>
RuntimeError: The size of tensor a (64) must match the size of tensor b (2048) at non-singleton dimension 2

直接乘不行,因?yàn)榫S數(shù)是不匹配的,那怎樣的維數(shù)才算匹配呢?

BROADCASTING SEMANTICS

以下內(nèi)容主要來(lái)源于自官方文檔

很多pytorch的運(yùn)算是支持broadcasting semantics的,而簡(jiǎn)單來(lái)說(shuō),如果運(yùn)算支持broadcast,則參與運(yùn)算的Tensor會(huì)自動(dòng)進(jìn)行擴(kuò)展來(lái)使得運(yùn)算符左右的Tensor維數(shù)匹配,而無(wú)需人手動(dòng)地去拷貝其中的某個(gè)Tensor,這就類似于我們開(kāi)頭的那個(gè)例子

>>> b = np.ones((4,5))
>>> a = np.arange(5)
>>> c = a + b
>>> c.shape
(4, 5)
>>> c
array([[1., 2., 3., 4., 5.],
? ? ? ?[1., 2., 3., 4., 5.],
? ? ? ?[1., 2., 3., 4., 5.],
? ? ? ?[1., 2., 3., 4., 5.]])

我們無(wú)需讓a的維數(shù)和b一樣,因?yàn)閚umpy自動(dòng)幫我們做了

這里的另一個(gè)重要的概念是broadcastable,如果兩個(gè)Tensor是broadcastable的,那么就可以對(duì)他倆使用支持broadcast的運(yùn)算,比如直接加減乘除

而兩個(gè)向量要是broadcast的話,必須滿足以下兩個(gè)條件

  • 每個(gè)tensor至少是一維的
  • 兩個(gè)tensor的維數(shù)從后往前,對(duì)應(yīng)的位置要么是相等的,要么其中一個(gè)是1,或者不存在

這是官方的例子解釋

>>> x=torch.empty(5,7,3)
>>> y=torch.empty(5,7,3)
# 相同維數(shù)的tensor一定是broadcastable的

>>> x=torch.empty((0,))
>>> y=torch.empty(2,2)
# 不是broadcastable的,因?yàn)槊總€(gè)tensor維數(shù)至少要是1

>>> x=torch.empty(5,3,4,1)
>>> y=torch.empty( ?3,1,1)
# 是broadcastable的,因?yàn)閺暮笸翱?,一定要注意是從后往前看?
# 第一個(gè)維度都是1,相等,滿足第二個(gè)條件
# 第二個(gè)維度其中有一個(gè)是1,滿足第二個(gè)條件
# 第三個(gè)維度都是3,相等,滿足第二個(gè)條件
# 第四個(gè)維度其中有一個(gè)不存在,滿足第二個(gè)條件

# 但是
>>> x=torch.empty(5,2,4,1)
>>> y=torch.empty( ?3,1,1)
# 不是broadcastable的,因?yàn)閺暮笸翱吹谌齻€(gè)維度是不match的 2!=3,且都不是1

如果x和y是broadcastable的,那么結(jié)果的tensor的size按照如下的規(guī)則計(jì)算

  • 如果兩者的維度不一樣,那么就自動(dòng)增加1維(也就是unsqueeze)
  • 對(duì)于結(jié)果的每個(gè)維度,它取x和y在那一維上的最大值

官方的例子

>>> x=torch.empty(5,1,4,1)
>>> y=torch.empty( ?3,1,1)
>>> (x+y).size()
torch.Size([5, 3, 4, 1])

>>> x=torch.empty(1)
>>> y=torch.empty(3,1,7)
>>> (x+y).size()
torch.Size([3, 1, 7])

>>> x=torch.empty(5,2,4,1)
>>> y=torch.empty(3,1,1)
>>> (x+y).size()
RuntimeError: The size of tensor a (2) must match the size of tensor b (3) at non-singleton dimension 1

此外,關(guān)于broadcast導(dǎo)致的就地(in-place)操作和梯度運(yùn)算的兼容性等問(wèn)題,可以自行參考官方文檔

解決問(wèn)題

上面我們看到,要想兩個(gè)Tensor支持element wise的運(yùn)算,需要它們是broadcastable的,而要想它們是broadcastable的,就需要它們的維度自后向前逐一匹配,回到我們?cè)瓉?lái)的問(wèn)題中,我們有兩個(gè)Tensor x(64 × 2048) att(5 × 64),為了讓它們broadcastable,我們只需要

>>> import torch
>>> bs = 10 # batch_size
>>> x = torch.randn(bs,64,2048)
>>> att = torch.randn(5,64)
>>> x = x.unsqueeze(1)
>>> att = att.view(1,*att.shape,1)
>>> x.shape
torch.Size([10, 1, 64, 2048])
>>> att.shape
torch.Size([1, 5, 64, 1])
>>> out = x * att
>>> out.shape
torch.Size([10, 5, 64, 2048])

最后我們來(lái)驗(yàn)證兩種方法是否結(jié)果相同

>>> import torch
>>> bs = 10?
>>> x = torch.randn(bs,64,2048)
>>> att = torch.randn(5,64)
>>> out1 = torch.matmul(att,x) ?# 直接矩陣相乘
>>> out.shape
torch.Size([10, 5, 2048])

>>> x = x.unsqueeze(1)
>>> att = att.view(1,*att.shape,1)
>>> out2 = x * att ?# element wise的方法
>>> out2 = out2.sum(dim=2)

>>> test = torch.sum((out1-out2)<0.00001) ?# 浮點(diǎn)數(shù)有微小的誤差
>>> test
tensor(102400)
>>> out1.numel() ?# 最后表明兩個(gè)out向量是相等的
102400

Reference

[1] https://docs.scipy.org/doc/numpy/user/basics.broadcasting.html#module-numpy.doc.broadcasting

[2] https://pytorch.org/docs/stable/notes/broadcasting.html#broadcasting-semantics

總結(jié)

以上為個(gè)人經(jīng)驗(yàn),希望能給大家一個(gè)參考,也希望大家多多支持腳本之家。

相關(guān)文章

最新評(píng)論

太保市| 林西县| 田阳县| 桑植县| 夹江县| 厦门市| 和田县| 浮梁县| 岑巩县| 普兰店市| 雷州市| 太原市| 夏津县| 古田县| 新邵县| 盐池县| 瓦房店市| 故城县| 达尔| 西平县| 陆川县| 乃东县| 岑溪市| 姜堰市| 蒙城县| 行唐县| 屏东市| 高州市| 通化县| 永平县| 昌邑市| 望都县| 马山县| 彩票| 德令哈市| 宁津县| 珲春市| 威信县| 公安县| 西安市| 迁安市|