制原理與PyTorch實(shí)現(xiàn)詳解)
自注意力機(jī)制本身并不復(fù)雜核心思想就是一句話讓序列里的每個(gè) token 都能和序列里的其他 token 做信息交互。但真正讓 Transformer 在 NLP、CV、多模態(tài)任務(wù)里全面站穩(wěn)腳跟的是那個(gè)看起來只加了一個(gè)字的模塊——多頭注意力機(jī)制Multi-Head Attention。它解決的是單頭注意力“只能有一組注意力分布”的表達(dá)瓶頸而方案不是增加復(fù)雜度而是把 Q、K、V 投影到多個(gè)低維子空間并行計(jì)算。這篇文章把多頭注意力的動(dòng)機(jī)、數(shù)學(xué)原理、PyTorch 實(shí)現(xiàn)、因果掩碼、MQA/GQA 變體以及它與殘差、層歸一化、FFN 的配合方式全部拆開講一遍。如果你正在學(xué) Transformer準(zhǔn)備從零復(fù)現(xiàn) BERT 或 GPT 系列模型如果你讀懂了論文里的公式但一到了代碼里就被view、transpose、contiguous繞暈或者你只是想知道為什么大模型推理時(shí)都在強(qiáng)調(diào) KV Cache而 GQA 能加速那么多——這篇文章就是給你準(zhǔn)備的。讀完你不僅能手寫一個(gè)可運(yùn)行的多頭注意力模塊還能說清楚它為什么有效、有哪些容易踩的坑。1. 多頭注意力機(jī)制核心速覽先把關(guān)鍵信息放在最前面。這一節(jié)不做推導(dǎo)只給結(jié)論。下面的表格基本覆蓋了多頭注意力機(jī)制的理解坐標(biāo)。項(xiàng)目說明機(jī)制名稱多頭注意力機(jī)制Multi-Head Attention, MHA提出論文Attention Is All You NeedTransformer 原論文解決的核心問題單頭自注意力表達(dá)能力有限難以同時(shí)建模多種依賴關(guān)系核心思路將 Q/K/V 投影到多個(gè)低維子空間并行做注意力計(jì)算再把結(jié)果拼接起來是否增加參數(shù)量標(biāo)準(zhǔn)實(shí)現(xiàn)下不增加Q/K/V 總參數(shù)量與單頭版本一致典型配置d_model512num_heads8每個(gè)頭維度 d_kd_v64前置知識(shí)縮放點(diǎn)積注意力、Softmax、線性投影、矩陣維度變換主要應(yīng)用Transformer、BERT、GPT、ViT、多模態(tài)模型等絕大多數(shù)現(xiàn)代架構(gòu)常見變體MHA、MQA、GQA以及 FlashAttention 等高頻實(shí)現(xiàn)學(xué)習(xí)門檻中等需要一定矩陣基礎(chǔ)但代碼復(fù)現(xiàn)并不難判斷自己是否真正理解多頭注意力可以拿下面三個(gè)問題自測當(dāng)d_model512、num_heads8時(shí)每個(gè)頭的維度是多少為什么是這個(gè)數(shù)多頭注意力的總參數(shù)量為什么和單頭注意力相同它到底“多”在哪里訓(xùn)練 GPT 這類自回歸模型時(shí)為什么在注意力分?jǐn)?shù)上要加一個(gè)上三角掩碼這三個(gè)問題如果在讀完后都能回答說明這章就真正通了。2. 為什么需要“多頭”單頭自注意力的局限單頭自注意力的局限主要體現(xiàn)在三個(gè)層面。第一表示能力單一。自注意力輸出是Attention(Q, K, V)它本質(zhì)上是在一組 Softmax 權(quán)重下對(duì)所有 Value 向量做加權(quán)求和。一個(gè)注意力頭只能輸出一種加權(quán)方式的結(jié)果。但語言中一個(gè)詞可能需要同時(shí)建模多種關(guān)系比如“蘋果”這個(gè)詞既和“紅色”有顏色關(guān)系又和“水果”有類別關(guān)系還和“喬布斯”有品牌關(guān)系。單頭注意力只能把這些關(guān)系全部揉在一起最終得到一個(gè)平均化的上下文表示。第二Softmax 存在“平均化”傾向。當(dāng)序列長度變長時(shí)注意力分?jǐn)?shù)經(jīng)過 Softmax 后很容易變得平緩尤其是每個(gè) token 的表示都比較接近的時(shí)候。這時(shí)候注意力頭實(shí)際上退化成了一種近似平均池化操作沒有真正突出某一個(gè)位置。解決思路有兩個(gè)方向一是降低溫度增大注意力分布的尖銳程度二是讓模型同時(shí)嘗試多組不同的注意力分布總有一組能學(xué)到關(guān)鍵依賴。第三優(yōu)化的計(jì)算路徑太單一。單頭自注意力從一個(gè)全量矩陣運(yùn)算中學(xué)習(xí)依賴關(guān)系一組 W_Q、W_K、W_V 只能覆蓋一種語義空間。模型把所有的語法、語義、指代、位置信息全部塞進(jìn)同一個(gè)低維投影里梯度更新時(shí)這些信息會(huì)互相干擾。多頭注意力解決問題的思路很直接既然一個(gè)頭不夠那就并行跑多個(gè)頭。每個(gè)頭使用獨(dú)立的投影矩陣把輸入映射到不同的子空間學(xué)習(xí)不同類型的依賴關(guān)系。最后把多個(gè)頭的輸出拼接起來再經(jīng)過一次線性投影融合成完整的表示。這樣做既保留了注意力的全局交互能力又增加了模型的表達(dá)自由度而且參數(shù)總量不漲。3. 多頭注意力機(jī)制原理拆解多頭注意力機(jī)制的輸入是一個(gè)序列表示矩陣X形狀為(batch_size, seq_len, d_model)。整個(gè)計(jì)算過程分四步。3.1 生成 Q、K、V 投影輸入通過三個(gè)可學(xué)習(xí)矩陣 W_Q、W_K、W_V 得到查詢、鍵、值矩陣。在標(biāo)準(zhǔn)實(shí)現(xiàn)中這三個(gè)矩陣的維度都是(d_model, d_model)Q XW_Q, K XW_K, V XW_V3.2 按頭拆分把 d_model 維度平均切分成 h 份每份維度 d_k d_model / h。拆分在代碼中常見做法是先經(jīng)過 Linear(d_model, d_model) 得到形狀 (batch, seq_len, d_model)再通過 view 和 transpose 重排為 (batch, h, seq_len, d_k)。這一操作等價(jià)于把一個(gè)大矩陣切成了 h 個(gè)子矩陣每個(gè)子矩陣代表一個(gè)子空間中的投影。3.3 縮放點(diǎn)積注意力每個(gè)頭獨(dú)立計(jì)算注意力分?jǐn)?shù)??s放點(diǎn)積注意力的標(biāo)準(zhǔn)公式為$$ \text{Attention}(Q, K, V) \text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}\right)V $$其中除以 sqrt(d_k) 是關(guān)鍵細(xì)節(jié)。當(dāng) d_k 較大時(shí)QK^T 的結(jié)果會(huì)有較大方差Softmax 的梯度會(huì)變得非常小訓(xùn)練不穩(wěn)定。除以 sqrt(d_k) 是為了把分?jǐn)?shù)拉回到合理的數(shù)值區(qū)間。3.4 拼接并輸出投影將 h 個(gè)頭的輸出在最后一個(gè)維度上拼接得到維度為 d_model 的向量再經(jīng)過輸出矩陣 W_O 完成一次線性變換$$ \text{MultiHead}(X) \text{Concat}(\text{head}_1, \ldots, \text{head}_h)W^O $$其中$$ \text{head}_i \text{Attention}(XW_i^Q, XW_i^K, XW_i^V) $$這里 W_i^Q、W_i^K、W_i^V 的維度是 (d_model, d_k)。從參數(shù)總量看h 個(gè)頭總共 h * (3 * d_model * d_k) 3 * d_model^2 個(gè)參數(shù)和單頭版本完全一致。區(qū)別在于單頭是一個(gè)大的線性投影多頭是把這個(gè)投影拆成了 h 份不同的子空間并通過輸出投影重新融合。4. 為什么多頭有效4 個(gè)關(guān)鍵原因多頭注意力之所以成為 Transformer 最核心的組件不是因?yàn)樗奥犉饋韽?fù)雜”而是因?yàn)樗谒膫€(gè)方面都有明確作用。4.1 子空間并行多頭各司其職Transformer 原論文通過在機(jī)器翻譯模型上的可視化實(shí)驗(yàn)觀察到不同的注意力頭確實(shí)在學(xué)習(xí)不同類型的關(guān)系有的頭關(guān)注句法依賴比如動(dòng)詞和主語有的頭關(guān)注指代關(guān)系比如代詞和先行詞有的頭關(guān)注相鄰位置的局部特征還有的頭關(guān)注長距離的跨段依賴。多頭機(jī)制本質(zhì)上是讓模型擁有 h 次機(jī)會(huì)去學(xué)習(xí)不同的注意力模式而不是強(qiáng)迫一個(gè)頭把所有關(guān)系都學(xué)會(huì)。4.2 打破單頭 Softmax 的“平均化”多頭機(jī)制相當(dāng)于把原來的一個(gè) Softmax 分布變成了 h 個(gè)獨(dú)立的 Softmax 分布。每個(gè)頭只需要在自己的子空間里找到最重要的位置不需要承擔(dān)所有信息的加權(quán)責(zé)任。即使某一個(gè)頭出現(xiàn)退化成“平均池化”的情況其他頭仍然可以保持尖銳的注意力分布。多個(gè)頭的組合讓模型更穩(wěn)定。4.3 參數(shù)效率極高很多第一次接觸多頭注意力的人會(huì)誤以為“多頭”意味著 h 倍參數(shù)量。事實(shí)并非如此。多頭通過拆分 d_model 維度來降低每個(gè)頭的維度總參數(shù)量和單頭完全一致。它增加的是“表征的自由度”而不是“參數(shù)的數(shù)量”。這也是為什么 Transformer 能在參數(shù)量不變的情況下獲得更高的模型容量。4.4 訓(xùn)練更穩(wěn)定梯度更平滑單頭注意力的輸出是一個(gè)大矩陣直接參與最后的加權(quán)求和所有信息集中在同一個(gè)路徑上。多頭輸出經(jīng)過拼接和線性投影后梯度可以通過多個(gè)分支回傳到不同的子空間避免了單個(gè)注意力頭的梯度主導(dǎo)整個(gè)模型更新的問題。多個(gè)頭還可以配合 Dropout 機(jī)制使用不同頭隨機(jī)丟棄部分注意力權(quán)重相當(dāng)于在注意力層面做了集成學(xué)習(xí)。5. 環(huán)境準(zhǔn)備與代碼復(fù)現(xiàn)理解公式之后必須動(dòng)手寫代碼。這里給出一套完全可運(yùn)行的 PyTorch 實(shí)現(xiàn)不需要 GPUCPU 環(huán)境即可驗(yàn)證維度邏輯。如果你本機(jī)已經(jīng)有 PyTorch直接跳過安裝步驟。首先確認(rèn) Python 版本建議 Python 3.8 以上然后安裝 PyTorch。pip install torch安裝完成后檢查是否可以正常導(dǎo)入。import torch import torch.nn as nn import torch.nn.functional as F import math print(torch.__version__)本文代碼的位置是在一個(gè)自包含的 Python 腳本里運(yùn)行不依賴額外項(xiàng)目結(jié)構(gòu)。建議把下面的代碼保存為multi_head_attention.py后續(xù)修改參數(shù)方便調(diào)試。6. 手寫多頭注意力模塊PyTorch這里給出一個(gè)最典型的實(shí)現(xiàn)方式。它嚴(yán)格按照上文公式展開重點(diǎn)在于理解 Q/K/V 的維度變換和 mask 的傳遞邏輯。import torch import torch.nn as nn import torch.nn.functional as F class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads, dropout0.1): super().__init__() assert d_model % num_heads 0, d_model must be divisible by num_heads self.d_model d_model self.num_heads num_heads self.d_k d_model // num_heads self.W_Q nn.Linear(d_model, d_model) self.W_K nn.Linear(d_model, d_model) self.W_V nn.Linear(d_model, d_model) self.W_O nn.Linear(d_model, d_model) self.dropout nn.Dropout(dropout) def forward(self, x, maskNone): batch_size, seq_len, _ x.size() # 1. 生成 Q、K、V并拆分成多頭形狀 # 目標(biāo)形狀: (batch, num_heads, seq_len, d_k) Q self.W_Q(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) K self.W_K(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) V self.W_V(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) # 2. 計(jì)算縮放點(diǎn)積注意力分?jǐn)?shù) # scores shape: (batch, num_heads, seq_len, seq_len) scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k) # 3. 如果傳入 mask則 mask 為 0 的位置置為負(fù)無窮 if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) # 4. Softmax 得到注意力權(quán)重再作用于 V attn F.softmax(scores, dim-1) attn self.dropout(attn) context torch.matmul(attn, V) # shape: (batch, num_heads, seq_len, d_k) # 5. 拼接所有頭恢復(fù) d_model 維度 context context.transpose(1, 2).contiguous().view(batch_size, seq_len, self.d_model) # 6. 輸出投影 output self.W_O(context) return output接下來做一個(gè)小測試驗(yàn)證輸出形狀是否正確。x torch.randn(2, 10, 512) # batch2, seq_len10, d_model512 mha MultiHeadAttention(d_model512, num_heads8) y mha(x) print(input shape:, x.shape) print(output shape:, y.shape)預(yù)期輸出input shape: torch.Size([2, 10, 512]) output shape: torch.Size([2, 10, 512])輸出形狀和輸入形狀完全一致這符合 Transformer 中殘差連接的使用前提。這里最容易踩坑的是view和transpose的配合。view是把最后兩個(gè)維度重組為(num_heads, d_k)transpose(1, 2)再把num_heads維度提前最終得到(batch, num_heads, seq_len, d_k)。這里的順序一旦寫錯(cuò)后續(xù)矩陣乘法的形狀就會(huì)全部錯(cuò)位。如果你熟悉 einsum也可以寫一個(gè)更緊湊的等價(jià)版本scores torch.einsum(bqhd,bkhd-bhqk, Q, K) / math.sqrt(self.d_k) context torch.einsum(bhqk,bkhd-bqhd, attn, V)兩種寫法計(jì)算邏輯完全一致。einsum 可讀性稍差但不容易出現(xiàn)維度順序錯(cuò)誤。7. 因果自注意力與 Mask 實(shí)現(xiàn)在 GPT 等自回歸模型里多頭注意力不能直接使用普通版本必須加一個(gè)因果掩碼causal mask所以這部分單獨(dú)拿出來講。在很多開源代碼中你會(huì)看到它被寫作 Causal Self-Attention也就是“因果自注意力”。因果掩碼的核心邏輯生成任務(wù)中token 在位置 t 只能看到位置 t 的 token不能看到未來的 token。否則模型在訓(xùn)練時(shí)“偷看”了未來信息推理時(shí)就沒有對(duì)應(yīng)的未來 token造成訓(xùn)練和推理不一致。掩碼的計(jì)算非常簡單。先用torch.tril生成一個(gè)下三角矩陣再將掩碼應(yīng)用到注意力分?jǐn)?shù)矩陣上。在 PyTorch 中實(shí)現(xiàn)如下def subsequent_mask(seq_len): mask torch.tril(torch.ones(seq_len, seq_len)).bool() return mask # shape: (seq_len, seq_len)測試一下mask subsequent_mask(5) print(mask)輸出tensor([[ True, False, False, False, False], [ True, True, False, False, False], [ True, True, True, False, False], [ True, True, True, True, False], [ True, True, True, True, True]])在多頭注意力 forward 中調(diào)用時(shí)mask 需要擴(kuò)展為和 scores 相同的維度也就是(batch, num_heads, seq_len, seq_len)seq_len x.size(1) mask subsequent_mask(seq_len).unsqueeze(0).unsqueeze(0) # (1, 1, seq_len, seq_len) mask mask.expand(x.size(0), mha.num_heads, -1, -1) # (batch, num_heads, seq_len, seq_len) output mha(x, maskmask)在 forward 內(nèi)部mask 為 False 的位置會(huì)被masked_fill替換為負(fù)無窮經(jīng)過 Softmax 后這些位置的權(quán)重趨近于 0。這里有一個(gè)實(shí)現(xiàn)上的關(guān)鍵點(diǎn)mask 必須在 softmax 之前加而不是在 softmax 之后把權(quán)重置零。如果 softmax 之后直接置零所有權(quán)重之和不再等于 1會(huì)破壞概率分布的語義而 softmax 之前加負(fù)無窮是標(biāo)準(zhǔn)做法。除了因果掩碼實(shí)際工程中還會(huì)用到 padding mask目的是讓注意力忽略掉 padding token。對(duì)于自回歸模型通常需要同時(shí)使用 padding mask 和 causal mask兩者取交集。實(shí)現(xiàn)上可以通過torch.logical_and將兩個(gè)掩碼合并成一個(gè)布爾矩陣再一次性傳給注意力模塊。8. 多頭注意力變體對(duì)比MHA / MQA / GQA隨著大模型推理部署的發(fā)展多頭注意力機(jī)制出現(xiàn)了幾個(gè)重要變體。理解這些變體能幫你理解為什么新一代大模型都在提“減少 KV Cache”。變體全稱核心思想?yún)?shù)量推理效率代表應(yīng)用情況MHAMulti-Head Attention每個(gè)頭都有自己的 K、V高一般Transformer、BERT、早期 GPTMQAMulti-Query Attention所有頭共享一組 K、V只有 Q 獨(dú)立低快部分早期大模型GQAGrouped-Query Attention若干個(gè)頭共享一組 K、V中較快Llama 2、Llama 3、Mistral 等在自回歸生成時(shí)模型每生成一個(gè)新 token 都需要用到之前所有 token 的 K、V 向量。如果不做緩存每一步都重新計(jì)算代價(jià)太高。因此引擎會(huì)將歷史 K、V 緩存到顯存中這部分緩存就是 KV Cache。MHA 因?yàn)槊總€(gè)頭都需要各自緩存 K、V顯存占用最高。MQA 讓所有頭共享一組 K、V緩存顯著減少但會(huì)犧牲一部分模型表達(dá)能力。GQA 是折中方案把頭分成若干組每組內(nèi)部共享 K、V。它既減少了緩存量又保留了一定程度的表達(dá)多樣性。這也是為什么 Llama 2 之后的很多開源模型都選擇 GQA。如果你自己在實(shí)現(xiàn) Transformer 推理可以從 MHA 開始跑通后再優(yōu)化為 GQA。優(yōu)化的第一步是理解“哪些 weight 需要緩存”Q 每次都重新生成不需要緩存K、V 需要跨 step 保留并拼接。9. 多頭注意力與殘差、層歸一化、FFN 的配合多頭注意力并不是單獨(dú)工作的。在 Transformer 中它總是和殘差連接、層歸一化、前饋網(wǎng)絡(luò)組合成一個(gè)完整的 Transformer Block。一個(gè)標(biāo)準(zhǔn) Transformer Block 的計(jì)算過程如下x x MultiHeadAttention(LayerNorm(x)) x x FeedForward(LayerNorm(x))其中 FeedForward 通常是一個(gè)兩層的多層感知機(jī)MLP先升維再降維中間用 ReLU 或 GELU 激活。這個(gè)結(jié)構(gòu)的兩個(gè)關(guān)鍵點(diǎn)第一多頭注意力輸出經(jīng)過殘差和 LayerNorm 后數(shù)值分布會(huì)更穩(wěn)定。多頭注意力內(nèi)部的矩陣乘法和 Softmax 操作會(huì)讓數(shù)值范圍波動(dòng)很大直接堆疊多層會(huì)出現(xiàn)訓(xùn)練不穩(wěn)定的情況。層歸一化LayerNorm在每個(gè) token 維度上做歸一化平均值拉到 0、方差拉到 1有效緩解梯度爆炸或消失。第二多頭注意力本質(zhì)上是線性投影和加權(quán)求和單靠它無法引入非線性。FFN 中的非線性激活函數(shù)承擔(dān)了這部分工作。多頭注意力負(fù)責(zé)在不同 token 之間交互信息FFN 負(fù)責(zé)在每個(gè) token 內(nèi)部做更高維的特征變換兩者分工明確。Pre-LN 和 Post-LN 是實(shí)現(xiàn)上的一個(gè)重要差別。上面給出的寫法是 Pre-LN先 LayerNorm 再進(jìn)注意力。GPT 系列模型普遍使用 Pre-LN因?yàn)樗梢宰屔顚泳W(wǎng)絡(luò)訓(xùn)練更穩(wěn)定。原始 Transformer 論文中的結(jié)構(gòu)更接近 Post-LN先注意力再 LayerNorm。理解這個(gè)差別有助于閱讀不同開源模型的源碼。10. 常見問題與排查方法實(shí)際寫代碼時(shí)最容易出的問題集中在維度變換和 mask 邏輯上。下面整理了一份排查清單。現(xiàn)象可能原因檢查方式解決方案矩陣乘法維度對(duì)不上d_model 無法被 num_heads 整除打印 Q、K、V 的 shape調(diào)整 num_heads或修改 d_model輸出 shape 與輸入不一致view和transpose順序?qū)懛丛?forward 中逐步打印 shape按view - transpose順序重新組織訓(xùn)練損失不下降或速度太慢忘記除以 sqrt(d_k)檢查 scores 計(jì)算代碼加上math.sqrt(self.d_k)mask 沒有生效mask 維度與 scores 不一致打印 mask 和 scores 的 shape將 mask 擴(kuò)展到 (batch, heads, seq_len, seq_len)Softmax 后取 mask 置零對(duì)概率直接置零分布不再歸一檢查 mask 是在 softmax 前還是后在 softmax 前通過負(fù)無窮屏蔽單頭輸出正常多頭后結(jié)果異常拼接后沒有調(diào)用 contiguous檢查.view前的報(bào)錯(cuò)拼接前調(diào)用.contiguous()長序列顯存溢出注意力分?jǐn)?shù)矩陣為 O(n^2)監(jiān)控顯存和序列長度使用 FlashAttention、稀疏注意力或梯度檢查點(diǎn)推理延遲過高M(jìn)HA 的 KV Cache 占用過高觀察顯存占用與緩存大小切換到 GQA 或 MQA其中contiguous的問題尤其隱蔽。transpose操作不會(huì)讓內(nèi)存連續(xù)此時(shí)直接調(diào)用view會(huì)報(bào)錯(cuò)。代碼中先transpose再contiguous再view這是正確順序。如果你在實(shí)現(xiàn)中遇到view size is not compatible with input tensors size大概率就是這里出了問題。11. 最佳實(shí)踐與使用建議學(xué)習(xí)多頭注意力機(jī)制不需要一開始就追求復(fù)雜實(shí)現(xiàn)。建議按下面的順序推進(jìn)。第一先跑通小規(guī)模測試。d_model 設(shè) 128、num_heads 設(shè) 4序列長度設(shè) 16先用隨機(jī)張量驗(yàn)證輸出形狀。形狀全部正確后再加 mask最后再接入 LayerNorm 和 FFN。第二結(jié)合 loss 曲線判斷實(shí)現(xiàn)是否正確。完全隨機(jī)初始化時(shí)Transformer 的 loss 應(yīng)該短暫下降且不會(huì)立即發(fā)散。如果 loss 在第一步就變成 NaN優(yōu)先檢查 scores 的縮放因子和 LayerNorm 的 eps 參數(shù)。第三頭數(shù)并不是越大越好。常見工程經(jīng)驗(yàn)是每個(gè)頭的維度在 64 附近例如 d_model512 對(duì)應(yīng) 8 個(gè)頭d_model768 對(duì)應(yīng) 12 個(gè)頭。頭數(shù)過少表達(dá)能力受限頭數(shù)過多單個(gè)頭維度太小能夠?qū)W到的特征有限且矩陣乘法形狀更碎GPU 利用率反而下降。第四長序列場景要主動(dòng)優(yōu)化。多頭注意力的計(jì)算復(fù)雜度是 O(n^2)序列長度從 512 提升到 2048計(jì)算量會(huì)增長 16 倍。實(shí)際工程中可以考慮 FlashAttention、稀疏注意力、局部窗口注意力等方案而不是盲目堆算力。第五推理階段要關(guān)注 KV Cache 的復(fù)用。自回歸模型生成時(shí)對(duì) QKV 的處理方式完全不同Q 只和當(dāng)前 token 有關(guān)不緩存K、V 需要保存歷史。如果只是拿模型做訓(xùn)練可以暫時(shí)忽略 KV Cache如果做部署和接口服務(wù)KV Cache 就是性能優(yōu)化的核心。第六任何涉及真實(shí)數(shù)據(jù)訓(xùn)練或應(yīng)用的項(xiàng)目要注意數(shù)據(jù)授權(quán)、隱私保護(hù)和內(nèi)容合規(guī)。模型訓(xùn)練使用他人文本、圖像、語音數(shù)據(jù)時(shí)需要確認(rèn)是否有合法使用權(quán)生成內(nèi)容對(duì)外發(fā)布前需要根據(jù)應(yīng)用場景做好安全審核。12. 總結(jié)與下一步多頭注意力機(jī)制的核心可以濃縮為一句話在參數(shù)總量不變的條件下把單一大矩陣投影拆成多個(gè)子空間并行學(xué)習(xí)再拼接融合讓模型獲得更多樣、更穩(wěn)定的注意力模式。它本身不是復(fù)雜機(jī)制但卻是理解和復(fù)現(xiàn)幾乎所有現(xiàn)代大模型的必經(jīng)之路。建議下一步動(dòng)手做三件事一是修改num_heads從 1 改成 4、8、16觀察輸出變化和顯存波動(dòng)二是給當(dāng)前模塊加入因果掩碼跑一個(gè)簡單的 n-gram 預(yù)測任務(wù)驗(yàn)證自回歸邏輯三是繼續(xù)學(xué)習(xí)位置編碼和 FlashAttention位置編碼解決的是“注意力本身不感知順序”的問題FlashAttention 解決的是長序列下顯存和速度的問題。把多頭注意力這一步踩扎實(shí)后面再看 BERT、GPT、ViT 的源碼會(huì)發(fā)現(xiàn)大量代碼都是這一章的重復(fù)與擴(kuò)展。