Hello! Welcome back to our course on Audio AI.
In the last lesson, we deconstructed the entire Transformer architecture, examining its key components: positional encoding, the encoder-decoder stacks, feed-forward networks, and multi-head attention. We focused on what these parts are and why they are essential.
Today, we transition from theory to practice. Our goal is to take one of those crucial building blocks and construct it from the ground up.
Your learning outcome for this lesson is to: Implement a multi-head self-attention layer in PyTorch, including support for causal masking.
We will start with the mathematical heart of attention—scaled dot-product attention—and progressively build it into a complete, reusable PyTorch nn.Module. This exercise is fundamental. A deep understanding of this implementation will empower you to move beyond using off-the-shelf models and start designing or modifying your own novel architectures, a key skill for your goal of becoming an audio researcher.
1. The Core Mechanism: Scaled Dot-Product Attention
Let's begin with the core formula that governs how attention is calculated. As a reminder from our previous lesson, attention can be described as a function of a Query (Q), a Key (K), and a Value (V).
The specific formula used in the Transformer is Scaled Dot-Product Attention:
Let's break this down:
- Score: We compute a score between each query and all keys using a dot product: . This matrix, often called the "attention logits" or "affinities," represents the raw similarity between every pair of elements in the sequence.
- Scale: We scale the scores by dividing by , where is the dimension of the key vectors. This is a crucial stabilization step that prevents the dot products from growing too large, which would push the softmax function into regions with very small gradients, stalling training.
- Mask (Optional): Before applying softmax, we can apply a mask to prevent certain positions from attending to others. We'll implement this for causal attention later.
- Weights: We apply a
softmaxfunction along the key dimension to normalize the scores into positive weights that sum to 1. These weights represent how much attention each query should pay to each key. - Output: We compute a weighted sum of the value vectors using the attention weights. The result is an output vector for each position that is a blend of all other value vectors, informed by the attention weights.

To see this process explained in a practical, code-oriented way, let's turn to a video by Andrej Karpathy.
Let's build GPT: from scratch, in code, spelled out.
In this segment, Andrej Karpathy explains the intuition behind Query, Key, and Value vectors and how their dot products create data-dependent affinities between tokens. He also introduces the concept of masking in the context of an autoregressive model.
Watch from 01:04:09 to 01:10:10. Focus on how the wei matrix is computed from the dot product of queries and keys, and why masking is applied to prevent future tokens from influencing the current one before the softmax operation.
2. Implementing a Single Head of Self-Attention
Now let's translate this into a PyTorch nn.Module. A self-attention layer takes a sequence of input vectors and produces the Q, K, and V vectors from it using linear projections. For efficiency, we can use a single large linear layer to create Q, K, and V at once.
The implementation steps are:
- Define learnable linear layers to project the input into Q, K, and V. A common and efficient strategy is to use one
nn.Linearlayer that projects frominput_dimto3 * head_dim. - In the
forwardpass, apply this projection and split the result into three separate tensors for Q, K, and V. - Implement the scaled dot-product attention formula using these tensors.
- Define a final linear layer, , to project the output of the attention mechanism back to the desired dimension.
This video segment walks through this exact process.
Let's build GPT: from scratch, in code, spelled out.
Continuing with the Karpathy video, this next part details the full implementation of a single self-attention head. It introduces the Value (V) vector and shows how all the pieces (Q, K, V, scaling, masking, softmax) come together to form the final output.
Watch from 01:10:10 to 01:19:46. Pay close attention to: How the value vector is created and used. The explanation of scaling by 1/sqrt(head_size) to control variance. How the logic is encapsulated within a Head module.
3. From Single Head to Multi-Head Attention
As we discussed in the previous lesson, Multi-Head Attention is more powerful than a single head. It allows the model to learn different relationships in parallel by running the attention mechanism multiple times with different, learned linear projections.
To implement this, we'll modify our single-head logic:
- Projections: We still use a single
nn.Linearfor efficiency, but now it projects frominput_dimto3 * embed_dim, whereembed_dimis the total dimension across all heads. - Reshaping for Parallelism: After the projection, we reshape and transpose the tensor to separate the heads. A tensor of shape
(batch, seq_len, 3 * embed_dim)becomes(batch, num_heads, seq_len, 3 * head_dim). This allows us to perform the attention calculation for all heads in parallel. - Parallel Attention: The scaled dot-product attention logic remains the same. PyTorch's
matmulwill automatically handle the batch and head dimensions. - Concatenation: After computing the attention output, we "concatenate" the heads back together by transposing and reshaping the tensor from
(batch, num_heads, seq_len, head_dim)back to(batch, seq_len, embed_dim). - Final Projection: We apply the final output projection
W^O.
The following resource provides a very clean implementation of this process.
Tutorial 6: Transformers and Multi-Head Attention
The UVADLC tutorial on Transformers offers a clear, well-structured PyTorch implementation of Multi-Head Attention. It's a great reference for seeing the theory translated into production-quality code.
Read the section "Multi-Head Attention". Study the MultiheadAttention class implementation closely. Focus on the forward pass and trace the tensor shapes as they are projected, reshaped for heads, passed through scaled_dot_product, and then reshaped back.
4. Implementing Causal Masking
For generative tasks like Text-to-Speech (TTS) or autoregressive speech recognition, we need to ensure that the model cannot "see into the future." The prediction for the current token (or audio sample) must only depend on the tokens that came before it. This is achieved with causal masking.
The implementation is straightforward:
- Create a square matrix where the upper triangle (excluding the diagonal) is
True(or 1) and the rest isFalse(or 0). This matrix represents the connections that are not allowed (i.e., attending to future tokens). - In the
forwardpass, before thesoftmax, use this mask to replace the attention scores for all future positions with a very large negative number (like-inf). - When
softmaxis applied, these positions will have a probability of effectively zero, preventing any information flow from the future.
For efficiency, this mask can be created once and stored as a buffer in the nn.Module. A buffer is a tensor that is part of the model's state but is not a trainable parameter.
This article provides a concise implementation of causal attention.
Creating a Transformer From Scratch - Part One: The Attention Mechanism
This article from Benjamin Warner provides a step-by-step guide to building an attention layer, culminating in a CausalAttention class that incorporates this crucial masking logic.
Read the section on "Causal Self-Attention". Pay special attention to how the causal_mask is created using torch.triu and registered as a buffer, and how it is applied in the forward method using masked_fill.
5. Putting It All Together: The Complete Implementation
Now, let's consolidate everything we've learned into a single, well-commented PyTorch module. This CausalSelfAttention class will be a reusable component that you can plug into a larger Transformer decoder.
Here is the complete implementation. Take your time to read through it, connecting each line of code to the concepts we've discussed.
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
class CausalSelfAttention(nn.Module):
"""
A full implementation of Multi-Head Causal Self-Attention.
"""
def __init__(self, embed_dim: int, num_heads: int, max_seq_len: int, bias: bool = False, dropout: float = 0.1):
"""
Args:
embed_dim (int): The total embedding dimension.
num_heads (int): The number of parallel attention heads.
max_seq_len (int): The maximum sequence length the model will see.
Used to create the causal mask buffer.
bias (bool): Whether to use bias in the linear layers.
dropout (float): Dropout probability.
"""
super().__init__()
assert embed_dim % num_heads == 0, "Embedding dimension must be divisible by number of heads."
self.embed_dim = embed_dim
self.num_heads = num_heads
self.head_dim = embed_dim // num_heads
# A single linear layer for Q, K, V projections for efficiency
self.qkv_proj = nn.Linear(embed_dim, 3 * embed_dim, bias=bias)
# Final output projection
self.o_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
# Dropout layers
self.attn_dropout = nn.Dropout(dropout)
self.resid_dropout = nn.Dropout(dropout)
# Causal mask to ensure attention is only applied to the left
# We register it as a buffer so it's part of the model's state,
# but not a parameter to be trained.
# `torch.triu` creates an upper-triangular matrix.
# We want to mask the upper triangle, so where the matrix is 1 (True).
mask = torch.triu(torch.ones(max_seq_len, max_seq_len), diagonal=1).bool()
self.register_buffer("causal_mask", mask)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Args:
x (torch.Tensor): Input tensor of shape (batch_size, seq_len, embed_dim)
Returns:
torch.Tensor: Output tensor of shape (batch_size, seq_len, embed_dim)
"""
batch_size, seq_len, _ = x.shape
# 1. Project to Q, K, V
# x: (B, S, C) -> qkv: (B, S, 3*C)
qkv = self.qkv_proj(x)
# 2. Reshape and transpose for multi-head attention
# (B, S, 3*C) -> (B, S, num_heads, 3*head_dim)
# -> (B, num_heads, S, 3*head_dim)
qkv = qkv.view(batch_size, seq_len, self.num_heads, 3 * self.head_dim).transpose(1, 2)
# 3. Split into Q, K, V
# q, k, v each have shape (B, num_heads, S, head_dim)
q, k, v = qkv.chunk(3, dim=-1)
# 4. Calculate attention scores (Scaled Dot-Product Attention)
# (B, H, S, D_h) @ (B, H, D_h, S) -> (B, H, S, S)
attn_scores = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(self.head_dim))
# 5. Apply causal mask
# We select the portion of the mask that matches the current sequence length
mask_slice = self.causal_mask[:seq_len, :seq_len]
attn_scores = attn_scores.masked_fill(mask_slice, float('-inf'))
# 6. Normalize scores to get attention weights
attn_weights = F.softmax(attn_scores, dim=-1)
attn_weights = self.attn_dropout(attn_weights)
# 7. Compute weighted sum of values
# (B, H, S, S) @ (B, H, S, D_h) -> (B, H, S, D_h)
output = attn_weights @ v
# 8. Concatenate heads and project
# (B, H, S, D_h) -> (B, S, H, D_h) -> (B, S, C)
output = output.transpose(1, 2).contiguous().view(batch_size, seq_len, self.embed_dim)
output = self.o_proj(output)
output = self.resid_dropout(output)
return output
To test it, you could instantiate it and pass a dummy tensor through:
# Example usage:
embed_dim = 512
num_heads = 8
max_seq_len = 1024
batch_size = 4
seq_len = 50 # current sequence length
attention_layer = CausalSelfAttention(embed_dim, num_heads, max_seq_len)
dummy_input = torch.randn(batch_size, seq_len, embed_dim)
output = attention_layer(dummy_input)
print(f"Input shape: {dummy_input.shape}")
print(f"Output shape: {output.shape}")
# Expected output:
# Input shape: torch.Size([4, 50, 512])
# Output shape: torch.Size([4, 50, 512])
This confirms our module correctly processes the input and maintains the tensor shape, as expected within a Transformer block.
Conclusion
In this lesson, you have moved from a theoretical understanding of attention to a practical, from-scratch implementation. By building this CausalSelfAttention module, you've mastered the core computational primitive of the Transformer decoder.
Key Takeaways:
- Scaled Dot-Product Attention is the core calculation, involving dot products of Q and K, scaling, masking, softmax, and a weighted sum of V.
- Multi-Head Attention is implemented efficiently through tensor reshaping and transposing, allowing parallel computation across heads.
- Causal Masking is essential for autoregressive models. It is implemented by setting attention scores for future tokens to
-infbefore the softmax operation. - A clean implementation uses a single linear layer for QKV projection and registers the causal mask as a non-trainable buffer for efficiency.
Preview of the Next Lesson:
We've built a powerful component, but it doesn't work in isolation. In our next lesson, we will assemble a Transformer encoder block combining multi-head attention and a feed-forward network with residual connections. We'll take the module we just built (in its non-causal form) and integrate it into the next level of the Transformer's hierarchy.