ARK · 理解深度學習CHAPTER 12 / 21

CHAPTER 12 / 21 · PART 3 · 結構的力量:四種主力架構

Transformer:注意力機制與語言模型的骨架

Transformers

把注意力拆成打分、正規化、加權三步,補上縮放、位置與遮罩,組成一層可畫出來的資料流。

§01學習重點

§02課程內容

一、把一段文字丟進全連接網路,會壞在哪裡

前面幾章的網路都有一個沒說出口的前提:輸入是一條長度固定的向量,而且第 3 個位置與第 40 個位置各配一組自己的權重。這個前提對表格資料成立,對一段文字則會連續壞三次。

第一次壞在長度。 句子有長有短,可是全連接層的輸入寬度在設計時就釘死了。你只能取一個最大長度,短的補空、長的截斷——補空的部分佔掉了大半個輸入,截斷掉的部分則是直接丟資訊。

第二次壞在共用。 一段文字裡的每個位置,本該用同一套規則處理。出現在第 3 個位置的某個詞,跟出現在第 40 個位置的同一個詞,是同一個詞;沒有任何理由要模型在兩個地方各學一次怎麼對付它。全連接層卻正好相反:它給每個位置各配一組權重,等於強迫模型把同一件事重學很多遍。

第三次壞在參數量。 假設每個位置用 300 個數字表示、一段最多容納 800 個位置,全部攤平就是 240,000 個數字。再接一層同樣寬的全連接層,光是那一層的權重矩陣就有 240,000 × 240,000 = 5.76 × 10¹⁰ 個參數。這還只是一層。

第 10 章的卷積網路(convolutional network)已經解決過同一型的前兩個毛病:同一組權重掃過每個位置,長度不再綁死,規則也自動共用。但它留下一個缺口。一個卷積核一次只覆蓋一小段,隔得遠的兩個位置要靠疊很多層才連得上;而文字裡需要建立關聯的兩個位置,往往隔了整整一段。

還缺最後一塊,而且這一塊才是關鍵:哪兩個位置該建立關聯,不能寫死在結構裡,得由內容當場決定。 卷積的連法是預先排好的——不管讀到什麼,第 n 個位置永遠只看左右那幾格。文字不吃這一套:一個代名詞該連到哪個名詞,答案在句子裡,不在結構裡。

三件事加起來,你要的東西輪廓就出來了:一種運算,處理規則對每個位置共用、輸入長度可以變,而且任何兩個位置之間的關聯強弱由當下的內容算出來。這就是自注意力(self-attention),本章其餘的篇幅都在把它拆開來看。

本章從頭到尾用同一個場景對照:把一段文字想成一場圓桌會議,每個位置是一位與會者,每位手上都有一份準備好要說的內容。自注意力要決定的,就是每一輪裡誰的話該佔多少分量。

二、先看合成:輸出是一份加權平均

自注意力有三步,我打算倒著講:先講最後一步「怎麼合成」,再回頭問「比重哪裡來」。理由很實際——後者難得多,而且只有先接受了前者,後者才有地方可放。

第一步單純。每個位置各自算出一份向量,叫它的值(value):

$$ \mathbf{v}_m \;=\; \boldsymbol\Omega_v \mathbf{x}_m + \mathbf{b}_v $$

逐項拆解:\(\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\) 個輸出是所有值向量的加權和:

$$ \mathbf{y}_n \;=\; \sum_{m=1}^{N} \alpha_{mn}\,\mathbf{v}_m $$

逐項拆解:\(N\) 是這一次送進去幾個位置;\(\alpha_{mn}\)(讀作 alpha)是一個純量,代表第 \(n\) 個輸出要分給第 \(m\) 個輸入多少比重——兩個下標的順序是「先來源、後去處」,寫的時候容易顛倒,值得記一下;求和跑遍所有 \(N\) 個位置,所以每一個輸出都用到了全部的值向量,只是比重不同。

這些比重有兩個約束:每一個都不小於零,而且固定住輸出位置 \(n\) 之後,它們加起來剛好是 1:

$$ \alpha_{mn} \;\ge\; 0, \qquad \sum_{m=1}^{N} \alpha_{mn} \;=\; 1 $$

這兩條約束合起來,說的就是「一整份被切開分掉」。權重不是拿來選一個,是拿來分配比例。 這句話值得多停一秒:不是「第 3 個輸出決定聽第 1 位的」,而是「第 3 個輸出分給第 1 位兩成、分給第 3 位將近六成、剩下的分給其他人」。

有三件事現在就該記住。第一,輸出與輸入的尺寸一模一樣,都是 \(D \times N\),所以這種區塊可以一層接一層疊上去,中間不需要改形狀。第二,如果把所有 \(\alpha_{mn}\) 都設成 \(1/N\),每個輸出就變成所有值向量的平均,而且 \(N\) 個輸出彼此完全相同——這是這個機制最沒用的一組權重,也剛好說明有用的部分全在權重的差異上。第三,到目前為止還沒有出現任何激活函數。

比喻: 一場圓桌會議正在討論某個議題,每位與會者手上都有一份準備好的發言稿,那就是他的值向量。輪到第三位整理自己這一輪的結論時,他不是挑一個人的話照抄,而是把發言權按比例分掉:第一位佔兩成、第二位幾乎不佔、他自己佔將近六成,剩下的分給後兩位——這一整份配額加起來是一,不能超發。他的結論就是這些發言稿按這組比例混出來的東西。這個比喻有一處明確失準:真人開會時「聽誰講」是互斥的,同一時間只能專心聽一個人;自注意力卻是同時按比例把所有人的內容都收進來,沒有先聽誰、後聽誰的問題。到了第五節談多頭時,這個差異還會再放大一次。

三、比重從哪裡來:查詢、鍵,以及 softmax 沿哪個方向做

現在回答上一節跳過的問題:\(\alpha_{mn}\) 是怎麼算出來的?

它得滿足兩個條件。一是由內容決定——第 \(n\) 個位置該分多少給第 \(m\) 個位置,要看這兩個位置各自帶著什麼。二是可訓練,也就是中間那些可調的數字得待在梯度算得動的地方。

做法是把同一批輸入再打成兩種角色,叫查詢(query)與鍵(key):

$$ \mathbf{q}_n \;=\; \boldsymbol\Omega_q \mathbf{x}_n + \mathbf{b}_q, \qquad \mathbf{k}_m \;=\; \boldsymbol\Omega_k \mathbf{x}_m + \mathbf{b}_k $$

逐項拆解:\(\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\))。

寫成式子:

$$ \alpha_{mn} \;=\; \frac{\exp\!\left[\,\mathbf{k}_m^{\mathsf{T}}\mathbf{q}_n / \sqrt{D_q}\,\right]} {\sum_{m'=1}^{N}\exp\!\left[\,\mathbf{k}_{m'}^{\mathsf{T}}\mathbf{q}_n / \sqrt{D_q}\,\right]} $$

逐項拆解:分子是第 \(m\) 個輸入對第 \(n\) 個輸出的分數取指數;分母把所有輸入位置的同一種東西加起來,加總用的啞變數寫成 \(m'\) 以免跟分子的 \(m\) 混淆。指數把可正可負的分數變成正數,除以總和讓整排加起來是 1——所以固定 \(n\) 時和為 1 是式子本身保證的,固定 \(m\) 時則沒有任何這樣的保證。分母裡的 \(\mathbf{q}_n\) 從頭到尾沒變,變的只有鍵,這就是「固定輸出、掃過所有輸入」的意思。至於為什麼要除以 \(\sqrt{D_q}\),下一節專門處理。

把所有比重排成一個 \(N \times N\) 的矩陣 \(\mathbf{A}\)(列是輸入位置 \(m\),欄是輸出位置 \(n\)),這件事就變成可以直接數的:每一欄加起來是 1,每一列不是。 整段運算壓成矩陣寫法是:

$$ \mathrm{SA}[\mathbf{X}] \;=\; \mathbf{V}\,\mathrm{softmax}\!\left[\frac{\mathbf{K}^{\mathsf{T}}\mathbf{Q}}{\sqrt{D_q}}\right] $$

逐項拆解:\(\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}\),等於拿權重矩陣的每一欄去混合各個值向量。

手刻一遍最能把方向釘住。底下這段程式造一場五個位置的小型會議,六維向量,查詢與鍵四維:

PYTHON
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)

實跑輸出的注意力矩陣是

CODE
[[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)\) 的最大值:

PYTHON
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}")

實跑輸出:

CODE
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 位的座位對調,輸出也只是把對應的兩欄對調,內容一個字都不變。原因就寫在式子裡:從頭到尾沒有任何一項提到「第幾個」,每個位置只認得內容。

PYTHON
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{MHSA}[\mathbf{X}] \;=\; \boldsymbol\Omega_c \begin{bmatrix} \mathrm{SA}_1[\mathbf{X}] \\ \mathrm{SA}_2[\mathbf{X}] \\ \vdots \\ \mathrm{SA}_H[\mathbf{X}] \end{bmatrix} $$

逐項拆解:\(\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、同樣的十個位置,只改頭數:

PYTHON
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}")

實跑輸出:

CODE
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 層之後才長出來的——長出上下文,正是自注意力在做的事。

語言模型的目標函數。 把整句話的機率拆成一串條件機率的乘積:

$$ Pr(t_1, t_2, \dots, t_N) \;=\; \prod_{n=1}^{N} Pr\bigl(t_n \mid t_1, \dots, t_{n-1}\bigr) $$

逐項拆解:\(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 了——加權平均變成加權「部分平均」,輸出的尺度會隨著位置一路縮水。要修就得再正規化一次,那等於白繞一圈。直接在分數上動手,乾淨得多。

PYTHON
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))

實跑輸出的矩陣是

CODE
[[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 章會正面處理「沒有標籤時到底該學什麼」這一整類問題。

最後,五件本章只點名不展開的事。

  1. 計算量隨長度平方成長。 每個位置都要跟每個位置打分,位置數翻倍,那張矩陣的格子就變成四倍。這是「上下文長度」為什麼是個大事的原因。有一整族方法在讓這張矩陣變稀疏或變小,本章不展開。
  2. 影像也能用同一套。 把圖切成小方塊,每一塊當成一個位置送進去就行。它與第 10 章的卷積是一組取捨:卷積把「鄰近的像素比較相關」這條歸納偏好(inductive bias)寫死在結構裡,注意力不寫,代價是得靠更多資料自己學出來。
  3. 生成時的取樣策略。 每一步都挑機率最大的詞元,容易吐出空轉、重複的句子,所以實務上有各種折衷的取樣做法。
  4. 這個機制之前的主力是循環網路(recurrent network)。 一次讀一個位置、把一個狀態往後傳,天然帶著順序,但隔得遠的資訊在傳遞途中會被沖淡。
  5. 這類模型的訓練不太穩定。 學習率的安排要特別設計,成因不只一個,本章不展開。

§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.2610.0270.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參考資料