Skip to main content
Create your own

Building a Decoder-Only Transformer in PyTorch

Introduction

Welcome to Module 3: "LLM Internals." In the previous module, we successfully set up a local environment, profiled a model's performance, and even served it behind a web API. Throughout that process, we treated the model from Hugging Face as a complete black box. Our goal now is to pry open that box and understand its inner workings from first principles.

This lesson is the first step on that journey. We will tackle the learning outcome: Implement a minimal decoder-only transformer from scratch in PyTorch, including embedding, attention, FFN, and output layers.

By the end of this lesson, you will have built the fundamental blueprint of nearly all modern large language models, from GPT-3 to Llama 3. This hands-on implementation is crucial, as it will form the basis for our upcoming lessons where we will calculate the exact memory and compute requirements of these architectures.

Architectural Overview

At its core, a decoder-only transformer is a neural network designed for a single task: predicting the next token in a sequence. It's composed of several key components stacked together. Before we dive into the code, let's look at a high-level schematic.

Decoder-Only Transformer Architecture Diagram
A high-level view of the decoder-only transformer architecture. It processes token IDs through an embedding layer, a series of identical "Decoder Blocks," and a final linear layer to produce output logits. Each decoder block contains self-attention and feed-forward network sub-layers.

As you can see, the architecture consists of three main parts:

  1. Embedding Layers: Converts input token IDs into dense vectors and adds positional information.
  2. Decoder Blocks: A stack of identical blocks, each performing communication (self-attention) and computation (feed-forward network). This is the heart of the model.
  3. Output Layer: A final linear layer that projects the processed vectors into logits over the entire vocabulary.

We will now implement each of these components in PyTorch.

1. The Building Blocks: Embedding Layers

The model doesn't work with raw text. First, text is converted to a sequence of integer token IDs. The embedding layers then convert these discrete IDs into continuous, high-dimensional vectors that the network can process.

Token Embedding

This is a learnable lookup table that maps each token ID in our vocabulary to a dense vector. In PyTorch, this is implemented with nn.Embedding.

import torch
import torch.nn as nn




# Configuration
vocab_size = 50257  # Example: GPT-2's vocabulary size
d_model = 768       # The dimensionality of the embedding vectors (and model's hidden state)




# Token Embedding Layer
token_embedding = nn.Embedding(num_embeddings=vocab_size, embedding_dim=d_model)




# Example: 4 tokens in a batch of 1
input_ids = torch.randint(0, vocab_size, (1, 4)) # (batch_size, seq_length)
token_embeds = token_embedding(input_ids)

print(f"Shape of input IDs: {input_ids.shape}")
print(f"Shape of token embeddings: {token_embeds.shape}")



# Expected output: torch.Size([1, 4, 768])

Positional Encoding

The self-attention mechanism, which we'll see next, is "permutation-invariant"—it treats the input as a set of vectors, with no inherent sense of order. To give the model information about the position of each token in the sequence, we must inject positional information. The original "Attention Is All You Need" paper proposed using a fixed set of sine and cosine functions of different frequencies.

Meet GPT, The Decoder-Only Transformer

The article 'Meet GPT, The Decoder-Only Transformer' provides a clean implementation of the classic sinusoidal positional encoding. Let's examine it.

Focus on 'Codeblock 4', which defines the PositionalEncoding class. This code directly implements the sin/cos formula from the original Transformer paper. Notice how it creates a fixed tensor that is added to the token embeddings.

Here is a simplified PositionalEncoding module based on that resource. We add these positional encodings directly to the token embeddings.

import math

class PositionalEncoding(nn.Module):
    def __init__(self, d_model: int, max_len: int = 512):
        super().__init__()
        position = torch.arange(max_len).unsqueeze(1)
        div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))
        
        pe = torch.zeros(max_len, d_model)
        pe[:, 0::2] = torch.sin(position * div_term)
        pe[:, 1::2] = torch.cos(position * div_term)
        pe = pe.unsqueeze(0) # Add batch dimension
        



        # register_buffer makes 'pe' a part of the module's state,
        # but not a parameter to be trained.
        self.register_buffer('pe', pe)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        x: Tensor, shape [batch_size, seq_len, d_model]
        """



        # Add positional encoding to the input tensor.
        # x.size(1) is the sequence length.
        x = x + self.pe[:, :x.size(1), :]
        return x




# Example usage:
pos_encoder = PositionalEncoding(d_model=d_model)
final_embeddings = pos_encoder(token_embeds)
print(f"Shape after adding positional encoding: {final_embeddings.shape}")

Note: Modern models like Llama often use more advanced methods like Rotary Positional Embeddings (RoPE), which we will explore in a later module.

2. The Core Logic: A Single Transformer Block

The main body of the transformer is a stack of identical blocks. Each block performs two key operations:

  1. Self-Attention: Allows tokens to "communicate" with each other and aggregate information.
  2. Feed-Forward Network (FFN): A standard multi-layer perceptron that processes each token's information independently (the "computation" or "thinking" part).

These two sub-layers are wrapped with residual connections and layer normalization to ensure stable training of deep networks.

Self-Attention: The Communication Hub

Self-attention is the mechanism that allows the model to weigh the importance of different tokens in the input sequence when processing a specific token. It works by projecting the input into three matrices: Query (Q), Key (K), and Value (V).

  • Query: What I am looking for.
  • Key: What I contain.
  • Value: What information I will provide if you find me relevant.

The core of self-attention is Scaled Dot-Product Attention:

For a deep and intuitive walk-through of how this works, there is no better resource than Andrej Karpathy's "Let's build GPT" video.

Let's build GPT: from scratch, in code, spelled out.

This video masterfully builds the concept of self-attention from scratch, starting with a simple average and culminating in the full mechanism. Watching this is key to building a strong mental model.

Please watch the following two segments: From 42:18 to 51:40: This section introduces a brilliant mathematical trick. It shows how a simple averaging of past token information can be expressed as a matrix multiplication with a lower-triangular matrix. This builds the foundation for the causal mask. From 01:01:31 to 01:11:00: This is the core of self-attention. It shows how to implement a single attention head, introducing Query, Key, and Value, the dot-product for affinity scores, and the final value aggregation.

Based on the principles from the video, here's the implementation of a single attention head. A crucial element for decoder-only models is the causal mask (or look-ahead mask). This prevents a token at a given position from attending to any future tokens, ensuring that predictions for position i can only depend on the known outputs at positions less than i.

class SelfAttentionHead(nn.Module):
    def __init__(self, d_model: int, d_head: int):
        super().__init__()
        self.d_head = d_head
        self.q_proj = nn.Linear(d_model, d_head, bias=False)
        self.k_proj = nn.Linear(d_model, d_head, bias=False)
        self.v_proj = nn.Linear(d_model, d_head, bias=False)

    def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:



        # x shape: (batch, seq_len, d_model)
        Q = self.q_proj(x) # (batch, seq_len, d_head)
        K = self.k_proj(x) # (batch, seq_len, d_head)
        V = self.v_proj(x) # (batch, seq_len, d_head)




        # Attention scores
        # K.transpose(-2, -1) results in (batch, d_head, seq_len)
        scores = Q @ K.transpose(-2, -1) / math.sqrt(self.d_head) # (batch, seq_len, seq_len)




        # Apply causal mask
        # The mask will have -inf on the upper triangle
        scores = scores.masked_fill(mask == 0, float('-inf'))
        
        attention_weights = torch.softmax(scores, dim=-1) # (batch, seq_len, seq_len)
        



        # Weighted sum of values
        output = attention_weights @ V # (batch, seq_len, d_head)
        return output




# Create a causal mask
seq_len = 4
causal_mask = torch.tril(torch.ones(seq_len, seq_len)).unsqueeze(0) # (1, seq_len, seq_len)

Note: We've implemented a single head. Multi-Head Attention involves running several such heads in parallel and concatenating their results. We'll explore this detail in later lessons.

Feed-Forward Network: The Computation Step

After tokens have exchanged information via attention, the Feed-Forward Network (FFN) processes each token's representation independently. It's typically a two-layer MLP that introduces non-linearity and allows for more complex computations.

class FeedForward(nn.Module):
    def __init__(self, d_model: int, d_ff: int):
        super().__init__()
        self.linear1 = nn.Linear(d_model, d_ff)
        self.activation = nn.GELU() # GELU is a common choice
        self.linear2 = nn.Linear(d_ff, d_model)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return self.linear2(self.activation(self.linear1(x)))

Typically, the intermediate dimension d_ff is a multiple (often 4) of d_model.

The Glue: Residual Connections and Layer Normalization

Stacking layers can lead to training instability (vanishing/exploding gradients). Two techniques are essential for building deep transformers:

  1. Residual Connections: A "shortcut" that adds the input of a sub-layer to its output (x + SubLayer(x)). This creates a direct path for gradients to flow, greatly improving optimization.
  2. Layer Normalization (LayerNorm): Normalizes the features for each token independently across the embedding dimension. This stabilizes the activation statistics during training.

Modern transformers typically use a pre-normalization scheme: x + SubLayer(LayerNorm(x)).

Build an LLM from Scratch 4: Implementing a GPT model from Scratch To Generate Text

Dr. Raschka's video provides a very clear, focused explanation of Layer Normalization and its implementation. This will clarify why and how we normalize the activations.

Watch the segment from 13:47 to 25:30. This covers the motivation for normalization, the difference between batch and layer norm, and a from-scratch implementation of LayerNorm in PyTorch. Pay attention to how it normalizes across the feature dimension for each token independently.

Now we can combine these components into a single TransformerBlock.

class TransformerBlock(nn.Module):
    def __init__(self, d_model: int, d_head: int, d_ff: int):
        super().__init__()
        self.attention = SelfAttentionHead(d_model, d_head)
        self.ffn = FeedForward(d_model, d_ff)
        self.norm1 = nn.LayerNorm(d_model)
        self.norm2 = nn.LayerNorm(d_model)

    def forward(self, x: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:



        # Pre-normalization for the attention layer
        attn_output = self.attention(self.norm1(x), mask)



        # Residual connection
        x = x + attn_output
        



        # Pre-normalization for the feed-forward network
        ffn_output = self.ffn(self.norm2(x))



        # Residual connection
        x = x + ffn_output
        
        return x

3. Assembling the Full Model

With all the components ready, we can assemble the final model. It will consist of the embedding layers, a stack of TransformerBlocks, and a final linear layer to produce logits.

Building a Decoder-Only Transformer Model Like Llama-2 and ...

The article 'Building a Decoder-Only Transformer Model' provides a very clean, complete PyTorch implementation that closely matches modern architectures like Llama. We'll use it as a reference for our final model structure.

First, review the DecoderLayer and TextGenerationModel classes in the main text. This shows the high-level assembly. Then, scroll down to the full code block and examine the TextGenerationModel class again. Notice how it initializes the embedding, the nn.ModuleList of DecoderLayers, and the final out layer, and how the forward method ties them all together.

Here is our complete minimal decoder-only transformer:




# --- Full Model Code ---
import torch
import torch.nn as nn
import math




# Use the classes defined above: PositionalEncoding, SelfAttentionHead, FeedForward, TransformerBlock

class MinimalDecoder(nn.Module):
    def __init__(self, vocab_size: int, d_model: int, n_layers: int, d_head: int, d_ff: int, max_len: int = 512):
        super().__init__()
        self.token_embedding = nn.Embedding(vocab_size, d_model)
        self.pos_encoding = PositionalEncoding(d_model, max_len)
        
        self.blocks = nn.ModuleList([
            TransformerBlock(d_model, d_head, d_ff) for _ in range(n_layers)
        ])
        
        self.final_norm = nn.LayerNorm(d_model)
        self.output_head = nn.Linear(d_model, vocab_size)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        batch_size, seq_len = x.shape
        



        # 1. Create Causal Mask
        mask = torch.tril(torch.ones(seq_len, seq_len, device=x.device)).unsqueeze(0)
        



        # 2. Input Embeddings + Positional Encoding
        x = self.token_embedding(x)
        x = self.pos_encoding(x)
        



        # 3. Transformer Blocks
        for block in self.blocks:
            x = block(x, mask)
        



        # 4. Final Normalization and Output Head
        x = self.final_norm(x)
        logits = self.output_head(x)
        
        return logits




# --- Test the model ---
# Configuration
vocab_size = 50257
d_model = 768
d_head = 64 # Typically d_model / num_heads
d_ff = d_model * 4
n_layers = 12

model = MinimalDecoder(vocab_size, d_model, n_layers, d_head, d_ff)




# Create a dummy input
input_ids = torch.randint(0, vocab_size, (1, 10)) # batch_size=1, seq_len=10




# Get the model output (logits)
output_logits = model(input_ids)

print(f"Input shape: {input_ids.shape}")
print(f"Output logits shape: {output_logits.shape}")



# Expected output shape: (batch_size, seq_len, vocab_size) -> (1, 10, 50257)

This final script ties everything together. You can run it to verify that the tensor shapes flow correctly from input token IDs to output logits.

Conclusion

Congratulations! You have just implemented a decoder-only transformer from scratch. While minimal, this model contains all the essential architectural DNA of today's most powerful LLMs.

Key Takeaways:

  • A decoder-only transformer is a stack of identical blocks processing token vectors.
  • Embeddings map discrete tokens to continuous space and add Positional Encoding to give a sense of order.
  • The Transformer Block is the core repeating unit, containing:
    • Self-Attention: A "communication" layer where tokens exchange information, governed by a causal mask to prevent looking into the future.
    • Feed-Forward Network: A "computation" layer where each token processes information independently.
    • Residual Connections and LayerNorm: Critical "glue" that enables stable training of deep networks.
  • The final Output Layer projects the result back to the vocabulary space to produce prediction logits.

Preview of the Next Lesson:

Now that you've built the model's structure layer by layer, you have the exact knowledge needed to reason about its size. In the next lesson, we will move on to the next outcome: "Write a function to programmatically compute the total parameter count of a transformer from its architectural hyperparameters." We will use the nn.Modules you've just implemented to derive a formula for a model's memory footprint based purely on its configuration.

Can't find a good explanation? Sign up and we'll make it for you

Sign up