§01學習重點
- 說出自注意力的三步——對所有位置打分、把分數壓成一組權重、再對值加權合成——並指出這三步裡的非線性是從哪裡來的
- 解釋查詢、鍵、值為什麼是同一批輸入被三組不同權重投影出來的三種角色,而不是三批不同的資料
- 指認 softmax 是沿哪個方向做的:固定住的是誰、被加總掉的是誰,並用一個具體的注意力矩陣說出「列和」與「欄和」的差別
- 說明分數為什麼要除以鍵維度的平方根,並親手量出不除的時候 softmax 的敏感度掉到什麼程度
- 親手驗證自注意力對位置重排是等變的,以及加上位置編碼之後這個性質怎麼被打破
- 說出多頭不是「同一件事做很多次」:維度怎麼被切開、拼接之後為什麼還要再乘一個混合矩陣
- 畫出一層 Transformer 的資料流,指認哪個子區塊負責橫向混合、哪個只做逐位置加工
- 解釋因果遮罩為什麼必須在 softmax 之前動手,並驗證加了遮罩之後前面的輸出不受後面輸入影響
§02課程內容
一、把一段文字丟進全連接網路,會壞在哪裡
前面幾章的網路都有一個沒說出口的前提:輸入是一條長度固定的向量,而且第 3 個位置與第 40 個位置各配一組自己的權重。這個前提對表格資料成立,對一段文字則會連續壞三次。
第一次壞在長度。 句子有長有短,可是全連接層的輸入寬度在設計時就釘死了。你只能取一個最大長度,短的補空、長的截斷——補空的部分佔掉了大半個輸入,截斷掉的部分則是直接丟資訊。
第二次壞在共用。 一段文字裡的每個位置,本該用同一套規則處理。出現在第 3 個位置的某個詞,跟出現在第 40 個位置的同一個詞,是同一個詞;沒有任何理由要模型在兩個地方各學一次怎麼對付它。全連接層卻正好相反:它給每個位置各配一組權重,等於強迫模型把同一件事重學很多遍。
第三次壞在參數量。 假設每個位置用 300 個數字表示、一段最多容納 800 個位置,全部攤平就是 240,000 個數字。再接一層同樣寬的全連接層,光是那一層的權重矩陣就有 240,000 × 240,000 = 5.76 × 10¹⁰ 個參數。這還只是一層。
第 10 章的卷積網路(convolutional network)已經解決過同一型的前兩個毛病:同一組權重掃過每個位置,長度不再綁死,規則也自動共用。但它留下一個缺口。一個卷積核一次只覆蓋一小段,隔得遠的兩個位置要靠疊很多層才連得上;而文字裡需要建立關聯的兩個位置,往往隔了整整一段。
還缺最後一塊,而且這一塊才是關鍵:哪兩個位置該建立關聯,不能寫死在結構裡,得由內容當場決定。 卷積的連法是預先排好的——不管讀到什麼,第 n 個位置永遠只看左右那幾格。文字不吃這一套:一個代名詞該連到哪個名詞,答案在句子裡,不在結構裡。
三件事加起來,你要的東西輪廓就出來了:一種運算,處理規則對每個位置共用、輸入長度可以變,而且任何兩個位置之間的關聯強弱由當下的內容算出來。這就是自注意力(self-attention),本章其餘的篇幅都在把它拆開來看。
本章從頭到尾用同一個場景對照:把一段文字想成一場圓桌會議,每個位置是一位與會者,每位手上都有一份準備好要說的內容。自注意力要決定的,就是每一輪裡誰的話該佔多少分量。
二、先看合成:輸出是一份加權平均
自注意力有三步,我打算倒著講:先講最後一步「怎麼合成」,再回頭問「比重哪裡來」。理由很實際——後者難得多,而且只有先接受了前者,後者才有地方可放。
第一步單純。每個位置各自算出一份向量,叫它的值(value):
逐項拆解:\(\mathbf{x}_m\) 是第 \(m\) 個位置的輸入向量,長度 \(D\);\(\boldsymbol\Omega_v\) 是一個 \(D \times D\) 的權重矩陣;\(\mathbf{b}_v\) 是一條長度 \(D\) 的偏置向量;\(\mathbf{v}_m\) 就是第 \(m\) 個位置的值向量。請特別注意一件事:\(\boldsymbol\Omega_v\) 與 \(\mathbf{b}_v\) 只有一組,所有位置共用。 位置有 \(N\) 個,這一步的參數量卻與 \(N\) 完全無關——上一節的第二、第三個毛病,光這一步就解掉了。
第二步是合成。第 \(n\) 個輸出是所有值向量的加權和:
逐項拆解:\(N\) 是這一次送進去幾個位置;\(\alpha_{mn}\)(讀作 alpha)是一個純量,代表第 \(n\) 個輸出要分給第 \(m\) 個輸入多少比重——兩個下標的順序是「先來源、後去處」,寫的時候容易顛倒,值得記一下;求和跑遍所有 \(N\) 個位置,所以每一個輸出都用到了全部的值向量,只是比重不同。
這些比重有兩個約束:每一個都不小於零,而且固定住輸出位置 \(n\) 之後,它們加起來剛好是 1:
這兩條約束合起來,說的就是「一整份被切開分掉」。權重不是拿來選一個,是拿來分配比例。 這句話值得多停一秒:不是「第 3 個輸出決定聽第 1 位的」,而是「第 3 個輸出分給第 1 位兩成、分給第 3 位將近六成、剩下的分給其他人」。
有三件事現在就該記住。第一,輸出與輸入的尺寸一模一樣,都是 \(D \times N\),所以這種區塊可以一層接一層疊上去,中間不需要改形狀。第二,如果把所有 \(\alpha_{mn}\) 都設成 \(1/N\),每個輸出就變成所有值向量的平均,而且 \(N\) 個輸出彼此完全相同——這是這個機制最沒用的一組權重,也剛好說明有用的部分全在權重的差異上。第三,到目前為止還沒有出現任何激活函數。
比喻: 一場圓桌會議正在討論某個議題,每位與會者手上都有一份準備好的發言稿,那就是他的值向量。輪到第三位整理自己這一輪的結論時,他不是挑一個人的話照抄,而是把發言權按比例分掉:第一位佔兩成、第二位幾乎不佔、他自己佔將近六成,剩下的分給後兩位——這一整份配額加起來是一,不能超發。他的結論就是這些發言稿按這組比例混出來的東西。這個比喻有一處明確失準:真人開會時「聽誰講」是互斥的,同一時間只能專心聽一個人;自注意力卻是同時按比例把所有人的內容都收進來,沒有先聽誰、後聽誰的問題。到了第五節談多頭時,這個差異還會再放大一次。
三、比重從哪裡來:查詢、鍵,以及 softmax 沿哪個方向做
現在回答上一節跳過的問題:\(\alpha_{mn}\) 是怎麼算出來的?
它得滿足兩個條件。一是由內容決定——第 \(n\) 個位置該分多少給第 \(m\) 個位置,要看這兩個位置各自帶著什麼。二是可訓練,也就是中間那些可調的數字得待在梯度算得動的地方。
做法是把同一批輸入再打成兩種角色,叫查詢(query)與鍵(key):
逐項拆解:\(\boldsymbol\Omega_q\) 與 \(\boldsymbol\Omega_k\) 是兩個 \(D_q \times D\) 的權重矩陣,\(\mathbf{b}_q\) 與 \(\mathbf{b}_k\) 是對應的偏置。\(D_q\) 是查詢與鍵的維度,兩者必須相同,因為下一步要把它們配對相乘。\(\mathbf{q}_n\) 代表「第 \(n\) 個位置這一輪想解決什麼」,\(\mathbf{k}_m\) 代表「第 \(m\) 個位置有什麼資格被找上」。這組名字取自查表的語彙,你不必想得太深,把它們當三個角色的代號就夠了。
三種角色,同一批資料。 學生最常在這裡卡住:以為查詢、鍵、值是三批不同的資料。不是。\(\mathbf{x}_1\) 到 \(\mathbf{x}_N\) 只有一批,是三組不同的投影權重把它們讀成三種身分——問問題的、被比對的、被搬運的。同一個位置同時扮演這三種身分,一點也不衝突。
為什麼是三種不是兩種? 因為「聽誰的」和「聽到什麼」是兩件事。打分只需要決定比重,值才是真正被搬走的內容。拆成三組之後,值的維度不必遷就打分的維度,兩邊可以各自調整;如果讓值兼任鍵,等於逼一個向量同時擅長被比對和被搬運。
打分那一步就是內積:\(s_{mn} = \mathbf{k}_m^{\mathsf{T}}\mathbf{q}_n\)。內積在這裡量的是「兩個方向有多一致」——同向時是正的大數,垂直時接近零,反向時是負數。
接下來是本章最大的卡點,所以我把它單獨拉出來寫死:
固定一個輸出位置 \(n\),把它對所有 \(N\) 個輸入位置的分數收成一整排,對這一整排做 softmax。被加總掉的是輸入位置那一維(下標 \(m\)),不是輸出位置那一維(下標 \(n\))。
寫成式子:
逐項拆解:分子是第 \(m\) 個輸入對第 \(n\) 個輸出的分數取指數;分母把所有輸入位置的同一種東西加起來,加總用的啞變數寫成 \(m'\) 以免跟分子的 \(m\) 混淆。指數把可正可負的分數變成正數,除以總和讓整排加起來是 1——所以固定 \(n\) 時和為 1 是式子本身保證的,固定 \(m\) 時則沒有任何這樣的保證。分母裡的 \(\mathbf{q}_n\) 從頭到尾沒變,變的只有鍵,這就是「固定輸出、掃過所有輸入」的意思。至於為什麼要除以 \(\sqrt{D_q}\),下一節專門處理。
把所有比重排成一個 \(N \times N\) 的矩陣 \(\mathbf{A}\)(列是輸入位置 \(m\),欄是輸出位置 \(n\)),這件事就變成可以直接數的:每一欄加起來是 1,每一列不是。 整段運算壓成矩陣寫法是:
逐項拆解:\(\mathbf{X}\) 是把 \(N\) 個輸入向量並排成的 \(D \times N\) 矩陣;\(\mathbf{Q}, \mathbf{K}, \mathbf{V}\) 同理,各自是查詢、鍵、值並排成的矩陣。\(\mathbf{K}^{\mathsf{T}}\mathbf{Q}\) 一次算完所有配對的內積,得到一個 \(N \times N\) 的分數矩陣,第 \(m\) 列第 \(n\) 欄就是 \(s_{mn}\)。這裡的 softmax 對每一欄各做一次;式子本身看不出這個方向,所以每次看到這種寫法都要自己補上一句「沿欄」。最後左乘 \(\mathbf{V}\),等於拿權重矩陣的每一欄去混合各個值向量。
手刻一遍最能把方向釘住。底下這段程式造一場五個位置的小型會議,六維向量,查詢與鍵四維:
import numpy as np
rng = np.random.default_rng(0)
D, N, Dq = 6, 5, 4 # 六維向量、五個位置、查詢與鍵四維
X = rng.normal(0, 1, (D, N)) # 第 n 欄=第 n 位的輸入向量
Wq = rng.normal(0, 0.5, (Dq, D)); bq = rng.normal(0, 0.5, (Dq, 1))
Wk = rng.normal(0, 0.5, (Dq, D)); bk = rng.normal(0, 0.5, (Dq, 1))
Wv = rng.normal(0, 0.5, (D, D)); bv = rng.normal(0, 0.5, (D, 1))
Q = Wq @ X + bq # (Dq, N):每欄一支查詢
K = Wk @ X + bk # (Dq, N):每欄一支鍵
V = Wv @ X + bv # (D, N):每欄一份值
def softmax_cols(S): # 固定一個輸出位置(一欄),對所有輸入位置正規化
E = np.exp(S - S.max(axis=0, keepdims=True))
return E / E.sum(axis=0, keepdims=True)
S = K.T @ Q / np.sqrt(Dq) # S[m, n]=k_m 與 q_n 的內積,再除以根號 Dq
A = softmax_cols(S) # A[m, n]=alpha_mn
Y = V @ A # (D, N):第 n 欄=第 n 個輸出
np.set_printoptions(precision=3, suppress=True)
print("注意力矩陣 A(欄=輸出位置,列=輸入位置):"); print(A)
print("每欄總和:", np.round(A.sum(axis=0), 6))
print("每列總和:", np.round(A.sum(axis=1), 3))
print("第 3 欄的分數:", np.round(S[:, 2], 3))
print("第 3 欄的權重:", np.round(A[:, 2], 3))
print("輸入形狀", X.shape, "輸出形狀", Y.shape)實跑輸出的注意力矩陣是
[[0.171 0.229 0.209 0.271 0.321]
[0.215 0.205 0.021 0.097 0.17 ]
[0.295 0.374 0.571 0.292 0.155]
[0.165 0.081 0.062 0.178 0.21 ]
[0.154 0.112 0.136 0.161 0.144]]每欄總和是 [1. 1. 1. 1. 1.],每列總和則是 [1.201 0.708 1.688 0.696 0.707]。這兩行並排看,卡點就消失了:欄和恆為 1,是 softmax 的方向決定的;列和五花八門,加起來只保證等於 5(因為五欄各湊 1)。列和大的位置,代表它的內容被很多輸出拿去用;列和小的位置,代表大家都不太理它。
第 3 欄的分數是 [ 1.309 -0.969 2.315 0.096 0.88 ],經過 softmax 之後變成權重 [0.209 0.021 0.571 0.062 0.136]。分數最高的第 3 位拿走將近六成,分數是負的第 2 位只剩兩個百分點——指數放大差距的效果在這裡看得很清楚。輸入形狀 (6, 5)、輸出形狀 (6, 5),尺寸確實沒變。
還有兩件事要在這裡講明白。
這一整段沒有激活函數,但它不是線性運算。 非線性有兩個來源:查詢與鍵的內積把輸入的兩個分量相乘,這本身就不是線性的;softmax 裡的指數與除法又是一層。所以不要因為看不到 ReLU 就以為這個區塊只是換個寫法的矩陣乘法。
權重是算出來的,不是存起來的。 第 3 到 11 章裡的權重都是參數:訓練完就固定,推論時原樣拿來用。\(\alpha_{mn}\) 不是。訓練完固定下來的是 \(\boldsymbol\Omega_q\)、\(\boldsymbol\Omega_k\)、\(\boldsymbol\Omega_v\) 這三組投影,\(\alpha_{mn}\) 每換一批輸入就整個重算一次。這是它跟前面所有層最大的體質差異,也是它能「由內容決定該連誰」的原因。
順帶提醒一個常見的過度解讀:注意力權重經常被拿來當成「模型在看哪裡」的說明。它是機制的一部分,不保證是可靠的解釋——某一格的權重大,不等於那個位置對最終答案的貢獻大,中間還隔著值向量與後面好幾層運算。
四、兩塊補丁:分數的尺度,與被弄丟的順序
上一節的機制有兩個洞,不補就不能用。兩個洞都出在打分那一步,所以放在一起講。
補丁一:分數會隨維度長大。 內積是 \(D_q\) 個乘積相加。如果兩個向量的各個分量彼此獨立、量級相當,那麼這 \(D_q\) 項加起來的典型大小會隨 \(\sqrt{D_q}\) 成長——項數變成一百倍,總和的典型大小就變成十倍。維度一高,分數就往兩邊跑得很遠。
分數一大,softmax 就飽和:最大那一項的指數把其他項壓成幾乎為零,權重變成「幾乎全押在一個位置上」。麻煩不在權重難看,而在梯度。softmax 對自己那一項分數的偏導數是 \(\alpha(1-\alpha)\),當 \(\alpha\) 逼近 1(或逼近 0),這個值就逼近零;往回傳的梯度被乘上一個接近零的數,這一層等於訓練不動了。這是第 7 章「梯度與初始化:讓深層網路真的訓練得動」處理過的那個現象,換一種方式發生。
補救很直接:把分數除以 \(\sqrt{D_q}\),把它拉回 softmax 還有反應的區間。底下這段程式把差別量出來——每種維度重抽兩千次,統計「五個權重裡最大的那個超過 0.99」的比例(我稱它飽和比例),以及 \(\alpha(1-\alpha)\) 的最大值:
import numpy as np
rng = np.random.default_rng(0)
N, R = 5, 2000 # 五個位置;每種維度重抽 2000 次
def softmax(v):
e = np.exp(v - v.max())
return e / e.sum()
def report(w): # 飽和=最大權重超過 0.99;敏感度=alpha(1-alpha) 的最大值
return (w.max(axis=1) > 0.99).mean(), (w * (1 - w)).max(axis=1).mean()
for Dq in (4, 50, 400):
s = np.array([rng.normal(0, 1, (N, Dq)) @ rng.normal(0, 1, Dq)
for _ in range(R)])
f_raw, d_raw = report(np.array([softmax(v) for v in s]))
f_scl, d_scl = report(np.array([softmax(v / np.sqrt(Dq)) for v in s]))
print(f"Dq={Dq:3d}|分數最大絕對值 {np.abs(s).max(axis=1).mean():5.2f}"
f"|未縮放:飽和比例 {f_raw:5.1%}、敏感度 {d_raw:.4f}"
f"|縮放後:飽和比例 {f_scl:4.1%}、敏感度 {d_scl:.4f}")實跑輸出:
Dq= 4|分數最大絕對值 2.98|未縮放:飽和比例 1.5%、敏感度 0.1981|縮放後:飽和比例 0.0%、敏感度 0.2253
Dq= 50|分數最大絕對值 11.04|未縮放:飽和比例 40.1%、敏感度 0.0774|縮放後:飽和比例 0.0%、敏感度 0.2310
Dq=400|分數最大絕對值 30.69|未縮放:飽和比例 73.4%、敏感度 0.0312|縮放後:飽和比例 0.0%、敏感度 0.2302四維的時候幾乎不飽和,四百維的時候有 73.4% 的抽樣一面倒,敏感度從 0.1981 掉到 0.0312。除以 \(\sqrt{D_q}\) 之後,三種維度的飽和比例全是 0.0%,敏感度穩定停在 0.23 附近——維度變了一百倍,softmax 的工作點卻沒有漂走。
這裡有個常見誤解要澄清:這一步不是在把向量長度正規化。向量本身完全沒被動到,被動到的只有分數的尺度,目的只有一個——別讓 softmax 卡在沒有反應的區域。
補丁二:順序被弄丟了。 把輸入的欄順序打亂,輸出會怎樣?先把兩個容易混的詞就地分清楚:不變(invariant)是「輸入怎麼動,輸出都不動」;等變(equivariant)是「輸入怎麼動,輸出就跟著怎麼動」。第 10 章「卷積網路:把影像的結構寫進模型」講過的平移等變性,就是後面這種。
自注意力對位置重排是等變的:你把第 1 位和第 3 位的座位對調,輸出也只是把對應的兩欄對調,內容一個字都不變。原因就寫在式子裡:從頭到尾沒有任何一項提到「第幾個」,每個位置只認得內容。
import numpy as np
rng = np.random.default_rng(0)
D, N, Dq = 6, 5, 4
X = rng.normal(0, 1, (D, N))
Wq = rng.normal(0, 0.5, (Dq, D)); Wk = rng.normal(0, 0.5, (Dq, D))
Wv = rng.normal(0, 0.5, (D, D))
def sa(Xin): # 無偏置版的自注意力,只為看清等變性
Q, K, V = Wq @ Xin, Wk @ Xin, Wv @ Xin
S = K.T @ Q / np.sqrt(Dq)
E = np.exp(S - S.max(axis=0, keepdims=True))
return V @ (E / E.sum(axis=0, keepdims=True))
perm = np.array([2, 0, 4, 1, 3]) # 把五個位置重排
Y, Yp = sa(X), sa(X[:, perm])
print("打亂輸入後的輸出,是否等於原輸出跟著換欄:", np.allclose(Yp, Y[:, perm]))
# 加上位置編碼:每個位置一組固定的正弦餘弦值
pos = np.arange(N)[None, :]
dim = np.arange(D)[:, None]
freq = 1.0 / (100.0 ** (2 * (dim // 2) / D))
P = np.where(dim % 2 == 0, np.sin(pos * freq), np.cos(pos * freq))
Z, Zp = sa(X + P), sa(X[:, perm] + P)
print("加了位置編碼後,同樣的檢查:", np.allclose(Zp, Z[:, perm]))
print("兩者第一欄的差距:", np.round(np.abs(Zp[:, 0] - Z[:, perm][:, 0]), 3))實跑輸出:第一問是 True,第二問是 False,加了位置編碼之後兩者第一欄的差距是 [0.629 0.201 0.319 0.091 0.46 0.068]。也就是說,重排前後的輸出真的不再只是換欄而已,內容本身變了。
為什麼這是致命傷?因為語言的意思有一大半就在順序裡。把一句話的詞打散重排,內容沒少一個字,意思卻可能整個翻過來。一個對順序等變的機制,看不出這兩句有什麼不同。
補救的辦法叫位置編碼(positional encoding),有兩條路。第一條是給每個位置一個獨有的向量 \(\mathbf{P}\)(尺寸與 \(\mathbf{X}\) 相同),在進第一層之前直接加進輸入。第二條是不動輸入,改在打分那一步多加一項,這一項只跟「兩個位置相隔多遠」有關——上面那段程式用的是第一條。
有人會問:為什麼是相加,不是把位置向量串在輸入後面?因為維度 \(D\) 通常遠大於實際用到的位置數,位置資訊完全可以擠在這個空間的一個角落裡,讓內容用其餘的部分。串接則要額外撥出維度來裝它,而後面每一層的參數量都跟維度掛鉤,這個成本會一路乘下去。
五、多頭,以及一層 Transformer 怎麼組起來
多頭不是把同一件事做很多次。 一組 \(\boldsymbol\Omega_q\) 與 \(\boldsymbol\Omega_k\) 只能對出一種「該聽誰」的判準。但一段文字裡同時有好幾條線索需要追:誰是這個動作的執行者、這個代名詞指的是誰、這一段的主題是什麼。一組投影要同時服務這幾條線索,只能折衷,折衷的結果是每條都做得不好。
多頭自注意力(multi-head self-attention)的做法是並排放 \(H\) 組。每個頭有自己的三組投影,各自跑一次完整的打分與加權,得到自己的一組輸出。\(H\) 個頭的輸出上下拼接起來之後,還要再乘一個混合矩陣 \(\boldsymbol\Omega_c\) 才算完:
逐項拆解:\(\mathrm{SA}_h[\mathbf{X}]\) 是第 \(h\) 個頭的輸出,尺寸是 \((D/H) \times N\);把 \(H\) 個這樣的矩陣上下疊起來,總高度回到 \(D\);\(\boldsymbol\Omega_c\) 是一個 \(D \times D\) 的矩陣。最後這一步很多人講到拼接就停了,可是它不能省:不乘 \(\boldsymbol\Omega_c\),各頭的結果只是並排堆著、彼此毫無往來,下一層拿到的每個維度都只來自單一個頭;乘上去之後,下一層看到的每個維度才可能同時吃到好幾個頭的資訊。
維度被切開了。 每個頭的查詢、鍵、值維度通常取 \(D/H\),不是各自都用滿 \(D\)。這一點是「多頭不等於做很多次」的關鍵:頭數翻倍,每個頭就變窄一半,總參數量與總計算量幾乎不動。底下這段程式把這件事量出來——同樣的總維度 48、同樣的十個位置,只改頭數:
import numpy as np
rng = np.random.default_rng(0)
D, N = 48, 10
X = rng.normal(0, 1, (D, N))
def softmax_cols(S):
E = np.exp(S - S.max(axis=0, keepdims=True))
return E / E.sum(axis=0, keepdims=True)
def multihead(H):
d = D // H # 每個頭分到的維度
heads, params, muls = [], 0, 0
for _ in range(H):
Wq = rng.normal(0, 0.2, (d, D)); Wk = rng.normal(0, 0.2, (d, D))
Wv = rng.normal(0, 0.2, (d, D))
Q, K, V = Wq @ X, Wk @ X, Wv @ X
heads.append(V @ softmax_cols(K.T @ Q / np.sqrt(d)))
params += 3 * d * D # 三組投影
muls += 3 * d * D * N + 2 * N * N * d # 投影+打分+加權
Wc = rng.normal(0, 0.2, (D, D)) # 拼接後的混合矩陣
params += D * D
muls += D * D * N
return Wc @ np.vstack(heads), params, muls
for H in (1, 2, 3, 6):
Y, params, muls = multihead(H)
print(f"H={H}|每頭維度 {D // H:2d}|投影與混合的參數 {params}"
f"|乘法次數 {muls}|輸出形狀 {Y.shape}")實跑輸出:
H=1|每頭維度 48|投影與混合的參數 9216|乘法次數 101760|輸出形狀 (48, 10)
H=2|每頭維度 24|投影與混合的參數 9216|乘法次數 101760|輸出形狀 (48, 10)
H=3|每頭維度 16|投影與混合的參數 9216|乘法次數 101760|輸出形狀 (48, 10)
H=6|每頭維度 8|投影與混合的參數 9216|乘法次數 101760|輸出形狀 (48, 10)四種頭數,參數 9,216 個、乘法 101,760 次,一個也沒多。多頭買到的是「同時追好幾條線索」,付出的不是算力,是每條線索能用的維度變窄了——所以頭數也不是越多越好,切太細之後每個頭的表達空間就不夠了。(另有一個常被提起的實務觀察:訓練完成之後移除一部分頭,效能下降得比想像中少。本課把它當成待查證的線索,不當定論。)
比喻: 同一場圓桌會議,改成先分三組並行討論,每組追一條不同的線索:一組盯預算、一組盯時程、一組盯法遵。每組各自重新分配一次發言權——盯預算那組會把比重壓在財務單位身上,盯時程那組壓在工務單位身上。三組討論完,紀錄不是各自歸檔了事,而是要再過一次彙整才變成這一輪的正式結論;彙整那一步就是混合矩陣,少了它,三份紀錄只是被釘在一起而已。這個比喻有一處失準:真人分三組要多找兩倍人手,總工時真的變三倍;多頭卻是把原本的維度切成三份、每組拿三分之一,總工作量幾乎沒變——上面那段程式的四行輸出就是這件事的證據。
一層 Transformer 長什麼樣。 到此為止拼出來的還不是一層,只是一層裡的半邊。一層有兩個子區塊,分工乾淨:
- 多頭自注意力:橫向混合。 位置與位置之間唯一交換資訊的地方就是這裡。
- 逐位置的小型全連接網路:縱向加工。 它對每個位置各做一遍,用的是同一組參數,而且各做各的——這個子區塊完全不讓位置之間通訊。學生常誤以為它會,其實它連隔壁是誰都不知道。
兩個子區塊都包在第 11 章講過的跳接裡,各自後面接一次正規化。這裡的正規化與第 11 章談的批次正規化不是同一種:統計量是在每個位置自己的特徵維度上算的,跟同一批裡的其他樣本無關。這個差別本章不展開,只要別把兩者混為一談就好;它有個現成的名字,叫層正規化(layer normalization)。
真正的網路是把這種層一個接一個疊上去。因為輸入與輸出的形狀相同,疊多少層都不必改接口——這也是為什麼「加深」在這個架構上特別容易。
六、遮罩:只准聽已經講過的話
先把文字接上去。 前五節談的都是「一疊向量進、一疊向量出」,還沒交代文字怎麼變成那疊向量。這一段快速帶過。
第一步是切分(tokenization),有一條「切多細」的軸。切到單一字元,表最小,但模型得從頭學怎麼把字元組成有意義的單位;切成整詞,同一個詞的各種變化形各佔一格,而且永遠會有沒收進表裡的詞。折衷是切成子詞(subword):常用的整詞原樣留著,罕見的詞由幾個片段拼出來。切完之後每個片段叫一個詞元(token),全部詞元的集合就是詞彙表 \(\mathcal{V}\),大小寫成 \(\lvert\mathcal{V}\rvert\)。
第二步是嵌入(embedding)。嵌入表 \(\boldsymbol\Omega_e\) 是一個 \(D \times \lvert\mathcal{V}\rvert\) 的矩陣,每個詞元對應其中一欄;這張表本身就是模型參數的一部分,跟著訓練一起調。(有些教材把嵌入表寫成 \(\mathbf{E}\),指的是同一個東西。)
這裡有一件事值得先講清楚:嵌入層出來的向量與上下文無關。 同一個詞元不管出現在哪一句、哪個位置,拿到的都是同一欄。上下文是進了 Transformer 層之後才長出來的——長出上下文,正是自注意力在做的事。
語言模型的目標函數。 把整句話的機率拆成一串條件機率的乘積:
逐項拆解:\(t_n\) 是第 \(n\) 個詞元;等號左邊是「這一整串詞元一起出現」的機率;右邊的每一項是「在看過前面所有詞元的條件下,第 \(n\) 個剛好是 \(t_n\)」的機率;\(\prod\) 是連乘符號,用法跟 \(\sum\) 一樣,只是把加號換成乘號。第一項 \(Pr(t_1)\) 前面沒有東西可看,就是它自己的邊際機率。
這個寫法帶來一個訓練上的大便宜:每個位置都是一道獨立的「猜下一個」的題目,而答案就是下一個詞元本身,不需要任何人標註。整句的損失就是這 \(N\) 道題的交叉熵(第 5 章)取平均。
問題來了。 訓練時我們當然希望一次把整句餵進去、一次算完 \(N\) 道題的損失,這樣才划算。但自注意力預設每個位置都看得到所有位置——第 3 個位置在猜第 4 個詞元的時候,第 4 個詞元就大剌剌地躺在它的輸入裡。答案就在輸入裡,這樣訓練等於白訓練。
解法不是把後面的詞元刪掉,那樣就退回一次只能算一個位置,前面那個便宜就沒了。解法是回到打分那一步動手:把「第 \(n\) 個輸出不該取用第 \(m\) 個輸入」的那些格子,在過 softmax 之前把分數壓成 \(-\infty\)。指數把 \(-\infty\) 送成 0,那些項的權重就是 0,而剩下的權重仍然自動加起來等於 1。這叫遮罩自注意力(masked self-attention),配上「只准看自己與前面」這條規則時也叫因果遮罩。
順序不能顛倒,這是這一節最要緊的一句。 如果先做完 softmax,再把不該看的權重歸零,剩下的那些加起來就不到 1 了——加權平均變成加權「部分平均」,輸出的尺度會隨著位置一路縮水。要修就得再正規化一次,那等於白繞一圈。直接在分數上動手,乾淨得多。
import numpy as np
rng = np.random.default_rng(0)
D, N, Dq = 6, 5, 4
X = rng.normal(0, 1, (D, N))
Wq = rng.normal(0, 0.5, (Dq, D)); bq = rng.normal(0, 0.5, (Dq, 1))
Wk = rng.normal(0, 0.5, (Dq, D)); bk = rng.normal(0, 0.5, (Dq, 1))
Wv = rng.normal(0, 0.5, (D, D)); bv = rng.normal(0, 0.5, (D, 1))
m_idx = np.arange(N)[:, None] # 輸入位置=列
n_idx = np.arange(N)[None, :] # 輸出位置=欄
M = m_idx <= n_idx # 只准取用自己與更前面的位置
def sa(Xin, mask=None):
Q, K, V = Wq @ Xin + bq, Wk @ Xin + bk, Wv @ Xin + bv
S = K.T @ Q / np.sqrt(Dq)
if mask is not None:
S = np.where(mask, S, -np.inf) # 過 softmax 之前先壓成負無窮
E = np.exp(S - S.max(axis=0, keepdims=True))
A = E / E.sum(axis=0, keepdims=True)
return V @ A, A
Y0, A0 = sa(X, M)
X2 = X.copy(); X2[:, 3] += 5.0 # 只動第 4 位的輸入
Y2, _ = sa(X2, M)
Yn0, _ = sa(X, None); Yn2, _ = sa(X2, None) # 對照組:不加遮罩
print("遮罩後的注意力矩陣(對角線以下全為 0):")
print(np.round(A0, 3))
print("每欄總和:", np.round(A0.sum(axis=0), 6))
print("有遮罩:前三欄輸出的最大變動", np.abs(Y2[:, :3] - Y0[:, :3]).max())
print("無遮罩:前三欄輸出的最大變動",
round(float(np.abs(Yn2[:, :3] - Yn0[:, :3]).max()), 4))實跑輸出的矩陣是
[[1. 0.528 0.261 0.324 0.321]
[0. 0.472 0.027 0.116 0.17 ]
[0. 0. 0.713 0.349 0.155]
[0. 0. 0. 0.212 0.21 ]
[0. 0. 0. 0. 0.144]]每欄總和依然是 [1. 1. 1. 1. 1.]。第一欄只剩一格,所以那一格必然是 1——第一個位置沒有前文可參考,只能全押自己。最後一欄什麼都沒被遮住,跟前一節那個未遮罩矩陣的最後一欄完全一樣(0.321 0.17 0.155 0.21 0.144),因為它本來就看得到全部。中間三欄的數字則整組變了:被砍掉的比重不會憑空消失,而是按原本的相對大小重新分給還看得見的那幾位。
最關鍵的是最後兩行。把第 4 位的輸入整個加 5、其餘不動:有遮罩時,前三欄輸出的最大變動是 0.0——一位小數都沒動;沒有遮罩時同樣的改動造成 14.2814 的變動。這就是「不會偷看後面」的機器證明。
生成的時候怎麼用。 訓練完之後要造句子,做法是把模型自己吐出來的詞元接回輸入的尾巴,序列長一格,再跑一次。因為每個位置只看得到自己與前面,前面那些位置算出來的鍵與值不會因為後面多了一格而改變,所以上一輪的結果可以整批留著重用,只算新加的那一格。這件事對生成的速度影響很大。
比喻: 把訓練想成一場圓桌會議的逐字稿覆盤。每位與會者要在自己那一輪開口之前,先猜出接下來會被說出的是什麼。整份逐字稿攤在會議桌上,所有輪次可以同時檢討,效率最高——但條件是每個人只准參考排在自己前面的發言,不准往下翻。遮罩做的就是這件事:不是把後面的內容從桌上收走,而是規定它們一律不計入發言權。這個比喻有一處明確失準:真人會議是有先後的,前面的人講了什麼會改變後面的人要講什麼;自注意力的一層裡沒有這種先後,所有位置是同時算完的,「一輪」指的是整層跑一次,不是一個人一個人輪流。
三種可見範圍。 這個架構的各種變體,用模型名字去記很快就亂了;用「一個位置看得到哪些位置」去記,就只有三種。
第一種,每個位置只看得到自己與前面。這正是剛才那種,可以一路往下猜、造出新內容,GPT 這一系屬於它。第二種,每個位置看得到所有位置,也就是完全不加遮罩。它不生成,它產出的是一組帶著上下文的表示,交給下游任務用,BERT 這一系屬於它。第三種,兩張桌子接起來:一張桌子先把來源那段讀成一組表示,另一張桌子一邊生成一邊回頭參考那組表示,用來把一段輸入映到另一段輸出——翻譯模型是最典型的例子,一張桌子讀源語言,另一張桌子產出目標語言。
第三種需要一個新接點,叫交叉注意力(cross-attention)。它用的還是同一套打分與加權,只換一件事:查詢來自其中一張桌子,鍵與值來自另一張。 「想解決什麼問題」由正在生成的那一邊提出,「有什麼內容可以拿」則由來源那一邊提供。看清楚這一點之後,這三種變體就不是三個新機制,而是同一個機制的三種可見範圍設定。
先大量無標註預訓練、再小量有標註微調。 這條路子之所以走得通,關鍵在上面那個目標函數:不管是「猜下一個詞元」還是「猜被遮住的詞元」,答案本來就在資料裡,不需要人標——這種做法叫自監督(self-supervised)。所以第一階段可以拿海量的純文字去訓練,這叫預訓練(pre-training)。第二階段才用少量有標註的資料,在後面接一層把輸出向量轉成任務要的形狀,這叫微調(fine-tuning)。第 14 章會正面處理「沒有標籤時到底該學什麼」這一整類問題。
最後,五件本章只點名不展開的事。
- 計算量隨長度平方成長。 每個位置都要跟每個位置打分,位置數翻倍,那張矩陣的格子就變成四倍。這是「上下文長度」為什麼是個大事的原因。有一整族方法在讓這張矩陣變稀疏或變小,本章不展開。
- 影像也能用同一套。 把圖切成小方塊,每一塊當成一個位置送進去就行。它與第 10 章的卷積是一組取捨:卷積把「鄰近的像素比較相關」這條歸納偏好(inductive bias)寫死在結構裡,注意力不寫,代價是得靠更多資料自己學出來。
- 生成時的取樣策略。 每一步都挑機率最大的詞元,容易吐出空轉、重複的句子,所以實務上有各種折衷的取樣做法。
- 這個機制之前的主力是循環網路(recurrent network)。 一次讀一個位置、把一個狀態往後傳,天然帶著順序,但隔得遠的資訊在傳遞途中會被沖淡。
- 這類模型的訓練不太穩定。 學習率的安排要特別設計,成因不只一個,本章不展開。
§03原書對照
原書第 12 章是全書篇幅最長的一章,橫跨三十三個印刷頁,從一個機制出發,一路鋪到語言與影像上的各種部署形態。以下按頁次指路,想深入的人可以直接翻。
開篇的 pp.207–208 用一段餐廳評論當引子,從中歸納出文字資料帶來的幾個難處,再據此推出這個機制該具備哪些性質。想知道動機是怎麼被逼出來的,這兩頁值得先看。
機制本身集中在 pp.208–212。原書先講一個輸出如何由多個輸入按比例合成,再回頭處理比例本身從哪裡來;p.209 有一張示意圖,把同一批輸入如何被拆成不同比例、組出三個不同的輸出畫了出來,pp.210–211 兩張圖則分別呈現這個運算裡的稀疏結構與打分的流程。想看矩陣寫法的人翻 p.212,那裡把整個運算壓成一行。
在實務上幾乎必備的三項補強寫在 pp.213–215:絕對與相對兩種位置資訊的處理方式在 pp.213–214,把內積除以維度平方根的理由在 p.214,多頭的維度配置與拼接後再作一次線性混合則在 pp.214–215,配圖在 p.215。完整的層結構——含兩段跳接與兩次正規化的四步流程——寫在 pp.215–216,p.216 的圖把資料流畫得很清楚。
文字管線的部分在 pp.216–219。子詞如何由字元逐步合併,原書用一段童謠做了六格圖解,在 p.217;嵌入表如何以獨熱向量取出對應欄位,圖在 p.219。
三種架構變體各佔一節。編碼器那一節以 BERT 為例,pp.219–222:模型規模與先預訓練後微調的兩段式在 p.219,遮住部分詞再預測的自監督任務在 p.220,微調後接上不同輸出層的三個下游任務範例在 pp.221–222。解碼器那一節以 GPT3 為例,pp.222–225:把一句話的機率拆成連乘的寫法在 p.222,遮罩如何切斷對未來的存取在 p.223,生成時的取樣策略在 pp.223–224,模型規模與少樣本現象在 pp.224–225,其中 p.225 附了一段模型實際續寫出來的文字,以及一組文法糾錯的示例。編碼器與解碼器接起來的翻譯架構在 pp.226–227,交叉注意力的資料流圖在 p.227。
序列一長,計算量就以平方成長。原書把各種稀疏化的互動樣式畫成一組矩陣圖,在 p.228,正文說明在 pp.227–228。影像上的做法在 pp.228–232:逐像素自迴歸的實驗在 p.229,切塊後送進編碼器的做法在 pp.229–231,多尺度與視窗位移的架構在 pp.230–232,各家在同一個影像基準上的錯誤率也列在這幾頁。
真正值得單獨翻的是章末註記,pp.232–238。那裡有循環網路的簡史與它遺忘長距離資訊的問題(p.233)、切分方法的族譜(p.234)、各種解碼演算法的比較(p.235)、注意力機制的變體家族以及它與其他模型的關係(pp.235–236)、位置編碼的研究綜述(p.236)、把序列拉長的三條技術路線(p.237),還有訓練為何不穩定、為何需要學習率暖身(pp.237–238)。影像、影片與圖文聯合模型的文獻索引在 p.238。
原書在這一章的頁緣另外標了四個可執行的示範筆記本,分別對應自注意力、多頭、切分與解碼策略,位置在 p.213、p.215、p.218 與 p.224,想照著原書動手跑的人可循這幾個標記找。
章末 p.239 有十道習題。其中一道要你親手算一組 softmax 的偏導數,算完就會明白為什麼分數不能太大;另一道要你證明打亂輸入的順序只會讓輸出跟著換位置,那正是位置編碼存在的理由。
原書第 12 章對應印刷頁 pp.207–239。
§04作業和解答
作業一:手算一組發言權,並看它怎麼被分數的尺度扭曲
會議桌上有三位與會者,鍵向量分別是 \(\mathbf{k}_1=(1,0)\)、\(\mathbf{k}_2=(0,1)\)、\(\mathbf{k}_3=(1,1)\),維度 \(D_q = 2\)。現在第二位提出一支查詢 \(\mathbf{q}=(2,0)\)。(a)算出三個未縮放的內積分數,再算出除以 \(\sqrt{D_q}\) 之後的分數與對應的三個 \(\alpha_{m2}\)。(b)如果不做縮放,直接對原始分數取 softmax,三個權重各是多少?(c)把查詢改成 \(\mathbf{q}=(10,0)\)(方向不變、長度變五倍),不做縮放時三個權重各是多少?從(a)到(c)說出你看到的趨勢,以及它對梯度的意義。
解答 SOLUTION
(a)內積分別是 \(\mathbf{k}_1\cdot\mathbf{q} = 2\)、\(\mathbf{k}_2\cdot\mathbf{q} = 0\)、\(\mathbf{k}_3\cdot\mathbf{q} = 2\)。除以 \(\sqrt{2} \approx 1.4142\) 之後是 1.4142、0、1.4142。取指數得 4.1133、1、4.1133,總和 9.2266;三個權重就是 0.4458、0.1084、0.4458。第一位與第三位分數相同,所以分到一樣多,這是式子的必然結果,不是巧合。
(b)不縮放時取指數得 7.3891、1、7.3891,總和 15.7781,權重是 0.4683、0.0634、0.4683。跟(a)比,兩端拉高、中間壓低——分數的尺度變大,分布就變尖。
(c)分數變成 10、0、10,權重是 0.5000、0.0000227、0.5000。中間那一位幾乎被歸零了。趨勢很清楚:分數的絕對尺度一放大,softmax 就從「分配比例」退化成「幾乎只挑一個」。對梯度的意義在於 \(\alpha(1-\alpha)\):(a)的中間項是 \(0.1084 \times 0.8916 \approx 0.0967\),(c)的中間項是 \(0.0000227 \times 0.9999773 \approx 0.0000227\),掉了三個數量級以上。這一項乘進反向傳播的鏈條裡,那一路的梯度就幾乎沒了。除以 \(\sqrt{D_q}\) 正是為了阻止分數的尺度隨維度失控。(以上數值以 numpy 重算核對過。)
作業二:為什麼遮罩一定要在 softmax 之前
第三節那個未遮罩的注意力矩陣,第 3 欄是 [0.209 0.021 0.571 0.062 0.136](完整精度為 0.208963、0.021426、0.571327、0.062148、0.136135)。現在要對它加上「只准看自己與前面」的遮罩,也就是第 4、5 兩項不該存在。(a)如果做法是「先算完 softmax,再把第 4、5 項直接設成 0」,剩下三項加起來是多少?這對輸出的尺度有什麼後果?(b)如果做法是「先算完 softmax,把第 4、5 項設成 0,再把剩下三項重新歸一化」,算出三個權重。(c)把(b)的答案,和第六節程式印出的那個遮罩矩陣的第 3 欄(由上往下讀是 0.261、0.027、0.713、0、0)比對,你發現什麼?據此說明「在 softmax 之前壓成負無窮」的真正好處是什麼。
解答 SOLUTION
(a)0.208963 + 0.021426 + 0.571327 = 0.801716。加起來不到 1,缺了將近兩成。後果是 \(\mathbf{y}_n\) 從「值向量的加權平均」變成「加權後又整體縮小了 0.8 倍」的東西,而且縮小的倍率每一欄都不一樣——越靠前的位置被砍掉的比重越多、縮得越兇。後面接的正規化層雖然會沖淡一部分影響,但這個尺度差本身沒有任何道理,純粹是做法錯了造成的。
(b)各項除以 0.801716:0.208963 / 0.801716 = 0.260645,0.021426 / 0.801716 = 0.026725,0.571327 / 0.801716 = 0.712630。
(c)第六節程式印出的遮罩矩陣第 3 欄是 0.261、0.027、0.713,跟(b)的 0.260645、0.026725、0.712630 完全對得上。所以「事前壓成負無窮」和「事後歸零再重新歸一化」在數學上是同一件事——真正錯的只有(a)那種「歸零卻不重新歸一化」的做法。事前做的好處因此不是數值不同,而是三點:一、不必多寫一次正規化,錯的機率低;二、被遮住的那些分數根本不會被算進指數,省掉一批計算;三、也是最重要的,遮罩變成「打分階段的一個設定」,於是第六節那三種可見範圍可以用同一份程式碼、只換一個布林矩陣就切換過去。
作業三:參數量與計算量,哪一個怕長句子
沿用第五節那段程式的口徑:一層多頭自注意力有 \(H\) 個頭,總維度 \(D\),每個頭維度 \(D/H\),序列長度 \(N\)。忽略偏置。(a)寫出投影與混合的總參數量,並說明它跟 \(N\) 有沒有關係。(b)寫出總乘法次數的表達式(含投影、打分、加權、混合四項),指出哪一項隨 \(N\) 平方成長。(c)取 \(D = 48\),算出 \(N\) 從 10 變成 40 時總乘法次數變成幾倍。(d)在什麼樣的 \(N\) 之下,平方那一項會第一次超過線性那些項?用 \(D\) 表示。
解答 SOLUTION
(a)每個頭有三組投影,各是 \((D/H) \times D\),一個頭就是 \(3D^2/H\);\(H\) 個頭合計 \(3D^2\),跟 \(H\) 無關。加上混合矩陣 \(\boldsymbol\Omega_c\) 的 \(D^2\),總參數量是 \(4D^2\)。它與 \(N\) 完全無關——這正是第二節說的「參數量與序列長度無關」。以 \(D = 48\) 代入是 \(4 \times 2304 = 9216\),與第五節四行輸出的 9216 相符。
(b)投影:\(3D^2 N\)(每個頭 \(3(D/H)DN\),乘上 \(H\) 個頭)。打分:\(N^2 D\)(每個頭 \(N^2 (D/H)\))。加權:同樣 \(N^2 D\)。混合:\(D^2 N\)。總計 \(4D^2N + 2N^2D\)。平方成長的是打分與加權那兩項,因為它們要跑遍所有位置配對。
(c)\(N = 10\):\(4 \times 2304 \times 10 + 2 \times 100 \times 48 = 92{,}160 + 9{,}600 = 101{,}760\),與第五節的實跑輸出一致。\(N = 40\):\(4 \times 2304 \times 40 + 2 \times 1600 \times 48 = 368{,}640 + 153{,}600 = 522{,}240\)。倍率是 \(522{,}240 / 101{,}760 \approx 5.13\)。注意兩項的成長速度完全不同:線性項變成 4 倍,平方項變成 16 倍。
(d)令 \(2N^2D > 4D^2N\),兩邊同除以 \(2ND\)(\(N, D\) 都是正的)得 \(N > 2D\)。也就是說序列長度超過總維度的兩倍之後,打分那一塊才開始主導計算量。以 \(D = 48\) 為例是 \(N = 96\),此時兩項都是 884,736,剛好相等。這解釋了一件容易被誤會的事:短序列時平方項根本不是瓶頸,是把上下文拉長之後它才變成主角。(本題數值以 numpy 重算核對過。)
§05參考資料
- Attention Is All You Need(原始論文) — 提出這個架構的論文,讀完本章之後再讀它,會發現大半的式子都認得
- The Illustrated Transformer — 用大量圖解走一遍同樣的機制,跟本章的敘事順序不同,適合當第二個角度
- The Annotated Transformer — 逐行代碼對照論文的實作導讀,想看完整版(含訓練迴圈)時的第一站
- Hugging Face NLP Course — 免費線上課,把切分、預訓練與微調這幾段本章只帶過的東西補齊
- nanoGPT — 一份刻意寫得很短的語言模型實作,本章第六節那套遮罩在裡面找得到對應的幾行
- NumPy 官方使用手冊 — 本章五段程式只依賴 numpy,查廣播與矩陣運算語法用它
- Understanding Deep Learning(MIT Press) — 本課課綱主題所本的原書出版頁(ISBN 9780262048644,2023-12 出版)
- udlbook 官方網站(作者釋出的 PDF、投影片與習題) — 原書作者維護的免費資源站(udlbook.com 會轉址到此)