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 還有其他不少的著名使用案例,如: