Hello! It's great to see you again.
In our last two lessons, we meticulously constructed the core components of the Transformer: the EncoderBlock and the DecoderBlock. You now understand how the encoder builds a rich representation of an input sequence and how the decoder uses masked self-attention and cross-attention to generate an output sequence conditioned on the encoder's memory.
We are now at the culminating point of this module. With all the building blocks in hand, we will assemble the full, end-to-end encoder-decoder Transformer model. This architecture is the foundation for countless state-of-the-art models in NLP and, as we'll see later, in audio processing for tasks like speech-to-text.
Your learning outcome for this lesson is to: Construct a full encoder-decoder Transformer model in PyTorch for sequence-to-sequence tasks.
We will integrate the token embeddings, positional encodings, encoder stack, decoder stack, and the final output layer into a single, cohesive nn.Module.
1. The Complete Picture
Let's start by visualizing our goal. We are building the entire structure shown in the diagram below. You've already built the components inside the "N x" boxes. Today, we connect everything around them.

A complete sequence-to-sequence Transformer consists of the following parts, executed in order:
- Source Embeddings: Convert the input tokens of the source sequence (e.g., German text) into vectors.
- Positional Encoding: Add positional information to the source embeddings.
- Encoder Stack: A stack of N
EncoderBlocks processes the embeddings to create a final contextual representation (often calledmemory). - Target Embeddings: Convert the input tokens of the target sequence (e.g., English text, shifted right) into vectors.
- Positional Encoding: Add positional information to the target embeddings.
- Decoder Stack: A stack of M
DecoderBlocks takes the target embeddings and the encoder'smemoryto generate an output sequence. - Final Projection: A linear layer followed by a softmax function converts the decoder's output vectors into probability distributions over the target vocabulary.
2. Building the Transformer Class in PyTorch
We will now construct a Transformer class that encapsulates this entire architecture. Your experience with object-oriented programming in Python and defining nn.Modules will be very useful here.
Our main reference will be a superbly written blog post from pi-tau, which provides a clean, from-scratch implementation. We'll build our understanding around its code structure.
An even more annotated Transformer
This article, "An even more annotated Transformer", provides one of the clearest from-scratch implementations available. We will focus on the final Transformer class, which brings together all the sub-modules we've discussed.
Please read the section titled "TRANSFORMER". Pay close attention to: The __init__ method: Notice how it initializes the TokenEmbedding layers for source and target, the encoder_stack and decoder_stack using nn.ModuleList, and the final normalization layers (enc_norm, dec_norm). Observe the logic for sharing embedding weights. The encode, decode, and forward methods: Trace the data flow. See how forward calls encode to get the mem tensor, then passes mem to the decode method. Finally, see how the output of the decoder is passed to a final linear layer to produce scores.
Let's break down the implementation based on that article.
The __init__ Method: Assembling the Components
The __init__ method is responsible for creating all the layers of our model.
import torch
import torch.nn as nn
import numpy as np
# Assuming EncoderBlock and DecoderBlock are defined as in previous lessons
# from encoder_block import EncoderBlock
# from decoder_block import DecoderBlock
# from token_embedding import TokenEmbedding
class Transformer(nn.Module):
def __init__(self,
src_vocab_size,
tgt_vocab_size,
max_seq_len,
d_model,
n_heads,
n_enc,
n_dec,
dim_mlp,
dropout):
super().__init__()
# --- 1. Embeddings & Positional Encoding ---
# Positional embeddings are shared between encoder and decoder
pos_embed = nn.Parameter(torch.randn(max_seq_len, d_model))
# Word embeddings
self.src_word_embed = nn.Embedding(src_vocab_size, d_model)
self.tgt_word_embed = nn.Embedding(tgt_vocab_size, d_model)
# Scale factor for embeddings as per the paper
self.scale = np.sqrt(d_model)
self.src_embed_dropout = nn.Dropout(dropout)
self.tgt_embed_dropout = nn.Dropout(dropout)
# We will manually add positional encodings in the forward pass
self.register_buffer("positions", torch.arange(max_seq_len).unsqueeze(0))
self.pos_embed = pos_embed
# --- 2. Encoder & Decoder Stacks ---
self.encoder_stack = nn.ModuleList([
EncoderBlock(d_model, n_heads, dim_mlp, dropout) for _ in range(n_enc)
])
self.decoder_stack = nn.ModuleList([
DecoderBlock(d_model, n_heads, dim_mlp, dropout) for _ in range(n_dec)
])
# --- 3. Final Layers ---
# Using Pre-LN architecture requires final normalization
self.enc_norm = nn.LayerNorm(d_model)
self.dec_norm = nn.LayerNorm(d_model)
# Final projection layer to map decoder output to vocab size
self.final_proj = nn.Linear(d_model, tgt_vocab_size)
# Optional: Weight sharing between target embedding and final projection
self.final_proj.weight = self.tgt_word_embed.weight
def encode(self, src, src_mask):
# src shape: (batch_size, src_seq_len)
# src_mask shape: (batch_size, 1, 1, src_seq_len) for attention
# Create embeddings and add positional encoding
src_emb = self.src_word_embed(src) * self.scale
pos = self.pos_embed[self.positions[:, :src.shape[1]], :]
z = self.src_embed_dropout(src_emb + pos)
# Pass through encoder stack
for encoder in self.encoder_stack:
z = encoder(z, src_mask)
return self.enc_norm(z)
def decode(self, tgt, mem, tgt_mask, mem_mask):
# tgt shape: (batch_size, tgt_seq_len)
# mem shape: (batch_size, src_seq_len, d_model)
# Create embeddings and add positional encoding
tgt_emb = self.tgt_word_embed(tgt) * self.scale
pos = self.pos_embed[self.positions[:, :tgt.shape[1]], :]
z = self.tgt_embed_dropout(tgt_emb + pos)
# Pass through decoder stack
for decoder in self.decoder_stack:
z = decoder(z, mem, tgt_mask, mem_mask)
return self.dec_norm(z)
def forward(self, src, tgt, src_mask, tgt_mask):
# Create memory from the encoder
mem = self.encode(src, src_mask)
# Decode using memory and the target sequence
out = self.decode(tgt, mem, tgt_mask, src_mask) # Note: src_mask is used as mem_mask
# Project to vocabulary space
tgt_scores = self.final_proj(out)
return tgt_scores
Note: This code snippet synthesizes the core ideas from the pi-tau article into a slightly different structure for clarity, manually handling the embedding and positional encoding steps. The core logic remains the same.
3. Practical Implementation with torch.nn.Transformer
While building from scratch is invaluable for understanding, in practice, you'll often use PyTorch's highly optimized, built-in modules. PyTorch provides torch.nn.Transformer, which encapsulates the entire encoder-decoder stack.
You are still responsible for:
- Creating the word embeddings.
- Creating and adding the positional encodings.
- Creating the final linear projection layer.
- Generating the appropriate masks.
This approach lets you focus on the data flow and task-specific parts while leveraging a fast, tested implementation of the core architecture. Let's see this in action.
Pytorch Transformers for Machine Translation
The YouTuber Aladdin Persson provides a fantastic tutorial on using the built-in torch.nn.Transformer for a machine translation task. This video demonstrates the practical, high-level way to assemble a Transformer model.
Please watch the following segments: Model Definition (00:05:41 - 00:09:44): See how the Transformer class is defined. Notice it still has embeddings, but the core is a single nn.Transformer instance. All the hyperparameters (heads, layers, etc.) are passed directly to it. Forward Pass (00:09:40 - 00:15:43): This is crucial. Observe how the source and target masks are created (make_source_mask, generate_square_subsequent_mask). Then, see how the embeddings are prepared and passed, along with the masks, to the self.transformer module in a single call. Training Loop (00:21:21 - 00:28:54): Skim this section to see how the model is used during training. Pay attention to how the target sequence is shifted (target[:-1, :]) for teacher forcing and how the loss function is configured with ignore_index to handle padding.
The key insight here is that nn.Transformer handles the internal looping through encoder and decoder stacks for you. Your job is to prepare the inputs (embeddings + positional encodings) and masks correctly.
4. Training and Inference
With our model constructed, how do we use it?
Training (Teacher Forcing):
During training, we have access to the complete ground-truth target sequence. To make training stable and parallelizable, we use a technique called teacher forcing.
- The entire source sequence is fed to the encoder.
- The entire target sequence, shifted one position to the right and with an
[SOS](start-of-sequence) token prepended, is fed to the decoder. - The model's predictions at each position are then compared against the original target sequence (with an
[EOS]end-of-sequence token) to calculate the loss.
This is why the forward pass takes both src and tgt as input.
Inference (Greedy/Beam Search):
During inference, we don't have the target sequence. We must generate it one token at a time.
- Feed the source sequence to the encoder to get the
memory. - Start the decoder with just the
[SOS]token. - Pass this through the decoder, get the logits for the next word, and pick the most likely one (this is greedy decoding).
- Append this new token to your target sequence and repeat from step 3 until an
[EOS]token is generated or a max length is reached.
The pi-tau article you read earlier contains a great implementation of this greedy_decode method.
An even more annotated Transformer
Let's revisit the pi-tau article to see how inference is handled.
Read the "INFERENCE" section and study the greedy_decode method. Notice the auto-regressive loop (for _ in range(max_len-1)) where the model's own output (next_idx) is concatenated to the input for the next step (tgt = torch.concat((tgt, next_idx), dim=1)). This is the fundamental difference between training and inference.
Conclusion
Congratulations! You have now assembled a complete Transformer model, from the individual attention heads to the full encoder-decoder architecture. You've seen both a from-scratch implementation that reveals the inner workings and a practical approach using PyTorch's optimized modules.
Key Takeaways:
- A full Transformer model is an
nn.Modulethat encapsulates embedding layers, positional encoding, an encoder stack (nn.ModuleList), a decoder stack (nn.ModuleList), and a final projection layer. - The
forwardpass orchestrates the entire sequence-to-sequence process: encoding the source to creatememory, then decoding the target conditioned on thatmemory. - During training, teacher forcing is used, where the ground-truth target is fed to the decoder.
- During inference, an auto-regressive loop is used to generate the output one token at a time.
- PyTorch's
nn.Transformerprovides a convenient and efficient way to implement the core architecture, abstracting away the encoder/decoder stacks.
Preview of the Next Lesson:
You've now reached the pinnacle of our module on sequence modeling. The Transformer is a powerful tool for transforming one sequence into another. In the next module, we'll shift our focus from transformation to generation. We'll explore the foundational principles of generative models, starting with Generative Adversarial Networks (GANs). This will lay the groundwork for understanding advanced audio synthesis models that can generate speech from scratch.