Understanding Transformer architectures requires getting comfortable with attention calculations. While PyTorch offers high-level layers like nn.MultiheadAttention, implementing Scaled Dot-Product Attention directly helps clarify matrix dimensions and numerical stability tricks like softmax scaling.
How does Scaled Dot-Product Attention work in PyTorch from scratch?
1 Answer
Scaled dot-product attention takes a query, compares it with every key, turns those comparison scores into weights, then uses the weights to combine the values. In PyTorch, the core is just matrix multiplication, scaling, and softmax.
For query Q, key K, and value V, the calculation is:
softmax((Q × Kᵀ) / √dₖ) × V
The division by √dₖ keeps scores from growing too large as the key dimension increases. Without it, softmax can become very peaked, which can make learning less effective.
A small PyTorch implementation
This version supports tensors with batch dimensions, such as [batch, sequence, features]. The returned attention weights show how strongly each query attends to each key.
import math
import torch
def scaled_dot_product_attention(query, key, value, mask=None):
# The final query dimension is the key feature size, d_k.
d_k = query.size(-1)
# Compare each query with each key, then scale the scores.
scores = torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(d_k)
if mask is not None:
# Boolean mask entries that are False are excluded from attention.
scores = scores.masked_fill(~mask, float("-inf"))
# Normalize over keys so each query's attention weights sum to one.
weights = torch.softmax(scores, dim=-1)
# Use those weights to form a weighted sum of the value vectors.
output = torch.matmul(weights, value)
return output, weights
# Example: one batch, four tokens, and eight features per token.
x = torch.randn(1, 4, 8)
# A lower-triangular mask prevents each token from seeing future tokens.
causal_mask = torch.tril(torch.ones(4, 4, dtype=torch.bool)).unsqueeze(0)
output, weights = scaled_dot_product_attention(
query=x,
key=x,
value=x,
mask=causal_mask,
)
print(output.shape) # torch.Size([1, 4, 8])
print(weights.shape) # torch.Size([1, 4, 4])
Here, each token supplies its own query, key, and value, as in self-attention. The mask has shape [1, 4, 4] and broadcasts across the batch. For unmasked attention, pass mask=None. Avoid masks that block every key for a query; softmax over an all-masked row can produce invalid values.
All operations in this implementation are regular PyTorch tensor operations, so gradients flow through them automatically. For production models, PyTorch also provides optimized attention functions, but writing out these steps makes the underlying operation easier to inspect.