PyTorchがGPUを利用できる場合はGPUを使用し、利用できない場合はCPUを使用する。今回はモデルの構造を理解することが目的なので、GPUがなくても実行できる。
import torch
import torch.nn as nn
import torch.nn.functional as F
device = torch.device(
"cuda" if torch.cuda.is_available() else "cpu"
)
print("使用デバイス:", device)
Encoder-Decoder型とDecoder-only型の両方でAttentionを使用する。
class MultiHeadAttention(nn.Module):
def __init__(self, hidden_dim=32, num_heads=4):
super().__init__()
assert hidden_dim % num_heads == 0
self.hidden_dim = hidden_dim
self.num_heads = num_heads
self.head_dim = hidden_dim // num_heads
# Q, K, V
self.W_Q = nn.Linear(hidden_dim, hidden_dim)
self.W_K = nn.Linear(hidden_dim, hidden_dim)
self.W_V = nn.Linear(hidden_dim, hidden_dim)
self.W_O = nn.Linear(hidden_dim, hidden_dim)
def forward(self, query, key=None, value=None, causal=False):
# Self-Attentionの場合
if key is None:
key = query
if value is None:
value = key
B, Nq, D = query.shape
Nk = key.shape[1]
# --------------------------------------------
# Q / K / V
# --------------------------------------------
Q = self.W_Q(query)
K = self.W_K(key)
V = self.W_V(value)
# --------------------------------------------
# Headに分割
# --------------------------------------------
Q = Q.view(B, Nq, self.num_heads, self.head_dim)
K = K.view(B, Nk, self.num_heads, self.head_dim)
V = V.view(B, Nk, self.num_heads, self.head_dim)
Q = Q.transpose(1, 2)
K = K.transpose(1, 2)
V = V.transpose(1, 2)
# --------------------------------------------
# Attention score
# --------------------------------------------
scores = torch.matmul(
Q,
K.transpose(-2, -1)
) / (self.head_dim ** 0.5)
# --------------------------------------------
# DecoderのSelf-Attentionでは
# 未来のトークンを見ない
# --------------------------------------------
if causal:
mask = torch.triu(
torch.ones(Nq, Nk, device=query.device),
diagonal=1
).bool()
scores = scores.masked_fill(
mask,
float("-inf")
)
# --------------------------------------------
# Softmax
# --------------------------------------------
weights = F.softmax(scores, dim=-1)
# --------------------------------------------
# Attention × V
# --------------------------------------------
context = torch.matmul(weights, V)
# --------------------------------------------
# Headを結合
# --------------------------------------------
context = context.transpose(1, 2)
context = context.contiguous().view(
B,
Nq,
D
)
return self.W_O(context)
Encoderは入力全体を処理して、「入力全体の表現」を作る。Encoderでは通常、入力系列の各トークンが他の入力トークンを自由に参照できる。そのため、ここではcausal=Trueにしない。
class Encoder(nn.Module):
def __init__(
self,
hidden_dim=32,
num_heads=4
):
super().__init__()
self.self_attention = MultiHeadAttention(
hidden_dim,
num_heads
)
# Attentionで得られた表現をさらに変換する
# Feed Forward Network
self.ffn = nn.Sequential(
nn.Linear(
hidden_dim,
hidden_dim * 4
),
nn.ReLU(),
nn.Linear(
hidden_dim * 4,
hidden_dim
)
)
def forward(self, x):
# ----------------------------------------------------
# Encoder Self-Attention
# ----------------------------------------------------
#
# Q / K / Vはすべて入力xから作られる。
x = x + self.self_attention(x)
# ----------------------------------------------------
# Feed Forward
# ----------------------------------------------------
x = x + self.ffn(x)
return x
Encoder-Decoder型では、Encoder → Encoder output → Decoder という構造になる。Decoderには、① Decoder Self-Attentionと② Encoder-Decoder Cross-Attentionの2種類のAttentionがある。
class EncoderDecoderTransformer(nn.Module):
def __init__(
self,
hidden_dim=32,
num_heads=4
):
super().__init__()
# ----------------------------------------------------
# Encoder
# ----------------------------------------------------
self.encoder = Encoder(
hidden_dim,
num_heads
)
# ----------------------------------------------------
# Decoder Self-Attention
# ----------------------------------------------------
self.decoder_self_attention = MultiHeadAttention(
hidden_dim,
num_heads
)
# ----------------------------------------------------
# Cross-Attention
# ----------------------------------------------------
#
# ここがEncoder-Decoder型の重要な部分。
#
# Query = Decoder
# Key = Encoder
# Value = Encoder
#
# つまりDecoderがEncoderの出力を参照する。
self.cross_attention = MultiHeadAttention(
hidden_dim,
num_heads
)
# ----------------------------------------------------
# Feed Forward
# ----------------------------------------------------
self.ffn = nn.Sequential(
nn.Linear(
hidden_dim,
hidden_dim * 4
),
nn.ReLU(),
nn.Linear(
hidden_dim * 4,
hidden_dim
)
)
def forward(self, source, target):
# ====================================================
# 4-1. Encoder
# ====================================================
encoder_output = self.encoder(source)
# ====================================================
# 4-2. Decoder Self-Attention
# ====================================================
#
# Decoder自身の過去の出力を参照する。
#
# causal=Trueなので、
# 未来のトークンは参照できない。
x = target + self.decoder_self_attention(
target,
causal=True
)
# ====================================================
# 4-3. Cross-Attention
# ====================================================
#
# DecoderのQueryから、
# EncoderのKey / Valueを参照する。
#
# Q = Decoder
# K = Encoder
# V = Encoder
#
# これによってDecoderは、
# 「入力側のどこを見るべきか」
# を判断できる。
x = x + self.cross_attention(
query=x,
key=encoder_output,
value=encoder_output
)
# ====================================================
# 4-4. Feed Forward
# ====================================================
x = x + self.ffn(x)
return x, encoder_output
GPT型ではEncoderが存在しない。入力をDecoderブロック自身で処理する。Input → Decoder → Output 。Decoder内部ではCausal Self-Attentionを使用する。Encoder-Decoder型に存在した「Encoder → Decoder Cross-Attention」は存在しない。
class DecoderOnlyTransformer(nn.Module):
def __init__(
self,
hidden_dim=32,
num_heads=4
):
super().__init__()
self.self_attention = MultiHeadAttention(
hidden_dim,
num_heads
)
self.ffn = nn.Sequential(
nn.Linear(
hidden_dim,
hidden_dim * 4
),
nn.ReLU(),
nn.Linear(
hidden_dim * 4,
hidden_dim
)
)
def forward(self, x):
# ====================================================
# 5-1. Causal Self-Attention
# ====================================================
#
# Q / K / Vはすべてxから作る。
#
# ただしcausal=Trueなので、
# 未来のトークンを見ることはできない。
x = x + self.self_attention(
x,
causal=True
)
# ====================================================
# 5-2. Feed Forward
# ====================================================
x = x + self.ffn(x)
return x
ここでは、実際の単語ではなく、Embedding後のベクトルを直接入力する。例えば、hidden_dim = 32 ならば各トークンが32次元のベクトルになっているとして乱数でダミーを作成。
BATCH_SIZE = 1
SOURCE_LEN = 5
TARGET_LEN = 4
HIDDEN_DIM = 32
NUM_HEADS = 4
source = torch.randn(
BATCH_SIZE,
SOURCE_LEN,
HIDDEN_DIM
).to(device)
target = torch.randn(
BATCH_SIZE,
TARGET_LEN,
HIDDEN_DIM
).to(device)
decoder_input = torch.randn(
BATCH_SIZE,
SOURCE_LEN,
HIDDEN_DIM
).to(device)
print("\n==============================")
print("Encoder-Decoder型")
print("==============================")
encoder_decoder = EncoderDecoderTransformer(
HIDDEN_DIM,
NUM_HEADS
).to(device)
with torch.no_grad():
output_ed, encoder_output = encoder_decoder(
source,
target
)
print("\nEncoder-Decoderの出力")
print("Source :", source.shape)
print("Encoder output :", encoder_output.shape)
print("Target :", target.shape)
print("Decoder output :", output_ed.shape) print("\n==============================")
print("Decoder-only型(GPT型)")
print("==============================")
decoder_only = DecoderOnlyTransformer(
HIDDEN_DIM,
NUM_HEADS
).to(device)
with torch.no_grad():
output_do = decoder_only(
decoder_input
)
print("\nDecoder-onlyの出力")
print("Input :", decoder_input.shape)
print("Decoder output :", output_do.shape)Mathematics is the language with which God has written the universe.