Beau
Alright, Jo. So we have our embedding and positional encoding sorted. The tokens are now vectors in a high-dimensional space. We're staring right at the core of the transformer block now... the self-attention mechanism.
Transcript
Beau
Alright, Jo. So we have our embedding and positional encoding sorted. The tokens are now vectors in a high-dimensional space. We're staring right at the core of the transformer block now... the self-attention mechanism.
Jo
Exactly. And we're going to treat it for what it is: a series of high-performance tensor operations. No metaphors about 'looking at other words'. This is pure linear algebra.
Beau
Okay, so let's start with the formula everyone sees: softmax of Q K-transpose over the square root of d_k. Simple enough. But implementing Q times K-transpose when you have batches and multiple heads... that's where the tensor ranks get tricky.
Jo
It is. You could do it with a series of transpositions and matmuls, but it's messy. The cleanest way, and frankly the most readable once you get it, is torch.einsum.
Beau
Ah, Einstein summation notation. I've used it in physics contexts, but less so in deep learning. So, let's walk through the notation for this specific case. Our Query and Key tensors... what's their shape entering this operation?
Jo
Good question. After we project our input into Q, K, and V and split them for multi-head attention, each will have a shape of, let's say: Batch Size, Number of Heads, Sequence Length, and Head Dimension. Let's call the dimensions b, h, s, and d.
Beau
So, for Q we have 'b, h, s, d'. For K, we also have 'b, h, s, d'. To do the dot product, we need to multiply along that last dimension, 'd', and transpose K's sequence dimension.
Jo
Exactly. So the einsum string would be 'bhsd,bhjd->bhsj'. We're telling it: take the tensors with these dimensions, multiply along the common dimension 'd', and produce an output tensor with dimensions 'b', 'h', 's', and 'j'. 'j' is just another name for the sequence length dimension from K.
Beau
And the output shape is Batch, Heads, Sequence Length, Sequence Length. That makes sense. It's the attention score of every token against every other token, for each head in each batch item. Super clean.
Jo
Much cleaner than reshaping and batching matrix multiplies. Now, before we apply softmax, we have two crucial steps. First, scaling.
Beau
The division by the square root of d_k, the head dimension. I've always understood this as a way to prevent the dot products from becoming too large, which would push the softmax into regions with tiny gradients.
Jo
That's the exact reason. It's a numerical stability trick. Without it, for larger head dimensions, the dot products could grow massive. Softmax would then sharpen to a one-hot distribution, and the gradients for all other tokens would vanish. Backpropagation would grind to a halt.
Beau
Okay, that's step one. What's step two? Ah... the causal mask. We can't let a token attend to future tokens in the sequence. That would be cheating.
Jo
Right. We need to enforce the autoregressive property. And we do this by creating a mask. The most efficient way in PyTorch is to use `torch.tril`, which gives you the lower triangular part of a matrix.
Beau
So you create a matrix of ones with the shape of our attention scores... Sequence Length by Sequence Length... and `tril` will zero out the upper triangle.
Jo
Almost. We then use that binary mask to fill the illegal, future positions in our attention scores matrix not with zero, but with a very large negative number. Negative infinity, essentially.
Beau
Right, because when you pass that through the softmax function, e to the power of negative infinity is zero. So those positions get an attention weight of exactly zero. Clever.
Jo
It's an elegant way to enforce causality at the compute level. So, recap of the block: einsum for `QK^T`, scale by `1/sqrt(d_k)`, apply the causal mask by adding negative infinity, and then finally apply the softmax along the last dimension.
Beau
And that gives us our attention weights. A matrix that says, for each token, how much it should 'pay attention' to every *previous* token, including itself.
Jo
Precisely. And the final step is to use these weights to create a weighted sum of the Value vectors, V. Which is another `einsum` operation, by the way.
Beau
Of course it is. Let me guess... something like 'bhsj,bhjd->bhsd'? We're combining the attention weights with the Values to get our final output with the original head dimension.
Jo
You got it. That's the entire computational flow of one attention head. The rest is just stacking these blocks, adding the feed-forward layers and the normalization. But this... this is the engine.