Transformers
1. Attention
Some feelings about how can these attention-powered transformers work so well in Gen AI:
Natural language, or any sequence data, is literally like a more advanced Solitaire game.
Each position can — and should — look way beyond just the previous position.
3️⃣ Roles Inside Context
To better look deeper beneath context, each word is treated in 3 ways:
- Query : which topic is searched for.
- Key : how each candidate is labeled.
- Value : what content each candidate holds.
So — if each pairs with a , then how about we use & matching results to decide how much each plays?
Stage ⓵: Scaled Dot-Product
.
.
☢️HIM — Heavy Incoming Math
Suppose all entries of and are i.i.d. from a distribution of mean 0 and variance 1, implying several things:
.
☝️ decouples into an implied fact — and are independent.
And this heaviest one:
=
So we can see that our distribution has:
I. Expectation: 0
=
= , so .
II. Variance:
=
= , so .
III. Assembly
Scale by to ensure :
When input follows a less extreme distribution, Softmax output tends to be less spiky.
Less spiky means less one-hot encoding — away from saturated vanishing-gradients areas.
Dot product comes afterward:
: with , is token’s relative proportion of attention to token.
each means token’s value representation.
yields a matrix of , where each row vector is a weighted sum of token vectors.
Weights are based on token’s attention to all tokens.
Be sure to set axis=-1 when doing Softmax. Along columns 😉😏
IV. 🤿Mask If Required
After doing , you have a square matrix , showing how much relative attention each token pays to other tokens.
But…..what if at times you aren’t totally free to pay attention?
Like, you can only pay attention to past and present. Not future.
Straightforward: given any , set to have .
Which is just to let Softmax produce 0, representing no attention, on disabled entries.
ones_matrix = torch.ones_like(qk_product) # Shape: (n, n).
# Diagonal = 1: don't mask entire diagonal as well, cuz present can be referenced.
mask = torch.triu(ones_matrix, diagonal=1).bool()
qk_product.masked_fill_(mask, float("-inf"))⬆️This code helps us mask out all upper triangular entries before entering Softmax.
Stage ⓶: MHA — Multi-Head
Simple intro: sets of 3 linear projections before scaled dot-product.
Philosophy behind: have more matrices track how each token behaves in a certain semantic space.
Detailed Maneuvers 🥘
Now Each head owns 3 weight matrices called:
— Help produce .
— Help produce .
— Help produce .
So all fit into expected input shapes of scaled dot-product.
scaled dot-product attention is . We will have such matrices.
With each row referring to a token, we of course can only concatenate attention along columns.
By doing so, we obtain .
Using another weight matrix , we can reach .
Yes, back to original shape of , but carrying learned attention.
🤔Subtle & Vital Guideline
For MHA, we need extra weight matrices:
First matrices — all that deal with .
Middle matrices — all that deal with .
Last 1 matrix — that brings back to original shape.
All these matrices share a common: dimensions contain no .
At the End of the Day
Attention is to grasp semantic relationship among tokens.
are all semantic spaces’ dimensions. isn’t.
These weight matrices are born to travel around semantic spaces. 😉😏
🕰️Time Complexity
- — pairs of doing inner product, with per pair.
Scaler & Softmax:
- They both perform constantly on all entries inside .
V Multiplication:
- — pairs of doing inner product, with per pair.
So total is
2. Generation 🏳️🌈
How to make good use of these learned attention?
Like writing a research paper:
- Ensure no topic deviation — look at already written parts.
- Extend relevant content — make reference to other papers.
Although I’ve never written a research paper 😎
🎭Masked Attention: No Deviation
Like I said, to prevent topic deviation, we must look at already written parts.
So in each step, already generated output is naturally the sources for our , and .
The adjective “masked” means only previous and current steps’ output is eligible.
⚔️Cross-Attention: Make Reference
Where can we let generated output make reference to incorporate relevant content?
Encoder’s learned attention, aka encoder stack. Thus, this time:
Encoder stack attention serves as sources of & , since it gets referenced.
We are allowed to look at every step. So no masks at all.
Tip: Output Is Topic Regardless
| Attention | source | & sources | Masked |
|---|---|---|---|
| Masked | Output. | Output. | Yes. |
| Cross | Output. | Encoder stack attention. | No. |
3. KV Cache
Now comes a discussion: can we optimize time or space complexities?
This brings up a quite fundamental topic indeed. Need a whole new page.
4. Positional Encoding
We chose this function because we hypothesized it would allow the model to easily learn to attend by relative positions, since for any fixed offset , can be represented as a linear function of .
“Attention Is All You Need” uses and to encode positions.
But there’s another interesting way to do so as well.
Second thought — decided to make an exclusive page for positional encoding only.