現(xiàn)Transformer:自注意力機(jī)制與編碼器-解碼器架構(gòu)詳解)
在自然語言處理乃至整個AI領(lǐng)域Transformer架構(gòu)無疑是過去幾年最具革命性的模型之一。然而許多學(xué)習(xí)者在初次接觸時(shí)往往會被其核心組件——自注意力機(jī)制Self-Attention——所吸引花費(fèi)大量精力理解Q、K、V矩陣和縮放點(diǎn)積注意力卻對Transformer作為一個完整系統(tǒng)如何運(yùn)作感到困惑。注意力機(jī)制固然關(guān)鍵但它只是整個架構(gòu)中的一個“零件”。本文將系統(tǒng)性地拆解Transformer的完整搭建過程從宏觀架構(gòu)到微觀實(shí)現(xiàn)結(jié)合代碼示例帶你真正理解這個“注意力機(jī)器”是如何被組裝起來并高效工作的。無論你是希望夯實(shí)基礎(chǔ)的學(xué)生還是需要在項(xiàng)目中應(yīng)用Transformer的開發(fā)者這篇文章都將提供從理論到實(shí)踐的完整路徑。1. Transformer 架構(gòu)全景不止于注意力在深入細(xì)節(jié)之前我們必須建立對Transformer的整體認(rèn)知。2017年Vaswani等人在論文《Attention Is All You Need》中提出了Transformer模型其核心思想是完全摒棄循環(huán)神經(jīng)網(wǎng)絡(luò)RNN和卷積神經(jīng)網(wǎng)絡(luò)CNN僅依賴注意力機(jī)制來構(gòu)建序列到序列的模型。1.1 宏觀架構(gòu)編碼器-解碼器范式Transformer遵循經(jīng)典的編碼器-解碼器Encoder-Decoder結(jié)構(gòu)但內(nèi)部組件全部換新。編碼器Encoder負(fù)責(zé)將輸入序列如一句英文編碼成一個富含上下文信息的中間表示。原始論文中編碼器由N6個完全相同的層堆疊而成。解碼器Decoder負(fù)責(zé)根據(jù)編碼器的輸出和已生成的部分輸出序列自回歸地一個接一個生成目標(biāo)序列如對應(yīng)的中文翻譯。解碼器同樣由N6個相同的層堆疊。每一層編碼器和解碼器都不是簡單的注意力模塊而是由更基礎(chǔ)的子層Sublayer通過殘差連接和層歸一化精巧組合而成。1.2 核心組件清單要搭建一個Transformer你需要準(zhǔn)備以下“零件”輸入嵌入Input Embedding將離散的符號單詞轉(zhuǎn)換為連續(xù)的向量。位置編碼Positional Encoding為序列注入順序信息彌補(bǔ)自注意力機(jī)制本身對位置不敏感的缺陷。多頭自注意力機(jī)制Multi-Head Self-Attention核心“零件”用于捕捉序列內(nèi)部的長距離依賴關(guān)系。前饋神經(jīng)網(wǎng)絡(luò)Position-wise Feed-Forward Network一個應(yīng)用于每個位置上的獨(dú)立全連接網(wǎng)絡(luò)用于進(jìn)行非線性變換。殘差連接Residual Connection將子層的輸入直接加到其輸出上緩解深層網(wǎng)絡(luò)訓(xùn)練中的梯度消失問題。層歸一化Layer Normalization對每個樣本的所有特征進(jìn)行歸一化穩(wěn)定訓(xùn)練過程。掩碼多頭注意力Masked Multi-Head Attention解碼器中使用的、防止當(dāng)前位置“看到”未來信息的注意力機(jī)制。編碼器-解碼器注意力Encoder-Decoder Attention解碼器中讓當(dāng)前生成位置關(guān)注整個輸入序列的注意力機(jī)制。線性層與SoftmaxLinear Softmax將解碼器的最終輸出映射到目標(biāo)詞匯表概率分布。理解了這些組件及其關(guān)系我們才能開始動手“搭建”。2. 環(huán)境準(zhǔn)備與基礎(chǔ)工具在開始編碼實(shí)現(xiàn)前我們需要配置開發(fā)環(huán)境。本文將以PyTorch框架為例進(jìn)行實(shí)現(xiàn)因?yàn)樗鼊討B(tài)圖的特點(diǎn)非常適合教學(xué)和原型設(shè)計(jì)。2.1 環(huán)境配置確保你已安裝Python推薦3.8版本和pip。然后安裝必要的庫# 創(chuàng)建虛擬環(huán)境可選但推薦 python -m venv transformer-env source transformer-env/bin/activate # Linux/Mac # transformer-env\Scripts\activate # Windows # 安裝核心依賴 pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu # 以CPU版本為例可根據(jù)CUDA版本調(diào)整 pip install numpy matplotlib tqdm2.2 項(xiàng)目結(jié)構(gòu)規(guī)劃一個清晰的項(xiàng)目結(jié)構(gòu)有助于管理復(fù)雜的模型代碼。建議如下transformer_from_scratch/ ├── model.py # Transformer模型核心架構(gòu)定義 ├── layers.py # 各個子層注意力、前饋網(wǎng)絡(luò)等的實(shí)現(xiàn) ├── embeddings.py # 詞嵌入和位置編碼的實(shí)現(xiàn) ├── utils.py # 工具函數(shù)如掩碼生成 ├── train.py # 訓(xùn)練腳本 ├── config.py # 超參數(shù)配置 └── data/ # 數(shù)據(jù)目錄接下來我們將從最底層的組件開始自底向上地構(gòu)建整個模型。3. 底層組件實(shí)現(xiàn)從嵌入到注意力3.1 詞嵌入與位置編碼首先實(shí)現(xiàn)embeddings.py。詞嵌入將單詞ID映射為d_model維的向量。import torch import torch.nn as nn import math class TokenEmbedding(nn.Module): 標(biāo)準(zhǔn)的詞嵌入層 def __init__(self, vocab_size, d_model): super().__init__() self.embedding nn.Embedding(vocab_size, d_model) self.d_model d_model def forward(self, x): # x: [batch_size, seq_len] # 乘以 sqrt(d_model) 是論文中的一種縮放有助于訓(xùn)練穩(wěn)定 return self.embedding(x) * math.sqrt(self.d_model)位置編碼是Transformer的精華之一它使用正弦和余弦函數(shù)來生成絕對位置信息。class PositionalEncoding(nn.Module): 位置編碼層 def __init__(self, d_model, max_len5000, dropout0.1): super().__init__() self.dropout nn.Dropout(pdropout) # 計(jì)算位置編碼矩陣 pe torch.zeros(max_len, d_model) position torch.arange(0, max_len, dtypetorch.float).unsqueeze(1) # [max_len, 1] div_term torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) pe[:, 0::2] torch.sin(position * div_term) # 偶數(shù)維度用sin pe[:, 1::2] torch.cos(position * div_term) # 奇數(shù)維度用cos pe pe.unsqueeze(0) # [1, max_len, d_model] 便于廣播 # 將pe注冊為buffer不參與訓(xùn)練的參數(shù) self.register_buffer(pe, pe) def forward(self, x): # x: [batch_size, seq_len, d_model] x x self.pe[:, :x.size(1), :] # 只取前seq_len個位置 return self.dropout(x)3.2 縮放點(diǎn)積注意力與多頭注意力這是Transformer的核心“零件”。我們先在layers.py中實(shí)現(xiàn)最基礎(chǔ)的縮放點(diǎn)積注意力。import torch.nn.functional as F def scaled_dot_product_attention(q, k, v, maskNone, dropoutNone): 計(jì)算縮放點(diǎn)積注意力。 參數(shù): q: query張量形狀 [batch_size, ..., seq_len_q, depth] k: key張量形狀 [batch_size, ..., seq_len_k, depth] v: value張量形狀 [batch_size, ..., seq_len_v, depth_v] (通常 seq_len_k seq_len_v) mask: 浮點(diǎn)數(shù)張量形狀可廣播到 [..., seq_len_q, seq_len_k] dropout: nn.Dropout層實(shí)例 返回: 輸出注意力權(quán)重 # 計(jì)算 Q K^T matmul_qk torch.matmul(q, k.transpose(-2, -1)) # [..., seq_len_q, seq_len_k] # 縮放 d_k q.size(-1) scaled_attention_logits matmul_qk / math.sqrt(d_k) # 應(yīng)用掩碼如果存在 if mask is not None: # 將mask中為1的位置需要被掩蓋設(shè)置為一個非常大的負(fù)數(shù)softmax后接近0 scaled_attention_logits scaled_attention_logits.masked_fill(mask 0, -1e9) # 計(jì)算softmax得到注意力權(quán)重 attention_weights F.softmax(scaled_attention_logits, dim-1) # [..., seq_len_q, seq_len_k] if dropout is not None: attention_weights dropout(attention_weights) # 加權(quán)求和 output torch.matmul(attention_weights, v) # [..., seq_len_q, depth_v] return output, attention_weights基于此我們實(shí)現(xiàn)多頭注意力。其思想是將d_model維的Q、K、V投影到h頭數(shù)個不同的、維度更低的子空間d_k,d_v在每個頭上并行計(jì)算注意力最后將結(jié)果拼接并投影回來。class MultiHeadAttention(nn.Module): 多頭注意力層 def __init__(self, d_model, num_heads, dropout0.1): super().__init__() assert d_model % num_heads 0, d_model必須能被num_heads整除 self.d_model d_model self.num_heads num_heads self.depth d_model // num_heads # 每個頭的維度 # 定義線性投影層 self.wq nn.Linear(d_model, d_model) # W^Q self.wk nn.Linear(d_model, d_model) # W^K self.wv nn.Linear(d_model, d_model) # W^V self.dense nn.Linear(d_model, d_model) # 最終輸出投影層 self.dropout nn.Dropout(dropout) def split_heads(self, x, batch_size): 將最后的d_model維度分割為(num_heads, depth)。 轉(zhuǎn)置后形狀變?yōu)?[batch_size, num_heads, seq_len, depth] x x.view(batch_size, -1, self.num_heads, self.depth) return x.permute(0, 2, 1, 3) def forward(self, v, k, q, maskNone): batch_size q.size(0) # 1. 線性投影并分頭 q self.wq(q) # [batch_size, seq_len_q, d_model] k self.wk(k) # [batch_size, seq_len_k, d_model] v self.wv(v) # [batch_size, seq_len_v, d_model] q self.split_heads(q, batch_size) # [batch_size, num_heads, seq_len_q, depth] k self.split_heads(k, batch_size) v self.split_heads(v, batch_size) # 2. 計(jì)算縮放點(diǎn)積注意力 scaled_attention, attention_weights scaled_dot_product_attention( q, k, v, mask, self.dropout ) # scaled_attention: [batch_size, num_heads, seq_len_q, depth] # 3. 合并多頭 scaled_attention scaled_attention.permute(0, 2, 1, 3).contiguous() # [batch_size, seq_len_q, num_heads, depth] concat_attention scaled_attention.view(batch_size, -1, self.d_model) # [batch_size, seq_len_q, d_model] # 4. 最終線性投影 output self.dense(concat_attention) # [batch_size, seq_len_q, d_model] return output, attention_weights3.3 前饋網(wǎng)絡(luò)與子層包裝每個編碼器和解碼器層中的前饋網(wǎng)絡(luò)是一個簡單的兩層全連接網(wǎng)絡(luò)中間有一個ReLU激活函數(shù)。它獨(dú)立地應(yīng)用于每個位置。class PositionwiseFeedForward(nn.Module): 位置前饋網(wǎng)絡(luò) def __init__(self, d_model, d_ff, dropout0.1): super().__init__() self.linear1 nn.Linear(d_model, d_ff) # 第一層擴(kuò)大維度 self.linear2 nn.Linear(d_ff, d_model) # 第二層投影回d_model self.dropout nn.Dropout(dropout) self.activation nn.ReLU() def forward(self, x): # x: [batch_size, seq_len, d_model] return self.linear2(self.dropout(self.activation(self.linear1(x))))現(xiàn)在我們需要一個通用的“子層”結(jié)構(gòu)它將核心操作如注意力或前饋網(wǎng)絡(luò)與殘差連接和層歸一化包裝起來。class SublayerConnection(nn.Module): 殘差連接后接層歸一化。 注意為了與原始論文一致歸一化在子層操作之前Pre-LN但有些實(shí)現(xiàn)放在之后Post-LN。 這里采用更常見的Pre-LN因?yàn)樗ǔS?xùn)練更穩(wěn)定。 def __init__(self, size, dropout): super().__init__() self.norm nn.LayerNorm(size) self.dropout nn.Dropout(dropout) def forward(self, x, sublayer): sublayer是一個函數(shù)它接受歸一化后的輸入并返回輸出。 # Pre-LN: 先歸一化再執(zhí)行子層操作最后加殘差和Dropout return x self.dropout(sublayer(self.norm(x)))4. 組裝編碼器與解碼器層有了這些基礎(chǔ)組件我們可以開始搭建編碼器和解碼器的單層結(jié)構(gòu)。4.1 編碼器層實(shí)現(xiàn)一個編碼器層包含兩個子層多頭自注意力層和前饋網(wǎng)絡(luò)層。class EncoderLayer(nn.Module): 單個編碼器層 def __init__(self, d_model, num_heads, d_ff, dropout): super().__init__() self.self_attn MultiHeadAttention(d_model, num_heads, dropout) self.feed_forward PositionwiseFeedForward(d_model, d_ff, dropout) self.sublayer nn.ModuleList([SublayerConnection(d_model, dropout) for _ in range(2)]) def forward(self, x, mask): x: [batch_size, seq_len, d_model] mask: 用于自注意力的掩碼形狀 [batch_size, 1, 1, seq_len]用于padding mask # 第一子層多頭自注意力自注意力意味著 qkvx x self.sublayer[0](x, lambda x: self.self_attn(x, x, x, mask)[0]) # 第二子層前饋網(wǎng)絡(luò) x self.sublayer[1](x, self.feed_forward) return x4.2 解碼器層實(shí)現(xiàn)解碼器層更復(fù)雜一些包含三個子層掩碼多頭自注意力層防止信息泄露。編碼器-解碼器注意力層關(guān)注編碼器輸出。前饋網(wǎng)絡(luò)層。class DecoderLayer(nn.Module): 單個解碼器層 def __init__(self, d_model, num_heads, d_ff, dropout): super().__init__() self.self_attn MultiHeadAttention(d_model, num_heads, dropout) self.cross_attn MultiHeadAttention(d_model, num_heads, dropout) # 編碼器-解碼器注意力 self.feed_forward PositionwiseFeedForward(d_model, d_ff, dropout) self.sublayer nn.ModuleList([SublayerConnection(d_model, dropout) for _ in range(3)]) def forward(self, x, encoder_output, src_mask, tgt_mask): x: 解碼器輸入即上一層的輸出或目標(biāo)序列嵌入[batch_size, tgt_seq_len, d_model] encoder_output: 編碼器最后一層的輸出 [batch_size, src_seq_len, d_model] src_mask: 源序列掩碼用于padding和/或編碼器-解碼器注意力[batch_size, 1, 1, src_seq_len] tgt_mask: 目標(biāo)序列掩碼用于padding和因果掩碼[batch_size, 1, tgt_seq_len, tgt_seq_len] # 第一子層掩碼多頭自注意力 x self.sublayer[0](x, lambda x: self.self_attn(x, x, x, tgt_mask)[0]) # 第二子層編碼器-解碼器注意力。q來自解碼器k和v來自編碼器輸出。 x self.sublayer[1](x, lambda x: self.cross_attn(x, encoder_output, encoder_output, src_mask)[0]) # 第三子層前饋網(wǎng)絡(luò) x self.sublayer[2](x, self.feed_forward) return x5. 構(gòu)建完整的Transformer模型現(xiàn)在我們將編碼器層、解碼器層、嵌入層等所有部件組裝成完整的Transformer模型。在model.py中實(shí)現(xiàn)。5.1 編碼器堆疊編碼器由N個EncoderLayer堆疊而成最前面是詞嵌入和位置編碼。class Encoder(nn.Module): 完整的編碼器嵌入 位置編碼 N個編碼器層 def __init__(self, vocab_size, d_model, num_layers, num_heads, d_ff, max_seq_len, dropout): super().__init__() self.d_model d_model self.token_embedding TokenEmbedding(vocab_size, d_model) self.positional_encoding PositionalEncoding(d_model, max_seq_len, dropout) self.layers nn.ModuleList([EncoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers)]) self.norm nn.LayerNorm(d_model) # 最終輸出前的層歸一化 def forward(self, src, src_mask): # src: [batch_size, src_seq_len] # src_mask: [batch_size, 1, 1, src_seq_len] # 1. 嵌入與位置編碼 x self.token_embedding(src) # [batch_size, src_seq_len, d_model] x self.positional_encoding(x) # 2. 通過N個編碼器層 for layer in self.layers: x layer(x, src_mask) # 3. 最終歸一化 return self.norm(x)5.2 解碼器堆疊解碼器結(jié)構(gòu)類似但多了編碼器-解碼器注意力層和因果掩碼。class Decoder(nn.Module): 完整的解碼器嵌入 位置編碼 N個解碼器層 最終歸一化 def __init__(self, vocab_size, d_model, num_layers, num_heads, d_ff, max_seq_len, dropout): super().__init__() self.d_model d_model self.token_embedding TokenEmbedding(vocab_size, d_model) self.positional_encoding PositionalEncoding(d_model, max_seq_len, dropout) self.layers nn.ModuleList([DecoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers)]) self.norm nn.LayerNorm(d_model) def forward(self, tgt, encoder_output, src_mask, tgt_mask): # tgt: 目標(biāo)序列訓(xùn)練時(shí)是右移一位的標(biāo)簽[batch_size, tgt_seq_len] # encoder_output: 編碼器輸出 [batch_size, src_seq_len, d_model] # src_mask: 源序列掩碼 [batch_size, 1, 1, src_seq_len] # tgt_mask: 目標(biāo)序列掩碼 [batch_size, 1, tgt_seq_len, tgt_seq_len] x self.token_embedding(tgt) x self.positional_encoding(x) for layer in self.layers: x layer(x, encoder_output, src_mask, tgt_mask) return self.norm(x)5.3 最終的Transformer類最后我們創(chuàng)建頂層的Transformer類它組合編碼器、解碼器并添加一個最終的線性投影層和softmax來生成詞匯表概率。class Transformer(nn.Module): 完整的Transformer模型 def __init__(self, src_vocab_size, tgt_vocab_size, d_model512, num_layers6, num_heads8, d_ff2048, max_seq_len5000, dropout0.1): super().__init__() self.encoder Encoder(src_vocab_size, d_model, num_layers, num_heads, d_ff, max_seq_len, dropout) self.decoder Decoder(tgt_vocab_size, d_model, num_layers, num_heads, d_ff, max_seq_len, dropout) self.final_linear nn.Linear(d_model, tgt_vocab_size) # 將解碼器輸出投影到目標(biāo)詞表大小 # 參數(shù)初始化重要 self._init_parameters() def _init_parameters(self): 使用Xavier均勻初始化參數(shù) for p in self.parameters(): if p.dim() 1: nn.init.xavier_uniform_(p) def forward(self, src, tgt, src_mask, tgt_mask): src: 源語言序列索引 [batch_size, src_len] tgt: 目標(biāo)語言序列索引訓(xùn)練時(shí)通常右移一位[batch_size, tgt_len] src_mask: 源序列掩碼 [batch_size, 1, 1, src_len] tgt_mask: 目標(biāo)序列掩碼 [batch_size, 1, tgt_len, tgt_len] 返回: 目標(biāo)詞表的logits [batch_size, tgt_len, tgt_vocab_size] encoder_output self.encoder(src, src_mask) decoder_output self.decoder(tgt, encoder_output, src_mask, tgt_mask) output_logits self.final_linear(decoder_output) return output_logits def encode(self, src, src_mask): 用于推理時(shí)編碼源序列 return self.encoder(src, src_mask) def decode(self, tgt, encoder_output, src_mask, tgt_mask): 用于推理時(shí)自回歸解碼 return self.decoder(tgt, encoder_output, src_mask, tgt_mask)6. 關(guān)鍵工具掩碼生成Transformer中掩碼Mask至關(guān)重要它有兩種主要類型填充掩碼Padding Mask在處理變長序列時(shí)忽略填充符pad的位置。前瞻掩碼Look-ahead Mask / Causal Mask在解碼器中防止當(dāng)前位置關(guān)注到未來的信息保證自回歸屬性。我們在utils.py中實(shí)現(xiàn)掩碼生成函數(shù)。def create_padding_mask(seq, pad_token_id0): 為序列創(chuàng)建填充掩碼。 參數(shù): seq: 整數(shù)張量形狀 [batch_size, seq_len] pad_token_id: 填充符的ID默認(rèn)為0 返回: mask: 浮點(diǎn)數(shù)張量形狀 [batch_size, 1, 1, seq_len] 其中pad_token_id的位置為0需要被掩蓋其他位置為1。 # 找出等于pad_token_id的位置 mask (seq ! pad_token_id).unsqueeze(1).unsqueeze(2) # [batch_size, 1, 1, seq_len] # 轉(zhuǎn)換為浮點(diǎn)數(shù)并調(diào)整值域1表示保留0表示掩蓋 return mask.float() def create_look_ahead_mask(size): 創(chuàng)建前瞻因果掩碼。 參數(shù): size: 目標(biāo)序列的長度 返回: mask: 形狀為 [size, size] 的下三角矩陣包含對角線 下三角部分包括對角線為1上三角部分為0。 # 創(chuàng)建一個全1矩陣 mask torch.ones(size, size) # 取上三角部分不包括對角線設(shè)置為0 mask torch.triu(mask, diagonal1) # 反轉(zhuǎn)需要被掩蓋的位置為0保留的位置為1 mask 1 - mask return mask # [size, size] # 組合掩碼在解碼器中需要同時(shí)應(yīng)用填充掩碼和前瞻掩碼 def create_decoder_mask(tgt_seq, pad_token_id0): 為解碼器創(chuàng)建組合掩碼。 參數(shù): tgt_seq: 目標(biāo)序列形狀 [batch_size, tgt_len] 返回: combined_mask: 形狀 [batch_size, 1, tgt_len, tgt_len] tgt_len tgt_seq.size(1) # 創(chuàng)建填充掩碼 [batch_size, 1, 1, tgt_len] padding_mask create_padding_mask(tgt_seq, pad_token_id) # 創(chuàng)建前瞻掩碼 [tgt_len, tgt_len] look_ahead_mask create_look_ahead_mask(tgt_len).to(tgt_seq.device) # 組合兩個掩碼都為1的位置才需要保留 # 將padding_mask擴(kuò)展維度以進(jìn)行廣播 [batch_size, 1, 1, tgt_len] - [batch_size, 1, tgt_len, tgt_len]? 需要調(diào)整 # 更常見的做法是combined_mask padding_mask look_ahead_mask (在需要的位置) # 但維度不匹配。標(biāo)準(zhǔn)做法是先創(chuàng)建前瞻掩碼然后與填充掩碼相乘廣播。 # 我們調(diào)整填充掩碼的維度使其能與前瞻掩碼逐元素相乘。 padding_mask padding_mask.squeeze(1).squeeze(1) # [batch_size, tgt_len] # 為了與 [tgt_len, tgt_len] 相乘需要擴(kuò)展維度 padding_mask padding_mask.unsqueeze(1) # [batch_size, 1, tgt_len] combined_mask look_ahead_mask.unsqueeze(0) # [1, tgt_len, tgt_len] combined_mask combined_mask * padding_mask # 廣播相乘 [batch_size, tgt_len, tgt_len] # 最后增加一個維度用于多頭注意力 combined_mask combined_mask.unsqueeze(1) # [batch_size, 1, tgt_len, tgt_len] return combined_mask7. 模型訓(xùn)練與推理示例7.1 配置與數(shù)據(jù)準(zhǔn)備在config.py中定義超參數(shù)并準(zhǔn)備一個簡單的模擬數(shù)據(jù)集用于演示。# config.py class Config: src_vocab_size 10000 # 源語言詞表大小 tgt_vocab_size 10000 # 目標(biāo)語言詞表大小 d_model 512 # 模型維度 num_layers 6 # 編碼器/解碼器層數(shù) num_heads 8 # 注意力頭數(shù) d_ff 2048 # 前饋網(wǎng)絡(luò)中間層維度 max_seq_len 100 # 最大序列長度 dropout 0.1 # Dropout率 batch_size 32 lr 1e-4 epochs 10 pad_token_id 07.2 訓(xùn)練循環(huán)骨架在train.py中我們展示一個簡化的訓(xùn)練循環(huán)。實(shí)際應(yīng)用中需要加載真實(shí)數(shù)據(jù)集、實(shí)現(xiàn)BPE/WordPiece分詞、構(gòu)建DataLoader等。# train.py (簡化版) import torch import torch.nn as nn from torch.optim import Adam from model import Transformer from utils import create_padding_mask, create_decoder_mask from config import Config def train_step(model, src_batch, tgt_batch, criterion, optimizer, config): model.train() optimizer.zero_grad() # 創(chuàng)建掩碼 src_mask create_padding_mask(src_batch, config.pad_token_id) # 解碼器的輸入是目標(biāo)序列去掉最后一個詞標(biāo)簽是目標(biāo)序列去掉第一個詞右移 tgt_input tgt_batch[:, :-1] tgt_labels tgt_batch[:, 1:] # 為解碼器輸入創(chuàng)建掩碼組合填充掩碼和前瞻掩碼 tgt_mask create_decoder_mask(tgt_input, config.pad_token_id) # 前向傳播 logits model(src_batch, tgt_input, src_mask, tgt_mask) # [batch, tgt_len-1, vocab] # 計(jì)算損失 loss criterion(logits.reshape(-1, config.tgt_vocab_size), tgt_labels.reshape(-1)) # 反向傳播與優(yōu)化 loss.backward() optimizer.step() return loss.item() def main(): config Config() device torch.device(cuda if torch.cuda.is_available() else cpu) # 初始化模型、損失函數(shù)、優(yōu)化器 model Transformer(config.src_vocab_size, config.tgt_vocab_size, config.d_model, config.num_layers, config.num_heads, config.d_ff, config.max_seq_len, config.dropout).to(device) criterion nn.CrossEntropyLoss(ignore_indexconfig.pad_token_id) # 忽略填充符的損失 optimizer Adam(model.parameters(), lrconfig.lr) # 模擬數(shù)據(jù)實(shí)際應(yīng)使用DataLoader for epoch in range(config.epochs): # 假設(shè)我們有一個數(shù)據(jù)生成器 # src_batch, tgt_batch next(data_iter) src_batch torch.randint(1, config.src_vocab_size, (config.batch_size, 20)).to(device) tgt_batch torch.randint(1, config.tgt_vocab_size, (config.batch_size, 25)).to(device) loss train_step(model, src_batch, tgt_batch, criterion, optimizer, config) print(fEpoch {epoch1}, Loss: {loss:.4f}) print(訓(xùn)練完成) if __name__ __main__: main()7.3 推理貪婪解碼示例訓(xùn)練完成后模型需要以自回歸方式進(jìn)行推理。def greedy_decode(model, src, src_mask, max_len, start_token_id, end_token_id, device): 貪婪解碼每次選擇概率最高的詞作為下一個輸入。 參數(shù): model: 訓(xùn)練好的Transformer模型 src: 源序列 [1, src_len] src_mask: 源序列掩碼 [1, 1, 1, src_len] max_len: 最大生成長度 start_token_id: 起始符如 s的ID end_token_id: 結(jié)束符如 /s的ID device: 設(shè)備 返回: result: 生成的目標(biāo)序列索引列表 model.eval() # 編碼源序列 memory model.encode(src, src_mask) # [1, src_len, d_model] # 初始化目標(biāo)序列以起始符開始 ys torch.ones(1, 1).fill_(start_token_id).type_as(src).to(device) # [1, 1] for i in range(max_len-1): # 為當(dāng)前已生成序列創(chuàng)建掩碼 tgt_mask create_decoder_mask(ys, pad_token_id0) # 假設(shè)0是pad這里ys沒有pad # 解碼 out model.decode(ys, memory, src_mask, tgt_mask) # [1, current_len, d_model] out model.final_linear(out[:, -1:, :]) # 只取最后一個位置的輸出 [1, 1, vocab] # 選擇概率最高的詞 prob torch.softmax(out, dim-1) next_word torch.argmax(prob, dim-1) # [1, 1] # 將新詞拼接到序列中 ys torch.cat([ys, next_word], dim1) # 如果生成結(jié)束符則停止 if next_word.item() end_token_id: break return ys.squeeze(0).tolist() # 返回列表8. 常見問題與排查思路在實(shí)現(xiàn)和訓(xùn)練Transformer時(shí)你可能會遇到以下典型問題問題現(xiàn)象可能原因排查思路與解決方案訓(xùn)練Loss為NaN或不下降1. 學(xué)習(xí)率過高。2. 梯度爆炸。3. 參數(shù)初始化不當(dāng)。4. 數(shù)據(jù)中存在異常值如非常大的數(shù)。1. 降低學(xué)習(xí)率如從1e-4降到1e-5。2. 使用梯度裁剪torch.nn.utils.clip_grad_norm_。3. 檢查并確保使用了正確的參數(shù)初始化如Xavier初始化。4. 檢查輸入數(shù)據(jù)范圍進(jìn)行歸一化或標(biāo)準(zhǔn)化。模型輸出全是同一個詞1. 學(xué)習(xí)率過低模型未更新。2. 損失函數(shù)忽略索引設(shè)置錯誤導(dǎo)致梯度無法回傳。3. 目標(biāo)序列掩碼因果掩碼錯誤導(dǎo)致模型無法有效學(xué)習(xí)序列依賴。1. 嘗試增大學(xué)習(xí)率。2. 檢查CrossEntropyLoss的ignore_index是否與pad_token_id一致。3. 可視化tgt_mask確保其是嚴(yán)格的下三角矩陣包括對角線。GPU內(nèi)存溢出OOM1. 批次大小Batch Size或序列長度Seq Len過大。2. 模型參數(shù)量過大d_model,num_layers等設(shè)置過高。3. 注意力權(quán)重大小為[batch, heads, seq_len, seq_len]序列很長時(shí)平方增長。1. 減小batch_size或使用梯度累積。2. 減小模型尺寸或使用模型并行。3. 對于長序列考慮使用稀疏注意力、線性注意力或分塊計(jì)算。推理速度非常慢1. 自回歸解碼時(shí)每次只生成一個詞序列長時(shí)循環(huán)次數(shù)多。2. 未使用緩存Key/Value緩存導(dǎo)致重復(fù)計(jì)算。1. 考慮使用束搜索Beam Search的優(yōu)化實(shí)現(xiàn)。2. 在推理時(shí)實(shí)現(xiàn)KV緩存將之前時(shí)間步計(jì)算過的K和V緩存起來避免重復(fù)計(jì)算這是生產(chǎn)級Transformer推理的必備優(yōu)化。位置編碼效果不佳1. 正弦/余弦位置編碼的波長設(shè)置不當(dāng)。2. 對于長于訓(xùn)練時(shí)max_seq_len的序列外推能力差。1. 確保div_term計(jì)算正確。2. 考慮使用可學(xué)習(xí)的位置編碼Learnable Positional Embedding或相對位置編碼如RoPE, ALiBi后者對長序列外推更友好。9. 最佳實(shí)踐與工程建議要將一個“玩具級”的Transformer實(shí)現(xiàn)用于實(shí)際項(xiàng)目需要考慮以下工程細(xì)節(jié)高效的批次處理與掩碼確保你的DataLoader能高效地生成批次數(shù)據(jù)并創(chuàng)建對應(yīng)的填充掩碼。對于變長序列通常需要先按長度排序再批次化以減少填充開銷。學(xué)習(xí)率調(diào)度與預(yù)熱WarmupTransformer模型通常受益于帶預(yù)熱的學(xué)習(xí)率調(diào)度策略如Noam調(diào)度器。在訓(xùn)練初期使用較小的學(xué)習(xí)率然后逐漸增大再按步數(shù)或輪次衰減。標(biāo)簽平滑Label Smoothing在計(jì)算交叉熵?fù)p失時(shí)使用標(biāo)簽平滑如nn.CrossEntropyLoss的label_smoothing參數(shù)可以防止模型對預(yù)測結(jié)果過于自信提升泛化能力。檢查點(diǎn)Checkpoint保存定期保存模型狀態(tài)和優(yōu)化器狀態(tài)以便從中斷處恢復(fù)訓(xùn)練并進(jìn)行模型選擇。使用現(xiàn)有庫進(jìn)行擴(kuò)展對于生產(chǎn)環(huán)境強(qiáng)烈建議基于成熟的深度學(xué)習(xí)框架如Hugging Face的Transformers庫、Fairseq、OpenNMT-py進(jìn)行開發(fā)。它們提供了高度優(yōu)化、經(jīng)過充分測試的Transformer實(shí)現(xiàn)以及豐富的預(yù)訓(xùn)練模型。理解不同的變體原始的Transformer只是起點(diǎn)。后續(xù)出現(xiàn)了眾多重要變體如BERT僅使用編碼器通過掩碼語言模型進(jìn)行預(yù)訓(xùn)練。GPT系列僅使用解碼器帶掩碼自注意力進(jìn)行自回歸語言建模預(yù)訓(xùn)練。T5統(tǒng)一的編碼器-解碼器框架將所有NLP任務(wù)轉(zhuǎn)化為文本到文本的格式。Vision Transformer (ViT)將圖像分割為圖塊視為序列輸入Transformer編碼器開創(chuàng)了視覺領(lǐng)域的Transformer時(shí)代。Swin Transformer引入分層設(shè)計(jì)和滑動窗口使ViT能高效處理高分辨率圖像。注意力機(jī)制的優(yōu)化標(biāo)準(zhǔn)自注意力的計(jì)算和內(nèi)存復(fù)雜度是序列長度的平方O(n2)這是處理長文本或高分辨率圖像的瓶頸。了解并適時(shí)使用如線性注意力Linear Attention、稀疏注意力Sparse Attention、局部窗口注意力Local Window Attention或Flash Attention等優(yōu)化技術(shù)至關(guān)重要。通過本文從零件到整機(jī)的逐步拆解與實(shí)現(xiàn)你應(yīng)該對Transformer架構(gòu)有了更立體、更深入的理解。注意力機(jī)制是它的心臟但殘差連接、層歸一化、位置編碼、前饋網(wǎng)絡(luò)等組件共同構(gòu)成了其健壯的軀體。理解這些組件如何協(xié)同工作是靈活運(yùn)用乃至改進(jìn)Transformer架構(gòu)的基礎(chǔ)。建議你親手運(yùn)行文中的代碼嘗試在小數(shù)據(jù)集如機(jī)器翻譯的IWSLT或文本生成的WikiText-2上訓(xùn)練一個迷你Transformer觀察其訓(xùn)練動態(tài)和生成效果這是將知識內(nèi)化的最佳途徑。