Skip to content

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 QQ: which topic is searched for.
  • Key KK: how each candidate is labeled.
  • Value VV: what content each candidate holds.

So — if each KK pairs with a VV, then how about we use QQ & KK matching results to decide how much each VV plays?

Stage ⓵: Scaled Dot-Product

  • Q=Rn×dk,K=Rn×dk,V=Rn×dv.Q = R^{n \times d_k}, K = R^{n \times d_k}, V = R^{n \times d_v}.

  • Attention(Q,K,V)=SoftmaxQKTdkV\text{Attention}(Q, K, V) = \text{Softmax} \frac{ QK^T }{ \sqrt{d_k} } V.

  • (QKT)i,j=Qi,:⋅(kj,:)T(QK^T)_{i, j} = Q_{i, :} \cdot (k_{j, :})^T.

☢️HIM — Heavy Incoming Math

Suppose all entries of QQ and KK are i.i.d. from a distribution of mean 0 and variance 1, implying several things:

  • E[Qi,h×Kj,h]=Cov(Qi,h,Kj,h)+E[Qi,h]×E[Kj,h]=0+0×0=0.E[Q_{i, h} \times K_{j, h}] = Cov(Q_{i, h}, K_{j, h}) + E[Q_{i, h}] \times E[K_{j, h}] = 0 + 0 \times 0 = 0.

  • ∀a,b∈R,E[Qi,ha×Kj,hb]=∫∫Qi,haKj,hbf(q,k)  dq  dk=∫∫Qi,haKj,hbfQ(q)fK(k)  dq  dk=∫Qi,hafQ(q)dq×∫Kj,hbfK(k)dk=E[Qi,ha]×E[Kj,hb]\forall a, b \in R, E[Q_{i, h}^a \times K_{j, h}^b] = \int \int Q_{i, h}^a K_{j, h}^b f(q, k) \; dq \; dk = \int \int Q_{i, h}^a K_{j, h}^b f_Q(q) f_K(k) \; dq \; dk = \int Q_{i, h}^a f_Q(q) dq \times \int K_{j, h}^b f_K(k) dk = E[Q_{i, h}^a] \times E[K_{j, h}^b].

  • ☝️f(q,k)=fQ(q)fK(k)f(q, k) = f_Q(q) f_K(k) decouples into an implied fact — Q2Q^2 and K2K^2 are independent.

And this heaviest one:

Var(Qi,h×Kj,h)=E[(Qi,h×Kj,h)2]−E[Qi,h×Kj,h]2Var(Q_{i, h} \times K_{j, h}) = E[(Q_{i, h} \times K_{j, h})^2] - E[Q_{i, h} \times K_{j, h}]^2

= E[Qi,h2×Kj,h2]−02=E[Qi,h2]×E[Kj,h2]=1×1=1.E[Q_{i, h}^2 \times K_{j, h}^2] - 0^2 = E[Q_{i, h}^2] \times E[K_{j, h}^2] = 1 \times 1 = 1.

So we can see that our (QKT)i,j(QK^T)_{i, j} distribution has:

I. Expectation: 0

E[(QKT)i,j]=E[Qi,:⋅(kj,:)T]=E[∑h=1dkQi,h×Kj,h]E[(QK^T)_{i, j}] = E[Q_{i, :} \cdot (k_{j, :})^T] = E[\sum_{h = 1}^{d_k} Q_{i, h} \times K_{j ,h}]

= ∑h=1dkE[Qi,h×Kj,h]=∑h=1dkE[Qi,h]×E[Kj,h]⏟All entries of Q and K are independent.\underbrace{ \sum_{h = 1}^{d_k} E[Q_{i, h} \times K_{j ,h}] = \sum_{h = 1}^{d_k} E[Q_{i, h}] \times E[K_{j ,h}] }_{\text{All entries of Q and K are } \textbf{independent.}}

= ∑h=1dk0×0=0\sum_{h = 1}^{d_k} 0 \times 0 = 0, so E[QKT]=0E[QK^T] = 0.

II. Variance: dkd_k

Var((QKT)i,j)=Var(Qi,:⋅(kj,:)T)Var((QK^T)_{i, j}) = Var(Q_{i, :} \cdot (k_{j, :})^T)

= Var(∑h=1dkQi,h×Kj,h)=∑h=1dkVar(Qi,h×Kj,h)⏟All entries of Q and K are independent so Cov(Qi,h,Kj,h)=0.\underbrace{ Var(\sum_{h = 1}^{d_k} Q_{i, h} \times K_{j ,h}) = \sum_{h = 1}^{d_k} Var(Q_{i, h} \times K_{j ,h}) }_{ \text{All entries of Q and K are } \textbf{independent so } Cov(Q_{i, h}, K_{j ,h}) = 0.}

= ∑h=1dk1=dk\sum_{h = 1}^{d_k} 1 = d_k, so Var(QKT)=dkVar(QK^T) = d_k.

III. Assembly

Scale QKTQK^T by 1dk\frac{1}{\sqrt{d_k}} to ensure Var(QKTdk)=1\text{Var}(\frac{ QK^T }{ \sqrt{d_k} }) = 1:

  • 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:

  • P=Softmax(QKTdk)=Rn×nP = \text{Softmax}(\frac{ QK^T }{ \sqrt{d_k} }) = R^{n \times n}: with ∑j=1nPi,j=1\sum_{j = 1}^n P_{i, j} = 1, Pi,jP_{i, j} is ithi^{th} token’s relative proportion of attention to jthj^{th} token.

  • V=Rn×dv:V = R^{n \times d_v}: each Vi,:=RdvV_{i, :} = R^{d_v} means ithi^{th} token’s value representation.

  • PVPV yields a matrix of Rn×dvR^{n \times d_v}, where each ithi^{th} row vector is a weighted sum of token vectors.

  • Weights are based on ithi^{th} token’s attention to all nn tokens.

Be sure to set axis=-1 when doing Softmax. Along columns 😉😏

IV. 🤿Mask If Required

After doing QKTdk\frac{ QK^T }{ \sqrt{d_k} }, you have a square matrix M=Rn×nM = R^{n \times n}, 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 ii, set Mi,j=−∞  ∀  i<jM_{i, j} = -\infty \; \forall \; i < j to have e−∞=0e^{-\infty} = 0.

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: hh 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 Q=K=V=Rn×dmodel.Q = K = V = R^{n \times d_{model}}. Each lthl^{th} head owns 3 weight matrices called:

  • WlQ=Rdmodel×dkW_l^Q = R^{d_{model} \times d_k} — Help produce Ql=QWlQ=Rn×dkQ_l = Q W_l^Q = R^{n \times d_k}.

  • WlK=Rdmodel×dkW_l^K = R^{d_{model} \times d_k} — Help produce Kl=KWlK=Rn×dkK_l = K W_l^K = R^{n \times d_k}.

  • WlV=Rdmodel×dvW_l^V = R^{d_{model} \times d_v} — Help produce Vl=VWlV=Rn×dvV_l = V W_l^V = R^{n \times d_v}.

So Ql,Kl,VlQ_l, K_l, V_l all fit into expected input shapes of scaled dot-product.

lthl^{th} scaled dot-product attention is Al=Rn×dvA_l = R^{n \times d_v}. We will have hh such matrices.

With each row referring to a token, we of course can only concatenate attention along columns.

By doing so, we obtain O=Concat(A1,...,Ah)=Rn×hdvO = \text{Concat}(A_1, ..., A_h) = R^{n \times h d_v}.

Using another weight matrix WO=Rhdv×dmodelW^O = R^{h d_v \times d_{model}}, we can reach OWO=Rn×dmodelO W^O = R^{n \times d_{model}}.

Yes, back to original shape of Q,K,VQ, K, V, but carrying learned attention.

🤔Subtle & Vital Guideline

For MHA, we need extra 3h+13h + 1 weight matrices:

  • First 2h2h matrices — all Rdmodel×dkR^{d_{model} \times d_k} that deal with Q,KQ, K.

  • Middle hh matrices — all Rdmodel×dvR^{d_{model} \times d_v} that deal with VV.

  • Last 1 matrix — Rhdv×dmodelR^{h d_v \times d_{model}} that brings back to original shape.

All these 3h+13h + 1 matrices share a common: dimensions contain no nn.

At the End of the Day

Attention is to grasp semantic relationship among nn tokens.

dmodel,dk,dv,hdvd_{model}, d_k, d_v, h d_v are all semantic spaces’ dimensions. nn isn’t.

These 3h+13h + 1 weight matrices are born to travel around semantic spaces. 😉😏

🕰️Time Complexity

QKT:O(n2×dk)QK^T: O(n^2 \times d_k)

  • Rn×dk⋅Rdk×nR^{n \times d_k} \cdot R^{d_k \times n} — n2n^2 pairs of RdkR^{d_k} doing inner product, with O(dk)O(d_k) per pair.

dk\sqrt{d_k} Scaler & Softmax: O(n2)O(n^2)

  • They both perform constantly on all n2n^2 entries inside QKTQK^T.

V Multiplication: O(n2×dk)O(n^2 \times d_k)

  • Rn×n⋅Rn×dkR^{n \times n} \cdot R^{n \times d_k} — n×dkn \times d_k pairs of RnR^n doing inner product, with O(n)O(n) per pair.

So total is O(n2×dk)+O(n2)+O(n2×dk)=O(n2×dk).O(n^2 \times d_k) + O(n^2) + O(n^2 \times d_k) = O(n^2 \times d_k).

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 QQ, KK and VV.

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 KK & VV, since it gets referenced.

  • We are allowed to look at every step. So no masks at all.

Tip: Output Is Topic Regardless

AttentionQQ sourceKK & VV sourcesMasked
MaskedOutput.Output.Yes.
CrossOutput.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 kk, PEpos+kPE_{pos + k} can be represented as a linear function of PEposPE_{pos}.

“Attention Is All You Need” uses sinsin and coscos 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.