PyTorch Building Transformer Model

Transformer is one of the most powerful models in modern machine learning.

The Transformer model is a deep learning architecture based on the self-attention mechanism. It has revolutionized the field of natural language processing (NLP) and become the foundation of modern deep learning models such as BERT, GPT, etc.

Transformer is a core architecture in modern NLP. With its powerful long-distance dependency modeling capability and efficient parallel computing advantages, it surpasses traditional Long Short-Term Memory (LSTM) networks in tasks such as language translation and text summarization.

If you don't yet understand Transformer, you can refer to:Introduction to the Transformer Model。

Building a Transformer Model with PyTorch

The steps to build a Transformer model are as follows:

1. Import the necessary libraries and modules

Import the PyTorch core library, neural network module, optimizer module, data processing tools, as well as math and object copying modules to provide support for defining the model architecture, managing data, and the training process.

import torch
import torch.nn as nn
import torch.optim as optim
import torch.utils.data as data
import math
import copy

Description:

  • torch: PyTorch's core library, used for tensor operations and automatic differentiation.

  • torch.nn: PyTorch's neural network module, containing various layers and loss functions.

  • torch.optim: Optimization algorithm module, such as Adam, SGD, etc.

  • math: Math function library, used for computing square roots, etc.

  • copy: Used for deep copying objects.

Define basic building blocks: multi-head attention, position-wise feed-forward network, positional encoding

Multi-Head AttentionBy computing the relationship between each pair of positions in the sequence through multiple "attention heads", it can capture different features and patterns of the input sequence.

The MultiHeadAttention class encapsulates the multi-head attention mechanism commonly used in Transformer models. It is responsible for splitting the input into multiple attention heads, applying attention to each head, and then combining the results, so that the model can capture various relationships in the input data at different scales and improve its expressive power.

Example

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, num_heads):
        super(MultiHeadAttention, self).__init__()
        assert d_model % num_heads == 0, "d_model must be divisible by num_heads"
       
        self.d_model = d_model    # Model dimension (e.g., 512)
        self.num_heads = num_heads # Number of attention heads (e.g., 8)
        self.d_k = d_model // num_heads # Dimension of each head (e.g., 64)
       
        # Define linear transformation layers (no bias)
        self.W_q = nn.Linear(d_model, d_model) # Query transformation
        self.W_k = nn.Linear(d_model, d_model) # Key transformation
        self.W_v = nn.Linear(d_model, d_model) # Value transformation
        self.W_o = nn.Linear(d_model, d_model) # Output transformation
       
    def scaled_dot_product_attention(self, Q, K, V, mask=None):
        """
Compute scaled dot-product attention
Input shape:
            Q: (batch_size, num_heads, seq_length, d_k)
K, V: same as Q
Output shape: (batch_size, num_heads, seq_length, d_k)
        """

        # Compute attention scores (dot product of Q and K)
        attn_scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
       
        # Apply mask (e.g., padding mask or future information mask)
        if mask is not None:
            attn_scores = attn_scores.masked_fill(mask == 0, -1e9)
       
        # Compute attention weights (softmax normalization)
        attn_probs = torch.softmax(attn_scores, dim=-1)
       
        # Weighted sum of value vectors
        output = torch.matmul(attn_probs, V)
        return output
       
    def split_heads(self, x):
        """
Split the input tensor into multiple heads
Input shape: (batch_size, seq_length, d_model)
Output shape: (batch_size, num_heads, seq_length, d_k)
        """

        batch_size, seq_length, d_model = x.size()
        return x.view(batch_size, seq_length, self.num_heads, self.d_k).transpose(1, 2)
       
    def combine_heads(self, x):
        """
Merge the outputs of multiple heads back to the original shape
Input shape: (batch_size, num_heads, seq_length, d_k)
Output shape: (batch_size, seq_length, d_model)
        """

        batch_size, _, seq_length, d_k = x.size()
        return x.transpose(1, 2).contiguous().view(batch_size, seq_length, self.d_model)
       
    def forward(self, Q, K, V, mask=None):
        """
Forward propagation
Input shape: Q/K/V: (batch_size, seq_length, d_model)
Output shape: (batch_size, seq_length, d_model)
        """

        # Linear transformation and split into multiple heads
        Q = self.split_heads(self.W_q(Q)) # (batch, heads, seq_len, d_k)
        K = self.split_heads(self.W_k(K))
        V = self.split_heads(self.W_v(V))
       
        # Compute attention
        attn_output = self.scaled_dot_product_attention(Q, K, V, mask)
       
        # Merge heads and apply output transformation
        output = self.W_o(self.combine_heads(attn_output))
        return output

Description:

  • Multi-head attention mechanism: Split the input into multiple heads, each head computes attention independently, and finally merge the results.

  • Scaled dot-product attention: Compute the dot product of queries and keys, scale it, use softmax to compute attention weights, and finally take a weighted sum of the values.

  • Mask: Used to mask invalid positions (e.g., padded parts).

Position-wise Feed-Forward Network

Example

class PositionWiseFeedForward(nn.Module):
    def __init__(self, d_model, d_ff):
        super(PositionWiseFeedForward, self).__init__()
        self.fc1 = nn.Linear(d_model, d_ff)  # First fully connected layer
        self.fc2 = nn.Linear(d_ff, d_model)  # Second fully connected layer
        self.relu = nn.ReLU()  # Activation function

    def forward(self, x):
        # Feed-forward network computation
        return self.fc2(self.relu(self.fc1(x)))

Feed-forward network:It consists of two fully connected layers and a ReLU activation function, used to further process the output of the attention mechanism.

Positional Encoding

Positional encoding is used to inject positional information of each token in the input sequence.

Use sine and cosine functions of different frequencies to generate positional encodings.

Example

class PositionalEncoding(nn.Module):
    def __init__(self, d_model, max_seq_length):
        super(PositionalEncoding, self).__init__()
        pe = torch.zeros(max_seq_length, d_model)  # Initialize the positional encoding matrix
        position = torch.arange(0, max_seq_length, dtype=torch.float).unsqueeze(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)  # Use sine function for even positions
        pe[:, 1::2] = torch.cos(position * div_term)  # Use cosine function for odd positions
        self.register_buffer('pe', pe.unsqueeze(0))  # Register as buffer
       
    def forward(self, x):
        # Add positional encoding to the input
        return x + self.pe[:, :x.size(1)]

Build the encoder block (Encoder Layer)

Encoder layer:Contains a self-attention mechanism and a feed-forward network, with each sublayer followed by residual connections and layer normalization.

Example

class EncoderLayer(nn.Module):
    def __init__(self, d_model, num_heads, d_ff, dropout):
        super(EncoderLayer, self).__init__()
        self.self_attn = MultiHeadAttention(d_model, num_heads)  # Self-attention mechanism
        self.feed_forward = PositionWiseFeedForward(d_model, d_ff)  # Feed-forward network
        self.norm1 = nn.LayerNorm(d_model)  # Layer normalization
        self.norm2 = nn.LayerNorm(d_model)
        self.dropout = nn.Dropout(dropout)  # Dropout
       
    def forward(self, x, mask):
        # Self-attention mechanism
        attn_output = self.self_attn(x, x, x, mask)
        x = self.norm1(x + self.dropout(attn_output))  # Residual connection and layer normalization
       
        # Feed-forward network
        ff_output = self.feed_forward(x)
        x = self.norm2(x + self.dropout(ff_output))  # Residual connection and layer normalization
        return x

Build the decoder module

Decoder layer:Contains a self-attention mechanism, a cross-attention mechanism, and a feed-forward network, with each sublayer followed by residual connections and layer normalization.

Example

class DecoderLayer(nn.Module):
    def __init__(self, d_model, num_heads, d_ff, dropout):
        super(DecoderLayer, self).__init__()
        self.self_attn = MultiHeadAttention(d_model, num_heads)  # Self-attention mechanism
        self.cross_attn = MultiHeadAttention(d_model, num_heads)  # Cross-attention mechanism
        self.feed_forward = PositionWiseFeedForward(d_model, d_ff)  # Feed-forward network
        self.norm1 = nn.LayerNorm(d_model)  # Layer normalization
        self.norm2 = nn.LayerNorm(d_model)
        self.norm3 = nn.LayerNorm(d_model)
        self.dropout = nn.Dropout(dropout)  # Dropout
       
    def forward(self, x, enc_output, src_mask, tgt_mask):
        # Self-attention mechanism
        attn_output = self.self_attn(x, x, x, tgt_mask)
        x = self.norm1(x + self.dropout(attn_output))  # Residual connection and layer normalization
       
        # Cross-attention mechanism
        attn_output = self.cross_attn(x, enc_output, enc_output, src_mask)
        x = self.norm2(x + self.dropout(attn_output))  # Residual connection and layer normalization
       
        # Feed-forward network
        ff_output = self.feed_forward(x)
        x = self.norm3(x + self.dropout(ff_output))  # Residual connection and layer normalization
        return x

Build the complete Transformer model

Example

class Transformer(nn.Module):
    def __init__(self, src_vocab_size, tgt_vocab_size, d_model, num_heads, num_layers, d_ff, max_seq_length, dropout):
        super(Transformer, self).__init__()
        self.encoder_embedding = nn.Embedding(src_vocab_size, d_model)  # Encoder word embedding
        self.decoder_embedding = nn.Embedding(tgt_vocab_size, d_model)  # Decoder word embedding
        self.positional_encoding = PositionalEncoding(d_model, max_seq_length)  # Positional encoding

        # Encoder and decoder layers
        self.encoder_layers = nn.ModuleList([EncoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers)])
        self.decoder_layers = nn.ModuleList([DecoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers)])

        self.fc = nn.Linear(d_model, tgt_vocab_size)  # Final fully connected layer
        self.dropout = nn.Dropout(dropout)  # Dropout

    def generate_mask(self, src, tgt):
        # Source mask: mask padding tokens (assume padding index is 0)
        # Shape: (batch_size, 1, 1, seq_length)
        src_mask = (src != 0).unsqueeze(1).unsqueeze(2)
   
        # Target mask: mask padding tokens and future information
        # Shape: (batch_size, 1, seq_length, 1)
        tgt_mask = (tgt != 0).unsqueeze(1).unsqueeze(3)
        seq_length = tgt.size(1)
        # Generate an upper triangular matrix mask to prevent seeing future information during decoding
        nopeak_mask = (1 - torch.triu(torch.ones(1, seq_length, seq_length), diagonal=1)).bool()
        tgt_mask = tgt_mask & nopeak_mask  # Combine the padding mask and the future information mask
        return src_mask, tgt_mask

    def forward(self, src, tgt):
        # Generate masks
        src_mask, tgt_mask = self.generate_mask(src, tgt)
       
        # Encoder part
        src_embedded = self.dropout(self.positional_encoding(self.encoder_embedding(src)))
        enc_output = src_embedded
        for enc_layer in self.encoder_layers:
            enc_output = enc_layer(enc_output, src_mask)
       
        # Decoder part
        tgt_embedded = self.dropout(self.positional_encoding(self.decoder_embedding(tgt)))
        dec_output = tgt_embedded
        for dec_layer in self.decoder_layers:
            dec_output = dec_layer(dec_output, enc_output, src_mask, tgt_mask)
       
        # Final output
        output = self.fc(dec_output)
        return output

Description:

  • Transformer modelContains encoder and decoder parts, each composed of multiple stacked layers.

  • Mask generationUsed to mask invalid positions and future information.

  • Forward propagationPasses through the encoder and decoder in sequence, and finally outputs through a fully connected layer.

Model initialization parameter description:

class Transformer(nn.Module):
    def __init__(
        self, 
        src_vocab_size,  # 源语言词汇表大小(如英文单词数)
        tgt_vocab_size,  # 目标语言词汇表大小(如中文单词数)
        d_model=512,     # 模型维度(每个词向量的长度)
        num_heads=8,     # 多头注意力的头数
        num_layers=6,    # 编码器/解码器的堆叠层数
        d_ff=2048,       # 前馈网络隐藏层维度
        max_seq_length=100, # 最大序列长度(用于位置编码)
        dropout=0.1      # Dropout概率
    ):

Train the PyTorch Transformer model

Use random data to train the model, compute the loss, and update parameters.

Example

# Hyperparameters
src_vocab_size = 5000  # Source vocabulary size
tgt_vocab_size = 5000  # Target vocabulary size
d_model = 512  # Model dimension
num_heads = 8  # Number of attention heads
num_layers = 6  # Number of encoder and decoder layers
d_ff = 2048  # Inner dimension of the feed-forward network
max_seq_length = 100  # Maximum sequence length
dropout = 0.1  # Dropout probability

# Initialize the model
transformer = Transformer(src_vocab_size, tgt_vocab_size, d_model, num_heads, num_layers, d_ff, max_seq_length, dropout)

# Generate random data
src_data = torch.randint(1, src_vocab_size, (64, max_seq_length))  # Source sequence
tgt_data = torch.randint(1, tgt_vocab_size, (64, max_seq_length))  # Target sequence

# Define loss function and optimizer
criterion = nn.CrossEntropyLoss(ignore_index=0)  # Ignore the loss of the padding part
optimizer = optim.Adam(transformer.parameters(), lr=0.0001, betas=(0.9, 0.98), eps=1e-9)

# Training loop
transformer.train()
for epoch in range(100):
    optimizer.zero_grad()  # Clear gradients to prevent accumulation
   
    # Remove the last word when inputting the target sequence (used to predict the next word)
    output = transformer(src_data, tgt_data[:, :-1])  
   
    # When computing the loss, the target sequence starts from the second word (i.e., predict the next word)
    # Output shape: (batch_size, seq_length-1, tgt_vocab_size)
    # Target shape: (batch_size, seq_length-1)
    loss = criterion(
        output.contiguous().view(-1, tgt_vocab_size),
        tgt_data[:, 1:].contiguous().view(-1)
    )
   
    loss.backward()        # Backpropagation
    optimizer.step()       # Update parameters
    print(f"Epoch: {epoch+1}, Loss: {loss.item()}")

Model evaluation

Evaluation process:Compute the loss on validation data to evaluate model performance.

Example

transformer.eval()
# Generate validation data
val_src_data = torch.randint(1, src_vocab_size, (64, max_seq_length))
val_tgt_data = torch.randint(1, tgt_vocab_size, (64, max_seq_length))
# Assume the input is a batch of English sentences and their corresponding Chinese translations (converted to indices)
# Example data:
# src_data: [[3, 14, 25, ..., 0, 0], ...] # English sentences (0 is the padding token)
# tgt_data: [[5, 20, 36, ..., 0, 0], ...] # Chinese translations (0 is the padding token)
# Note: In practical applications, text needs preprocessing such as tokenization, encoding, and padding
with torch.no_grad():
    val_output = transformer(val_src_data, val_tgt_data[:, :-1])
    val_loss = criterion(val_output.contiguous().view(-1, tgt_vocab_size), val_tgt_data[:, 1:].contiguous().view(-1))
    print(f"Validation Loss: {val_loss.item()}")
Other extensions