Transformer 是一種完全基於注意力機制(attention mechanism)的神經網路架構,最初由 Vaswani 等人於 2017 年提出,並被廣泛應用於自然語言處理、語音辨識、影像辨識等需要捕捉長距離依賴關係的任務。雖然在應用面上與 LSTM 相似,但不同的是,Transformer 不依賴遞迴結構,而是透過自注意力機制(self-attention)讓序列中的每個位置都能直接與其他所有位置互動,因此在平行運算上具有顯著優勢。Transformer 的完整架構包含 encoder 與 decoder,本篇章將主要基於原始版本進行介紹。Encoder 部分的介紹如下。
Transformer encoder 的核心元件,是由多頭自注意力(Multi-Head Self-Attention,MHSA)與前饋網路(Feed-Forward Network,FFN)所組成。對於給定的輸入序列 X,我們首先分別透過三個線性投影,得到查詢(Query)、鍵(Key)、值(Value)三組矩陣:
Q = XWQ,K = XWK,V = XWV
接著計算注意力輸出:
Attention(Q, K, V) = softmax(QKT / dk1/2)V
其中 dk 為 K 的維度,除以其平方根是為了避免乘積過大,導致 softmax 梯度消失。多頭注意力則是將上述運算沿特徵維度平均分割成多份來執行,每份使用不同的投影矩陣,最後將結果並排後再做一次線性投影:
headi = Attention(XWQi, XWKi, XWVi),
MultiHead(X) = Concat(head1, ..., headh)WO
其中,MultiHead(X) 代表 Q、K、V 皆來自 X。而一個完整的 encoder block,還包含殘差連接(residual connection)與層正規化(Layer Normalization,LN),以及一個逐位置的前饋網路:
Z = LN(X + MultiHead(X))
Y = LN(Z + FFN(Z))
其中,FFN 通常是兩層全連接層加上激勵函數:
FFN(Z) = ReLU(ZW1 + b1)W2 + b2
此外,由於 Transformer 沒有遞迴結構,無法從運算過程中得知各位置的順序資訊,因此需要在輸入端加入位置編碼(positional encoding)。原始論文使用正弦與餘弦函數:
PE[pos, 2i] = sin(pos / 100002i/d)
PE[pos, 2i+1] = cos(pos / 100002i/d)
其中 pos 為時間軸的索引,d 為特徵的總維度,i 為 0 到 d/2 之間的整數。依此填完的 PE 矩陣,其 shape 會與 X 相同,因此我們只要把 X 替換為 X + PE,再餵入 Transformer 即可。
上述的詳細運算步驟,若要用 PyTorch class 的方式表示,則寫法如下:
import torch import torch.nn as nn class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads): super().__init__() 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) def forward(self, x): N, T, d = x.shape Q = self.W_q(x).view(N, T, self.num_heads, self.d_k).transpose(1, 2) # (N, h, T, d_k) K = self.W_k(x).view(N, T, self.num_heads, self.d_k).transpose(1, 2) V = self.W_v(x).view(N, T, self.num_heads, self.d_k).transpose(1, 2) scores = Q @ K.transpose(-2, -1) / self.d_k ** 0.5 # (N, h, T, T) attn = torch.softmax(scores, dim=-1) out = (attn @ V).transpose(1, 2).contiguous().view(N, T, d) # (N, T, d) return self.W_o(out) class PositionalEncoding(nn.Module): def __init__(self, d_model): super().__init__() self.d_model = d_model def forward(self, x): T = x.size(1) pe = torch.zeros(T, self.d_model, device=x.device) pos = torch.arange(0, T, device=x.device).unsqueeze(1) div = torch.pow(10000.0, torch.arange(0, self.d_model, 2, device=x.device) / self.d_model) pe[:, 0::2] = torch.sin(pos / div) pe[:, 1::2] = torch.cos(pos / div) return x + pe.unsqueeze(0) class TransformerEncoderBlock(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout=0.1): super().__init__() self.attn = MultiHeadAttention(d_model, num_heads) self.ff = nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model) ) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.dropout = nn.Dropout(dropout) def forward(self, x): attn_out = self.attn(x) x = self.norm1(x + self.dropout(attn_out)) ff_out = self.ff(x) x = self.norm2(x + self.dropout(ff_out)) return x if __name__ == '__main__': d_model = 64 num_heads = 4 d_ff = 128 T = 10 N = 2 pe = PositionalEncoding(d_model) block = TransformerEncoderBlock(d_model, num_heads, d_ff) x = torch.randn(N, T, d_model) x = pe(x) out = block(x) print('Input shape:', x.shape) print('Output shape:', out.shape)Transformer 的 decoder block 與 encoder block 結構相似,但多了一個交叉注意力(Cross-Attention)層,其使用的 Q 來自 decoder 自身,K 與 V 則來自 encoder 的輸出。一個完整的 decoder block,依序包含三個子層:
Z1 = LN(Xdec + MultiHead(Xdec))
Z2 = LN(Z1 + CrossAttn(Z1, Xenc))
Y = LN(Z2 + FFN(Z2))
其中 CrossAttn(Z1, Xenc) 表示 Q 來自 Z1,K 與 V 來自 encoder 輸出 Xenc:
CrossAttn(Z1, Xenc) = Attention(Z1WQ, XencWK, XencWV)
上述步驟的 PyTorch class 的寫法如下:
class MultiHeadCrossAttention(nn.Module): def __init__(self, d_model, num_heads): super().__init__() 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) def forward(self, x_dec, x_enc): N, T_dec, d = x_dec.shape T_enc = x_enc.size(1) Q = self.W_q(x_dec).view(N, T_dec, self.num_heads, self.d_k).transpose(1, 2) K = self.W_k(x_enc).view(N, T_enc, self.num_heads, self.d_k).transpose(1, 2) V = self.W_v(x_enc).view(N, T_enc, self.num_heads, self.d_k).transpose(1, 2) scores = Q @ K.transpose(-2, -1) / self.d_k ** 0.5 attn = torch.softmax(scores, dim=-1) out = (attn @ V).transpose(1, 2).contiguous().view(N, T_dec, d) return self.W_o(out) class TransformerDecoderBlock(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout=0.1): super().__init__() self.self_attn = MultiHeadAttention(d_model, num_heads) self.cross_attn = MultiHeadCrossAttention(d_model, num_heads) self.ff = nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model) ) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.norm3 = nn.LayerNorm(d_model) self.dropout = nn.Dropout(dropout) def forward(self, x_dec, x_enc): x = self.norm1(x_dec + self.dropout(self.self_attn(x_dec))) x = self.norm2(x + self.dropout(self.cross_attn(x, x_enc))) x = self.norm3(x + self.dropout(self.ff(x))) return x上述的 encoder 和 decoder 都各自可以堆疊多層,此時 decoder 的每一層都是從 encoder 最後一層的輸出取 K 和 V,而不是對應層的輸出;也就是說 encoder 跑完所有層之後,只會把最終輸出傳給 decoder 的每一層使用。
這個完整的 encoder-decoder 架構,通常會用於生成類的問題,例如將文字從一種語言翻譯成另一種語言。需要注意的是,由於在生成任務的推論過程中,是無法看到未來資訊的;而若要在訓練階段模擬此種行為,則會需要在 decoder 端,利用遮罩(mask)來避免之;對於 encoder 端,若你在特定應用中,不希望讓每個位置取用到距離它太遠的資訊,則也會使用 mask 來處理。關於 mask 的細節和使用方式,暫且不列入本篇章的介紹範圍內。
與 LSTM 相仿,PyTorch 也有內建的 class,實務上不需要自己寫。而如果只是要用 Transformer 來進行分類任務,則通常只會取用 encoder 的部分,如下:
import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu') train_set = datasets.MNIST(root='./data', train=True, download=True, transform=transforms.ToTensor()) test_set = datasets.MNIST(root='./data', train=False, download=True, transform=transforms.ToTensor()) print('Data shapes:', train_set.data.shape, test_set.data.shape) train_loader = DataLoader(train_set, batch_size=128, shuffle=True) test_loader = DataLoader(test_set, batch_size=128) class PositionalEncoding(nn.Module): def __init__(self, d_model): super().__init__() self.d_model = d_model def forward(self, x): T = x.size(1) pe = torch.zeros(T, self.d_model, device=x.device) pos = torch.arange(0, T, device=x.device).unsqueeze(1) div = torch.pow(10000.0, torch.arange(0, self.d_model, 2, device=x.device) / self.d_model) pe[:, 0::2] = torch.sin(pos / div) pe[:, 1::2] = torch.cos(pos / div) return x + pe.unsqueeze(0) class TransformerClassifier(nn.Module): def __init__(self, d_model=28, num_heads=2, d_ff=64, dropout=0.1): super().__init__() self.pe = PositionalEncoding(d_model) self.encoder = nn.TransformerEncoderLayer( d_model=d_model, nhead=num_heads, dim_feedforward=d_ff, dropout=dropout, batch_first=True ) self.linear = nn.Linear(d_model, 10) def forward(self, x): x = x.squeeze(1) x = self.pe(x) x = self.encoder(x) return self.linear(x[:, -1, :]) model = TransformerClassifier().to(device) optim = torch.optim.Adam(model.parameters()) criterion = nn.CrossEntropyLoss() model.train() for i in range(10): print(f'Epoch {i+1}/10') for X_batch, y_batch in train_loader: loss = criterion(model(X_batch.to(device)), y_batch.to(device)) optim.zero_grad() loss.backward() optim.step() model.eval() correct = sum( (model(X.to(device)).argmax(dim=1) == y.to(device)).sum().item() for X, y in test_loader ) print('Accuracy: {:.2f}%'.format(100 * correct / len(test_set)))在上述範例中,我們仿照 LSTM 篇章的做法,只取最後一個時間步的輸出來進行後續的分類,你當然還是可以嘗試其他做法。
Transformer 還有其他不少的著名使用案例,如:
- BERT (Bidirectional Encoder Representations from Transformers):由 Google 於 2018 年提出,僅使用 encoder 部分,透過將訓練語句的一部分遮住,來訓練 Transformer encoder 預測被遮住的內容,使模型學到豐富的雙向語境表示,並讓接續使用的研究者,可以再針對下游任務進行 fine-tuning。
- GPT (Generative Pre-trained Transformer):由 OpenAI 提出,僅使用 decoder 部分,以自迴歸方式逐步預測下一個 token,並在透過大規模語料預訓練後,可用於文字生成、問答、摘要等多種任務。GPT 系列歷經多個版本的演進,ChatGPT 即以此為基礎。
- Whisper:由 OpenAI 於 2022 年提出的語音辨識模型,使用完整的 encoder-decoder 架構,encoder 處理音訊的頻譜特徵,decoder 自迴歸地生成對應的文字,並透過大規模多語言資料訓練,在多種語言的辨識任務上表現穩健。
- Vision Transformer (ViT):由 Google 於 2020 年提出,將影像切成固定大小的 patch 並展平後視為序列,送入 Transformer encoder 進行影像分類,證明了純 Transformer 架構在電腦視覺任務上同樣具有競爭力。