JOVANA
Explore Library Glossary Getting Started Three Levels Fields How it works Mission
Join the mission
All guides

機器學習加速器內部:脈動陣列

上一篇的四條設計守則,到這裡就不再抽象了。我們打開一顆深度學習加速器,發現它做的幾乎所有事都是一場龐大的矩陣乘法,並看著一格格小小乘法器組成的網格——脈動陣列——把數字泵過自己,讓每個值一旦被取進來,就在離開之前做上數百次乘法。

工作負載只有一個形狀:矩陣乘法

上一篇給了你設計領域專用架構的四條守則:把電晶體花在更多算術單元上、用專用且由軟體管理的記憶體而非快取去餵它們、榨乾這個領域交到你手上的平行性,並降到應用能容忍的最低精度。本篇就看著這四條守則在那個招牌例子——深度學習加速器——上一次到位,而它們之所以落得如此乾淨,靠的是關於這個工作負載的一個單一事實。一張神經網路所做的幾乎每件事,無論訓練還是推論,都歸結為把矩陣相乘。一層把一個輸入向量變成一個輸出向量,就是一次矩陣對向量的乘積;一次處理一整批,它就變成矩陣對矩陣的乘積。卷積、注意力、那些龐大的全連接層——掀開蓋子,它們全是同一種算術。

這恰恰是領域專用守則所盼望的禮物:一個(每層數百萬次乘法)、穩定(同一個運算,一層接一層、一個模型接一個模型)、且大規模平行(每個輸出元素都是一次獨立的點積)的工作負載。一顆通用中央處理器把那些乘法的每一次都當成一條要去取指、解碼、排程的獨立指令;一顆 圖形處理器靠著數千條通道做得好上太多,卻仍要為了把運算元搬進搬出暫存器檔而花掉實打實的能量。加速器的賭注是:如果工作就只是矩陣乘法,你就能造出一塊除此之外什麼都不做的矽——並且做每一次乘法時幾乎沒有額外開銷。

真正的敵人是資料搬運,不是算術

這裡有個塑造了整個設計的微妙之處。在現代矽裡,一次乘法便宜得驚人——一個小乘法器很小、燒掉的能量極少。昂貴的是把要相乘的那兩個數字取進來。從晶片外的 DRAM 讀一個運算元,耗掉的能量可能是它所餵的那次乘法的數百倍,就算從晶片內的儲存讀,也讓那點算術相形見絀。這就是你很久以前見過的記憶體牆,如今戴的是能量的帽子,而非延遲的帽子。所以一顆加速器的優劣指標,不是它有幾個乘法器,而是它的算術強度——它每從記憶體拉進一個位元組,就執行幾次運算。

這樣框定目標,整個設計就會自己改寫一遍:訣竅是把每個數字只取進來一次,然後在它離開晶片之前,盡你所能從它身上榨出越多次乘法越好。在一場矩陣乘法裡,這種重用是巨大的——每個輸入列都餵進數百個輸出行,所以每個值應該被乘上數百次。問題純粹是結構性的:你要怎麼把乘法器接起來,好讓一個值被取進陣列一次之後,就自然地流經每一個需要它的乘法器,而完全不必再回記憶體跑一趟?那個結構性的答案,就是脈動陣列。

這片網格實際上是怎麼相乘的

想像一片二維網格,由一格格相同的小單元組成——Google 第一代 張量處理器那著名的網格是 256 乘 256,共 65,536 個乘法單元。每格在每個時脈滴答上只做一件卑微的事:它把兩個數字相乘、把乘積加進自己保有的一個累加總和裡,再把它的兩個輸入傳給鄰居——一個往右、一個往下。這個「相乘、累加、傳下去」就是一次乘加運算(multiply-accumulate,MAC),是點積的原子,也是這格單元唯一懂的運算。沒有取指、沒有解碼、沒有暫存器更名、沒有分支預測——前面那些階花了好久去打造的機械,一樣都沒有。這格單元幾乎是純算術。那份簡樸正是重點所在:因為沒有額外開銷,你才付得起 65,536 個。

現在說流動。加速器先把兩個矩陣中的一個——比方說這層的權重——載進網格,每格釘住一個權重,在整場矩陣乘法裡都待著不動。然後它把另一個矩陣,也就是活化值,從左邊緣串流進來:每一列輸入向右行軍橫越網格,每個時脈滴答前進一行單元,像一道波。當一個輸入值穿過一格單元時,那格把它乘上自己常駐的權重、加進往下流動的部分和裡,再把這個輸入交給它右邊的下一格。因為資料是錯位進場的(每一列都比上一列晚一個滴答開始),所以陣列一旦填滿,每一格在每個滴答上都在忙,而每個週期都有一個算完的行和從底部邊緣掉出來。

A 3x3 systolic array, weights W pinned in each cell.
Activations a enter from the LEFT, skewed by one tick per row.
Partial sums flow DOWN. Each cell does: sum += a * W; pass a right.

   a0 ->[W00]->[W01]->[W02]      (row 0 enters at tick 0)
         |     |     |
   a1 ->[W10]->[W11]->[W12]      (row 1 enters at tick 1)
         |     |     |
   a2 ->[W20]->[W21]->[W22]      (row 2 enters at tick 2)
         v     v     v
      out0  out1  out2   <- finished dot-products fall out the bottom

Each activation, fetched ONCE into the left edge, is reused by every
cell in its row. Each weight, loaded ONCE, is reused by every
activation that streams past it. One fetch -> hundreds of multiplies.
權重不動,活化值流過去;一個值取進來一次,就被它經過的每一格相乘,所以幾乎沒有任何運算元會跑第二趟記憶體。

追蹤一個輸入值,去感受這份贏面。數字 a0 進入左上那格、在那裡被乘一次,然後往右流、在下一格又被乘一次,再下一格——一次取進,接著是一整列的乘法,全程不再碰記憶體。同時,每個被釘住的權重都被流經它的每一個活化值重用。這就是把算術強度那個目標化為實體:一個值進陣列一次,就在它路徑上的每一格參與一次乘法。原本主宰能量預算的資料搬運,被相鄰單元之間的短跳取代了——是皮米級的導線,而不是跑一趟 DRAM。

更低的精度:第四條守則加倍回本

最後一條設計守則是「用應用能容忍的最低精度」,而神經網路能容忍的不少。一個訓練好的模型是統計性的、帶雜訊的;它不需要科學計算所要求的那種完整 32 位元浮點數。所以加速器以降精度運行:bfloat16,一種 16 位元的浮點數,保留了 32 位元浮點完整的 8 位元指數範圍(所以在大的那個不會溢位的地方,它也絕不溢位),卻丟掉大部分網路根本不會想念的尾數位元;以及 int8,一種 8 位元整數,在模型從浮點被量化下來後用於推論。

這回本回了兩次,而那很容易被漏看。第一,給 8 位元數字用的乘法器,遠比 32 位元的更小、更省能——乘法器的面積大致隨位元寬度的平方成長,所以從 32 位元降到 8 位元,能把每格單元縮小一個數量級,讓你在同一塊矽裡塞進多得多的乘加單元。第二,現在每個運算元只剩四分之一或八分之一的位元組,於是同樣寶貴的記憶體頻寬,每秒能載運多上四到八倍的數字。降精度同時攻擊算術成本資料搬運成本——這兩件事,正是整個架構存在所要去極小化的。

軟體是這台機器的一半,它的極限也是

脈動陣列是硬到極點的簡單硬體,這意味著它把大量工作推給了軟體。陣列內部沒有任何東西決定要取什麼、何時載入權重、或要怎麼把一層大到塞不進網格的層,切成裝得下的小塊——那塊暫存記憶體是由軟體管理的,所以編譯器必須調度每一個位元組的到來與離去。這正是最純粹形態的軟硬體協同設計:晶片被刻意做笨,好讓編譯器可以聰明,而兩者是當成一個系統一起設計出來的。一個弱編譯器讓那 65,536 格單元閒著;一個好的則讓陣列吃飽、嗡嗡作響。這顆加速器,真的就只是這項產品的一半而已。

也要對這一切專門化的代價誠實。脈動陣列在稠密矩陣乘法上美得令人屏息,在其餘幾乎所有事情上卻無能為力——它跑不了你的作業系統、排不了一串清單、也剖析不了文字。更糟的是,它的效率取決於讓那片網格保持滿載:餵它小矩陣,或餵它無法整齊鋪到 256 乘 256 陣列上的形狀,大多數單元就閒著,而你仍得付它們的功率。這恰恰是把領域專用的那筆交易重講一遍——加速器只有在工作負載又大、又穩定、又平行時才贏。一旦工作變小、變不規則、或變得多分支,一顆 圖形處理器或一顆樸素的中央處理器才是更好的工具。沒有萬用的加速器,只有一個對得起某個配得上它的工作負載的加速器。