制:從核回歸到Transformer的權(quán)重化信息聚合)
1. 從“看哪里”到“學(xué)哪里”注意力機(jī)制的核心直覺在機(jī)器學(xué)習(xí)和深度學(xué)習(xí)的實(shí)踐中我們常常面臨一個(gè)根本性的挑戰(zhàn)如何處理海量的輸入信息無論是處理一張高分辨率圖片中的千萬像素還是分析一篇長(zhǎng)文檔中的每個(gè)詞語(yǔ)模型如果對(duì)每個(gè)輸入單元都“一視同仁”地投入同等計(jì)算資源不僅效率低下而且容易淹沒在噪聲中無法抓住關(guān)鍵信息。這就像我們?nèi)祟愒陂喿x時(shí)不會(huì)逐字逐句以相同的精力去分析而是會(huì)快速掃視將注意力集中在標(biāo)題、關(guān)鍵詞和核心段落上。這種“選擇性聚焦”的能力正是注意力機(jī)制試圖賦予模型的。注意力機(jī)制的核心思想可以概括為“權(quán)重化的信息聚合”。它不是一個(gè)具體的模型而是一種設(shè)計(jì)范式一種資源分配策略。其目標(biāo)是為輸入序列中的不同部分分配不同的重要性權(quán)重然后根據(jù)這些權(quán)重對(duì)信息進(jìn)行加權(quán)匯總從而得到一個(gè)更能代表當(dāng)前任務(wù)需求的上下文表示。簡(jiǎn)單來說它教會(huì)模型在“看”的時(shí)候知道“哪里更重要”。這種思想并非憑空而來其數(shù)學(xué)根源可以追溯到統(tǒng)計(jì)學(xué)中的非參數(shù)回歸方法特別是核回歸。核回歸為我們提供了一種優(yōu)雅的框架如何根據(jù)查詢點(diǎn)我們當(dāng)前關(guān)注的問題與一系列鍵值對(duì)歷史經(jīng)驗(yàn)或輸入數(shù)據(jù)的相似度來動(dòng)態(tài)地計(jì)算一個(gè)加權(quán)平均的預(yù)測(cè)值。注意力機(jī)制尤其是其最基礎(chǔ)的“注意力池化”形式可以看作是核回歸在深度學(xué)習(xí)語(yǔ)境下的一個(gè)神經(jīng)化、參數(shù)化的擴(kuò)展。理解了這個(gè)連接我們就能從更堅(jiān)實(shí)的統(tǒng)計(jì)基礎(chǔ)出發(fā)而不僅僅是把注意力當(dāng)作一個(gè)“魔法模塊”。在接下來的內(nèi)容里我們將從最直觀的“注意力提示”概念入手逐步深入到其數(shù)學(xué)實(shí)現(xiàn)——注意力池化并揭示其與核回歸的血緣關(guān)系。我們會(huì)用具體的例子和代碼片段展示如何從零構(gòu)建一個(gè)最簡(jiǎn)單的注意力模型并討論其在現(xiàn)代深度學(xué)習(xí)架構(gòu)如Transformer中的核心地位。無論你是剛?cè)腴T的新手還是希望鞏固基礎(chǔ)的老手理解這個(gè)“從統(tǒng)計(jì)到神經(jīng)網(wǎng)絡(luò)”的演進(jìn)路徑都將大有裨益。2. 注意力提示從生物本能到算法框架在深入公式之前讓我們先建立一個(gè)牢固的直覺。注意力本質(zhì)上是一種資源分配方案。在計(jì)算資源有限的前提下將更多的“算力”分配給更重要的輸入部分。2.1 生活中的注意力提示想象一下你在一個(gè)嘈雜的雞尾酒會(huì)上。房間里充滿了各種對(duì)話聲、音樂聲和杯盤碰撞聲。此時(shí)你的朋友叫了你的名字。盡管環(huán)境音的總音量可能遠(yuǎn)大于朋友的聲音但你的大腦會(huì)瞬間將“聽覺注意力”聚焦到朋友聲音傳來的方向抑制其他背景噪音。這里的“你的名字”就是一個(gè)強(qiáng)大的非自主性提示非自主性線索它基于刺激本身的突出性顯著性自動(dòng)捕獲了你的注意力。另一種情況是你正在聚精會(huì)神地閱讀一份復(fù)雜的項(xiàng)目報(bào)告尋找關(guān)于預(yù)算的部分。此時(shí)你的注意力是由你內(nèi)心的任務(wù)和目標(biāo)驅(qū)動(dòng)的這是一種自主性提示自主性線索。你主動(dòng)地、有意識(shí)地將認(rèn)知資源導(dǎo)向與“預(yù)算”相關(guān)的章節(jié)、表格和數(shù)字。在機(jī)器學(xué)習(xí)模型中這兩種提示都有其對(duì)應(yīng)物非自主性提示顯著性例如在圖像中一個(gè)像素與其周圍像素差異巨大高對(duì)比度邊緣、明亮斑點(diǎn)這個(gè)區(qū)域本身就具有視覺顯著性容易吸引模型的“注意”。在序列中一個(gè)出現(xiàn)頻率極低或極高的詞如專業(yè)術(shù)語(yǔ)或停用詞也可能具有統(tǒng)計(jì)顯著性。自主性提示任務(wù)驅(qū)動(dòng)這是更強(qiáng)大、更常用的方式。模型根據(jù)當(dāng)前要解決的具體任務(wù)例如“翻譯這句話”、“回答這個(gè)問題”、“檢測(cè)圖中的貓”生成一個(gè)查詢Query。這個(gè)查詢就像我們大腦中的“任務(wù)指令”用于在輸入數(shù)據(jù)鍵Key中尋找最相關(guān)的內(nèi)容并提取對(duì)應(yīng)的值Value。2.2. 查詢、鍵與值注意力機(jī)制的三元組這是理解注意力機(jī)制最關(guān)鍵的抽象。我們可以將其類比于信息檢索系統(tǒng)查詢Query代表當(dāng)前模型“想知道什么”或“關(guān)注什么”。例如在翻譯任務(wù)中當(dāng)模型在生成目標(biāo)語(yǔ)言的第t個(gè)詞時(shí)它需要一個(gè)查詢來回顧源語(yǔ)言句子中哪些部分最相關(guān)。鍵Key代表輸入數(shù)據(jù)中每個(gè)元素的“標(biāo)識(shí)”或“索引”。它用于與查詢進(jìn)行匹配計(jì)算相似度。鍵和查詢通常存在于同一個(gè)向量空間以便進(jìn)行相似度比較。值Value代表輸入數(shù)據(jù)中每個(gè)元素實(shí)際包含的“信息內(nèi)容”。一旦通過查詢-鍵匹配找到了相關(guān)的元素我們就需要提取這些元素所承載的具體信息值。一個(gè)簡(jiǎn)單的比喻你Query去圖書館輸入數(shù)據(jù)找一本關(guān)于“深度學(xué)習(xí)注意力機(jī)制”Query的內(nèi)容的書。圖書館的圖書檢索系統(tǒng)Key里存有每本書的標(biāo)題和關(guān)鍵詞Key。你輸入查詢系統(tǒng)返回一系列相似度高的書名Key-Query匹配。最后你根據(jù)這個(gè)列表去書架上找到對(duì)應(yīng)的書籍Value并閱讀其中的內(nèi)容聚合Value。在絕大多數(shù)注意力實(shí)現(xiàn)中鍵和值通常來源于同一個(gè)輸入序列甚至是相同的向量但分別經(jīng)過不同的線性變換層W_K,W_V投影到不同的空間以承擔(dān)不同的角色。查詢則可能來自另一個(gè)序列如解碼器狀態(tài)或同一序列的不同位置自注意力。注意這種“鍵-值”分離的設(shè)計(jì)是精妙的。它允許模型學(xué)習(xí)到根據(jù)什么特征Key去檢索和檢索到之后提取什么信息Value這兩者可以是不同的。例如在基于內(nèi)容的推薦系統(tǒng)中Key可以是電影的類型、演員用于匹配用戶興趣Query而Value可以是電影的詳細(xì)描述、評(píng)分用于最終生成推薦列表。3. 注意力池化核回歸的神經(jīng)化詮釋現(xiàn)在我們將直覺轉(zhuǎn)化為數(shù)學(xué)。注意力池化是注意力機(jī)制最基礎(chǔ)的計(jì)算單元其目標(biāo)就是根據(jù)查詢q對(duì)一組鍵值對(duì){(k1, v1), (k2, v2), ..., (kn, vn)}進(jìn)行加權(quán)求和得到輸出。3.1 從平均池化到加權(quán)池化假設(shè)我們有一組數(shù)據(jù)點(diǎn)(x1, y1), (x2, y2), ..., (xn, yn)想要預(yù)測(cè)在位置x查詢對(duì)應(yīng)的y值。最簡(jiǎn)單的方法是平均池化f(x) mean(yi)。這顯然不合理因?yàn)榫嚯xx遠(yuǎn)近不同的xi應(yīng)該對(duì)預(yù)測(cè)有不同貢獻(xiàn)。更合理的方法是Nadaraya-Watson核回歸。它的預(yù)測(cè)公式為f(x) Σ_i [α(x, xi) * yi]其中α(x, xi)是權(quán)重由核函數(shù)K計(jì)算得出α(x, xi) K(x - xi) / Σ_j K(x - xj)。這里x是查詢Query。xi是鍵Key即數(shù)據(jù)點(diǎn)的位置。yi是值Value即數(shù)據(jù)點(diǎn)的標(biāo)簽。K(·)是一個(gè)核函數(shù)如高斯核用于度量查詢x與鍵xi之間的相似度。相似度越高權(quán)重α越大。分母是一個(gè)歸一化項(xiàng)通常稱為注意力權(quán)重確保所有權(quán)重之和為1使得輸出f(x)是值yi的凸組合。這就是最原始的注意力池化模型在預(yù)測(cè)點(diǎn)x的輸出是所有訓(xùn)練樣本yi的加權(quán)平均權(quán)重取決于x與每個(gè)xi的相似度。3.2 引入可學(xué)習(xí)參數(shù)從非參數(shù)到參數(shù)化經(jīng)典的核回歸是非參數(shù)的核函數(shù)K是固定的如高斯函數(shù)。深度學(xué)習(xí)中的注意力池化則對(duì)其進(jìn)行了參數(shù)化改造使其能夠從數(shù)據(jù)中學(xué)習(xí)如何計(jì)算相似度。最常見的做法是使用加性注意力或縮放點(diǎn)積注意力來計(jì)算相似度分?jǐn)?shù)。以縮放點(diǎn)積注意力為例 假設(shè)查詢q、鍵k、值v都是向量。我們計(jì)算查詢與每個(gè)鍵的點(diǎn)積衡量相似度然后進(jìn)行縮放和歸一化Softmax得到權(quán)重最后對(duì)值進(jìn)行加權(quán)求和。# 偽代碼示意 def attention_pooling(query, keys, values): # 計(jì)算相似度分?jǐn)?shù) scores[i] query · keys[i] scores torch.matmul(query, keys.transpose(-2, -1)) # 縮放為了梯度穩(wěn)定并計(jì)算注意力權(quán)重 weights F.softmax(scores / sqrt(d_k), dim-1) # d_k是鍵向量的維度 # 加權(quán)求和 output torch.matmul(weights, values) return output, weights在這個(gè)框架下核函數(shù)K的角色被點(diǎn)積相似度或加性網(wǎng)絡(luò)替代并且查詢、鍵、值都可以通過神經(jīng)網(wǎng)絡(luò)W_Q, W_K, W_V從原始輸入學(xué)習(xí)得到。歸一化項(xiàng)Softmax對(duì)應(yīng)核回歸公式中的分母Σ_j K(x - xj)確保權(quán)重和為1。通過這種參數(shù)化注意力池化不再依賴于預(yù)設(shè)的、固定的距離度量如高斯核的歐氏距離而是可以學(xué)習(xí)適應(yīng)特定任務(wù)的最優(yōu)相似度計(jì)算方式。例如在文本任務(wù)中它可以學(xué)習(xí)到“蘋果”公司”和“水果”蘋果”在與不同查詢交互時(shí)應(yīng)有不同的相似度。3.3 一個(gè)簡(jiǎn)單的NumPy實(shí)現(xiàn)理解計(jì)算流程讓我們拋開深度學(xué)習(xí)框架用最基礎(chǔ)的NumPy來實(shí)現(xiàn)一個(gè)最簡(jiǎn)版的注意力池化加深理解。import numpy as np def nadaraya_watson_kernel_regression(x_train, y_train, x_query, bandwidth1.0): 簡(jiǎn)單的Nadaraya-Watson核回歸高斯核 x_train: 訓(xùn)練鍵 (n_samples,) y_train: 訓(xùn)練值 (n_samples,) x_query: 查詢點(diǎn) (1,) bandwidth: 高斯核的帶寬參數(shù) # 計(jì)算查詢與所有訓(xùn)練鍵的歐氏距離負(fù)數(shù)因?yàn)楦咚购耸蔷嚯x的減函數(shù) distances x_query - x_train # (n_samples,) # 使用高斯核計(jì)算非歸一化權(quán)重 unnormalized_weights np.exp(-distances**2 / (2 * bandwidth**2)) # (n_samples,) # 歸一化得到注意力權(quán)重 attention_weights unnormalized_weights / np.sum(unnormalized_weights) # (n_samples,) # 加權(quán)池化 y_pred np.sum(attention_weights * y_train) # (1,) return y_pred, attention_weights # 生成一些非線性數(shù)據(jù) np.random.seed(42) x_train np.linspace(-5, 5, 50) y_train np.sin(x_train) 0.2 * np.random.randn(50) # 正弦函數(shù)加噪聲 # 在多個(gè)查詢點(diǎn)上進(jìn)行預(yù)測(cè) x_queries np.linspace(-5, 5, 200) predictions [] all_weights [] for xq in x_queries: yp, aw nadaraya_watson_kernel_regression(x_train, y_train, xq, bandwidth0.5) predictions.append(yp) all_weights.append(aw) # 可視化此處省略繪圖代碼但概念上我們會(huì)看到一條平滑曲線 # 對(duì)于每個(gè)x_query模型都“注意”到了附近x_train的點(diǎn)并給出了預(yù)測(cè)。這個(gè)例子清晰地展示了注意力池化的流程計(jì)算相似度高斯核- 歸一化權(quán)重Softmax的連續(xù)類比- 加權(quán)求和池化。帶寬參數(shù)bandwidth控制了注意力的“聚焦”程度。帶寬小則注意力集中只關(guān)注非常近的點(diǎn)預(yù)測(cè)曲線波動(dòng)大帶寬大則注意力分散平滑效應(yīng)強(qiáng)。實(shí)操心得帶寬的選擇在核回歸或類似注意力中帶寬或縮放因子sqrt(d_k)是一個(gè)超參數(shù)其作用類似于卷積神經(jīng)網(wǎng)絡(luò)中的感受野。它決定了模型關(guān)注的范圍大小。在實(shí)踐中通常通過驗(yàn)證集來調(diào)整這個(gè)參數(shù)。在Transformer的縮放點(diǎn)積注意力中縮放因子sqrt(d_k)是為了防止點(diǎn)積結(jié)果過大導(dǎo)致Softmax梯度消失這是一個(gè)重要的工程技巧。4. 注意力機(jī)制的全景從池化到現(xiàn)代架構(gòu)基礎(chǔ)的注意力池化是一個(gè)強(qiáng)大的模塊但將其嵌入到完整的神經(jīng)網(wǎng)絡(luò)中并規(guī)?;耪嬲尫帕似錆摿?。4.1 注意力機(jī)制的幾種基本形態(tài)加性注意力早期RNN編碼器-解碼器架構(gòu)中常用。它通過一個(gè)小的前饋網(wǎng)絡(luò)來計(jì)算查詢和鍵的相似度score(q, k) v^T * tanh(W_q * q W_k * k)。這種方式更靈活但計(jì)算量稍大。點(diǎn)積注意力查詢和鍵直接做點(diǎn)積score(q, k) q^T * k。計(jì)算高效但要求查詢和鍵的維度相同且當(dāng)維度d_k較高時(shí)點(diǎn)積值的方差會(huì)變大容易將Softmax推入梯度極小的區(qū)域??s放點(diǎn)積注意力點(diǎn)積注意力的改進(jìn)版score(q, k) q^T * k / sqrt(d_k)??s放操作使得點(diǎn)積值的方差穩(wěn)定在1左右有利于訓(xùn)練。這是Transformer中使用的標(biāo)準(zhǔn)形式。自注意力當(dāng)查詢、鍵、值都來自同一個(gè)序列時(shí)稱為自注意力。它允許序列中的每個(gè)位置與序列中所有位置包括自身進(jìn)行交互從而捕捉序列內(nèi)部的長(zhǎng)期依賴關(guān)系。這是Transformer的核心。交叉注意力查詢來自一個(gè)序列如解碼器而鍵和值來自另一個(gè)序列如編碼器。常用于機(jī)器翻譯、問答等需要跨序列對(duì)齊的任務(wù)。4.2 多頭注意力并行化的注意力“專家”單一的注意力池化在每次計(jì)算時(shí)只能建立一種類型的依賴關(guān)系。為了讓模型同時(shí)關(guān)注來自不同表示子空間的信息提出了多頭注意力。其思想很簡(jiǎn)單將查詢、鍵、值通過不同的線性投影矩陣投影到h個(gè)不同的低維子空間頭。在每個(gè)頭上獨(dú)立地執(zhí)行縮放點(diǎn)積注意力得到h個(gè)輸出。最后將這些輸出拼接起來再通過一個(gè)線性投影得到最終結(jié)果。# 偽代碼概念 class MultiHeadAttention(nn.Module): def forward(self, Q, K, V): # 1. 線性投影拆分成h個(gè)頭 Q_heads split(self.W_q(Q)) # (batch, h, seq_len, d_k) K_heads split(self.W_k(K)) V_heads split(self.W_v(V)) # 2. 每個(gè)頭獨(dú)立計(jì)算注意力 head_outputs [] for i in range(h): output_i, _ scaled_dot_product_attention(Q_heads[i], K_heads[i], V_heads[i]) head_outputs.append(output_i) # 3. 拼接所有頭的輸出 concat_output concatenate(head_outputs) # (batch, seq_len, h*d_v) # 4. 最終線性投影 final_output self.W_o(concat_output) return final_output這相當(dāng)于讓h個(gè)不同的“注意力專家”同時(shí)工作一個(gè)可能專注于語(yǔ)法結(jié)構(gòu)一個(gè)可能專注于語(yǔ)義相似另一個(gè)可能專注于指代關(guān)系。最后綜合所有專家的意見做出更穩(wěn)健的決策。4.3 注意力機(jī)制在模型中的位置與作用在現(xiàn)代架構(gòu)中注意力機(jī)制通常不是孤立存在的而是與其它層交織在一起Transformer Block標(biāo)準(zhǔn)Transformer塊包含一個(gè)多頭自注意力層和一個(gè)前饋神經(jīng)網(wǎng)絡(luò)層每個(gè)層周圍都有殘差連接和層歸一化。這種設(shè)計(jì)使得注意力能夠被深度堆疊。編碼器-解碼器結(jié)構(gòu)在編碼器中使用自注意力來理解源序列的內(nèi)部結(jié)構(gòu)。在解碼器中使用掩碼自注意力防止看到未來信息和交叉注意力關(guān)注編碼器輸出來生成目標(biāo)序列。視覺Transformer將圖像分割成 patches每個(gè) patch 視為一個(gè) token然后直接應(yīng)用 Transformer 編碼器。其中的注意力機(jī)制讓模型能夠建立圖像塊之間的全局依賴超越了CNN局部感受野的限制。踩坑實(shí)錄注意力權(quán)重的可視化與解釋。我們常常想通過可視化注意力權(quán)重來理解模型“在看哪里”。但這需要謹(jǐn)慎Softmax的競(jìng)爭(zhēng)性Softmax使得權(quán)重是相對(duì)的。一個(gè)位置權(quán)重高不一定是因?yàn)樗^對(duì)重要可能只是因?yàn)槠渌恢酶恢匾?。特別是在長(zhǎng)序列中權(quán)重分布可能非常均勻。多頭注意力的分散不同頭的注意力模式可能差異很大簡(jiǎn)單平均可能沒有意義。需要分別檢查每個(gè)頭。不能直接等價(jià)于重要性高注意力權(quán)重表明該位置的信息被大量用于計(jì)算當(dāng)前輸出但這不一定是“因果性”的重要。有時(shí)模型可能通過注意力機(jī)制“忽略”某些噪聲給低權(quán)重這也是一種重要的能力。 我的經(jīng)驗(yàn)是將注意力權(quán)重作為理解模型內(nèi)部工作機(jī)理的一種輔助工具而不是“金標(biāo)準(zhǔn)”。結(jié)合梯度類方法如Grad-CAM或擾動(dòng)測(cè)試能獲得更可靠的解釋。5. 超越基礎(chǔ)注意力機(jī)制的變體與優(yōu)化基礎(chǔ)的縮放點(diǎn)積注意力雖然強(qiáng)大但在處理長(zhǎng)序列時(shí)面臨O(n^2)計(jì)算和內(nèi)存復(fù)雜度的瓶頸因?yàn)樾枰?jì)算所有查詢-鍵對(duì)。為此研究者提出了多種高效注意力變體。5.1 局部注意力與稀疏注意力思想并非所有查詢都需要和所有鍵交互。強(qiáng)制每個(gè)查詢只關(guān)注一個(gè)局部窗口如前后w個(gè)位置或一種預(yù)定義的稀疏模式如固定步長(zhǎng)、塊狀模式。局部注意力類似CNN的局部感受野計(jì)算復(fù)雜度降至O(n*w)。在圖像或某些具有強(qiáng)局部相關(guān)性的序列上很有效。稀疏Transformer設(shè)計(jì)固定的稀疏注意力模式例如Stride模式關(guān)注固定間隔的位置、Fixed模式關(guān)注某些固定位置。這需要先驗(yàn)知識(shí)。軸向注意力在多維數(shù)據(jù)如圖像中沿高度和寬度兩個(gè)軸分別進(jìn)行注意力計(jì)算將O(h^2 * w^2)復(fù)雜度降為O(h^2 w^2)。5.2 線性化注意力核心思路通過數(shù)學(xué)變換將計(jì)算注意力權(quán)重的順序進(jìn)行交換從而避免計(jì)算顯式的n x n注意力矩陣。一個(gè)著名的代表是Linformer和Linear Transformer。它們的基本思想是將標(biāo)準(zhǔn)的Softmax注意力公式Attention(Q, K, V) softmax(QK^T/sqrt(d)) V進(jìn)行重寫。通過使用核函數(shù)近似或低秩投影將K和V投影到低維空間使得QK^T的計(jì)算不再需要顯式的n x n矩陣。例如Linear Transformer使用elu(x)1作為核函數(shù)使得注意力可以寫成(Q * (K^T V))的形式從而實(shí)現(xiàn)線性復(fù)雜度。這類方法在長(zhǎng)序列推理中能極大節(jié)省內(nèi)存和時(shí)間。5.3 內(nèi)存壓縮與分塊計(jì)算內(nèi)存高效的注意力如FlashAttention通過精妙的IO感知算法在GPU顯存層次結(jié)構(gòu)HBM - SRAM中重新組織計(jì)算順序避免存儲(chǔ)龐大的中間注意力矩陣從而在幾乎不改變算法的情況下大幅降低內(nèi)存占用并提升速度。分塊注意力將長(zhǎng)序列分成塊在塊內(nèi)進(jìn)行精確注意力計(jì)算在塊間使用一種簡(jiǎn)化的注意力機(jī)制如平均池化后的全局向量。這是一種工程上的折中方案。技術(shù)選型思考如何選擇注意力變體這取決于你的具體任務(wù)和資源約束任務(wù)特性如果你的數(shù)據(jù)具有強(qiáng)烈的局部性如圖像、音頻局部注意力或軸向注意力是很好的起點(diǎn)。如果需要完全的全局交互如某些文檔級(jí)NLP任務(wù)則需考慮線性注意力或內(nèi)存優(yōu)化方法。序列長(zhǎng)度這是決定性因素。對(duì)于n512的序列標(biāo)準(zhǔn)注意力通??梢猿惺堋?duì)于n1024就必須考慮高效注意力變體。硬件資源如果GPU內(nèi)存有限FlashAttention是必選項(xiàng)。它現(xiàn)在已被集成進(jìn)主流的深度學(xué)習(xí)框架如PyTorch 2.0的scaled_dot_product_attention高效實(shí)現(xiàn)。精度要求有些線性化或稀疏化方法會(huì)引入近似誤差。在關(guān)鍵任務(wù)上需要通過實(shí)驗(yàn)驗(yàn)證其對(duì)最終性能的影響。 我的建議是優(yōu)先使用經(jīng)過充分優(yōu)化的標(biāo)準(zhǔn)注意力實(shí)現(xiàn)如PyTorch的F.scaled_dot_product_attention它內(nèi)部可能已經(jīng)集成了FlashAttention等優(yōu)化。只有當(dāng)序列長(zhǎng)度成為明確瓶頸時(shí)再著手研究和引入特定的高效注意力變體。6. 實(shí)戰(zhàn)構(gòu)建一個(gè)用于回歸任務(wù)的注意力層理論說了這么多我們來動(dòng)手實(shí)現(xiàn)一個(gè)可以嵌入全連接網(wǎng)絡(luò)的、最簡(jiǎn)單的注意力池化層并用于一個(gè)簡(jiǎn)單的回歸任務(wù)。我們將實(shí)現(xiàn)一個(gè)“通用的”注意力池化層它接受一組鍵值對(duì)和一個(gè)查詢輸出加權(quán)后的值。然后將其用于擬合一個(gè)一維的非線性函數(shù)。import torch import torch.nn as nn import torch.nn.functional as F import numpy as np import matplotlib.pyplot as plt class SimpleAttentionPooling(nn.Module): 一個(gè)簡(jiǎn)單的注意力池化層使用縮放點(diǎn)積注意力。 def __init__(self, d_k, d_v): super().__init__() # 通常我們會(huì)有關(guān)聯(lián)的線性層來投影Q, K, V。這里為了簡(jiǎn)單假設(shè)輸入已經(jīng)投影好。 # 或者我們內(nèi)置投影層。這里我們選擇內(nèi)置更通用。 self.d_k d_k # 注意這個(gè)例子中我們讓查詢、鍵、值的維度可以不同但計(jì)算注意力時(shí)q和k維度需相同。 # 我們假設(shè)輸入是原始特征用線性層投影到指定維度。 self.W_q nn.Linear(d_k, d_k, biasFalse) # 查詢投影 self.W_k nn.Linear(d_k, d_k, biasFalse) # 鍵投影 self.W_v nn.Linear(d_v, d_v, biasFalse) # 值投影 def forward(self, queries, keys, values): Args: queries: (batch_size, num_queries, d_k) keys: (batch_size, num_keys, d_k) values: (batch_size, num_keys, d_v) Returns: output: (batch_size, num_queries, d_v) attn_weights: (batch_size, num_queries, num_keys) Q self.W_q(queries) # (B, Nq, d_k) K self.W_k(keys) # (B, Nk, d_k) V self.W_v(values) # (B, Nk, d_v) # 計(jì)算縮放點(diǎn)積注意力 scores torch.matmul(Q, K.transpose(-2, -1)) / (self.d_k ** 0.5) # (B, Nq, Nk) attn_weights F.softmax(scores, dim-1) # (B, Nq, Nk) output torch.matmul(attn_weights, V) # (B, Nq, d_v) return output, attn_weights # 構(gòu)建一個(gè)使用注意力池化的簡(jiǎn)單回歸模型 class AttentionRegressionModel(nn.Module): def __init__(self, input_dim1, hidden_dim64, output_dim1): super().__init__() # 我們將整個(gè)訓(xùn)練集視為“記憶”鍵值對(duì)。 # 但實(shí)際上我們需要?jiǎng)討B(tài)處理。這里我們用一個(gè)網(wǎng)絡(luò)來生成“記憶”的鍵和值。 self.memory_net nn.Sequential( nn.Linear(input_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, hidden_dim), ) # 注意力池化層d_k和d_v都設(shè)為hidden_dim self.attention_pool SimpleAttentionPooling(d_khidden_dim, d_vhidden_dim) # 輸出層從池化后的特征預(yù)測(cè)輸出 self.output_layer nn.Linear(hidden_dim, output_dim) def forward(self, x_query, x_memory, y_memory): x_query: 要預(yù)測(cè)的查詢點(diǎn) (B, Nq, 1) x_memory: 作為記憶的訓(xùn)練點(diǎn)坐標(biāo) (B, Nm, 1) y_memory: 作為記憶的訓(xùn)練點(diǎn)標(biāo)簽 (B, Nm, 1) 注意這是一個(gè)“非參數(shù)”風(fēng)格的使用記憶數(shù)據(jù)作為輸入的一部分。 更常見的參數(shù)化方式是將記憶編碼到模型參數(shù)中這里僅為演示注意力機(jī)制。 # 1. 將記憶的x編碼為鍵和值 # 鍵由x_memory編碼得到 K self.memory_net(x_memory) # (B, Nm, hidden_dim) # 值我們想讓值包含y的信息。一種簡(jiǎn)單做法是將y_memory與編碼后的特征結(jié)合。 # 這里為了極端簡(jiǎn)化我們直接用y_memory作為值的一部分或者也通過一個(gè)網(wǎng)絡(luò)。 # 更合理的做法值也應(yīng)該是一個(gè)學(xué)習(xí)到的表示。這里我們偷懶用K作為值即鍵值相同。 V K # (B, Nm, hidden_dim) # 2. 將查詢點(diǎn)x_query編碼為查詢向量 Q self.memory_net(x_query) # (B, Nq, hidden_dim) # 3. 注意力池化 context, attn_weights self.attention_pool(Q, K, V) # context: (B, Nq, hidden_dim) # 4. 輸出預(yù)測(cè) y_pred self.output_layer(context) # (B, Nq, 1) return y_pred, attn_weights # 生成模擬數(shù)據(jù) def generate_data(num_samples100): x np.linspace(-3, 3, num_samples) y np.sin(x) * np.exp(-0.1 * x**2) 0.1 * np.random.randn(num_samples) # 一個(gè)衰減振蕩信號(hào) return torch.FloatTensor(x).view(-1, 1), torch.FloatTensor(y).view(-1, 1) # 訓(xùn)練和評(píng)估 def train_and_evaluate(): # 數(shù)據(jù) x_all, y_all generate_data(200) # 劃分“記憶”集訓(xùn)練集和查詢集測(cè)試集 indices np.random.permutation(len(x_all)) train_idx, test_idx indices[:150], indices[150:] x_train, y_train x_all[train_idx], y_all[train_idx] x_test, y_test x_all[test_idx], y_all[test_idx] model AttentionRegressionModel(input_dim1, hidden_dim32, output_dim1) optimizer torch.optim.Adam(model.parameters(), lr0.01) criterion nn.MSELoss() epochs 500 batch_size 32 num_train len(x_train) for epoch in range(epochs): model.train() perm torch.randperm(num_train) total_loss 0 for i in range(0, num_train, batch_size): idx perm[i:ibatch_size] batch_x_mem x_train[idx].unsqueeze(0) # (1, B, 1) - 這里簡(jiǎn)化假設(shè)batch內(nèi)記憶相同 batch_y_mem y_train[idx].unsqueeze(0) # 在這個(gè)batch中我們用記憶數(shù)據(jù)來預(yù)測(cè)記憶數(shù)據(jù)本身自回歸這只是一個(gè)演示。 # 更合理的設(shè)置是從記憶集中采樣一部分作為支持集另一部分作為查詢。 # 這里我們簡(jiǎn)單地將batch內(nèi)的點(diǎn)既作記憶又作查詢。 y_pred, _ model(batch_x_mem, batch_x_mem, batch_y_mem) loss criterion(y_pred.squeeze(0), batch_y_mem.squeeze(0)) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() if (epoch1) % 100 0: print(fEpoch [{epoch1}/{epochs}], Loss: {total_loss/(num_train//batch_size):.4f}) # 評(píng)估在測(cè)試集上用全部訓(xùn)練集作為記憶 model.eval() with torch.no_grad(): # 將全部訓(xùn)練集作為記憶 x_mem x_train.unsqueeze(0) # (1, N_train, 1) y_mem y_train.unsqueeze(0) # 預(yù)測(cè)測(cè)試集 y_pred_test, attn_weights model(x_test.unsqueeze(0), x_mem, y_mem) test_loss criterion(y_pred_test.squeeze(0), y_test) print(fTest MSE: {test_loss.item():.4f}) # 可視化 plt.figure(figsize(12, 5)) plt.subplot(1, 2, 1) plt.scatter(x_train.numpy(), y_train.numpy(), alpha0.6, labelTrain (Memory)) plt.scatter(x_test.numpy(), y_test.numpy(), alpha0.6, labelTest (Query)) # 生成平滑曲線用于繪制預(yù)測(cè) x_plot torch.linspace(-3, 3, 300).view(-1, 1) y_plot_pred, _ model(x_plot.unsqueeze(0), x_mem, y_mem) plt.plot(x_plot.numpy(), y_plot_pred.squeeze().numpy(), r-, linewidth2, labelModel Prediction) plt.legend() plt.title(Regression Fit with Attention) # 可視化某個(gè)測(cè)試查詢點(diǎn)的注意力權(quán)重 plt.subplot(1, 2, 2) query_idx 25 # 選擇一個(gè)測(cè)試點(diǎn) sample_attn attn_weights[0, query_idx].cpu().numpy() # (N_train,) plt.bar(x_train.squeeze().numpy(), sample_attn, alpha0.7, width0.05) plt.axvline(xx_test[query_idx].item(), colorr, linestyle--, labelfQuery x{x_test[query_idx].item():.2f}) plt.xlabel(Memory x) plt.ylabel(Attention Weight) plt.title(fAttention Weights for a Test Query) plt.legend() plt.tight_layout() plt.show() if __name__ __main__: train_and_evaluate()這個(gè)例子雖然簡(jiǎn)單但完整展示了如何將注意力池化作為一個(gè)可微分的神經(jīng)網(wǎng)絡(luò)層來構(gòu)建和使用。模型通過學(xué)習(xí)能夠?yàn)槊總€(gè)查詢點(diǎn)動(dòng)態(tài)地從“記憶”訓(xùn)練集中檢索并聚合信息。可視化注意力權(quán)重可以看到對(duì)于某個(gè)查詢點(diǎn)x模型確實(shí)會(huì)給附近的x_train點(diǎn)分配更高的權(quán)重這與核回歸的直覺一致但這里的相似度度量通過memory_net學(xué)習(xí)比預(yù)設(shè)的高斯核更加靈活。注意事項(xiàng)與擴(kuò)展記憶集的處理上面的例子中記憶集是作為模型輸入動(dòng)態(tài)傳入的這更像是一種“非參數(shù)”或“基于記憶”的學(xué)習(xí)方式。更常見的參數(shù)化方式是將知識(shí)固化在網(wǎng)絡(luò)的權(quán)重中注意力用于處理序列輸入本身。計(jì)算效率在實(shí)際應(yīng)用中如果記憶集很大這種每次計(jì)算所有查詢-鍵對(duì)的方式開銷巨大。這就需要用到我們前面提到的高效注意力機(jī)制。鍵與值的分離本例中為了簡(jiǎn)化令VK。在更復(fù)雜的任務(wù)中V應(yīng)該獨(dú)立學(xué)習(xí)以承載與K不同的信息。位置信息對(duì)于序列數(shù)據(jù)輸入本身沒有順序信息。需要額外加入位置編碼如正弦余弦編碼、可學(xué)習(xí)編碼來讓注意力機(jī)制感知位置。這是Transformer成功的關(guān)鍵之一。注意力機(jī)制從核回歸的統(tǒng)計(jì)思想出發(fā)通過神經(jīng)網(wǎng)絡(luò)的參數(shù)化改造已成為深度學(xué)習(xí)中最核心的構(gòu)件之一。理解其從“提示”到“池化”再到“架構(gòu)”的演進(jìn)脈絡(luò)不僅能幫助我們?cè)趯?shí)踐中更好地應(yīng)用它例如選擇合適的變體、調(diào)試注意力權(quán)重更能讓我們洞察其本質(zhì)——一種動(dòng)態(tài)的、數(shù)據(jù)驅(qū)動(dòng)的資源分配策略。無論是處理自然語(yǔ)言、圖像還是其他序列化數(shù)據(jù)當(dāng)你希望模型學(xué)會(huì)“有選擇地聚焦”時(shí)注意力機(jī)制幾乎總是你的第一選擇。