一個頭不夠:多頭注意力
在第 2 篇指南中,我們建立了單一的自注意力運算:每個 patch 都產生一個 query、一個 key 與一個 value,再讓每個 patch 自行決定要多認真聆聽其他每個 patch。這個運算很強大,但它逼迫所有比較都發生在同一個共享的「空間」裡。想像你要求一個人同時追蹤整張影像的顏色相似度、空間鄰近度「以及」材質——他做得到,但最後會把這些非常不同的「相關」概念混成一個模糊的平均判斷。
多頭注意力透過「並行」執行多個注意力運算來解決這個問題,這些並行的運算稱為頭(head)。每個頭都有自己學到的 Q/K/V 投影矩陣,因此每個頭都透過自己的鏡頭去看這些 patch。某個頭可能學會關注顏色相近的 patch,另一個關注空間相鄰的 patch,第三個關注重複出現的材質。比喻是:與其只用一個過勞的通才,不如召集一個由專家組成的小組,每位專家從自己的角度檢視同一個場景,最後再把他們的筆記彙整起來。
一張圖顯示一個輸入張量分支成數個並行的注意力頭,每個頭都有自己的 Q、K、V 投影,其輸出被串接後再通過最終的輸出投影。
多頭注意力即為並行的注意力頭,經串接後再投影。
讓我們拆解每個符號。共有 h 個頭(也就是專家的數量)。對第 i 個頭而言,矩陣 W_Q^i, W_K^i, W_V^i 是學到的投影,把完整的 D 維 token 投影到一個較小的維度 d_k = D/h。在每個頭內部,\mathrm{Attention}(\cdot) 正是第 2 篇的縮放點積公式:\mathrm{Attention}(Q,K,V)=\mathrm{softmax}\!\big(\tfrac{QK^\top}{\sqrt{d_k}}\big)V——唯一的改變是 Q,K,V 現在是該頭自己較小的投影。每個頭輸出一個 N \times d_k 的結果。接著我們把這 h 個頭並排 `Concat`(串接),將 h 段寬度為 d_k 的片段重新拼回寬度 h \cdot d_k = D。最後 W_O(一個 D \times D 的矩陣)把各個頭的發現混合成每個 token 一致的輸出。
用 ViT-Base 來看具體數字。模型維度為 D = 768,我們用 h = 12 個頭。那麼每個頭在維度 d_k = 768 / 12 = 64 中運作。於是這個龐大、寬度 768 的向量被切成 12 條寬度為 64 的條帶;每條各自執行自己的注意力;這 12 個寬度 64 的輸出再被串接回 12 \times 64 = 768;最後由 W_O(768×768)混合。注意,成本大致和一個 768 維的大注意力相同——我們並沒有付出 12 倍的計算量,只是把同樣的預算重新組織成 12 個專門化的子空間。
我在哪裡?位置編碼
我們目前建立的注意力有一個微妙但關鍵的問題:它具有排列不變性(permutation-invariant)。注意力純粹根據每個 patch 向量的內容來計算它與其他每個 patch 的相關程度——它從不去看一個 patch 原本位於影像中的何處。如果你把這些 patch 打亂成隨機順序,注意力會產生完全相同的一組關係,只是順序被重排。對注意力而言,這些 patch 是一個無序的袋子,而不是一個網格。
對視覺而言,這是一場災難。左上角的 patch 和右下角的 patch 是非常不同的東西,即使它們碰巧包含相似的像素;鼻子在嘴巴上方構成一張臉,相同的部件被打亂後就不是了。比喻是:有人把一格格的漫畫打亂成一疊交給你。圖畫都在,但沒有格子編號,你就無法重組故事。位置編碼(positional encoding)把這些編號加回去——它把一個與位置相關的向量注入每個 patch embedding,讓模型知道每個 patch 來自哪裡。
一排 patch embedding 向量,每個都被逐元素加上一個獨特的位置向量,產生具有位置感知的 token。
送入編碼器的輸入序列:patch embedding 加上 class token,再加上位置編碼。
由左到右閱讀:x_p^j 是第 j 個攤平後的 patch(回想第 1 篇,我們把影像切成 N 個 patch),乘上patch embedding 矩陣 E 後把它投影成一個 D 維的 token。我們在最前面加上一個額外的可學習向量 x_{\text{class}}——也就是 class token——它本身不帶任何像素,但稍後會彙整整張影像的摘要(我們在第 5 節會用到它)。方括號只是表示「把這 N+1 個 token 疊成一個序列」。接著我們加上 E_{\text{pos}},這是一個形狀為 (N{+}1) \times D 的可學習表格:表格的第 k 列就是第 k 個位置的位置向量。關鍵細節是我們把位置向量相加(逐元素),而不是串接。相加會把維度維持在 D,因此後續每一層仍看到寬度為 D 的 token;串接則會增加寬度,迫使其他每個矩陣都得改變形狀。
填入 E_{\text{pos}} 有兩種常見方式。ViT 通常使用學習得到的一維位置:E_{\text{pos}} 就像其他任何參數一樣從頭訓練,模型自行找出有用的位置向量(「一維」是指我們以掃描順序把 patch 編號為 0、1、2……,而不是分別編碼它們的二維列/行)。另一種來自原始 Transformer,是一種固定的正弦(sinusoidal)編碼,如下所示。
固定的正弦替代方案:每個維度都是一個不同頻率的正弦或餘弦波。
這裡 pos 是 token 在序列中的整數位置(0、1、2……),而 i 是 d 維向量內部維度的索引(因此 2i 與 2i{+}1 是成對的正弦/餘弦位置)。訣竅在於低維度使用長波長(隨位置變化緩慢),而高維度使用短波長(變化快速),所以所有正弦與餘弦的組合就給了每個位置一個獨一無二的「指紋」——很像二進位數字組合起來計數。正弦編碼不需要訓練,且能外推到更長的序列;但對於 ViT 處理的固定解析度影像,學習得到的表格通常表現至少一樣好且更簡單,這就是 ViT 預設採用它的原因。
再想一想:前饋 MLP 區塊
注意力很擅長在 token「之間」混合資訊——它讓每個 patch 從其他 patch 蒐集脈絡。但在這番蒐集之後,每個 token 持有一個剛更新過的向量,值得進行一些更深入、私下的處理。這正是前饋 MLP 區塊的工作。它被獨立地套用到每個 token 的向量上——同一個小型神經網路分別在每個 token 上執行,token 之間不再有任何交流。比喻是:注意力是大家分享所聽聞內容的小組討論;而 MLP 是討論結束後每個人各自回家,安靜地把一切想清楚。
一個寬度為 D 的 token 向量被投影到寬度 4D,通過 GELU 非線性,再投影回寬度 D。
兩層線性層,中間夾著一個 GELU 非線性。
拆解一下:z 是一個 token 寬度為 D 的向量。第一層線性層 W_1(含偏置 b_1)把它從 D 擴展到一個更寬的隱藏維度,慣例是 4D。接著逐元素套用平滑非線性 \mathrm{GELU}——它的行為像是「保留正值、抑制負值」的柔化版本,而這種平滑(相較於硬性的 ReLU)往往有助於訓練。最後第二層線性層 W_2(含偏置 b_2)把 4D 的隱藏向量收縮回 D。兩個偏置 b_1, b_2 只是可學習的位移量。由於兩層線性層在每個 token 上共享相同的權重,這個區塊沒有任何位置或順序的概念——所有跨 token 的混合早已在注意力中完成。
ViT-Base 的具體維度:768 \to 3072 \to 768。所以 W_1 是 768 \times 3072,W_2 是 3072 \times 768。為什麼要先把寬度放大 4 倍再縮回去?這個寬大的隱藏層是草稿空間:它讓網路有餘裕並行計算許多中間特徵並重新組合,很像你在寫下一行答案之前,先在一大張紙上粗略演算。一個直接 768 \to 768 的層在所能表達的轉換上會受限得多。事實上,MLP 區塊持有 Transformer 大多數的參數,而模型大部分原始的「思考能力」就住在這裡。
拼起來:編碼器區塊
我們現在有了兩個引擎——多頭注意力(MHA)與前饋 MLP 區塊。一個 Transformer 編碼器區塊把它們接在一起,再加上兩種成分:每個引擎「之前」的 LayerNorm(LN),以及每個引擎周圍的殘差(跳接)連接。ViT 採用前置正規化(pre-norm)的排列方式,意思是 LayerNorm 出現在子層「之前」,而非之後。這個區塊是:LN → MHA → 把輸入加回來;接著 LN → MLP → 再加一次。
一條垂直的資料路徑:輸入分支,一條分支經過 LayerNorm 再經過多頭注意力後加回輸入;結果再次分支,經過 LayerNorm 再經過 MLP 後加回。
前置正規化編碼器區塊的兩個子層,每個都包覆在一個殘差連接中。
讓我們走過第 \ell 個區塊的一次完整流程。輸入是 z_{\ell-1},也就是上一個區塊輸出的 token 序列(形狀 N{\times}D)。第一個子層:我們用 \mathrm{LN} 將它正規化,送入多頭注意力,然後把原始的 z_{\ell-1} 加回來——那個 `+ z_{\ell-1}` 就是殘差連接。結果 z'_\ell 是在輸入之上疊加了一層基於注意力的更新。第二個子層:我們正規化 z'_\ell,讓它通過 MLP,再把 z'_\ell 加回來,得到 z_\ell。這裡 \ell 只是區塊索引(區塊 1、2、3……),\mathrm{LN} 是層正規化,每個 + 都是一次殘差相加。關鍵在於:輸入與輸出的形狀都是 N{\times}D——這個區塊吃進 N 個寬度為 D 的 token,再交還 N 個寬度為 D 的 token,只是內容更豐富了。
那為什麼要 LayerNorm?當訊號流經許多層時,激活值的尺度可能會漂移——有些變得巨大,有些縮到趨近於零——這會使訓練不穩定。\mathrm{LN} 在向量進入子層之前,把每個 token 的向量重新縮放成受控的平均值與變異數(跨它的 D 個特徵),讓數值維持在健康的範圍內。由於 LN、殘差、注意力與 MLP 全都保持 N{\times}D 的形狀,一個區塊的輸出可以直接當作下一個區塊的輸入——這正是我們在第 5 節得以堆疊相同區塊的原因。
def encoder_block(z, params):
# z has shape (N, D) and stays (N, D) throughout
# --- sub-layer 1: multi-head attention with residual ---
a = multi_head_attention(layer_norm(z)) # LN, then MHA
z = z + a # residual add
# --- sub-layer 2: feedforward MLP with residual ---
m = mlp(layer_norm(z)) # LN, then MLP (GELU inside)
z = z + m # residual add
return z # still (N, D)堆疊深度,組成完整 ViT
由於每個編碼器區塊都把 N{\times}D 對映到 N{\times}D,我們可以把 L 個相同的區塊一層疊一層,把每個區塊的輸出直接餵進下一個區塊。每個區塊都把 token 表示再精煉一點:較早的區塊傾向捕捉局部、低層次的關係,較晚的區塊則把它們組合成全域、語意層次的關係。ViT-Base 使用 L = 12 個區塊;像 ViT-Large 這類更大的變體則用 24 個。這就是這裡「深度」的意思——一輪又一輪的蒐集再思考。
在最後一個區塊之後,我們如何把 N{+}1 個 token 向量轉成單一預測?我們只取出 class token——也就是我們在第 2 節最前面加上的那個特殊位置。在全部 L 個區塊中,注意力讓這個 token 從每個 patch 蒐集資訊,因此到最後它持有整張影像的摘要。我們把這一個向量通過一個小型分類頭——一個線性層接著一個 softmax——以產生跨各標籤類別的機率。這就完成了視覺 Transformer。
分類頭:取最終的 class token,正規化,投影成類別 logits,再 softmax。
逐符號來看:z_L^0 是最終(第 L 個)區塊之後 class token 的向量——上標 0 取的是第 0 個位置,也就是 class token,下標 L 表示「在全部 L 個區塊之後」。我們對它套用最後一個 \mathrm{LN},然後乘上 W_{\text{head}},這是形狀為 D \times C 的分類權重矩陣,其中 C 是類別數量;這會把寬度為 D 的摘要對映成 C 個原始分數,稱為 logits。最後 \mathrm{softmax} 對這 C 個 logits 取指數並正規化,使它們成為一個總和為 1 的機率分布——也就是模型對每個標籤的信心。舉例來說,若 D{=}768、C{=}1000 個 ImageNet 類別,則 W_{\text{head}} 是 768 \times 1000,而 y 是一個長度 1000 的機率向量,其最大的元素就是預測的類別。
- 分塊(第 1 篇):把影像切成 N 個固定大小的 patch,將每個攤平,再用 patch embedding 矩陣 E 投影,得到 N 個寬度為 D 的 token。
- 在最前面加上 class token,並加上學習得到的位置編碼 E_pos,得到形狀為 (N+1)×D 的輸入序列 z_0。
- 讓 z_0 通過 L 個相同的編碼器區塊;每個執行 LN → 多頭注意力 → 殘差,接著 LN → MLP → 殘差,並維持形狀 (N+1)×D。
- 取出最終區塊的 class token,z_L^0,套用 LayerNorm、分類頭 W_head 與 softmax,讀出預測的標籤機率 y。
這就是完整的前向傳遞,從原始像素到預測標籤,跨越第 1 到第 3 篇建構而成:patch 與 embedding、多頭自注意力、位置編碼、MLP、LayerNorm 與殘差、L 個堆疊的編碼器區塊,以及 class token 分類頭。這個架構現在已經完整,而且令人安心地相當一致——幾乎一切都是同一個區塊的重複。我們「尚未」討論的是如何真正把這東西訓練好。ViT 以渴求資料聞名,而這正是第 4 篇(DeiT、蒸餾與混合架構)將要處理的問題;之後第 5 篇會用階層式的 Swin Transformer 重新改造架構本身,以擴展到大型、真實世界的影像。