At a glance
Sasha Rush opens by saying this one is going to be different. Normally he talks about new research, brings in a lot of citations, and goes into hardcore technical detail. Today he is giving a tutorial, it is a bit casual, and the goal is intuition: "I'm going to do this by presenting intuition about five core formulas that help me understand language models in general."
The framing question is not how large language models work. It is narrower and more honest than that. "My goal is to develop a language that lets me reason about how large language models work. Now the problem with this is that we do not really even understand how small language models work. I'm not really an optimist at heart so I'm not going to tell you that we're going to figure this out very soon. However, I can say that there are some specific areas where we can reason about the behavior of large language models relatively precisely."
Five such areas, five formulas. He names them in the first ninety seconds so nobody is left hanging, and then spends an hour building each one from the ground up: the toy language with 10,000 words, the binary codes, the lower triangular matrices, the exact global memory read counts for a 3x3 matrix multiply done two ways, the parameter and token counts of three real models, and a handful of RASP programs that compile to transformer weights. Every section ends with him saying where that formula stops being useful.
At the time of the talk he was an associate professor at Cornell Tech and a researcher at Hugging Face, and the tutorial began life as an invited session for the Harvard Data Science Initiative. He has since moved to Cursor. This page rebuilds the whole hour in his order.
The five, named up front
At [0:01:32] he puts all five on one slide, with their mysterious names and their plain English translations side by side:
| Formula | What it governs |
|---|---|
| Perplexity | Generation |
| Attention | Memory |
| GEMM | Efficiency |
| Chinchilla | Scaling |
| RASP | Reasoning |
"Now if you haven't seen these before those will be pretty mysterious names. If it helps, these correspond to generation, memory, efficiency, scaling and reasoning."
The one caveat he puts before everything else
There is exactly one disclaimer in the talk and it arrives at [0:01:32], before any content: "I'm going to be focusing on conceptual understanding so I'm going to simplify a lot of details and probably get things wrong. You should think about this as acting in a kind of frictionless environment. This is a kind of idealized version of language modeling that focuses more on understanding how to think about the system than on the specifics."
Worth holding on to. Every simplification he makes after this is deliberate, and he flags most of them as he goes.
Part one: generation, and the formula for perplexity
The toy language
He builds a miniature world first, at [0:02:04], so the numbers stay concrete for the whole section:
- A collection of documents, each made up of exactly 1,000 word tokens.
- A vocabulary of 10,000 word types. Think of this as the dictionary.
A language model is then a probabilistic model of a document. It gives the probability of the tokens x_1 through x_T, and it uses a set of parameters theta:
p(x_1, x_2, ..., x_T ; theta)
He likes this notation for one specific reason: "it lets us isolate the parameters theta separate from the probabilistic model." Since theta is going to be a giant neural network, writing it this way lets him defer the entire problem of defining that network to section two. Part one is about the probability; part two is about the network.
Factoring the document left to right
The probability of the document is the joint probability of all its word tokens, so standard rules of probability let you write it however you like. The useful way is as a product of conditionals, factored left to right:
p(x_1, ..., x_T) = p(x_1) * p(x_2 | x_1) * p(x_3 | x_1, x_2) * ... * p(x_T | x_1, ..., x_{T-1})
Parameterize those individual conditionals and you have an autoregressive language model: a predictive model where you predict the next token conditioned on the previous ones. He unpacks the name at [0:03:35]: "we call it autoregressive where auto refers to the fact that we're feeding back in previous predictions and regressive refers to the fact that we are predicting the next token."
More tangibly: each conditional produces a distribution over 10,000 different choices. Every word in the dictionary gets a probability. He draws it as a histogram over all possible next word choices, and that histogram is the picture he reuses for the rest of the hour.
The nice thing about having the joint distribution is that you can sample a document from it. Sample x_1, feed it in as the condition for the next step, sample x_2, feed that in, and keep going. That is the autoregressive process made concrete.
Everything so far is a hundred years old
At [0:04:35] he stops and makes a point of how little he has said: "we're about 5 minutes in and honestly I haven't told you anything new. In fact basically everything I've said so far was figured out by Markov about 100 years ago."
Andrey Markov and the Markov chain are the whole foundation here, and the reason to walk through the basics anyway is that it lets him name two assumptions people used to make in language modeling that no longer hold in modern systems.
Assumption one: fixed history. Until recently it was very common to assume that the probability of x_t really only depended on a few of the previous words. The most aggressive version is that x_t depends only on x_{t-1}:
p(x_t | x_1, ..., x_{t-1}) ~= p(x_t | x_{t-1})
He is fair to it: "this seems like an aggressive assumption but x_{t-1} definitely provides the most information about predicting x_t and words further away outweigh less."
Assumption two: the categorical assumption. That the probability of the next word could be modeled roughly with a categorical distribution, which is to say a lookup table of counts. "This assumption went away when people started using neural networks to model the probability of next word prediction."
Both assumptions are present in Claude Shannon's pioneering 1948 work, A Mathematical Theory of Communication, where he develops some of the first language models: a one step categorical language model, sampled exactly the way Rush described a minute earlier. "This model is actually not so different from a lot of the language models that were developed before the modern resurgence of neural networks."
Why you cannot just measure accuracy
Shannon's model is not great, Rush says, but that raises the measurement problem, because "language modeling is in some sense unsupervised learning, it's kind of hard to quantify when a model is good or when it's bad." The procedure has to be: give it unseen text, check how close its predictions are to what the text actually says, compute a metric, use that metric to compare models.
The naive thing to try is accuracy. Take the running example he uses all hour:
the dog walked to the ____
Look at the mode of the distribution. The model predicted park. The true answer was lawn. Zero out of one.
"Which is, I guess, right, but a bit unsatisfying. We were pretty close, we almost got the right answer, but we get zero points."
It gets worse, because words in language follow a Zipfian distribution. Very common words make up a very large portion of the probability mass, but very uncommon words are seen pretty frequently. So not every prediction is equal: "a lot of the time we're going to be predicting pretty common words like the or a, but not too infrequently will we have to predict very challenging words like pizza or raincoat." An accuracy number averages those together as if they were the same event, and they are not.
From probabilities to bits
The fix, at [0:07:41], is to stop scoring guesses and start counting bits. Convert the probability distribution into a string of binary values, one per word in the vocabulary. The conversion can be done deterministically, but the intuition he offers is that you are placing a bet on each word, and you only have so many values you can allocate, so you use the probabilities to choose the length of the string you assign to each word.
In his example:
| Word | Probability | Code | Length |
|---|---|---|---|
| park | high | 101 |
3 bits |
| lawn | lower | 100101 |
6 bits |
| totally irrelevant words | tiny | very long strings | many bits |
And you can show that the optimal length of these strings is roughly
length(w) ~= -log2( p(w) )
which "explains why words with very low probability will have very long strings, and the closer we are to one the closer we are to a short string."
He is careful to say this is not an abstraction: "this is not just a theoretical conversion, these strings literally give us a way to compress the underlying language and communicate it to a party that has access to our language model." A language model is a compressor. The code lengths are the compressed message.
The formula
That conversion between probability and bits gives the main metric used in language modeling. Perplexity, for historical reasons, is two to the average number of bits per word on a held out test set:
perplexity = 2 ^ ( -(1/T) * sum_{t=1..T} log2 p(x_t | x_1, ..., x_{t-1} ; theta) )
Read the formula from the inside out and it is exactly the recipe he describes at [0:09:14]: "we simply compute the probability of each of the true next words, take a log base 2, average, negate, and then send to a power of two." The exponent is the average code length. Raising two to it "puts it in a slightly easier form to work with for users of language models."
Every symbol, since this is the one formula everything downstream rests on:
| Symbol | Meaning |
|---|---|
x_t |
the true next token at position t in the held out text |
T |
number of tokens scored |
theta |
the model parameters, deferred to part two |
p(x_t | ...) |
probability the model assigned to the correct next token |
log2 |
base two, so the units are bits |
-(1/T) * sum |
average bits per word, negated so it is positive |
2 ^ (...) |
the historical convention that turns bits back into an effective vocabulary size |
What the extreme values mean
He commits the formula to intuition by walking the ends of its range at [0:09:44].
Perplexity of 1. This implies the number of bits needed per word is zero. "How's that possible, don't we need to communicate something about the next word? Well in this case we actually don't. If our perplexity was one, the person we're talking with basically knows the next word. It's always whatever had the highest probability. We don't actually need to communicate anything."
Perplexity of 10,000, the vocabulary size. This implies you need a string of log2(10000) bits to communicate any word in the dictionary, which "means we're not really getting any advantage from language modeling at all. It implies that our model is basically uniform. If we were kind of trying to guess the next word we would basically have to roll a 10,000 sided dice."
Perplexity above 10,000. This is possible, and he flags it specifically: it happens "if our model was overly confident about the wrong prediction. It might assign a very short code to the wrong next word and an extremely long code to the correct next word. This would lead to a very bad perplexity value as we would be spending a huge amount of bits communicating the next word which we thought could never actually happen." Confident and wrong is worse than uncertain.
That is the whole intuition: perplexity is the effective number of equally likely choices the model is still deciding between. One means certainty. Ten thousand means it learned nothing.
The Wall Street Journal ladder
For a long time, people studied language modeling on the Wall Street Journal corpus, which Rush describes as "using a couple years of newspaper articles as your training data and then trying to assess your perplexity on today's newspaper." He then gives the whole history of the field as a single descending ladder of numbers, at [0:11:14]:
The numbers, in his words:
- A uniform distribution gives roughly 10,000, as the formula predicts.
- Give yourself the previous word and you get down to around 600.
- The previous two words gets you nearly to 200.
- The previous four words, plus about 40 years of clever tricks, gets you all the way down to about 140.
- When deep learning started looking at the problem, early Markov neural networks got to a similar performance of around 140.
- When people started developing early non Markov neural networks that looked at longer ranges, they could get down to around 100 or even lower.
"This is greatly simplifying a very rich and interesting literature, but it gives you a rough sense of about where things stood around 2015."
And then, much later in the hour, the comparison point he throws in at [0:14:50]: "while it's not totally a fair comparison, we can go back and apply GPT-3 to the challenging Wall Street Journal test corpus, you get a perplexity of 20.5, a major leap from some of the earlier language models applied to this task."
Why anyone cared about perplexity in 2015
Rush asks the question directly at [0:12:15]: in 2015, why did people care about this problem at all? Two answers, and they arrive in order.
First answer: it was a usable proxy for tasks people did care about. Machine translation, for instance. Translating a sentence from French to English can be posed as a conditional language modeling problem: instead of conditioning only on the previous English words, you also condition on the French sentence, and then you measure the conditional perplexity of generating the English words given the French input.
What researchers found is that this correlated very strongly with downstream performance. He shows a table of perplexity values for a translation experiment against the corresponding translation accuracy, measured with "a bespoke measurement known as BLEU score", which is the BLEU metric from Papineni and colleagues. As the perplexity went down, the BLEU score went up.
Second answer, and the one that mattered. "However if this was all that happened, very few people would care outside of NLP. The major result that people found next was that perplexity on the task of language modeling by itself could be used to produce models that would be really good at tasks that the model had never seen or had just seen a few examples of."
The table he shows for this is specific. Raw language modeling perplexity goes from 5.84 down to 3.23 across a series of experiments. At the same time, three other very different downstream tasks all get significantly better. And the key detail: "these tasks were not included as part of the original training data but instead were given as a small set of examples after the fact."
Just by reducing the perplexity on general purpose language, the model became usable on tasks nobody trained it for. "This idea of course now underlies all of modern large language modeling research." The field that result opened up is the one documented in Language Models are Few-Shot Learners, the GPT-3 paper.
Llama 2, and trusting perplexity on its own
How far that trust now extends is his closing point for part one, at [0:14:18]: "in modern papers like Llama 2, sometimes they really just show the perplexity and people trust that things will work really well."
The Llama 2 table he displays has four models of varying sizes, each with better perplexity than the last:
- Final perplexity goes from about 1.8 down to about 1.5 across the four sizes.
- The bottom model, Llama 2 70B, "still remains one of the best open source large language models, and we can trust that because it has an extremely good general perplexity."
- "As we've seen earlier, a perplexity of 1.5 means that the model is really capturing the distribution of English well."
Hold that against the extremes he established ten minutes earlier. A perplexity of 1.5 is half a bit of surprise per token. On his own scale, 1 is a reader who already knows what you are going to say.
Where part one leaves off
"But I'm getting ahead of myself. I still haven't told you how you actually get from Shannon's model to GPT-3, and to do this we're going to need to remove the two major assumptions. We're first going to have to move to the use of neural networks, and then we're going to need to figure out how we can take into account all previous word tokens."
Assumption two falls in part two's first five minutes. Assumption one takes the rest of the section.
Part two: memory, and the attention equation
Into theta
Part two opens at [0:15:20] by picking up the thing part one deliberately set aside: "we're going to dive into the theta, we're going to try to understand better how to create a neural network that can power the probabilistic model that predicts the next word."
He does not drop both assumptions at once. He keeps the Markov assumption for now, looking at only the last two words, and swaps the categorical table for a neural network:
p(x_t | x_{t-2}, x_{t-1}) = softmax( NN( x_{t-2}, x_{t-1} ) )
Run a neural network over x_{t-2} and x_{t-1}, get a vector out, send it to the softmax function, and the softmax converts that vector into a distribution over the 10,000 possible classes. That distribution is the language model.
One hot in, 10,000 out
"Well it took a while to figure out the best neural networks to use. Looking back, the form of all these networks is relatively straightforward." The pipeline, at [0:16:22]:
- Encode
x_{t-2}andx_{t-1}as one hot vectors: "vectors that have zeros for every position and a one for the position of the word token they represent." Each is length 10,000. - Feed both vectors into a neural network, which processes them in the standard form.
- The network outputs another vector of size 10,000.
- Learn the whole thing on language.
And the deflationary aside he cannot resist: "all in all it's probably about five or 10 lines of PyTorch."
Step four is the softmax, and he spells out what it actually does: "this is done by applying the softmax function, which ensures that the output is positive and sums to one. We do this by exponentiating each element of the vector and then normalizing." Which produces exactly the histogram over 10,000 next word choices from part one.
softmax(z)_i = exp(z_i) / sum_j exp(z_j)
He is explicit about how much he is skating over: "obviously I'm being pretty casual about this process. There were many forms of early neural network language models and they all had innovations that made it possible to get to this point." The famous one he names is word2vec, from Mikolov and colleagues: "this model came out exactly 10 years ago and it demonstrated a lot of techniques that became foundational to later models." The talk is from January 2024, word2vec from 2013, so "exactly 10 years" is a precise claim.
One forward reference to part three, planted here: "it was much harder to build models like this at that time, particularly because the compute infrastructure and hardware was less developed."
Easy bits and hard bits
Once you have the infrastructure you can push more data through it and build bigger networks. You can take a Markovian neural language model and "simply replace the internals with a much larger neural network." At [0:18:24] he says why that alone is not enough.
"One of the problems is that while the most recent words are particularly important, there are many bits that are hard to recover without looking at longer term context." Even with a very large network, the last two words may not contain enough information.
Back to the running example. Predict the next word in the dog walked to the ____ using only the last two words, and you are predicting the blank from to the. "This tells us that it's a location and a noun, but it really doesn't tell us much about the semantics of the sentence itself. We don't know who is going or what the verb was, and that information can really help us get some of the harder bits in this prediction problem."
Then the example that drives the rest of the talk, and the one he returns to in part five. A newspaper article mentions a person's name in the intro paragraph. Later it brings up the same person and you need to recall what their last name is.
"In theory this is relatively rare, and mostly you can get the easy bits just by saying some proper noun. However to get the hard bits of exactly who that person was requires a very long term memory."
This is the argument for fully autoregressive models, which is to say models with the ability to use all previous tokens. The split between "easy bits" and "hard bits" is doing real work in that argument: the easy bits are cheap, high probability, and recoverable from local context. The hard bits are the ones a two word window cannot reach, and they are the ones that make a model feel like it understands the document.
Why a longer window does not fix it
The obvious objection, which he raises himself at [0:20:26]: can you not just take the Markov neural network and make it much longer?
"The problem is that if you do it this way you end up learning very specific information about absolute positioning. For instance you might learn some particular information about position 7, but language doesn't really work that way. There is no real specific information about position 7. It's going to depend on context and the dynamic structure that gets constructed in the document."
That is the whole reason a learned fixed window fails: it binds knowledge to slots instead of to content. What you want is addressing by content, not by index.
Attention as a lookup table
"The way to think about attention is to think about a neural network version of random access memory, or even simpler as a neural network version of a lookup table. We're going to save all the previous information and then refer back to it as we need it."
Three pieces of information, at [0:21:28]:
| Piece | What it is |
|---|---|
| query | one vector that "has looked at the whole sentence so far" |
| key | one per previous position, what that position offers to be matched against |
| value | one per previous position, what that position hands back if matched |
And three steps: "the query matches the key, the best match is selected, and then we return the value of that key." The returned value is passed to a neural network, which predicts the next word.
The problem with argmax
"However this process has a foundational issue. The problem is that we'd like the whole thing to be embedded within a neural network. Neural networks learn through the use of derivatives. The problem is that the argmax operation that would be used to select the best key, does it have a useful derivative?"
It does not. "If we write it as a one dimensional function we can see that we get a flat structure with a derivative of zero." A flat function carries no gradient signal, so nothing upstream of a hard selection can be trained.
The fix is to replace argmax with softmax, which he has already introduced for a different purpose. Used here, "this softmax, instead of producing a distribution over word types, produces a distribution over previous token positions. This distribution is computed softly and has a non trivial derivative at every location." In two dimensions, "this softmax function can be drawn with a sigmoid shape which has a nice derivative at every location."
The three steps, softened
With that substitution, at [0:23:01], the process becomes:
- "We use the query to score our keys."
- "Instead of picking the highest scoring value we instead use a softmax to normalize the scores."
- "Instead of producing a single value we utilize the softmax to average over the different values, weighting them by how well their key matched the query."
Walked through on the example: start with the same query, keys and values. Match the query to the keys to get a score for each location. Compute a softmax, which gives a histogram over key locations. Use that histogram to average the values together. "This produces a new vector that is some average of the different values weighted by the histogram." Final step unchanged: use the weighted value to predict the next word.
"This process is fully differentiable and is a good way to learn a neural network that can decide which previous words are useful for the next prediction."
Attention Is All You Need
"This attention operation was central to the key work in language modeling known as Attention Is All You Need. This paper written in 2017 introduces a neural net architecture known as a Transformer that relies heavily on this attention step."
He waves at the gap between his version and the real one: "there are many extensions beyond the simple version that I've shown earlier, but roughly the idea holds. You have some way of computing keys, queries and values, and then you use attention repeatedly until you're ready to predict the next word."
On the architecture diagram, at [0:25:05]: "the diagram of the Transformer architecture has become quite iconic in the field. It roughly consists of two stages. The first stage is the attention that we just previously saw and the second stage is a rather large standard neural network. These two stages are repeated many times before the final prediction is made."
And then the line that is the whole reason to teach it this way: "this diagram looks a bit complex, but note that this describes basically the entire GPT system. Given how important that is as a large language model, it's actually surprisingly simple."
The real reason attention won
Here is where he refuses the easy answer, at [0:25:35]. "I've told you what attention is but not why it's the best way to do this sort of long form language model. The tempting answer is to say that attention is kind of like memory, and so it makes sense that this sort of architecture would actually win out in practice. The real answer though is a bit different."
"It turns out that the sort of attention that's used in the Transformer architecture happens to be very efficient and parallelizable. It's a nice combination of long range dependency and something that runs fast on modern hardware."
Attention won on hardware, not on cognitive plausibility. That claim is the hinge between part two and part three, and he spends the last two minutes of part two proving the parallelizable half of it.
Batching the queries into matrices
In practice you compute several queries simultaneously, so group them into a matrix. The memory that holds the keys and values can be written as two more matrices. Then, at [0:26:35]:
- Combine the query matrix and the key matrix and compute the softmax, and "we end up with many different histograms representing all the combinations of the queries and the keys."
- That whole set of histograms "can be computed by running a softmax over the matrix product between the queries and the keys."
- "Similarly these can then be combined with the value matrix to compute the set of weighted averages simply by taking a matrix multiply between the softmax of the queries and the keys and the value matrix."
"The punchline of this process is that the entire attention step can be written as two matrix multiplies with a softmax around the inner one."
That is the second formula, in the form he presents it:
attention(Q, K, V) = softmax( Q K^T ) V
| Symbol | Meaning |
|---|---|
Q |
the queries, one row per position being predicted, stacked into a matrix |
K |
the keys, one row per previous position |
V |
the values, one row per previous position |
Q K^T |
every query scored against every key, in one matrix multiply |
softmax(...) |
applied per row, turning each row of scores into a histogram over positions |
... V |
the weighted averages, one per query, in a second matrix multiply |
He reads it back against his own three steps so nothing is left implicit: "the queries score the keys through a matrix multiply, a softmax then normalizes the scores to produce the histograms, and then the weighted average is taken with the values by computing another matrix multiply. This whole thing is just a series of softmaxes and matrix multiplies."
One thing his simplified form drops: the published equation divides Q K^T by the square root of the key dimension, sqrt(d_k), to keep the scores from saturating the softmax at large dimensions. He is working in the frictionless environment he declared at the top, and the scaling factor is not what the section is about.
The handoff
"So far we've now described a full autoregressive language model. This model is known as a generative Transformer. But what I haven't told you yet is why this actually runs fast in practice. In the next section we'll dive deeper into this matrix multiply operation and see how do we actually make it run fast on GPUs."
Part three: efficiency, and the GEMM
The honest answer is GPUs
Part three opens at [0:28:08] with him answering a question he has been dodging all talk: "I've been bouncing around this question of why language models suddenly got so much better, and if I'm honest the answer is pretty simple, it's mostly because of GPUs."
Then the joke, which is also the slide: "I can show you a graph of the speed of GPUs over recent years, but in some sense showing you the graph of Nvidia stock price gets the same point across."
"Over the last several years we've seen the core centrality of hardware in the process of building bigger and more powerful large language models. As a deep learning practitioner, the rise of general purpose programming on GPUs has fundamentally altered what sorts of models were possible to be built."
A research literature that evaporated
His first demonstration of what "altered what sorts of models were possible" means is a small piece of field history, at [0:28:39], and it is the softmax again.
To compute a softmax you have to normalize, which means taking a sum over every word type in the vocabulary. "The sum is surprisingly big, it can be upwards of 10,000 different choices."
"If we think about language models circa 2010, there was lots of research into effectively approximating the denominator of the softmax function. If we could come up with some way to efficiently approximate the denominator, we could compute the softmax more efficiently on CPU hardware. But after the introduction and widespread use of GPUs this totally changed. It turns out that this denominator is pretty trivial to compute on GPUs, and all of a sudden we no longer had to figure out fancy ways of approximating this function."
An entire subfield existed to work around one sum, and a hardware change made the problem go away. That is the cleanest statement in the talk of why the efficiency formula deserves a slot next to the mathematical ones.
Matrix multiplication is the same story with higher stakes: "matrix multiplication is central to every part of deep learning. It's used within neural networks itself, and we've seen that it's central to the computation that's necessary for attention. If we can map an operation into some form of matrix multiplication then we can certainly run it fast on new GPUs."
So he uses it as a running example to teach GPU programming from scratch, which is the most hands on stretch of the hour.
What a GPU actually is
"So first off, what is a GPU? At a kind of high level approximation you can think about this as just being a parallel computer. A GPU has many threads and they all run the same code simultaneously. To make things less intimidating we'll think of each GPU thread as being a little robot. The robot can do mathematical operations and it can read and write from memory."
Three levels, at [0:30:11]:
| Level | What it is | Memory it can touch |
|---|---|---|
| thread | one little robot, running the same code as all the others | its own registers |
| block | a group of threads. His example picture has 12 threads in one block | block memory, shared by the block, "quite efficient" |
| grid | the whole GPU, all the blocks | global memory, shared by the entire grid, "quite inefficient to read and write from" |
"The main rules of GPU programming are that we're only able to have a limited number of threads per each block, but the blocks will be quite important. The blocks are essential because reading and writing from global memory is much much slower than utilizing our block memory. We're going to want to do as many operations as possible within the block as opposed to resorting to the global memory."
The anti pattern and the pattern, stated plainly at [0:31:45]:
- Bad: "it would be very bad if each of the individual threads was reading and writing to global memory by themselves. You'd get parallelism but it would be really slow to calculate things."
- Good: "in the ideal world we first load from the global memory into our local block memory, we then do some computation where we compute things within the block itself, maybe read and write from the local memory several times, and then when we're done write back out to global memory with our final answers."
"This is the main trick for GPU programming, but it can be a little bit counterintuitive, and seeing it applied in practice for the first time can be quite challenging."
So he counts the reads.
The 3x3 multiply, done naively
Two square matrices, A and B, both 3x3. Each element of AB is computed by multiplying a column in A with a row in B and summing the results. He works the second row, first column entry on the slide.
Done naively, one thread computes each of the nine outputs. To compute one value, a thread does:
- Six global reads: the column of
Aand the row ofB, three values each. - Each of the multiplications.
- Sum them up inside the thread.
- Write the answer back out to global memory.
"Note that doing this requires six global reads for each thread."
naive 3x3: 6 global reads per thread x 9 threads = 54 global reads
The same multiply, done in blocks
"A better method is going to be to first read from global memory into the block's memory. We can then calculate important intermediate results within the block itself and then finally write back out to the final value."
Step by step, at [0:33:16]:
- Read from global memory into block memory, "reading the whole matrix into our block memory, yielding 2 x 9 global reads."
- "Once we have the matrix in our block memory we can have each thread do the same operations we saw before, but now the reads are from the shared memory, not from the global memory. This is much much faster in practice."
- "Since the memory is shared, different threads can reuse the same rows and columns to compute different positions in the output matrix."
- When done, copy the shared
ABmatrix back out to global memory for use in the next operation.
blocked 3x3: 2 x 9 = 18 global reads, all in the first stage
"If we look at the number of reads, with the naive method each thread did six global reads and there were nine threads total, yielding 54 global reads. If we do our block based method all of the reads happen in the first step, which yields 9 x 2 reads for a total of 18."
Point three is where the win actually comes from. Reuse is the whole mechanism: the naive version reads the same row of B three separate times from the slow memory because three different threads each need it. Loading it once into shared memory and letting all three threads read it from there is the entire optimization.
6x6, when the matrix does not fit
"This is many fewer per thread, but that's the case where the entire matrix fit into our block. Recall that I mentioned that blocks have to be a fixed size, and so we can't scale this approach to arbitrarily large matrices."
So: how would you do a 6x6 matrix multiply with blocks of the same size? "The answer is that you end up having to do it in multiple steps." At [0:34:51]:
- Pass one. Instead of reading in the entire matrix, read in a 3x3 block of each of
AandB. Again2 x 9reads into the 3x3 blocks. - Calculate a part of the final value by multiplying together the three values of the column of the top matrix and the three values of the row of the bottom matrix. "We then use a single thread reading from block memory to compute the cell."
- Pass two. "Once this part is done we copy in a new part of the two original matrices into our shared memory. For the top matrix we do the bottom part, and for the bottom matrix we do the right part. We then use our threads to multiply together these components and sum them into the final value."
- "Between the first and the second part of this process we now have computed the full multiply between the row on the bottom and the column on the top. This gives us the correct answer for the 3x3 block in the A times B matrix."
- Write that 3x3 block back out to global memory. "This gives us a 3x3 part of the full A by B output matrix."
- "While this is happening, other blocks are completing the rest of the A times B matrix. Each of these are again only doing 2 x 9 reads each time."
blocked 6x6: 2 passes x (2 x 9) = 36 global reads per block
"In this case we end up doing 36 total global reads per block in order to compute the full final matrix."
That is the real algorithm, and it is the reason the section exists. A matrix multiply that does not fit in fast memory is not a different algorithm, it is the same algorithm run in tiles, with a partial sum accumulated across passes. Everything about how these models are actually executed comes out of that one restructuring.
Why it is called GEMM
"And that's the main operation for this section. I probably should have just called it matrix multiplication."
He names it properly at [0:36:24]: "in practice this operation is often called GEMM when applied on GPUs. The GPU operation lets us do a generalized version of this matrix multiply that also allows us to add in an additional term and to scale the operation, but you get the idea. We can do this sort of low level efficient matrix multiplication by exploiting all the power of GPUs."
The generalized form he is describing, the one the BLAS interface has had for decades and the one cuBLAS implements:
C <- alpha * (A @ B) + beta * C
The extra term and the two scalars are exactly the "additional term" and "scale the operation" he mentions, and they are there because fusing a scale and an accumulate into the multiply saves another round trip through memory.
How much of this he is leaving out
"As I mentioned earlier, it's really really hard to underestimate the importance of this operation for modern neural networks. It's used in basically all the main parts of the system. It's particularly important for the calculation of attention as well as the core neural network blocks that are utilized in a Transformer."
And the caveat for part three: "I've also only shown you the most basic form of matrix multiplication. In modern GPUs there's all sorts of specialized hardware that is continued to make the calculation of matrix multiplies even faster and faster with each release."
If this section is the one you want to actually practice, Rush maintains the teaching materials for it: GPU-Puzzles builds up exactly this reasoning as a sequence of small problems, and Triton-Puzzles does the same one level up.
The handoff
"But you might wonder why we're focusing so much on speed. Isn't enough enough? You can run a language model on your laptop, isn't that good enough? Why do we have to optimize it so intensely? You might wonder where all this compute is going and why people are fighting over buying up all the newest latest GPUs."
Part four: scaling, and Chinchilla
Two knobs, one multiplicative budget
Part four is about pretraining, and specifically about "the question of scaling these sorts of language models to be trained on lots and lots of data with very large neural networks."
"Unlike some of the other sections of this talk, the key decisions in scaling seem quite simple." There are two:
- How big a neural network. "That means how many parameters should we try to fit in our neural network layers."
- How much training data. "Roughly how many documents or tokens should we train our model on. Training on more tokens allows the model to fit the data better and potentially have a better perplexity."
And then the structural fact that makes the problem interesting, at [0:38:57]: "what's interesting about these two variables is that they form a multiplicative relationship. If we take the neural network size and we take the training data, the amount of total compute we need to dedicate to pretraining scales as the product of their two sizes."
compute ~ parameters x tokens (N x D)
The intuition he gives for why it is a product, not a sum: "each token that goes to the neural network needs to touch every one of the neural network parameters. This is what forms the multiplicative relationship and forms the total compute of the system."
Three famous models, two with the numbers on screen
At [0:39:29] he reads the relationship off real models.
| Model | Year | Parameters | Training tokens | Total compute |
|---|---|---|---|---|
| BERT Base | 2018 | about 109 million | 250 billion | about 1.6e20 |
| PaLM | 2022 | 540 billion | 780 billion | about 2.5e24 |
"So models are getting bigger, they're being trained on more data, and more importantly they're utilizing more and more compute. If we can utilize the compute to the best purpose we can get better language models."
Two things worth doing with those four numbers, both of which the page is adding rather than quoting.
First, the constant he left out. He says compute scales as the product and does not give the factor. His own figures pin it down exactly, because 6 x 109e6 x 250e9 = 1.6e20 and 6 x 540e9 x 780e9 = 2.5e24. Both land on his stated values, so the slide is using the standard forward plus backward pass count:
C = 6 * N * D N = parameters, D = training tokens, C in FLOPs
Second, the gap between those two rows. 2.5e24 / 1.6e20 is about 15,600 times more compute in four years. But look at the allocation rather than the total:
- BERT Base:
250e9 / 109e6= about 2,294 training tokens per parameter. - PaLM:
780e9 / 540e9= about 1.4 training tokens per parameter.
Those two models are at opposite extremes of the exact trade off the rest of the section is about, and the ratio between their ratios is roughly 1,600 to one. He gets to that point about four minutes later.
The power law, and why "make everything bigger" is not an answer
"So you might ask which of these variables we want to change." In the OpenAI Scaling Laws for Neural Language Models paper, "researchers demonstrated that for each of these quantities the perplexity of the model is going to improve as a power law. Roughly this means that if we make a log log plot of perplexity versus each of these individual powers we get a linear line showing the decrease in perplexity as we make a large increase in compute, parameters or data size."
The naive conclusion follows immediately: "a natural conclusion is just to increase all of these parameters as much as possible. They all seem to help performance so let's just make them all as big as we can."
And the problem with it, at [0:40:34]: "even if you're extremely GPU rich you still have a compute budget and you have to determine how to best utilize the compute that you have available."
The same area, two shapes
His slide for this is two rectangles. "These two diagrams both utilize roughly the same amount of compute, or area in the diagram, but the one on the left allocates more compute to utilizing more tokens to train on whereas the one on the right utilizes that compute for a model with more parameters. How do we determine which one would yield the best perplexity in the end?"
The rectangle picture is doing real work: because compute is N x D, a fixed budget is a fixed area, and the only question is the aspect ratio. Tall and thin is parameter heavy. Wide and flat is token heavy.
"This is not really a theoretical problem. For example the PaLM model which I mentioned earlier utilized a very large amount of parameters and actually relatively few tokens. You might ask if this was the best thing they could have done or if they could have done better."
The formula, term by term
"The approach we'll use to study this problem is to write down the formula for a power law and then try to fit this formula to empirical curves that show perplexity. If we can get a good fit we can maybe extrapolate onto how we should train a new model, basically whether we should use more parameters or more tokens."
He admits the formula looks a little bit complicated and draws it as a picture instead, but the picture he describes is exactly this:
L(N, D) = A / N^alpha + B / D^beta + E
| Term | What he says it is |
|---|---|
L(N, D) |
the perplexity you get for a given model size and data size |
A / N^alpha |
"some value A over the number of parameters we use to the alpha term" |
B / D^beta |
"a second term which is B over the amount of data we use to a beta term" |
E |
"an additional E term which acts as a bias and corresponds to the best possible perplexity you could get for the language" |
alpha, beta |
"the key terms of interest are the exponents on the model size and the data size. This will tell us roughly how to scale our model." |
The E term is the one worth sitting with. It says there is an irreducible floor, a perplexity below which no amount of parameters or data will take you, because language itself is not deterministic. That is Shannon's entropy of English from part one, showing up as a constant in a scaling law.
The result: the exponents are roughly equal
"This formula and its fit are explored in a paper known as Chinchilla. The main result of Chinchilla is that the exponents for the model and the data, the blue and the green box, are roughly the same. This implies that the best perplexity can be achieved with an equal scaling formula. That is, if we scale the data and the model in roughly equal proportions, we'll be able to get the best perplexity for the least amount of compute."
alpha ~= beta => scale N and D together
That is the whole finding, stated as he states it. The paper's own phrasing is that "the model size and the number of training tokens should be scaled equally", and that for every doubling of model size the number of training tokens should also be doubled. The paper's tables put the compute optimal allocation at roughly 20 training tokens per parameter, and its demonstration model, Chinchilla, is a 70 billion parameter model trained on 4x more data than the 280 billion parameter Gopher it beats. Rush does not quote the 20 to 1 ratio or the Chinchilla and Gopher sizes in the talk; he stops at "roughly equal exponents", which is the part that is actually a result rather than a fitted constant.
What that says about PaLM, GPT-3 and Megatron-Turing
"The consequence of this result is that we can go back and look at some previous models. So in particular we can note that the PaLM model, which used a scaling that favored parameters over data, was maybe the wrong decision. They ended up spending too much compute for the final perplexity and outcome of their model. A better strategy would be to scale the parameters and the data equally. You end up with something looking more like a square than a rectangle. This allows you to get a very good model with less compute and often less dollars."
Then the figure from the paper, at [0:43:38]: "the notable figure of this work shows a graph with the continuing scaling of compute represented by FLOPs in the line in the center of the graph. The Chinchilla model, which we've seen uses equal scaling, ends up being on this line, where other models such as GPT-3 or approaches like Megatron-Turing, which is 530 billion parameters, end up overparameterized for the amount of tokens that they use."
"There are of course lots of complexities here and other factors, but in general it provides a very nice key for how people can produce better models. A lot of the conversation in building large language models has centered around this notion of Chinchilla scaling and how to take advantage of various constraints."
The caveat that broke the rule in practice
This is the best caveat in the talk, because it is the one that actually changed what the industry shipped. At [0:44:39]:
"There are some really important caveats to this scaling property though. The main one is that there is an asymmetry between model parameters and training data. The training data is only utilized during the training process and is therefore a one time cost. However the model parameters have to be used for every inference of the actual large language model. So every time you actually call ChatGPT it has to basically use all the parameters of the model."
"Because of this it may actually be beneficial to use some of your training time compute to instead produce a model that's smaller and maybe utilizes more tokens. This can lead to weird asymmetries where people produce models that may not be the best model they could compute for their constraints but end up leading to better actual applications."
Read that against the formula. Chinchilla optimizes one thing: the perplexity you get per unit of training compute. It says nothing about what the model costs to run, and the model runs forever. Once you put serving cost in the objective, the optimum moves toward fewer parameters and more tokens, past the point the scaling law would pick.
LLaMA, deliberately suboptimal
"The most notable example of this was the original LLaMA model. This model was purposely suboptimal in terms of training compute but ended up producing a model that had many fewer parameters than some of the larger models with similar compute. To do this they simply paid more upfront cost to train it on more tokens and then they distributed the smaller version of the model."
The paper makes the case in one line of its abstract: LLaMA-13B outperforms GPT-3 at 175 billion parameters on most benchmarks, and LLaMA-65B is competitive with Chinchilla-70B and PaLM-540B. A model 13 times smaller than GPT-3 beating it is the asymmetry he is describing, cashed out.
What part four concludes
"So just to summarize, the key way we produce large language models is through this initial pretraining stage, particularly when working with generative pretrained Transformers. Compute is really the constraining factor and the best use of this compute can produce the model with the best perplexity. Even with GPUs we are still limited to the amount of compute we have, and so deciding on the best allocation, both during training and for downstream inference, is a core problem."
And the dissatisfaction that sets up part five: "and yet this is still a bit unsatisfactory. We know that these models learn a ton from this pretraining stage, but we're really just relying on the fact that perplexity goes down and therefore all sorts of tasks get better. What people really seem to want to know now is what the models are actually doing. What do they learn, and how come they're able to perform so well on such hard tasks?"
Part five: reasoning, and RASP
Back to the name in the first paragraph
He opens part five at [0:46:39] by calling it "part five, algorithms", which is the working title for what the chapter list calls reasoning. And he goes straight back to the example from part two:
"Let's jump all the way back to this question of memory. We discussed this example earlier in the talk where we have a New York Times article that mentions a person's name and then later refers to him in the article. In order to solve this kind of problem you need a model that's both going to be really good at the easy stuff, so able to figure out the syntax of the language and its structure, but also able to do complex algorithmic tasks such as remember a previous name and then be able to use it based on context clues."
"For many years this type of problem was actually considered very very very challenging for language models, but we've seen that as perplexity goes down with GPT based models it is able to answer these questions correctly almost all of the time."
So the question for part five is not whether models do this. They demonstrably do. The question is what mechanism could possibly be doing it.
The synthetic version: associative recall
"Sometimes the complication of natural language can make these problems more challenging to study, so people have proposed a bunch of synthetic tasks that simulate some of the interesting properties that language models display."
The one he picks, at [0:47:40], is associative memory. Strip the newspaper article down to its skeleton: a sequence of symbols in which some marker symbol appears more than once, and the task is, on seeing the marker again, to produce whatever was next to it the first time. In his example the answer is the letter A, because A is the symbol sitting next to the first occurrence of the comparison symbol.
"This is obviously an abstraction of the original task, but it is a good representation of what makes the problem hard."
The hardness is specific: the distance between the two occurrences is unbounded. No fixed window reaches it. You need content addressing, which is exactly what part two built.
"Honestly, I actually have no idea"
Then the most useful thirty seconds in the talk, at [0:48:12]:
"So you might ask, how do language models learn to do this sort of task? How are they able to make it work so well for so many different domains? And honestly I actually have no idea. There are a lot of papers written about this topic but I still don't feel like I've gotten a conclusive clear answer."
"I'm a bit more comfortable with a version of this question that looks into how models MIGHT do this. So in particular they look to build abstractions that could represent similar properties, just to show that it's either possible or impossible for a language model to accomplish certain goals."
That reframing is the whole methodology of part five. He is not going to tell you what trained transformers do. He is going to tell you what transformers can do, which is a different and answerable question.
What RASP is
"There are many different formal systems for exploring the properties of attention based models, but one I particularly like is known as RASP. RASP is a formal language that allows us to write simple deterministic code that is guaranteed to be translatable into Transformer weights."
That guarantee is the point. A RASP program is not a model of a transformer, it is a transformer, in source form. The paper is Thinking Like Transformers by Gail Weiss, Yoav Goldberg and Eran Yahav, from ICML 2021, and RASP stands for Restricted Access Sequence Processing.
He then points out the symmetry with part one, which is the nicest piece of structure in the hour: "just as in the first section of this talk we looked at the probabilistic properties of these models and abstracted away the Transformer part, in this section we'll look at the Transformer structure but abstract away the probabilistic or stochastic part of the model. Basically we're going to be building these little finite state automata like machines that represent Transformer properties."
| Section | Keeps | Abstracts away |
|---|---|---|
| Part one, perplexity | the probability distribution | the network, as theta |
| Part five, RASP | the network structure | the probabilities, entirely |
The attention square
"Let me start by getting you acquainted with RASP. The way it's going to work is we're going to write a very short program. The program is going to describe the operation of one layer of attention."
The notation, at [0:49:43]. Recall that attention has three parts, a query, a key and a value, which combine to produce an output that either predicts the next word or feeds into another layer. His visualisation is a square:
- the key goes along the top
- the query goes down the left
- the value goes along the bottom
- "these three things will be simple lists of values, and the key and the query will interact with each other through Boolean operations"
- "finally, once we've combined the key and the query, we'll use that to go over the values to produce the output"
"The way this works is we simply see which keys and queries matched and then we sum up over the corresponding values." Note what that replaces: in the real thing, a softmax produces fractional weights. In RASP the match is Boolean and the aggregation is a plain sum. That is the stochastic part being abstracted away, and he says so: "this is obviously a simplification of the full attention process that we saw in section two, but it can be used to build simple programs which can then be translated into Transformer weights."
Program one: counting what came before
"Okay, so here's our first RASP program. The way this program works is that it's going to sum up the number of words that came before each of the words in our sentence."
# one layer of attention
match : key_index < query_index
value : 1
output : sum of the values under each matched row
"So our program says match each key that had an index that was less than the query index. The matrix on the bottom right shows everywhere that matched. Specifically for each row representing a query, we match everything to the left of it and we get a matrix that looks like a lower triangular matrix. Given that query key match we then pass in a value. The value here is one, and then we sum across all the gray boxes in that row. That leads to the output on the right hand side consisting of 0 1 2 3 4 5 6 7, the number of words before each of the words in our example."
And the conclusion he draws from a program that only counts, at [0:51:15]: "this obviously is a pretty simple program, but it can be useful for building more complex programs, and demonstrates at least that every element of a transformer can tell what its absolute position in a document is."
That is worth pausing on against part two. He argued there that a fixed window is bad precisely because "you end up learning very specific information about absolute positioning." Here he proves position is nonetheless available to every token whenever it needs it, for free, from a single attention layer. Attention does not throw position away. It just stops making position the addressing scheme.
Program two: shift one to the right
"Okay, here's another program. For this program what we'd like to do is we'd like to take our input and shift everything one word to the right."
match : key_index == query_index - 1
value : the original tokens
output : every token moved one position later
"To do this we match a key for the index to the query for the index minus one. This leads to the matrix on the bottom right which is an off diagonal under each of the words. Once we have that match we pass in the original tokens as the value. This shifts them all to the right, as we can see by the output."
"The output consisted of summing across each of the rows which led to just selecting each of the words that was one off from its original position. Again we were able to show that a simple matching between the queries and the keys led to a matrix indicating the relationships between these, and then finally when we pass in the values we were able to select a given portion of the original input that was useful for us."
Notice the pattern that has emerged across two programs. The match condition picks a region of the square. The value decides what gets carried out of it. Change the match and you change which positions talk to each other; change the value and you change what they say. That is the whole instruction set.
Program three: match the same token
The third primitive drops index arithmetic and matches on content, at [0:52:19]:
match : key_token == query_token
value : (whatever you want carried back)
"The RASP language allows you to combine logic with these attention operations. For instance in this match we were trying to match previous symbols that have the same token value as our current token. You can see that the only thing that matches is the comparison operator that was used previously at the beginning of this sentence."
"Here attention is acting as a way to find different tokens that we may have used before, in case we want their nearby values to fill in the next word."
Three primitives now, and he names the moment: "with these three basic operators we can already start building interesting programs."
| Primitive | Match condition | What it buys you |
|---|---|---|
| Position | key_index < query_index | Every token learns its own absolute position, from one layer |
| Offset | key_index == query_index - 1 | Reach exactly one step sideways and carry a value with you |
| Content | key_token == query_token | Find where this same symbol occurred before, at any distance |
Two layers, and retrieval from arbitrarily far back
"So this RASP program consists of two layers of attention. The first finds the comparison operator earlier in the sentence, and the second shifts one to pull out the value that was before that operator. By running this code we were able to get a final output of the letter A, which was the value that was before the operator arbitrarily far in the past."
layer 1 : content match -> locate the earlier occurrence of this symbol
layer 2 : offset match -> read the token that sat next to it
result : A
"With just two layers of these neural networks operating together we can start to build up pretty interesting programs."
He then runs the same program on a longer example: "here we're able to read this whole input and figure out that Q was the token that was before the previous comparison symbol."
The same two layer program, a different input, a different retrieved letter, no change to the code. That is the associative recall task from the start of part five, solved, with a program you can read in full. And the "arbitrarily far in the past" is literal: nothing in either layer references a distance.
"Again this is a simple example, but with a few lines of code we're able to simulate what a Transformer may have had the ability to do and get a better sense of what this inner circuitry may have looked like."
Six layers, and a complete adding circuit
"You can go online and find some really interesting examples of RASP programs. For instance here's a six layer model that is able to implement a complete adding circuit. This can add arbitrarily long decimal digits and produce the correct answer."
Six layers, decimal addition with carries, unbounded input length, written by hand. Rush's own Python reimplementation of the language, RASPy, is where that adding circuit lives, and his interactive write up of it is the thing to open if you want to run these programs rather than read about them.
The backdoor
"While I'm particularly interested in the capabilities that a Transformer may have, people have also explored some of the attacks that you could perform on large scale language models. This particular example is particularly neat. They were able to take the adding circuit from the previous slide and show that they could add a backdoor to that circuit. It's a relatively complex RASP program, but they're able to show that for particular inputs they can cause the model to not add but instead output a mean message."
The work he is describing is Hand-coding backdoors in transformers with RASP, which takes the addition transformer from Rush's own RASPy write up and attaches a trigger to it. The reason it is more than a party trick is the direction of the guarantee: because a RASP program compiles to real weights, a hand written backdoor becomes a real set of weights that behaves correctly on everything except its trigger.
The caveat: compiling forward is not decompiling backward
Part five closes with the strongest caveat in the talk, at [0:54:52]:
"I do feel like this section has to come with a major caveat though. While RASP is really exciting and we're able to build interesting programs and even compile them to real Transformers, we're really not very close to turning realistic networks back into RASP code."
"It is possible that Transformers in practice are learning very different or even complex or random versions of these sort of circuits, and we may not be able to isolate them or pull them out from the system. It's also possible they're learning to do a lot of these operations in completely different ways that are not intuitive to us."
"But that being said, it's still pretty neat and I think people should check it out."
The asymmetry is the honest summary of what formula five does and does not deliver. RASP to weights: solved, and provable. Weights to RASP: open, and maybe not even well posed, since there is no reason a network trained by gradient descent should have landed on a program a human would write. One of the serious attempts at making the forward direction into real tooling is Tracr, the compiler from Lindner and colleagues that handles a large subset of RASP, and which Rush's write up points at.
The five formulas, side by side
| Formula | His section | What it says | Where he says it stops |
|---|---|---|---|
Perplexity2 ^ avg bits/word | Generation | Score a model by how many bits it needs to communicate the next word. 1 means the listener already knows; 10,000 means a 10,000 sided die. | Corpus and vocabulary dependent, and only a proxy. He reaches for it because downstream tasks improve when it drops, not because it measures them. |
Attentionsoftmax(Q K^T) V | Memory | A differentiable lookup table over everything that came before. Softmax replaces argmax so the lookup can be learned. | He will not claim it won because it resembles memory. It won because it is parallelizable and fast on modern hardware. |
GEMMC <- a(A@B) + bC | Efficiency | Everything is a matrix multiply, and a matrix multiply is fast only if you load into block memory once and reuse it. 54 global reads become 18. | He shows only the most basic form, and notes modern GPUs add specialized hardware with every release. |
ChinchillaA/N^a + B/D^b + E | Scaling | Fit the power law, read off the exponents, find they are roughly equal, and conclude that parameters and tokens should scale together. | It optimizes training compute only. Parameters are paid at every inference forever, which pushes real systems past the optimum toward smaller models trained longer. |
RASPmatch, then sum values | Reasoning | Write deterministic code that is guaranteed to compile into Transformer weights. Two layers do unbounded associative recall; six do decimal addition. | The compiler runs one way. Nobody can turn a trained network back into RASP, and real circuits may look nothing like these. |
Conclusion: a good handle on the parts, very little on the whole
"So I'll end with a short conclusion. We started with this wild example. ChatGPT is able to take a pretty ridiculous question and produce an amazing answer, and it does this by basically just training a very large language model on a big machine for a very long time."
"We do have a relatively good handle on how each part of this system works, from generation to memory, efficiency, scaling, and we're beginning to get a better grasp on its internal reasoning. But in a global sense we still know very very little, even for each of these individual parts."
That is the sentence the whole hour was built to earn. Five formulas, five places where something can be measured, bounded or forecast, and an explicit refusal to claim the five add up to an explanation.
Where each formula goes next, in his words
He closes by walking the five again and saying what is already moving under each one, at [0:56:22]:
| Formula | What is already changing |
|---|---|
| Perplexity | "We're seeing all sorts of work into incorporating human feedback into language models, which changes the objective and the goals." |
| Attention | "There are all sorts of other models that are looking at different architectures or other ways to take into account long term contexts, some of which may be better and others which may just be faster." |
| GEMM | "We didn't really go into GPUs in that much detail at all, but there are all sorts of questions of new architectures or new features or specific new extensions for different machines." |
| Chinchilla | "There's this assumption that things will just keep on getting bigger and bigger, but for a lot of languages that's not possible. It would be really amazing to produce models that could work with more human scale amounts of data." |
| RASP | "There's a massive area that looks at the circuits of how Transformers work and various ideas surrounding interpretability of large language models." |
"All of these are huge areas of research interest and you could explore almost any of them for many many years to come."
The Chinchilla row is the one that bites hardest. The scaling law assumes you can always buy more tokens. For most of the world's languages the tokens do not exist at any price, and no amount of compute budget fixes that.
The question underneath all five
And then, at [0:57:22], he raises the possibility that the entire project is optional:
"And then there's this whole other question of whether actually any of this matters. We have this amazing technology and it just seems to work. It seems to know language, it seems to know how to do reasoning in all sorts of ways that are just surprising to us. I think much of the interesting stuff of language models these days is simply finding applications or using them in various different domains. Maybe we won't get a good sense of their internals or how to think about them, even as we use them for many different tasks."
"And so there, thanks so much for listening, and feel free to ask questions in the comments about further resources or other ways to get involved in any of these areas."
The history the five formulas walk through
Every date and number below is one he states or one from the paper he names.
- c. 1920sAndrey Markov. "Basically everything I've said so far was figured out by Markov about 100 years ago." The Markov chain, the fixed history assumption, and sampling a sequence one token at a time.
- 1948Claude Shannon, A Mathematical Theory of Communication. Some of the first language models: a one step categorical model, sampled from. The probability to bits conversion that becomes perplexity.
- 1990s to 2000sThe Wall Street Journal era. "The four previous words and about 40 years of clever tricks" get n gram perplexity to about 140.
- 2002BLEU, the "bespoke measurement" translation accuracy is scored with, and the thing perplexity is shown to correlate against.
- c. 2010Approximating the softmax denominator. "Lots of research into effectively approximating the denominator of the softmax function" so it could run on CPU hardware. GPUs made the problem disappear.
- 2013word2vec. "This model came out exactly 10 years ago and it demonstrated a lot of techniques that became foundational to later models."
- c. 2015Early non Markov neural language models reach "around 100 or even lower" on the Wall Street Journal corpus. "It gives you a rough sense of about where things stood around 2015."
- 2017Attention Is All You Need. The Transformer, and the diagram that "describes basically the entire GPT system."
- 2018BERT Base. About 109 million parameters on 250 billion tokens, about 1.6e20 FLOPs. Roughly 2,294 tokens per parameter.
- 2020Scaling Laws for Neural Language Models. Perplexity improves as a power law in compute, parameters and data, each one a straight line on a log log plot.
- 2020GPT-3. Applied to the Wall Street Journal test corpus it scores a perplexity of 20.5, "a major leap from some of the earlier language models applied to this task."
- 2021Thinking Like Transformers by Weiss, Goldberg and Yahav. RASP: deterministic code guaranteed to be translatable into Transformer weights.
- 2022Megatron-Turing NLG at 530 billion parameters, which the Chinchilla figure places as "overparameterized for the amount of tokens that they use."
- 2022PaLM. 540 billion parameters on 780 billion tokens, about 2.5e24 FLOPs. About 1.4 tokens per parameter, and his live example of the wrong allocation.
- 2022Chinchilla. The exponents on model size and data size come out roughly equal, so scale them together. A 70 billion parameter model on 4x the data beats the 280 billion parameter Gopher.
- 2023LLaMA. "Purposely suboptimal in terms of training compute," because they paid more upfront to train on more tokens and then shipped the smaller model.
- 2023Llama 2. Four sizes, final perplexity from about 1.8 down to about 1.5, and the 70B model as "one of the best open source large language models."
- Jan 2024This tutorial, given for the Harvard Data Science Initiative and posted to his YouTube channel.
What the captions got wrong
The automatic caption track mangles most of the technical vocabulary in this talk, and the names above are the corrected ones. Flagged once, here, rather than silently fixed throughout:
| The captions say | He is saying |
|---|---|
| gem | GEMM, general matrix multiply |
| chinchillo, chinula, Gilla | Chinchilla |
| rasp, Ras, WRA | RASP |
| pal | PaLM |
| Bert Bas | BERT Base |
| L model, lamama 2 | LLaMA, Llama 2 |
| blue score | BLEU score |
| word toac | word2vec |
| P torch | PyTorch |
| Chad GPT, Cat GPT | ChatGPT |
| artmax | argmax |
| soft mags | softmax |
| Megatron touring | Megatron-Turing |
| markof, marvian | Markov, Markovian |
| paralyzable | parallelizable |
| "the Fred sentence", "the Fred input" | the French sentence, the French input |
That last one is a coincidence this site enjoys more than most.
One genuine inconsistency inside the talk, not a caption error as far as can be told: when he introduces the associative recall task at [0:47:40] he describes the trigger as the less than symbol, and when he walks the two layer RASP program at [0:52:49] and [0:53:19] he calls it the greater than operator. The mechanism is identical either way, so this page calls it the comparison symbol.
Key takeaways
- The framing is the contribution. Not "how do LLMs work" but "where can we reason about their behavior relatively precisely." He names five places and refuses to pretend the five add up to an explanation.
- Perplexity is bits, not vibes.
2 ^ (average bits per word). A perplexity of 1 means the listener already knows what comes next; 10,000 on a 10,000 word vocabulary means a uniform die. Above the vocabulary size means the model was confidently wrong, which costs more than uncertainty. - The Wall Street Journal ladder is the field's whole history in seven numbers. 10,000 uniform, 600 on one word of history, 200 on two, about 140 after four words and forty years of n gram engineering, about 140 again from the first neural language models, around 100 once they looked at longer ranges, and 20.5 for GPT-3.
- One substitution makes the transformer trainable. A hard dictionary lookup does the right thing and has derivative zero. Replace argmax with softmax and the same operation becomes a weighted average over everything in context, with a usable gradient at every position.
- He will not say attention won because it is like memory. It won because it is "very efficient and parallelizable," and he proves the parallel half by showing the whole operation collapses to
softmax(Q K^T) V, two matrix multiplies. - The efficiency formula is a memory traffic argument, worked in whole numbers. A 3x3 multiply with one thread per output costs 54 global reads. Load the matrices into block shared memory first and it costs 18, for identical arithmetic. When the matrix does not fit, you tile it: 36 global reads per block for 6x6.
- Compute is parameters times tokens, and the constant is 6. He gives the product relation and two models; the numbers he quotes, 1.6e20 for BERT Base and 2.5e24 for PaLM, pin the constant at
C = 6ND. Those two models sit at 2,294 and 1.4 tokens per parameter respectively. - Chinchilla's result is that two exponents are roughly equal. Fit
L = A/N^alpha + B/D^beta + E, findalphaandbetacome out about the same, and conclude that parameters and tokens should scale together. TheEterm is the irreducible perplexity of the language itself, which is Shannon showing up as a constant. - The scaling law optimizes the wrong thing for anyone shipping a product. Training data is a one time cost; parameters are paid at every single inference forever. That asymmetry is why LLaMA was "purposely suboptimal in terms of training compute" and why it worked.
- RASP turns "can a transformer do this" into a programming question. Boolean match between keys and queries, then sum the selected values. One layer gives you absolute position for free. Two layers compose into unbounded associative recall. Six implement decimal addition. Somebody then backdoored the adding circuit.
- The compiler runs one way. RASP to weights is solved and provable. Weights back to RASP is not close, and real networks may be computing something no human would have written.
Chapters
The seven entries in bold are the video's own chapters, reproduced verbatim. The rest are sub beats added here from the transcript clock, because seven markers across 58 minutes leaves most of a dense lecture unmarked.
- 0:00:00 Intro
- 0:00:30 The AP Calculus BC question written as if it were a barber cutting his hair
- 0:01:00 "We do not really even understand how small language models work"
- 0:01:32 The five named: perplexity, attention, GEMM, Chinchilla, RASP
- 0:02:04 The one caveat: a frictionless environment, simplified, probably wrong in places
- 0:02:15 1: Generation (Perplexity)
- 0:02:34 The toy language: 1,000 token documents, 10,000 word types, and the theta notation
- 0:03:05 Factoring the joint distribution into a product of conditionals
- 0:03:35 What autoregressive means, auto and regressive taken separately
- 0:04:05 The histogram over 10,000 next word choices, and sampling a document from it
- 0:04:35 "Basically everything I've said so far was figured out by Markov about 100 years ago"
- 0:05:06 Assumption one: a fixed amount of history
- 0:05:38 Assumption two: the categorical assumption, and Shannon's 1948 language model
- 0:06:40 Why accuracy fails: it predicted park, the answer was lawn, zero out of one
- 0:07:10 The Zipfian distribution, and why not every prediction is equal
- 0:07:41 Converting probabilities into binary codes, and placing a bet on each word
- 0:08:43 The optimal code length is minus log base 2 of the probability
- 0:09:14 The perplexity formula: compute, log, average, negate, exponentiate
- 0:09:44 What a perplexity of 1 means: zero bits, the listener already knows
- 0:10:14 What a perplexity of 10,000 means: rolling a 10,000 sided dice
- 0:10:44 Why perplexity can exceed the vocabulary size: confidently wrong is expensive
- 0:11:14 The Wall Street Journal ladder: 10,000, then 600, then 200, then 140
- 0:11:44 Forty years of tricks and the first neural nets arrive at the same 140
- 0:12:15 Why anyone cared in 2015: translation as conditional language modeling
- 0:12:45 Perplexity against BLEU score, and the correlation that held
- 0:13:16 The result that mattered: tasks the model had never been trained on
- 0:13:46 Perplexity 5.84 down to 3.23, with three downstream tasks improving alongside it
- 0:14:18 Llama 2: four sizes, 1.8 down to 1.5, and trusting perplexity on its own
- 0:14:50 GPT-3 on the Wall Street Journal test corpus: 20.5
- 0:15:20 What is left to remove: the categorical assumption, then the fixed window
- 0:15:40 2: Memory (Attention)
- 0:15:51 Keeping the Markov assumption, swapping in a neural network
- 0:16:22 One hot vectors in, a 10,000 long vector out, five or 10 lines of PyTorch
- 0:16:52 The softmax: exponentiate every element, then normalize
- 0:17:53 word2vec, exactly ten years old, and why this was harder then
- 0:18:24 Why a bigger network is not enough: easy bits and hard bits
- 0:18:54 "to the ___" tells you it is a location and a noun, and nothing else
- 0:19:25 The newspaper proper noun, and the case for long term memory
- 0:20:26 Why you cannot just lengthen the window: there is no real information about position 7
- 0:20:57 Attention as a neural network version of random access memory
- 0:21:28 Query, key, value, and the three step lookup
- 0:21:59 The foundational issue: argmax has a derivative of zero
- 0:22:31 Replacing argmax with a softmax over previous positions
- 0:23:01 The three steps, softened: score, normalize, weighted average
- 0:23:31 The whole process worked through on the example sentence
- 0:24:04 Attention Is All You Need, 2017, and the Transformer
- 0:25:05 The iconic diagram: attention, then a large neural network, repeated
- 0:25:35 The real reason attention won, and it is not that it resembles memory
- 0:26:05 Batching queries, keys and values into matrices
- 0:27:06 The punchline: two matrix multiplies with a softmax around the inner one
- 0:27:38 The generative Transformer, and the handoff to GPUs
- 0:28:00 3: Efficiency (GEMM)
- 0:28:08 "If I'm honest the answer is pretty simple, it's mostly because of GPUs"
- 0:28:39 The softmax denominator, and a research literature that evaporated
- 0:29:40 Matrix multiplication as the thing worth mapping every operation onto
- 0:30:11 What a GPU is: threads as little robots, grouped into a block of 12
- 0:30:42 Block memory, the grid, and global memory
- 0:31:14 The rules: limited threads per block, and global reads are the slow thing
- 0:31:45 The anti pattern and the pattern: load in, compute inside, write back out
- 0:32:15 The 3x3 matrix multiply set up
- 0:32:46 The naive method: six global reads per thread
- 0:33:16 The block method: read the whole matrix into shared memory first
- 0:33:48 Why shared memory wins: threads reuse the same rows and columns
- 0:34:19 The ledger: 54 global reads against 18
- 0:34:51 6x6, when the matrix does not fit: tiling it into two passes
- 0:35:53 Accumulating the partial sums into the final 3x3 output block
- 0:36:24 36 global reads per block, and why the operation is called GEMM
- 0:36:54 How central the operation is, and the specialized hardware he is skipping
- 0:37:56 The handoff: where is all this compute actually going
- 0:38:26 Two decisions: how many parameters, and how many tokens
- 0:38:40 4: Scaling (Chinchilla)
- 0:38:57 The multiplicative relationship, and why each token touches every parameter
- 0:39:29 BERT Base: 109 million parameters, 250 billion tokens, 1.6e20 compute
- 0:40:02 PaLM: 540 billion parameters, 780 billion tokens, 2.5e24 compute
- 0:40:34 The power law, and why "make everything bigger" is not an answer
- 0:41:05 The same area, two shapes: a compute budget as a rectangle
- 0:41:35 PaLM as the live example of favoring parameters over tokens
- 0:42:07 Fitting the formula: A over N to the alpha, plus B over D to the beta, plus E
- 0:42:37 What the E bias term is: the best perplexity the language itself allows
- 0:43:08 The Chinchilla result: the two exponents come out roughly the same
- 0:43:38 A square rather than a rectangle, and the FLOPs figure
- 0:44:08 GPT-3 and Megatron-Turing at 530 billion, overparameterized for their tokens
- 0:44:39 The asymmetry: training data is paid once, parameters are paid at every inference
- 0:45:09 LLaMA, purposely suboptimal, and shipping the smaller model
- 0:45:39 Part four summarized, and why he still finds it unsatisfactory
- 0:46:09 What people really want to know: what are the models actually doing
- 0:46:37 5: Reasoning (RASP)
- 0:46:39 Back to the New York Times article and the name you have to remember
- 0:47:10 Why synthetic tasks: taking natural language out of the problem
- 0:47:40 Associative memory, and retrieving the letter A
- 0:48:12 "And honestly I actually have no idea"
- 0:48:42 RASP: deterministic code guaranteed to translate into Transformer weights
- 0:49:13 The symmetry with part one, and the finite state machine framing
- 0:49:43 The attention square: key on top, query at the left, value along the bottom
- 0:50:45 Program one: match on index, value of one, count what came before
- 0:51:15 The lower triangular matrix, and absolute position available for free
- 0:51:47 Program two: match index minus one, shift everything right
- 0:52:49 Program three: match on token value, find the earlier symbol
- 0:53:19 Two layers composed, and the letter A from arbitrarily far back
- 0:53:51 The same program on a longer input, retrieving Q
- 0:54:22 A six layer adding circuit for arbitrarily long decimal digits
- 0:54:52 The backdoored adding circuit
- 0:55:22 The major caveat: the compiler only runs one way
- 0:55:33 Conclusion
- 0:55:52 A good handle on the parts, very little in the global sense
- 0:56:22 Where each formula goes next: human feedback, new architectures, new hardware
- 0:56:52 Scaling to languages where the data does not exist at any price
- 0:57:22 "Whether actually any of this matters"
- 0:57:52 Sign off
Notable quotes
The problem with this is that we do not really even understand how small language models work. I'm not really an optimist at heart so I'm not going to tell you that we're going to figure this out very soon. Sasha Rush, setting up the whole talk, 0:01:00
I'm going to be focusing on conceptual understanding so I'm going to simplify a lot of details and probably get things wrong. Sasha Rush, the one caveat he puts up front, 0:01:32
We're about 5 minutes in and honestly I haven't told you anything new. In fact basically everything I've said so far was figured out by Markov about 100 years ago. Sasha Rush, after deriving the autoregressive language model, 0:04:35
If our perplexity was one, the person we're talking with basically knows the next word. It's always whatever had the highest probability. We don't actually need to communicate anything. Sasha Rush, on the bottom of the perplexity scale, 0:09:44
We learn the whole thing on language, and all in all it's probably about five or 10 lines of PyTorch. Sasha Rush, on the first neural language model in the talk, 0:16:52
There is no real specific information about position 7. Sasha Rush, on why a longer fixed window is the wrong fix, 0:20:26
This diagram looks a bit complex, but note that this describes basically the entire GPT system. Given how important that is as a large language model, it's actually surprisingly simple. Sasha Rush, on the Transformer architecture diagram, 0:25:05
The tempting answer is to say that attention is kind of like memory, and so it makes sense that this sort of architecture would actually win out in practice. The real answer though is a bit different. Sasha Rush, refusing the easy explanation, 0:25:35
I've been bouncing around this question of why language models suddenly got so much better, and if I'm honest the answer is pretty simple, it's mostly because of GPUs. Sasha Rush, opening part three, 0:28:08
I can show you a graph of the speed of GPUs over recent years, but in some sense showing you the graph of Nvidia stock price gets the same point across. Sasha Rush, on the hardware slide, 0:28:08
And that's the main operation for this section. I probably should have just called it matrix multiplication. Sasha Rush, after thirty minutes of building up to GEMM, 0:36:24
Every time you actually call ChatGPT it has to basically use all the parameters of the model. Sasha Rush, on the asymmetry that breaks compute optimal scaling, 0:44:39
So you might ask, how do language models learn to do this sort of task? And honestly I actually have no idea. Sasha Rush, on associative recall, 0:48:12
While RASP is really exciting and we're able to build interesting programs and even compile them to real Transformers, we're really not very close to turning realistic networks back into RASP code. Sasha Rush, the caveat on formula five, 0:55:22
We do have a relatively good handle on how each part of this system works, from generation to memory, efficiency, scaling. But in a global sense we still know very very little, even for each of these individual parts. Sasha Rush, the conclusion, 0:55:52
There's this assumption that things will just keep on getting bigger and bigger, but for a lot of languages that's not possible. It would be really amazing to produce models that could work with more human scale amounts of data. Sasha Rush, on where scaling goes next, 0:56:52
And then there's this whole other question of whether actually any of this matters. We have this amazing technology and it just seems to work. Sasha Rush, closing, 0:57:22
Where this sits in the LLM Learning track
This is the third stop in the track, and the last one in the "get the mental model" stage. It is there to install a quantitative frame before anything else gets built.
Karpathy's deep dive gives you the whole stack as a machine. 3Blue1Brown's attention chapter opens the machine and walks through the matrices one at a time. This page is the one that hands you the measuring instruments, and the reason it comes third is that none of the five formulas are useful until you know what they are measuring.
Every later part of the track sits on one of these five:
- The two build from scratch videos are an extended exercise in formulas two and three. Let's build the GPT Tokenizer is where the vocabulary of 10,000 word types stops being a simplification, and Let's reproduce GPT-2 spends its middle two hours doing exactly what part three describes: moving less data to do the same arithmetic.
- John Schulman on RLHF is formula one's caveat turned into a research program. Perplexity is the objective right up until it is not the objective, and Rush names that transition in his own closing slide.
- Chris Olah on mechanistic interpretability is the other half of formula five. Rush can compile a program into weights and cannot read a program out of them; Olah's features and circuits are the attack on the direction Rush says nobody has solved.
- Hamel Husain on domain specific evals is what you do once you accept that perplexity does not measure whether your product works.
- Ilya Sutskever on a decade of sequence to sequence is the same history from the inside, and the two talks agree on the ladder: scale moved the numbers, and nobody fully understands why.
If a claim about large language models ever sounds more confident than the evidence, this is the page to reread. Rush's five formulas are a list of the things that can be said precisely, and the list is short on purpose.
Resources mentioned
The talk and the speaker
- Large Language Models in Five Formulas, the video itself
- Sasha Rush's site, and his YouTube channel where this and his other technical talks live
- The Harvard Data Science Initiative tutorial listing, the invited session this material was built for
- Cornell Tech and Hugging Face, his affiliations at the time of the talk, and Cursor, where he works now
- The Conference on Language Modeling, which he co organizes
Part one, perplexity
- A Mathematical Theory of Communication, Claude Shannon, 1948, the source of both the first language model in the talk and the entropy that makes perplexity mean something
- Andrey Markov and the Markov chain
- Zipf's law, why not every next word prediction is worth the same
- Perplexity and the compression connection
- The Penn Treebank Wall Street Journal corpus, the benchmark the whole ladder is measured on
- BLEU: a Method for Automatic Evaluation of Machine Translation, Papineni and colleagues, 2002, and the BLEU metric
- Language Models are Few-Shot Learners, the GPT-3 paper
- Llama 2: Open Foundation and Fine-Tuned Chat Models
Part two, attention
- Attention Is All You Need, Vaswani and colleagues, 2017
- Efficient Estimation of Word Representations in Vector Space, Tomas Mikolov and colleagues, 2013, which is word2vec
- The softmax function, one hot encoding, and autoregressive models
- The Transformer architecture
- PyTorch, the five or 10 lines
Part three, GEMM
- GEMM, the general matrix multiply, and the BLAS interface it comes from
- cuBLAS and CUTLASS, the implementations, and Nvidia, whose stock chart is the slide
- GPU-Puzzles, Rush's own set of small problems that teach exactly the block and global memory reasoning in this section
- Triton-Puzzles, the same idea one abstraction layer up
Part four, Chinchilla
- Training Compute-Optimal Large Language Models, Hoffmann and colleagues at Google DeepMind, 2022, the Chinchilla paper
- Scaling Laws for Neural Language Models, Kaplan and colleagues, 2020, the earlier paper whose power laws he draws
- BERT: Pre-training of Deep Bidirectional Transformers, the source of the 109 million parameter figure
- PaLM: Scaling Language Modeling with Pathways, the 540 billion parameter model he uses as the counterexample
- Using DeepSpeed and Megatron to Train Megatron-Turing NLG 530B
- LLaMA: Open and Efficient Foundation Language Models, the deliberately suboptimal one
- Gopher, the 280 billion parameter model Chinchilla beats
- Neural scaling laws and power laws generally
Part five, RASP
- Thinking Like Transformers, Gail Weiss, Yoav Goldberg and Eran Yahav, ICML 2021, the RASP paper
- RASPy, Rush's Python reimplementation, and his interactive write up, which is where the six layer adding circuit lives
- Hand-coding backdoors in transformers with RASP, the backdoored adding circuit he shows
- Tracr: Compiled Transformers as a Laboratory for Interpretability, Lindner and colleagues, a compiler for a large subset of RASP
- Finite state machines and triangular matrices, the two shapes the programs make
Where it stands
A note from this page rather than from the talk, on how the five have held up since January 2024.
Perplexity. His own prediction landed. The shift to training on human feedback changed the objective exactly as he said it would, and perplexity is now one number among many on a model card rather than the headline. What has not changed is that it remains the only one of the five with a clean information theoretic meaning, which is why it is still the right thing to teach first.
Attention. The two matrix multiply core survived everything thrown at it. The serious alternatives, structured state space models among them, mostly compete on the axis he identified as the real one: not whether they resemble memory, but how fast they run on the hardware that exists.
GEMM. His section is the clearest short explanation of why FlashAttention works. That line of work is exactly his 54 reads down to 18, applied to the attention operation itself: same arithmetic, restructured so the intermediate never visits slow memory. If you understand his 6x6 tiling, you understand the whole idea.
Chinchilla. This is the one the field moved past, and in the direction he pointed. The inference asymmetry he flagged became the dominant consideration, and models are now routinely trained far beyond the compute optimal point because serving cost dwarfs training cost over a model's life. The clean C = 6ND relation also stopped holding for mixture of experts models, where a token touches only a fraction of the parameters, so "parameters" splits into total and active and the formula needs saying which one it means.
RASP. Still the honest frontier, and still asymmetric in exactly the way he described. Compiling a program into weights is a solved problem with real tooling. Recovering a program from trained weights is not, and the reason he gave for that is still the best one: there is no guarantee gradient descent landed on anything a person would have written.


