注意力機制
一個字要理解自己的意思,得看看前面都有哪些字。從最簡單的"平均前面所有的字"出發,一步步加上查詢、鍵、值,因果掩碼和縮放,寫出一個完整的注意力頭,並和 PyTorch 的實現核對。
- 約 50 分鐘
- 難度:深入
- 實測:2026-09-15 torch 2.14,CPU,固定隨機種子
程式碼和執行結果保留原樣(簡體中文),與實際執行時完全一致。
"明月"的"明"和"明天"的"明",意思不一樣。一個字該怎麼理解,要看它周圍是什麼字。
我們要訓練的模型,每個位置的任務是預測下一個字。它在"床前明月"之後要寫出"光",就必須知道前面出現過"明月"。所以每個位置都需要一個辦法,從前面的字那裡收集資訊。
注意力(attention)就是這個辦法,它是 Transformer 的核心。這一課從最簡單的做法開始,一步一步把它搭出來。
python attention.py
最簡單的辦法:取平均
每個字先用一個向量表示(第 01 模組第 5 課講過向量嵌入)。最簡單的"收集前面的資訊",就是把自己和前面所有字的向量取平均。
這件事可以用一次矩陣乘法完成。構造一個下三角的矩陣,每一行只有前面若干個位置有值,而且加起來是 1:
x = torch.randn(T, 4) # 5 个词,每个词一个 4 维向量(先随便取)
weights = torch.ones(T, T).tril() # 下三角:第 i 行只有前 i+1 个位置是 1
weights = weights / weights.sum(dim=1, keepdim=True) # 每行除以个数,变成平均
out = weights @ x
== 1. 用一个矩阵乘法,让每个位置得到自己和前面所有位置的平均
tensor([[1.00, 0.00, 0.00, 0.00, 0.00],
[0.50, 0.50, 0.00, 0.00, 0.00],
[0.33, 0.33, 0.33, 0.00, 0.00],
[0.25, 0.25, 0.25, 0.25, 0.00],
[0.20, 0.20, 0.20, 0.20, 0.20]])
第 3 个位置('明')的结果 = 前三个向量的平均?True
第 i 行表示第 i 個字從每個位置拿多少:第 1 個字只能看自己,第 3 個字平均地看前三個,第 5 個字平均地看全部五個。weights @ x 一次就算出了所有位置的結果。
這個矩陣乘法的寫法請記住,後面的注意力也是這個形式:一個權重矩陣乘以一堆向量。不同的只是權重怎麼來。
平均的問題很明顯:所有字一視同仁。"光"想知道的是前面有沒有"月","床"和"前"對它來說沒那麼重要。權重應該由內容決定。
查詢和鍵:讓每個字決定該看誰
辦法是給每個字準備兩個向量:
- 查詢(query,q):我想找什麼樣的資訊。
- 鍵(key,k):我這裡有什麼樣的資訊。
一個字對另一個字的"關注程度",就用自己的 q 和對方的 k 做點積來打分。兩個向量方向越接近,點積越大。
先用手工造的向量看看效果。設向量有兩維,第一維表示"和天空有關",第二維表示"和位置有關"。"光"的查詢是 [1, 0],表示它想找和天空有關的東西:
k = torch.tensor([[0.1, 0.9], # 床:一个地点
[0.0, 1.0], # 前:一个方位
[0.9, 0.1], # 明:和天空、光有关
[1.0, 0.0], # 月:天空里的东西
[0.8, 0.2]]) # 光
q_guang = torch.tensor([1.0, 0.0]) # "光"想找的是:和天空有关的东西
scores = k @ q_guang
== 2. 点积打分:查询(q)和每个键(k)越像,分数越高
'光' 对每个词的分数: {'床': 0.1, '前': 0.0, '明': 0.9, '月': 1.0, '光': 0.8}
softmax 之后的权重: {'床': 0.12, '前': 0.11, '明': 0.26, '月': 0.29, '光': 0.23}
分数放大 5 倍再 softmax: {'床': 0.01, '前': 0.0, '明': 0.3, '月': 0.5, '光': 0.18}
分數經過 softmax 變成權重(第 08 模組第 5 課講過 softmax),加起來是 1。"月"的分數最高,拿到的權重最大。
注意最後一行:分數整體放大 5 倍,softmax 之後的權重就變得集中得多,"月"一個拿走了一半。分數的大小決定了注意力是"雨露均霑"還是"只盯著一個"。這一點在後面講縮放時很重要。
真實的模型裡,q 和 k 不是手工造的,而是用兩個矩陣從每個字的向量算出來的,矩陣裡的參數通過訓練學到。模型會自己學會,什麼樣的字該去找什麼樣的資訊。
值:真正傳遞的內容
打完分,決定了從每個位置拿多少,那拿的是什麼?是第三個向量:值(value,v),表示"如果你關注我,我給你這些資訊"。
為什麼不直接拿原來的向量?因為"用來匹配的資訊"和"用來傳遞的資訊"不一定一樣。把它們分開,模型更靈活。
三個向量合在一起,就是注意力的全部:
权重 = softmax(q 和每个 k 的点积)
输出 = 用这些权重,把每个位置的 v 加起来
因果掩碼:不能偷看後面
訓練時,模型要在每個位置預測下一個字。如果第 3 個位置能看到第 4 個字,它就可以直接"抄答案",什麼都學不到。所以每個位置只能看自己和前面的位置。
做法是在 softmax 之前,把後面位置的分數設成負無窮。負無窮取指數是 0,softmax 之後權重就是 0:
scores = q @ k.T / math.sqrt(8)
mask = torch.ones(T, T, dtype=torch.bool).tril()
scores = scores.masked_fill(~mask, float("-inf"))
== 3. 因果掩码:把后面位置的分数设成负无穷,softmax 之后权重就是 0
tensor([[1.00, 0.00, 0.00, 0.00, 0.00],
[0.96, 0.04, 0.00, 0.00, 0.00],
[0.24, 0.26, 0.50, 0.00, 0.00],
[0.09, 0.39, 0.43, 0.09, 0.00],
[0.02, 0.14, 0.02, 0.09, 0.73]])
和第一節的平均矩陣對比:形狀一樣,都是下三角,每行加起來都是 1。區別是現在每行的權重不再平均,而是由 q 和 k 決定的。這裡的 q 和 k 是隨機的,所以權重看起來沒有規律;訓練之後,它們會變得有意義。
因為有這個掩碼,這種注意力叫因果自注意力:自注意力是指 q、k、v 都來自同一段文字,因果是指只能看過去、不能看未來。GPT 這類"一個字一個字往後寫"的模型用的都是它。
為什麼要除以根號維度
上面的程式碼裡有一個 / math.sqrt(8),8 是 q 和 k 的維度。這一步叫縮放,看看不做會怎樣:
== 4. 为什么要除以 √d:维度越大,点积的数值越大,softmax 会变得非常极端
d= 16 点积的标准差 3.89(√d = 4.00) 每行最大权重的平均:不缩放 0.79,除以 √d 后 0.53
d= 64 点积的标准差 8.17(√d = 8.00) 每行最大权重的平均:不缩放 0.86,除以 √d 后 0.47
d= 256 点积的标准差 16.08(√d = 16.00) 每行最大权重的平均:不缩放 0.95,除以 √d 后 0.47
d= 1024 点积的标准差 32.53(√d = 32.00) 每行最大权重的平均:不缩放 0.92,除以 √d 后 0.39
兩個隨機向量的點積,是 d 個乘積加起來。加的項越多,結果的波動越大,標準差大約是根號 d:d=256 時點積的標準差是 16。
分數的波動大了,就像第二節"放大 5 倍"那樣,softmax 會變得極端:不縮放時,每一行最大的那個權重平均在 0.8 到 0.95,幾乎所有注意力都給了一個位置。這在訓練剛開始時很糟糕:權重接近 0 或 1 的地方,softmax 的梯度非常小,模型很難學習。
除以根號 d,分數的標準差回到 1 左右,每行最大的權重在 0.4 到 0.5 之間,注意力是分散的,梯度正常,模型可以慢慢學會該關注誰。
一個完整的注意力頭
把前面的東西合在一起:
d_model, d_head = 16, 8
x = torch.randn(T, d_model)
W_q, W_k, W_v = (torch.randn(d_model, d_head) / math.sqrt(d_model) for _ in range(3))
q, k, v = x @ W_q, x @ W_k, x @ W_v
scores = (q @ k.T / math.sqrt(d_head)).masked_fill(~mask, float("-inf"))
out = F.softmax(scores, dim=-1) @ v
== 5. 一个完整的注意力头:q、k、v 都由同一个输入经过不同的矩阵得到
输入 (5, 16) → q、k、v 各 (5, 8) → 输出 (5, 8)
和 PyTorch 自带的 scaled_dot_product_attention 比,最大差别 2.4e-07
五行程式碼:三個矩陣把輸入變成 q、k、v,打分、縮放、掩碼,softmax,加權求和。結果和 PyTorch 自帶的 scaled_dot_product_attention 一致(差別在小數點後 7 位,是浮點數計算順序不同帶來的誤差)。
W_q、W_k、W_v 就是這個注意力頭要學的參數。訓練時,它們和其他參數一樣,通過反向傳播和梯度下降更新。
這個公式最早出自 2017 年的論文《Attention Is All You Need》,寫成數學形式是 softmax(QKᵀ/√d)V。現在你知道其中每一個符號是做什麼的了。
注意力的代價
每個位置都要和前面所有的位置打分。序列長度是 T,分數表就是 T×T 的大小。長度翻一倍,計算量和這張表佔的記憶體就變成四倍。
這是長上下文很貴的根本原因之一(第 01 模組第 4 課講過上下文和成本)。各家模型為了支援幾十萬、上百萬詞元的上下文,在注意力的計算方式上做了大量最佳化,但基本的思路還是這一課的這幾行。
練習
- 在第二節裡,給"光"換一個查詢向量
[0, 1](想找和位置有關的東西),權重會怎麼變? - 在第五節裡,去掉因果掩碼,再和
scaled_dot_product_attention(..., is_causal=False)對比。 - 把第五節的序列長度從 5 改成 1000、2000、4000,用
time.time()測一下每種長度要多少時間,是不是大約按平方增長?
自測
1. 查詢、鍵、值分別起什麼作用?
查詢表示"我想找什麼",鍵表示"我這裡有什麼"。用一個位置的查詢和所有位置的鍵做點積,得到它對每個位置的關注分數,經過 softmax 變成權重。值表示"關注我時能拿到的資訊",輸出就是按這些權重把所有位置的值加起來。
2. 因果掩碼是怎麼實現的?為什麼需要它?
在 softmax 之前把每個位置後面的分數設成負無窮,softmax 之後這些位置的權重就是 0。需要它是因為模型要學習預測下一個字,如果能看到後面的字,就能直接抄答案,什麼也學不到。
3. 注意力分數為什麼要除以根號維度?
d 維的兩個向量的點積,標準差大約是根號 d,維度越大分數的波動越大,softmax 之後的權重會極端地集中在一個位置,梯度變得很小,模型難以訓練。除以根號 d 讓分數的標準差回到 1 左右。
提問與討論
這一課沒看懂的地方,在這裡問。看到別人的問題,也歡迎你來回答。
提問 +3 點,回答別人 +6 點。內容經審核後公開。
正在載入討論…