How to Implement Residual Connections in Transformer Blocks
Residual connections in Transformer blocks are implemented by adding the input tensor to the output of each sub-layer (attention and MLP) after applying layer normalization, typically using the pattern x = x + sublayer(norm(x)) to enable stable gradient flow through deep networks.
Implementing residual connections correctly is essential when building Transformer architectures from scratch. In the FareedKhan-dev/train-llm-from-scratch repository, residual connections are handled inside the Block class and propagated through the stacked Transformer model. This guide explains exactly how to implement residual connections in transformer blocks using the source code patterns found in this educational LLM implementation.
Where Residual Connections Are Implemented
Residual connections appear in two primary locations within the codebase. The ** Block class** in src/models/transformer_block.py contains the core residual logic, while the ** Transformer class** in src/models/transformer.py orchestrates how these blocks stack.
According to the source code analysis, the specific implementations are:
src/models/transformer_block.pylines 43-46: TheBlock.forwardmethod adds the input tensor to the attention output, then adds that result to the MLP outputsrc/models/transformer_block.pylines 59-61: TheBlock.forward_embeddingmethod returns both the transformed tensor and the residual after the attention sub-layersrc/models/transformer.pylines 68-70: TheTransformer.forwardmethod chains blocks that internally handle their own residualssrc/models/transformer.pylines 90-94: TheTransformer.forward_embeddingmethod propagates residuals upstream from each block
The Pre-Norm Residual Architecture
The repository follows the Pre-Norm design pattern, where layer normalization occurs before the sub-layer rather than after. This approach differs from the original Transformer paper but provides superior training stability for deep networks.
Attention Sub-Layer with Residual
Inside Block.forward, the first residual connection wraps the multi-head attention mechanism. The implementation first normalizes the input, computes attention, then adds the original input back to the result:
x = x + self.attn(self.ln1(x))
This pattern appears at line 43 in src/models/transformer_block.py. The original input x is preserved through the addition operation, creating the skip connection that allows gradients to flow directly back to earlier layers.
MLP Sub-Layer with Residual
Immediately following the attention residual, the block applies a second residual connection around the feed-forward network:
x = x + self.mlp(self.ln2(x))
This implementation at line 45 maintains the same additive pattern. By applying layer normalization (self.ln2) before the transformation and adding the pre-transformation input to the output, the network can learn identity functions easily if a layer does not improve the representation.
Stacking Blocks in the Full Transformer
When building a complete model, the Transformer class in src/models/transformer.py instantiates multiple Block instances and processes input through them sequentially. Each Block independently manages its own residual connections, so the full model inherits a hierarchy of skip connections spanning all layers.
The Transformer.forward method (lines 68-70) iterates through self.attn_blocks, calling each block which internally executes the residual logic described above. After processing through all blocks, a final layer normalization (self.layer_norm) prepares the representations for the language model head.
For embedding extraction tasks, Transformer.forward_embedding (lines 90-94) propagates residuals upstream from each block, mirroring the behavior of Block.forward_embedding which returns the residual state after the attention sub-layer.
Why Residual Connections Matter
Understanding the implementation requires understanding the motivation behind these architectural choices.
Gradient flow preservation. The additive shortcut allows gradients to bypass attention and MLP sub-layers during backpropagation, directly addressing the vanishing gradient problem in deep stacks. Without residuals, training networks with 12, 24, or more blocks would be computationally infeasible.
Training stability. The Pre-Norm configuration combined with residual connections mitigates the "exploding activation" issue common in deep Transformers. By normalizing before the transformation and adding the unnormalized residual, the network maintains stable activation ranges throughout training.
Identity function learning. If a particular layer cannot improve the representation, the residual path enables it to default to an identity function (output = input). This allows the optimizer to allocate capacity to layers that actually contribute to the task while preserving information flow through less useful layers.
Code Implementation Examples
The following examples demonstrate how to instantiate blocks and verify residual connections are active.
Single Block Forward Pass
To create a single Transformer block and observe residual behavior:
import torch
from src.models.transformer_block import Block
# Configuration
batch_size = 2
seq_len = 5
embed_dim = 32
num_heads = 4
context_len = seq_len
# Random input tensor (B, T, C)
x = torch.randn(batch_size, seq_len, embed_dim)
# Initialize block
block = Block(n_head=num_heads, n_embed=embed_dim, context_length=context_len)
# Standard forward (residuals applied internally)
output = block(x)
print(f"Output shape: {output.shape}") # (2, 5, 32)
# Access intermediate residual state
out_emb, residual = block.forward_embedding(x)
print(f"Post-attention residual shape: {residual.shape}") # (2, 5, 32)
Building a Multi-Layer Model
To stack multiple blocks with residual connections in a full Transformer:
from src.models.transformer import Transformer
vocab_size = 100
n_blocks = 3
model = Transformer(
n_head=num_heads,
n_embed=embed_dim,
context_length=context_len,
vocab_size=vocab_size,
N_BLOCKS=n_blocks,
)
# Dummy input
idx = torch.randint(0, vocab_size, (batch_size, seq_len))
# Forward pass (residuals handled automatically in each block)
logits, loss = model(idx, targets=idx)
print(f"Logits shape: {logits.shape}") # (2, 5, 100)
In this multi-layer configuration, residual connections exist at two points within each of the three blocks, creating six total skip connections through which gradients can flow.
Summary
- Residual connections in this implementation use the Pre-Norm pattern:
x = x + sublayer(norm(x)) - Core implementation resides in
src/models/transformer_block.pyat lines 43-46, handling both attention and MLP residuals - Dual pathways exist:
Block.forwardfor standard inference andBlock.forward_embeddingfor extracting intermediate residual states - Automatic propagation occurs when stacking blocks via
Transformerclass insrc/models/transformer.py - Key benefits include stabilized gradient flow, mitigation of vanishing gradients, and the ability for layers to learn identity functions when necessary
Frequently Asked Questions
What is the difference between Pre-Norm and Post-Norm in Transformers?
Pre-Norm applies layer normalization before the sub-layer (attention or MLP), while Post-Norm applies it after the residual addition. The train-llm-from-scratch repository uses Pre-Norm (lines 43-46 in transformer_block.py), which provides more stable gradients during training but may require slightly different learning rate schedules compared to the original Post-Norm Transformer design.
Why does the Block class return residuals in forward_embedding?
The forward_embedding method (lines 59-61) returns the residual state after the attention sub-layer to support analysis and specialized training workflows. This allows researchers to inspect the pre-MLP representation or propagate specific residual states through custom training loops while maintaining the standard residual logic internally.
How do residual connections prevent vanishing gradients?
Residual connections create additive shortcut paths that allow gradients to flow directly from later layers back to earlier ones without passing through activation functions or weight matrices. In the implementation at src/models/transformer_block.py, the x + ... operations ensure that the gradient with respect to the input x receives a direct contribution from the output gradient, preventing the multiplicative decay that occurs when gradients pass through many layers sequentially.
Can I modify the residual connection strength?
While the repository uses standard additive residuals (x = x + ...), you could implement gated residuals by modifying lines 43-45 in transformer_block.py to use a learnable parameter: x = x + alpha * self.attn(...). However, the current implementation follows the standard Transformer design where the residual weight is implicitly 1.0, which empirical research shows provides sufficient gradient flow for training deep language models.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →