現(xiàn):從正弦編碼到RoPE)
1. 項(xiàng)目概述位置編碼——讓模型“看見(jiàn)”序列的秩序在深度學(xué)習(xí)的序列建模領(lǐng)域無(wú)論是處理自然語(yǔ)言、音頻還是時(shí)間序列數(shù)據(jù)模型本身通常是“無(wú)序”的。一個(gè)經(jīng)典的Transformer模型其自注意力機(jī)制Self-Attention在處理一個(gè)句子時(shí)對(duì)于“我愛(ài)北京”和“北京愛(ài)我”這兩個(gè)詞序完全不同的輸入如果不做任何處理它會(huì)計(jì)算出幾乎相同的注意力權(quán)重因?yàn)樗魂P(guān)心詞與詞之間的語(yǔ)義關(guān)聯(lián)而忽略了它們?cè)谛蛄兄械慕^對(duì)位置和相對(duì)順序。這顯然不符合我們的認(rèn)知。位置編碼Positional Encoding, PE就是為了解決這個(gè)問(wèn)題而誕生的核心組件它像給序列中的每個(gè)元素貼上一個(gè)“坐標(biāo)標(biāo)簽”告訴模型“這個(gè)詞在第幾個(gè)位置”。這個(gè)項(xiàng)目標(biāo)題“07-位置編碼 ”暗示了這是一個(gè)系列教程或筆記中的第七部分聚焦于位置編碼這一關(guān)鍵技術(shù)。這個(gè)表情符號(hào)直觀地表達(dá)了“定位”的概念。從相關(guān)熱詞來(lái)看它緊密關(guān)聯(lián)著Transformer架構(gòu)、PyTorch實(shí)現(xiàn)、正弦編碼、可學(xué)習(xí)編碼以及RoPE、ViT位置編碼等前沿變體。理解位置編碼不僅是理解Transformer的基石也是掌握當(dāng)下眾多基于Transformer的視覺(jué)ViT、語(yǔ)音乃至多模態(tài)大模型的關(guān)鍵。本文將從一個(gè)實(shí)踐者的角度深入拆解位置編碼的為什么、是什么和怎么做。我會(huì)結(jié)合PyTorch代碼帶你從零實(shí)現(xiàn)經(jīng)典的正弦位置編碼探討可學(xué)習(xí)位置編碼的優(yōu)劣并分析像RoPE旋轉(zhuǎn)位置編碼這樣的現(xiàn)代方案為何能成為大語(yǔ)言模型的寵兒。無(wú)論你是剛接觸Transformer的新手還是希望深化對(duì)模型細(xì)節(jié)理解的中級(jí)開(kāi)發(fā)者這篇文章都將提供可直接復(fù)現(xiàn)的代碼和背后深刻的原理剖析。2. 位置編碼的核心原理與設(shè)計(jì)思路2.1 自注意力機(jī)制的“位置盲”問(wèn)題要理解位置編碼的必要性必須回到自注意力機(jī)制本身。自注意力通過(guò)計(jì)算查詢(xún)Query、鍵Key、值Value向量之間的相似度來(lái)聚合全局信息。其計(jì)算過(guò)程本質(zhì)上是置換等變Permutation Equivariant的。簡(jiǎn)單來(lái)說(shuō)如果你把輸入序列的順序打亂輸出的序列順序也會(huì)相應(yīng)打亂但每個(gè)輸出位置所聚合的信息內(nèi)容不考慮位置是相似的。用一個(gè)簡(jiǎn)單的例子說(shuō)明假設(shè)我們有一個(gè)包含詞嵌入的序列X [x1, x2, x3]。自注意力層計(jì)算輸出Z Attention(Q, K, V)其中QKVXWW是可學(xué)習(xí)的權(quán)重矩陣。由于點(diǎn)積注意力softmax((QK^T)/√d_k)V的計(jì)算只依賴(lài)于向量間的點(diǎn)積而點(diǎn)積運(yùn)算與向量的絕對(duì)位置無(wú)關(guān)。因此對(duì)于輸入X‘ [x2, x1, x3]交換了x1和x2其輸出Z‘將會(huì)是Z的相應(yīng)行被交換后的結(jié)果。模型無(wú)法區(qū)分“貓追老鼠”和“老鼠追貓”。2.2 位置編碼的注入方式為了解決這個(gè)問(wèn)題我們需要將位置信息顯式地注入到模型中。主流的方法是將位置編碼向量與詞嵌入向量進(jìn)行相加。設(shè)輸入序列長(zhǎng)度為L(zhǎng)詞嵌入維度為d_model。詞嵌入矩陣為E ∈ R^(L×d_model)位置編碼矩陣為P ∈ R^(L×d_model)。那么Transformer的輸入就是X E P這個(gè)簡(jiǎn)單的加法操作是經(jīng)過(guò)精心設(shè)計(jì)的。它假設(shè)位置信息和語(yǔ)義信息存在于同一個(gè)向量空間的不同“子空間”或通道中模型可以通過(guò)后續(xù)的線(xiàn)性變換和注意力機(jī)制學(xué)習(xí)到如何同時(shí)利用這兩種信息。注意為什么是相加而不是拼接相加保持了輸入維度不變?nèi)允莇_model避免了參數(shù)量的顯著增加同時(shí)實(shí)踐表明模型能夠有效學(xué)習(xí)到這種混合表示。拼接雖然信息分離更徹底但會(huì)改變輸入維度需要調(diào)整后續(xù)所有層的權(quán)重維度不夠優(yōu)雅且效率未必更高。2.3 絕對(duì)位置編碼 vs. 相對(duì)位置編碼根據(jù)編碼方式所蘊(yùn)含的信息位置編碼可以分為兩大類(lèi)絕對(duì)位置編碼Absolute Positional Encoding為序列中的每個(gè)絕對(duì)位置如第1個(gè)詞、第2個(gè)詞分配一個(gè)獨(dú)特的編碼向量。經(jīng)典的正弦/余弦編碼和可學(xué)習(xí)位置編碼都屬于此類(lèi)。它直接告訴模型“這是第幾個(gè)位置”。相對(duì)位置編碼Relative Positional Encoding不關(guān)心絕對(duì)位置而是編碼序列中任意兩個(gè)元素之間的相對(duì)距離或相對(duì)位置關(guān)系。例如它編碼“當(dāng)前詞”和“前一個(gè)詞”、“后兩個(gè)詞”之間的關(guān)系。RoPE旋轉(zhuǎn)位置編碼和Transformer-XL中使用的編碼是這類(lèi)方法的杰出代表。它更符合語(yǔ)言的內(nèi)在規(guī)律因?yàn)槲覀兝斫庖粋€(gè)詞的意義往往更依賴(lài)于它與其他詞的相對(duì)關(guān)系而非它在句子中的絕對(duì)序號(hào)。近年來(lái)相對(duì)位置編碼因其更好的長(zhǎng)度外推性處理比訓(xùn)練時(shí)更長(zhǎng)的序列和理論上的優(yōu)越性在大型語(yǔ)言模型中逐漸成為主流。3. 經(jīng)典位置編碼方案詳解與PyTorch實(shí)現(xiàn)3.1 正弦/余弦位置編碼Sinusoidal Positional Encoding這是原版Transformer論文《Attention Is All You Need》提出的方法也是最具標(biāo)志性的位置編碼。它并非可學(xué)習(xí)參數(shù)而是一個(gè)基于正弦和余弦函數(shù)的確定性公式。公式解析對(duì)于位置pos從0開(kāi)始計(jì)數(shù)和維度索引ii0,1,...,d_model-1位置編碼向量P(pos, 2i)和P(pos, 2i1)的計(jì)算公式如下P(pos, 2i) sin(pos / 10000^(2i / d_model))P(pos, 2i1) cos(pos / 10000^(2i / d_model))為什么設(shè)計(jì)成這樣唯一性與連續(xù)性每個(gè)位置都有唯一的編碼。同時(shí)由于正弦函數(shù)的性質(zhì)相鄰位置的編碼是平滑變化的模型可以更容易地學(xué)習(xí)到位置之間的鄰近關(guān)系。可擴(kuò)展性長(zhǎng)度外推對(duì)于訓(xùn)練時(shí)未見(jiàn)過(guò)的更長(zhǎng)序列pos很大由于公式是定義好的我們可以直接計(jì)算其位置編碼而不需要重新訓(xùn)練。盡管外推效果可能下降但至少是可行的。相對(duì)位置的可表達(dá)性一個(gè)關(guān)鍵的性質(zhì)是對(duì)于一個(gè)固定的偏移量kP(posk)可以表示為P(pos)的線(xiàn)性函數(shù)。這意味著模型可能僅通過(guò)注意力機(jī)制中的線(xiàn)性變換就能學(xué)會(huì)關(guān)注相對(duì)位置信息。這是其設(shè)計(jì)精妙之處。PyTorch實(shí)現(xiàn)import torch import torch.nn as nn import math class SinusoidalPositionalEncoding(nn.Module): def __init__(self, d_model: int, max_len: int 5000): super().__init__() # 創(chuàng)建一個(gè)形狀為 (max_len, d_model) 的零矩陣來(lái)存儲(chǔ)位置編碼 pe torch.zeros(max_len, d_model) # 生成位置索引 (max_len, 1) position torch.arange(0, max_len, dtypetorch.float).unsqueeze(1) # 計(jì)算分母項(xiàng)10000^(2i/d_model)使用對(duì)數(shù)空間計(jì)算避免數(shù)值過(guò)大 div_term torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) # 對(duì)偶數(shù)維度應(yīng)用正弦函數(shù) pe[:, 0::2] torch.sin(position * div_term) # 對(duì)奇數(shù)維度應(yīng)用余弦函數(shù) pe[:, 1::2] torch.cos(position * div_term) # 增加一個(gè)批次維度最終形狀為 (1, max_len, d_model)便于廣播相加 pe pe.unsqueeze(0) # 將其注冊(cè)為緩沖區(qū)buffer而不是可訓(xùn)練參數(shù)parameter # 這意味著它會(huì)被保存和加載但不會(huì)被優(yōu)化器更新 self.register_buffer(pe, pe) def forward(self, x: torch.Tensor) - torch.Tensor: Args: x: Tensor, shape [batch_size, seq_len, embedding_dim] Returns: Tensor: 添加了位置編碼的輸入形狀不變 # 將位置編碼加到輸入張量上。pe[:, :x.size(1)] 是為了適配可變序列長(zhǎng)度 x x self.pe[:, :x.size(1)] return x # 使用示例 d_model 512 seq_len 100 batch_size 4 embedding torch.randn(batch_size, seq_len, d_model) # 模擬詞嵌入 pos_encoder SinusoidalPositionalEncoding(d_model) output pos_encoder(embedding) print(f輸入形狀: {embedding.shape}) print(f輸出形狀: {output.shape})實(shí)操心得register_buffer是關(guān)鍵。這確保了位置編碼矩陣pe會(huì)隨著模型一起被保存state_dict和加載但不會(huì)被梯度更新。如果你錯(cuò)誤地將其定義為nn.Parameter優(yōu)化器會(huì)嘗試更新它這違背了正弦編碼“固定不變”的設(shè)計(jì)初衷。在實(shí)際的Transformer模型中位置編碼通常加在嵌入層之后進(jìn)入編碼器堆疊之前。對(duì)于非常長(zhǎng)的序列接近或超過(guò)max_len雖然可以計(jì)算但高頻維度i較大的維度的波長(zhǎng)會(huì)非常長(zhǎng)可能導(dǎo)致位置信息區(qū)分度下降。這是所有絕對(duì)位置編碼面臨的共同挑戰(zhàn)。3.2 可學(xué)習(xí)位置編碼Learnable Positional Encoding這是一種更簡(jiǎn)單直觀的方法將位置編碼直接視為可訓(xùn)練的模型參數(shù)。實(shí)現(xiàn)方式創(chuàng)建一個(gè)形狀為(max_len, d_model)的nn.Embedding層或nn.Parameter。在 forward 過(guò)程中根據(jù)輸入序列的長(zhǎng)度取出對(duì)應(yīng)位置的可學(xué)習(xí)向量加到詞嵌入上。PyTorch實(shí)現(xiàn)class LearnablePositionalEncoding(nn.Module): def __init__(self, d_model: int, max_len: int 5000): super().__init__() # 定義一個(gè)可學(xué)習(xí)的位置嵌入層 self.pe nn.Parameter(torch.zeros(1, max_len, d_model)) # 通常使用較小的標(biāo)準(zhǔn)差進(jìn)行初始化如0.02或0.01 nn.init.normal_(self.pe, mean0.0, std0.02) def forward(self, x: torch.Tensor) - torch.Tensor: x x self.pe[:, :x.size(1)] return x優(yōu)點(diǎn)與缺點(diǎn)分析特性可學(xué)習(xí)位置編碼正弦位置編碼靈活性高。模型可以從數(shù)據(jù)中學(xué)習(xí)最適合任務(wù)的位置表示。低。形式固定無(wú)法根據(jù)數(shù)據(jù)調(diào)整。長(zhǎng)度外推差。只能處理訓(xùn)練時(shí)見(jiàn)過(guò)的位置≤max_len。對(duì)于更長(zhǎng)的序列沒(méi)有對(duì)應(yīng)的學(xué)習(xí)過(guò)的向量。較好。可通過(guò)公式計(jì)算任意位置但外推性能會(huì)衰減。訓(xùn)練穩(wěn)定性可能需要更仔細(xì)的初始化和調(diào)優(yōu)。非常穩(wěn)定無(wú)需擔(dān)心初始化。常見(jiàn)應(yīng)用場(chǎng)景在訓(xùn)練和推理序列長(zhǎng)度固定或變化不大的任務(wù)中表現(xiàn)良好如早期的BERT、一些機(jī)器翻譯模型。需要處理可變長(zhǎng)度或希望有理論保障外推能力的場(chǎng)景原版Transformer。注意事項(xiàng)初始化很重要可學(xué)習(xí)位置編碼的初始化會(huì)影響訓(xùn)練收斂。通常使用較小的隨機(jī)初始化如正態(tài)分布N(0, 0.02)避免初始值過(guò)大淹沒(méi)詞嵌入信號(hào)。過(guò)擬合風(fēng)險(xiǎn)在數(shù)據(jù)量較小的任務(wù)上可學(xué)習(xí)參數(shù)可能無(wú)法充分學(xué)習(xí)到有效的位置模式反而容易過(guò)擬合。max_len的選擇這是一個(gè)超參數(shù)。設(shè)置過(guò)小會(huì)限制模型處理長(zhǎng)序列的能力設(shè)置過(guò)大會(huì)增加不必要的參數(shù)并可能使模型難以學(xué)習(xí)到遠(yuǎn)處位置的有效表示因?yàn)槟切┪恢迷谟?xùn)練數(shù)據(jù)中很少出現(xiàn)。4. 進(jìn)階位置編碼方案RoPE與相對(duì)位置編碼4.1 旋轉(zhuǎn)位置編碼RoPE原理淺析RoPE是近年來(lái)在大型語(yǔ)言模型如LLaMA、GPT-NeoX中廣泛使用的相對(duì)位置編碼方法。它的核心思想非常巧妙通過(guò)旋轉(zhuǎn)矩陣將絕對(duì)位置信息注入到注意力分?jǐn)?shù)的計(jì)算中從而間接地實(shí)現(xiàn)相對(duì)位置編碼的效果。直觀理解想象一下我們把詞嵌入向量看作高維空間中的點(diǎn)。RoPE為每個(gè)位置分配一個(gè)特定的“旋轉(zhuǎn)角度”。在計(jì)算注意力得分Query和Key的點(diǎn)積時(shí)先將Query向量和Key向量根據(jù)它們各自的位置進(jìn)行旋轉(zhuǎn)然后再做點(diǎn)積。神奇的是旋轉(zhuǎn)后的點(diǎn)積結(jié)果只依賴(lài)于兩個(gè)向量的原始內(nèi)容以及它們之間的相對(duì)位置差而與它們的絕對(duì)位置無(wú)關(guān)。數(shù)學(xué)表達(dá)簡(jiǎn)化對(duì)于位置m的Query向量q_m和位置n的Key向量k_nRoPE通過(guò)一個(gè)復(fù)數(shù)旋轉(zhuǎn)操作在代碼中通常用實(shí)數(shù)矩陣實(shí)現(xiàn)將它們轉(zhuǎn)換為q_m’和k_n’使得注意力分?jǐn)?shù)滿(mǎn)足q_m‘, k_n’ g(q_m, k_n, m-n)這里表示點(diǎn)積g是一個(gè)只依賴(lài)于原始向量和相對(duì)位置m-n的函數(shù)。這就實(shí)現(xiàn)了將相對(duì)位置信息編碼到注意力機(jī)制中。為什么RoPE如此受歡迎相對(duì)性直接建模了相對(duì)位置關(guān)系更符合語(yǔ)言建模的直覺(jué)。長(zhǎng)度外推性?xún)?yōu)秀由于其數(shù)學(xué)形式RoPE在處理遠(yuǎn)長(zhǎng)于訓(xùn)練序列的文本時(shí)性能下降相對(duì)平滑外推能力顯著優(yōu)于絕對(duì)位置編碼。兼容自注意力實(shí)現(xiàn)上非常優(yōu)雅只需在計(jì)算Q和K之后、計(jì)算注意力分?jǐn)?shù)之前對(duì)Q和K應(yīng)用旋轉(zhuǎn)變換即可不改變模型的其他部分。4.2 RoPE的PyTorch核心實(shí)現(xiàn)RoPE的實(shí)現(xiàn)涉及一些線(xiàn)性代數(shù)操作。以下是其核心部分的一個(gè)簡(jiǎn)化示例幫助理解其流程import torch import torch.nn as nn import torch.nn.functional as F import math def precompute_freqs_cis(dim: int, end: int, theta: float 10000.0): 預(yù)計(jì)算復(fù)數(shù)旋轉(zhuǎn)因子cis cos i*sin。 Args: dim: 詞嵌入維度必須是偶數(shù)。 end: 最大序列長(zhǎng)度。 theta: 用于控制波長(zhǎng)的基礎(chǔ)值。 Returns: freqs_cis: 復(fù)數(shù)張量形狀 (end, dim//2) # 計(jì)算頻率theta^(-2i/dim) for i in [0, 1, ..., dim//2 -1] freqs 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)) # 生成位置序列 t [0, 1, ..., end-1] t torch.arange(end, devicefreqs.device) # 計(jì)算外積freqs * t形狀 (end, dim//2) freqs torch.outer(t, freqs).float() # 將其轉(zhuǎn)換為復(fù)數(shù)形式cis(freqs) cos(freqs) i*sin(freqs) freqs_cis torch.polar(torch.ones_like(freqs), freqs) # 幅度為1相位為freqs return freqs_cis def apply_rotary_emb( xq: torch.Tensor, xk: torch.Tensor, freqs_cis: torch.Tensor, ): 將旋轉(zhuǎn)位置編碼應(yīng)用到Query和Key上。 Args: xq, xk: Query和Key張量形狀均為 (batch_size, seq_len, num_heads, head_dim) freqs_cis: 預(yù)計(jì)算的旋轉(zhuǎn)因子形狀 (seq_len, head_dim//2) Returns: 旋轉(zhuǎn)后的xq, xk形狀不變 # 將xq和xk的最后一維head_dim視為復(fù)數(shù)對(duì) (x0, x1, x2, x3, ...) - (x0ix1, x2ix3, ...) xq_ torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1, 2)) xk_ torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1, 2)) # 調(diào)整freqs_cis形狀以進(jìn)行廣播 (seq_len, head_dim//2) - (1, seq_len, 1, head_dim//2) freqs_cis freqs_cis.unsqueeze(0).unsqueeze(2) # 復(fù)數(shù)乘法實(shí)現(xiàn)旋轉(zhuǎn) (abi) * (cosθ i*sinθ) (a cosθ - b sinθ) i(a sinθ b cosθ) xq_out torch.view_as_real(xq_ * freqs_cis).flatten(3) xk_out torch.view_as_real(xk_ * freqs_cis).flatten(3) return xq_out.type_as(xq), xk_out.type_as(xk) # 在Transformer注意力模塊中的使用示例偽代碼 class AttentionWithRoPE(nn.Module): def __init__(self, args): super().__init__() self.n_heads args.n_heads self.head_dim args.dim // args.n_heads # ... 其他初始化Wq, Wk, Wv投影層等 # 預(yù)計(jì)算旋轉(zhuǎn)因子假設(shè)最大序列長(zhǎng)度為args.max_seq_len self.freqs_cis precompute_freqs_cis(self.head_dim, args.max_seq_len) def forward(self, x: torch.Tensor): batch_size, seq_len, _ x.shape # 1. 計(jì)算Q, K, V q self.wq(x) # (B, L, dim) k self.wk(x) v self.wv(x) # 2. 重塑為多頭形式 (B, L, n_heads, head_dim) q q.view(batch_size, seq_len, self.n_heads, self.head_dim) k k.view(batch_size, seq_len, self.n_heads, self.head_dim) v v.view(batch_size, seq_len, self.n_heads, self.head_dim) # 3. 應(yīng)用旋轉(zhuǎn)位置編碼僅對(duì)Q和K # 取出當(dāng)前序列長(zhǎng)度對(duì)應(yīng)的旋轉(zhuǎn)因子 freqs_cis self.freqs_cis[:seq_len] q, k apply_rotary_emb(q, k, freqs_cis) # 4. 轉(zhuǎn)置以進(jìn)行批量矩陣乘法 (B, n_heads, L, head_dim) q, k, v q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2) # 5. 計(jì)算縮放點(diǎn)積注意力分?jǐn)?shù) (B, n_heads, L, L) # 此時(shí)Q和K已包含相對(duì)位置信息 scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim) # ... 后續(xù)mask, softmax, 與V相乘等操作實(shí)操心得與避坑指南數(shù)值穩(wěn)定性在計(jì)算freqs時(shí)theta ** (2i/dim)可能導(dǎo)致數(shù)值上溢當(dāng)theta很大時(shí)或下溢當(dāng)i/dim很大時(shí)。使用對(duì)數(shù)空間計(jì)算exp(-log(theta) * 2i / dim)是更穩(wěn)定的做法如上文SinusoidalPositionalEncoding所示。RoPE的實(shí)現(xiàn)中也常采用此技巧。精度問(wèn)題旋轉(zhuǎn)操作涉及三角函數(shù)對(duì)數(shù)值精度敏感。在混合精度訓(xùn)練如AMP中確保關(guān)鍵計(jì)算如apply_rotary_emb在足夠的精度如float32下進(jìn)行或者使用經(jīng)過(guò)數(shù)值穩(wěn)定性?xún)?yōu)化的庫(kù)如xformers庫(kù)中的apply_rotary_emb函數(shù)。因果注意力在自回歸語(yǔ)言模型中需要結(jié)合因果掩碼Causal Mask使用確保當(dāng)前位置只能看到之前的位置。RoPE本身不提供掩碼它只改變了Q和K的計(jì)算方式。5. 位置編碼在視覺(jué)TransformerViT等領(lǐng)域的應(yīng)用與變體5.1 ViT中的位置編碼從1D到2D視覺(jué)Transformer將圖像切分為一系列圖像塊Patches然后將這些塊視為一個(gè)序列進(jìn)行處理。因此它也需要位置編碼來(lái)區(qū)分不同空間位置的圖像塊。1D位置編碼原版ViT最簡(jiǎn)單直接的方式將二維空間位置行列展平為一維序列索引。例如一個(gè)14x14的網(wǎng)格按行優(yōu)先展開(kāi)成0, 1, 2, ..., 195的序列然后使用標(biāo)準(zhǔn)的可學(xué)習(xí)1D位置編碼。這種方法忽略了二維空間的鄰近性例如第13行的最后一個(gè)塊和第14行的第一個(gè)塊在1D序列中相鄰但在2D空間中卻相隔甚遠(yuǎn)。2D位置編碼為了更好保留空間結(jié)構(gòu)可以為行和列分別分配位置編碼然后合并。可學(xué)習(xí)2D編碼定義兩個(gè)可學(xué)習(xí)的嵌入表row_embed和col_embed形狀分別為(num_rows, d/2)和(num_cols, d/2)。對(duì)于一個(gè)位于(i, j)的塊其位置編碼為concat(row_embed[i], col_embed[j])或row_embed[i] col_embed[j]。2D正弦編碼將正弦公式擴(kuò)展到二維。為行坐標(biāo)pos_x和列坐標(biāo)pos_y分別計(jì)算正弦編碼然后拼接或相加。這能更好地建模二維空間中的相對(duì)位置關(guān)系。PyTorch實(shí)現(xiàn)2D可學(xué)習(xí)位置編碼示例class Learnable2DPositionalEncoding(nn.Module): def __init__(self, d_model: int, grid_size: tuple): grid_size: (height, width) 圖像塊網(wǎng)格的高度和寬度 super().__init__() self.height, self.width grid_size # 為行和列分別創(chuàng)建可學(xué)習(xí)嵌入 self.row_embed nn.Parameter(torch.randn(self.height, d_model // 2)) self.col_embed nn.Parameter(torch.randn(self.width, d_model // 2)) nn.init.normal_(self.row_embed, std0.02) nn.init.normal_(self.col_embed, std0.02) def forward(self, x: torch.Tensor) - torch.Tensor: x: (B, L, d_model), L應(yīng)該等于 height * width 假設(shè)x中塊的順序是行優(yōu)先展開(kāi)的。 batch_size, seq_len, d_model x.shape assert seq_len self.height * self.width, 序列長(zhǎng)度必須等于網(wǎng)格大小 # 生成所有位置索引 rows torch.arange(self.height).repeat_interleave(self.width) # [0,0,...,1,1,..., H-1] cols torch.arange(self.width).repeat(self.height) # [0,1,...,W-1,0,1,...] # 獲取對(duì)應(yīng)的行、列嵌入并拼接 row_emb self.row_embed[rows] # (L, d_model//2) col_emb self.col_embed[cols] # (L, d_model//2) pos_emb torch.cat([row_emb, col_emb], dim-1) # (L, d_model) # 廣播并相加 pos_emb pos_emb.unsqueeze(0) # (1, L, d_model) x x pos_emb return x5.2 無(wú)需位置編碼探索位置感知的替代方案近年來(lái)也有一些研究嘗試完全摒棄顯式的位置編碼讓模型從數(shù)據(jù)中隱式地學(xué)習(xí)位置信息。相對(duì)注意力偏置Relative Attention Bias不向輸入添加位置向量而是在計(jì)算注意力分?jǐn)?shù)時(shí)直接加上一個(gè)基于查詢(xún)鍵相對(duì)位置的偏置項(xiàng)b(i-j)。這個(gè)偏置矩陣B是可學(xué)習(xí)的。Swin Transformer中就使用了這種相對(duì)位置偏置。它參數(shù)更少且天然是平移不變的對(duì)于圖像分類(lèi)等任務(wù)有益。卷積或池化預(yù)處理在將圖像塊輸入Transformer之前先使用輕量的卷積層或池化層進(jìn)行處理。卷積操作本身具有平移等變性并能捕獲局部空間關(guān)系可以在一定程度上提供位置信息。條件位置編碼Conditional Positional Encoding, CPECPE不是固定的或可學(xué)習(xí)的查找表而是根據(jù)輸入內(nèi)容動(dòng)態(tài)生成的。例如使用一個(gè)深度可分離卷積Depthwise Convolution作用于輸入序列圖像塊其輸出作為位置編碼。這樣位置編碼能適應(yīng)輸入內(nèi)容更具靈活性。選擇建議對(duì)于自然語(yǔ)言處理RoPE因其優(yōu)秀的外推性和理論性質(zhì)已成為大語(yǔ)言模型的事實(shí)標(biāo)準(zhǔn)。對(duì)于計(jì)算機(jī)視覺(jué)ViT可學(xué)習(xí)的1D或2D位置編碼仍是主流且有效的選擇簡(jiǎn)單可靠。Swin Transformer的相對(duì)偏置方法在層次化設(shè)計(jì)中表現(xiàn)優(yōu)異。對(duì)于音頻或時(shí)間序列正弦位置編碼或可學(xué)習(xí)編碼都是常見(jiàn)選擇需根據(jù)序列長(zhǎng)度是否固定、是否需要外推來(lái)決定。當(dāng)追求極致的平移不變性如圖像分類(lèi)或處理非網(wǎng)格數(shù)據(jù)如圖、點(diǎn)云時(shí)相對(duì)注意力偏置或動(dòng)態(tài)位置編碼如CPE值得嘗試。6. 位置編碼的常見(jiàn)問(wèn)題、調(diào)試技巧與實(shí)戰(zhàn)經(jīng)驗(yàn)6.1 長(zhǎng)度外推Length Extrapolation難題與應(yīng)對(duì)長(zhǎng)度外推是指模型在推理時(shí)處理比訓(xùn)練時(shí)更長(zhǎng)的序列的能力。這是位置編碼面臨的一大挑戰(zhàn)。問(wèn)題表現(xiàn)模型在長(zhǎng)序列上性能急劇下降生成無(wú)意義的文本或預(yù)測(cè)準(zhǔn)確率暴跌。根本原因絕對(duì)位置編碼正弦/可學(xué)習(xí)對(duì)于正弦編碼雖然能計(jì)算但高頻維度在長(zhǎng)序列下波長(zhǎng)過(guò)長(zhǎng)區(qū)分度下降。對(duì)于可學(xué)習(xí)編碼模型根本沒(méi)見(jiàn)過(guò)長(zhǎng)位置對(duì)應(yīng)的向量。注意力模式變化隨著序列變長(zhǎng)注意力權(quán)重的分布可能發(fā)生變化模型未學(xué)習(xí)過(guò)這種模式。應(yīng)對(duì)策略訓(xùn)練時(shí)使用更長(zhǎng)序列最直接有效的方法。在資源允許的情況下盡量用更長(zhǎng)的序列訓(xùn)練模型。位置插值Position Interpolation對(duì)于已經(jīng)用短序列訓(xùn)練好的模型特別是使用RoPE的模型可以將位置索引進(jìn)行縮放。例如訓(xùn)練時(shí)最大位置為2048推理時(shí)需要4096。我們可以將推理時(shí)的位置索引pos除以一個(gè)縮放因子s如s2即使用pos/s來(lái)查詢(xún)位置編碼。這相當(dāng)于將位置編碼的“頻率”降低使其能覆蓋更長(zhǎng)的范圍。LLaMA等模型的外推就采用了此類(lèi)技術(shù)。NTK-aware Scaled RoPE這是一種更聰明的RoPE外推方法它不是在推理時(shí)簡(jiǎn)單縮放而是在訓(xùn)練時(shí)就不均勻地縮放不同維度的頻率。高頻維度對(duì)應(yīng)i大的維度縮放得多一些低頻維度縮放得少一些。這樣能更好地保持模型在訓(xùn)練長(zhǎng)度內(nèi)的性能同時(shí)提升外推能力。使用外推性更好的編碼從一開(kāi)始就選擇RoPE這類(lèi)相對(duì)位置編碼其天然的外推性?xún)?yōu)于絕對(duì)位置編碼。6.2 位置編碼的初始化與融合策略初始化可學(xué)習(xí)位置編碼務(wù)必使用小標(biāo)準(zhǔn)差初始化如0.02。過(guò)大的初始化會(huì)干擾詞嵌入的語(yǔ)義信息導(dǎo)致訓(xùn)練初期不穩(wěn)定。與詞嵌入的尺度協(xié)調(diào)位置編碼的幅度應(yīng)與詞嵌入的幅度相匹配。通常在相加之前會(huì)對(duì)詞嵌入乘以一個(gè)縮放因子sqrt(d_model)以控制其方差。確保位置編碼的初始化幅度與之協(xié)調(diào)。融合策略除了簡(jiǎn)單的加法也有研究嘗試其他融合方式如拼接后通過(guò)一個(gè)線(xiàn)性層投影增加參數(shù)或使用門(mén)控機(jī)制動(dòng)態(tài)調(diào)整位置信息的權(quán)重。但在大多數(shù)實(shí)踐中加法已被證明是簡(jiǎn)單且有效的應(yīng)作為首選。6.3 調(diào)試與驗(yàn)證技巧可視化位置編碼繪制位置編碼矩陣的熱力圖plt.imshow(pe.squeeze().T)觀察其模式。正弦編碼應(yīng)呈現(xiàn)清晰的條紋狀周期模式。可學(xué)習(xí)編碼在訓(xùn)練初期可能是雜亂的訓(xùn)練后應(yīng)呈現(xiàn)出一定的結(jié)構(gòu)如平滑變化。檢查梯度在訓(xùn)練初期監(jiān)控位置編碼參數(shù)的梯度。如果梯度始終為零或異常大可能意味著它與模型其他部分的交互有問(wèn)題。設(shè)計(jì)簡(jiǎn)單測(cè)試構(gòu)建一個(gè)極簡(jiǎn)任務(wù)如“輸出序列中每個(gè)元素的位置索引”。用一個(gè)只有位置編碼作為輸入詞嵌入設(shè)為零的小Transformer來(lái)學(xué)習(xí)這個(gè)任務(wù)。如果模型無(wú)法快速學(xué)會(huì)說(shuō)明位置編碼的信息注入可能有問(wèn)題。對(duì)比消融實(shí)驗(yàn)在你自己任務(wù)的驗(yàn)證集上嘗試去掉位置編碼、使用不同種類(lèi)的位置編碼觀察性能變化。這是最直接的驗(yàn)證方式。一個(gè)常見(jiàn)的坑序列長(zhǎng)度不一致的批處理在訓(xùn)練時(shí)我們常使用動(dòng)態(tài)填充padding來(lái)組成批次。位置編碼應(yīng)該只加到真實(shí)的 token 上而不是 padding 部分。通常在注意力機(jī)制中會(huì)使用注意力掩碼Attention Mask來(lái)屏蔽 padding 位置。位置編碼的加法操作本身不需要特殊處理因?yàn)楹罄m(xù)的注意力掩碼會(huì)阻止模型關(guān)注這些加了位置編碼的 padding 位置。但是如果你使用了像RNN這樣的遞歸網(wǎng)絡(luò)則需要小心處理。位置編碼雖是一個(gè)“小”組件卻是Transformer系列模型不可或缺的“靈魂”之一。理解其背后的原理根據(jù)任務(wù)需求選擇合適的方案并能在實(shí)踐中調(diào)試和優(yōu)化是構(gòu)建高效Transformer模型的關(guān)鍵一步。從確定性的正弦波到可學(xué)習(xí)的參數(shù)表再到精巧的旋轉(zhuǎn)操作位置編碼的發(fā)展也體現(xiàn)了深度學(xué)習(xí)從手工設(shè)計(jì)特征到數(shù)據(jù)驅(qū)動(dòng)學(xué)習(xí)再到尋求更優(yōu)歸納偏置的演進(jìn)路徑。