制原理與實(shí)戰(zhàn):從QKV到Transformer應(yīng)用詳解)
如果你經(jīng)常用大模型聊天大概率遇到過(guò)這個(gè)場(chǎng)景上一輪還在討論周末去哪家火鍋店下一輪問(wèn)一句“那家店需要提前預(yù)約嗎”模型居然知道“那家店”指的就是剛才那家店。這個(gè)能力并不來(lái)自某種神奇的“記憶”而是來(lái)自一個(gè)在深度學(xué)習(xí)中被廣泛使用的組件——注意力機(jī)制Attention Mechanism。本文會(huì)從原理、變體、代碼到工程實(shí)踐把注意力機(jī)制完整拆開(kāi)講解。內(nèi)容包括QKV 到底是什么、縮放點(diǎn)積注意力是怎么算的、多頭注意力為什么有效、CV 領(lǐng)域常見(jiàn)的 SE/CBAM/ECA/CA 注意力、上下文窗口與注意力之間的關(guān)系以及一個(gè)基于 PyTorch 的完整可運(yùn)行示例。適合剛接觸深度學(xué)習(xí)的同學(xué)也適合想在業(yè)務(wù)模型中加入注意力模塊的開(kāi)發(fā)者。1. 背景與核心概念注意力機(jī)制解決什么問(wèn)題1.1 從 RNN 的長(zhǎng)距離依賴(lài)?yán)Ь痴f(shuō)起在 Transformer 出現(xiàn)之前處理文本序列主要靠 RNN循環(huán)神經(jīng)網(wǎng)絡(luò)和 LSTM。RNN 的思路是“按順序讀”當(dāng)前時(shí)刻的輸出依賴(lài)上一時(shí)刻的隱藏狀態(tài)。這種串行結(jié)構(gòu)有兩個(gè)明顯問(wèn)題長(zhǎng)距離依賴(lài)難以捕捉。句子開(kāi)頭的信息傳到句子結(jié)尾時(shí)經(jīng)過(guò)多次非線性變換后容易衰減或丟失。串行計(jì)算難以并行。一個(gè) token 一個(gè) token 地往后走GPU 的并行能力得不到充分發(fā)揮。舉個(gè)例子在英文句子 “The cat that chased the mouse was tired” 中真正決定was單復(fù)數(shù)的是最前面的cat而不是緊挨著它的mouse。RNN 雖然理論上能處理這種依賴(lài)但實(shí)際訓(xùn)練時(shí)效果和效率都不理想。注意力機(jī)制的出現(xiàn)改變了游戲規(guī)則。它允許模型在計(jì)算某個(gè)位置的表示時(shí)直接“查看”序列中任意其他位置并按相關(guān)性分配權(quán)重。這樣一來(lái)長(zhǎng)距離依賴(lài)變成了一個(gè)加權(quán)求和問(wèn)題不同位置之間也不再是嚴(yán)格的串行依賴(lài)可以進(jìn)行并行計(jì)算。1.2 注意力機(jī)制的直觀理解注意力機(jī)制的本質(zhì)可以概括為八個(gè)字按相關(guān)性加權(quán)匯總信息。想象你在讀一篇技術(shù)論文目光并不是平均掃過(guò)每一個(gè)字而是會(huì)在關(guān)鍵詞、公式、圖表附近停留更久。注意力機(jī)制做的就是這樣一件事對(duì)當(dāng)前需要理解的內(nèi)容計(jì)算它與其他位置信息的匹配程度再根據(jù)匹配程度從這些位置中提取內(nèi)容。相關(guān)性高的信息權(quán)重高貢獻(xiàn)大相關(guān)性低的信息權(quán)重低貢獻(xiàn)小。在深度學(xué)習(xí)里這種“匹配程度”通常被建模成一個(gè)概率分布所有位置的權(quán)重加起來(lái)等于 1模型選擇把“注意力”更多地放在哪些內(nèi)容上。1.3 為什么上下文理解離不開(kāi)注意力“上下文”這個(gè)詞在 AI 領(lǐng)域有兩層常見(jiàn)含義。在對(duì)話場(chǎng)景中它指前幾輪聊天內(nèi)容在文本場(chǎng)景中它指當(dāng)前詞周?chē)钠渌~。兩層含義的共同點(diǎn)是語(yǔ)義依賴(lài)于更廣的范圍而不只是當(dāng)前詞本身。注意力機(jī)制通過(guò)兩兩計(jì)算 token 之間的相關(guān)性把“遠(yuǎn)距離上下文”直接融入當(dāng)前 token 的表示。比如“銀行”前面出現(xiàn)“河”時(shí)注意力機(jī)制會(huì)讓“銀行”更傾向于理解為河岸而不是金融機(jī)構(gòu)如果前后文提到“存款”“利率”則更傾向于金融機(jī)構(gòu)。這種動(dòng)態(tài)、按輸入內(nèi)容變化的相關(guān)性計(jì)算正是模型理解上下文的關(guān)鍵。這里也順帶解釋了一個(gè)常見(jiàn)困惑模型并沒(méi)有顯式的“記憶”模塊它靠的是在每一層 Transformer 中對(duì)上下文所有 token 做注意力加權(quán)把相關(guān)信息融合進(jìn)當(dāng)前位置的向量。所謂“理解上下文”本質(zhì)上是多層注意力疊加后的結(jié)果。2. 核心原理QKV 與縮放點(diǎn)積注意力2.1 Query、Key、Value 三個(gè)角色要理解注意力機(jī)制繞不開(kāi) Q、K、V 三個(gè)概念。它們是一套檢索系統(tǒng)的抽象Query查詢(xún)表示“我現(xiàn)在想找什么”。Key鍵表示“我這里有什么可以被匹配”。Value值表示“如果匹配上了提供什么內(nèi)容”??梢灶?lèi)比搜索引擎用戶(hù)輸入的關(guān)鍵詞是 Query網(wǎng)頁(yè)標(biāo)題和標(biāo)簽是 Key網(wǎng)頁(yè)正文是 Value。搜索引擎先計(jì)算 Query 與 Key 的相關(guān)性再按相關(guān)程度返回 Value 中的內(nèi)容。在 Transformer 中Q、K、V 由輸入向量分別乘上三個(gè)可學(xué)習(xí)的權(quán)重矩陣得到# 偽代碼展示 Q/K/V 的生成方式 Q X W_q K X W_k V X W_v其中X是輸入序列的表示W(wǎng)_q、W_k、W_v都是可訓(xùn)練參數(shù)。通過(guò)訓(xùn)練模型會(huì)學(xué)會(huì)“什么樣的 Key 能匹配上當(dāng)前 Query”。2.2 縮放點(diǎn)積注意力的計(jì)算過(guò)程最常用的注意力計(jì)算方式是縮放點(diǎn)積注意力Scaled Dot-Product Attention公式可以寫(xiě)成Attention(Q, K, V) softmax(Q * K^T / sqrt(d_k)) * V逐步拆解Q與K做點(diǎn)積得到兩個(gè) token 之間的相似度分?jǐn)?shù)。除以sqrt(d_k)做縮放。d_k是 Key 向量的維度當(dāng)維度較大時(shí)點(diǎn)積數(shù)值可能非常大softmax 會(huì)趨向于 one-hot 分布梯度很小不利于訓(xùn)練縮放可以把數(shù)值拉回平緩區(qū)間。對(duì)分?jǐn)?shù)做 softmax 歸一化得到“注意力權(quán)重”每一行所有位置權(quán)重之和為 1。用權(quán)重對(duì)V做加權(quán)求和得到當(dāng)前位置的輸出向量。這里最關(guān)鍵的一步是 softmax。它決定了模型“關(guān)注誰(shuí)、忽略誰(shuí)”也是注意力機(jī)制可解釋性的來(lái)源——你可以直接觀察某個(gè) token 對(duì)哪些其他 token 分配了較高的權(quán)重。2.3 用 NumPy 跑一個(gè)最小例子先不看 Transformer我們用 NumPy 把核心計(jì)算完整實(shí)現(xiàn)一遍。假設(shè)有一個(gè)長(zhǎng)度為 3 的序列每個(gè) token 用 4 維向量表示這里為了演示讓 Q、K、V 來(lái)自同一份輸入。import numpy as np def scaled_dot_product_attention(Q, K, V): d_k K.shape[-1] scores np.dot(Q, K.T) / np.sqrt(d_k) # 數(shù)值穩(wěn)定的 softmax exp_scores np.exp(scores - scores.max(axis-1, keepdimsTrue)) weights exp_scores / exp_scores.sum(axis-1, keepdimsTrue) output np.dot(weights, V) return output, weights Q K V np.array([ [1.0, 0.0, 1.0, 0.0], # token_1 [0.0, 1.0, 0.0, 1.0], # token_2 [1.0, 1.0, 0.0, 0.0], # token_3 ]) output, weights scaled_dot_product_attention(Q, K, V) print(注意力權(quán)重:\n, weights) print(輸出:\n, output)運(yùn)行結(jié)果大致如下注意力權(quán)重: [[0.5065 0.1863 0.3072] [0.1863 0.5065