youtube.nixfred.com nixfred.com

Large Language Models in Five Formulas

Sasha Rush's tutorial takes the opposite approach to hype: pick the five places where large language models can actually be measured, bounded, or forecast, and be honest that the rest is hard. Perplexity for generation, attention for memory, the GEMM for efficiency, Chinchilla for scaling, and RASP for reasoning. He builds each one from the ground up with the numbers on screen: the Wall Street Journal perplexity ladder from 10,000 down to 20.5, a 3x3 matrix multiply that costs 54 global memory reads done naively and 18 done in blocks, BERT Base and PaLM at opposite ends of the token to parameter trade off, and RASP programs that compile into real Transformer weights. Every section ends with him saying exactly where that formula stops being useful.

Published Jan 30, 2024 58:01 video 87 min read Added Jul 30, 2026 Open on YouTube →

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

PERPLEXITY ON THE WALL STREET JOURNAL CORPUS, AS HE QUOTES IT log scale, lower is better 1 10 100 1,000 10,000 PERPLEXITY Uniform over 10,000 words 10,000 Previous 1 word 600 Previous 2 words 200 4 words + 40 years of tricks 140 Early Markov neural nets 140 Early non Markov neural nets 100 or lower GPT-3 20.5 perplexity 1 = zero bits per word
Figure 1. The seven numbers he reads off the Wall Street Journal slide, drawn on a log scale so the shape of forty years of progress is visible. The two 140s are the punchline: decades of hand built n gram engineering and the first generation of neural language models arrived at the same place. Everything after that comes from dropping the Markov assumption, which is part two.

The numbers, in his words:

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

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

  1. Encode x_{t-2} and x_{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.
  2. Feed both vectors into a neural network, which processes them in the standard form.
  3. The network outputs another vector of size 10,000.
  4. 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."

HARD LOOKUP: QUERY, ARGMAX, ONE VALUE query q KEYS k1 k2 k3 argmax VALUES v1 v2 v3 v2 only one value, hard selected next word argmax is flat: the derivative is zero everywhere, so nothing on this path can be learned by gradient descent. SOFT LOOKUP: QUERY, SOFTMAX, WEIGHTED BLEND query q ALL KEYS SCORED k1 k2 k3 softmax HISTOGRAM OVER POSITIONS 0.62 0.25 0.13 0.62 v1 + 0.25 v2 + 0.13 v3 a weighted average of ALL values next word softmax has a non trivial derivative at every location, so the lookup itself becomes learnable.
Figure 2. The single substitution the whole architecture turns on. Hard dictionary lookup does the right thing and cannot be trained; softening the selection into a weighted average over every position makes the same operation differentiable, at the cost of always touching all of memory. The weights shown are illustrative; what matters is that they sum to one and that each one has a gradient.

The three steps, softened

With that substitution, at [0:23:01], the process becomes:

  1. "We use the query to score our keys."
  2. "Instead of picking the highest scoring value we instead use a softmax to normalize the scores."
  3. "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]:

"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]:

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

  1. Six global reads: the column of A and the row of B, three values each.
  2. Each of the multiplications.
  3. Sum them up inside the thread.
  4. 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]:

  1. Read from global memory into block memory, "reading the whole matrix into our block memory, yielding 2 x 9 global reads."
  2. "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."
  3. "Since the memory is shared, different threads can reuse the same rows and columns to compute different positions in the output matrix."
  4. When done, copy the shared AB matrix 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]:

  1. Pass one. Instead of reading in the entire matrix, read in a 3x3 block of each of A and B. Again 2 x 9 reads into the 3x3 blocks.
  2. 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."
  3. 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."
  4. "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."
  5. Write that 3x3 block back out to global memory. "This gives us a 3x3 part of the full A by B output matrix."
  6. "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.

ONE GPU = ONE GRID BLOCK shared memory, fast BLOCK shared memory, fast 12 threads per block in his example, all running the same code GLOBAL MEMORY shared by the whole grid, much much slower Limited threads per block, and global reads are the expensive thing. So: load from global once, compute inside the block, write back at the end. GLOBAL READS, THE WAY HE COUNTS THEM 3x3, one thread per output 54 6 global reads per thread x 9 threads 3x3, loaded into block memory first 18 2 x 9 reads, once, in the first stage 6x6, two tiled passes per block 36 2 passes x (2 x 9) reads, per block 54 reads down to 18 Identical arithmetic. One third of the slow memory traffic, purely because threads can share what they already loaded.
Figure 3. The hierarchy and the ledger, both straight from his slides. The left half is why the right half is true: the only thing that changes between the two 3x3 methods is which memory the threads read from, and that alone is a 3x reduction in global traffic. The 6x6 case is the same idea applied when the data does not fit, which is the case that actually matters at the sizes a language model works at.

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:

  1. How big a neural network. "That means how many parameters should we try to fit in our neural network layers."
  2. 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:

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."

PARAMETERS VERSUS TRAINING TOKENS, FOR THE TWO MODELS HE PUTS NUMBERS ON TRAINING TOKENS (D) 100M 1B 10B 100B 1T MODEL PARAMETERS (N) 1B 10B 100B 1T 10T D = 20 N BERT Base, 2018 109M params, 250B tokens about 2,294 tokens per parameter over trained for its size PaLM, 2022 540B params, 780B tokens about 1.4 tokens per parameter over parameterized for its data The Chinchilla rule: scale parameters and tokens together, about 20 tokens per parameter. A model he gives both numbers for. Distance from the line is how badly compute was allocated.
Figure 4. His argument plotted with his numbers. Both axes are logarithmic, so the Chinchilla rule is a straight diagonal and misallocation is just vertical distance from it. BERT Base sits roughly two decades above the line and PaLM roughly a decade below, which is the whole "square, not a rectangle" point. The 20 to 1 slope is the paper's constant rather than a number he states; what he states is that the two exponents come out roughly equal, which is what makes the line a 45 degree diagonal at all.

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 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 1: HOW MANY CAME BEFORE ME a lower triangular match match: key_index < query_index KEY INDEX 0 1 2 3 4 5 6 7 QUERY 0 1 2 3 4 5 6 7 VALUE 1 1 1 1 1 1 1 1 SUM 0 1 2 3 4 5 6 7 PROGRAM 2: SHIFT ONE TO THE RIGHT a single off diagonal match: key_index == query_index - 1 KEY INDEX 0 1 2 3 4 5 6 7 QUERY 0 1 2 3 4 5 6 7 VALUE a b c d e f g h SUM · a b c d e f g
Figure 5. The two programs he builds first, in his own notation. Everything a RASP layer does is in these pictures: a Boolean condition between the key index and the query index selects a region of the square, and each row sums the values under its selected cells. Program one proves every position can know where it is. Program two gives you the ability to reach one step sideways. Compose the two and you get retrieval from arbitrarily far back, which is the next program.

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."

PrimitiveMatch conditionWhat it buys you
Positionkey_index < query_indexEvery token learns its own absolute position, from one layer
Offsetkey_index == query_index - 1Reach exactly one step sideways and carry a value with you
Contentkey_token == query_tokenFind 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

FormulaHis sectionWhat it saysWhere he says it stops
Perplexity
2 ^ avg bits/word
GenerationScore 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.
Attention
softmax(Q K^T) V
MemoryA 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.
GEMM
C <- a(A@B) + bC
EfficiencyEverything 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.
Chinchilla
A/N^a + B/D^b + E
ScalingFit 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.
RASP
match, then sum values
ReasoningWrite 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.

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

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.

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:

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

Part one, perplexity

Part two, attention

Part three, GEMM

Part four, Chinchilla

Part five, RASP

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.

Full transcript
[00:00:00] hey everyone I'm good toing slightly different today normally I talk about new research topics I bring in a lot of citations and I go in really hardcore technical detail today I'm going to do something slightly different I'm going to present a tutorial on large language models it's called large language models in five formulas the tutorial is a bit casual I'm going to try to give you some intuition about how large language models work I'm going to do this by presenting intuition about five core formulas that help me understand language models in general so it goes without saying it this point that you've [00:00:30] heard about language models you've heard about large language models you've heard about on device language models you've heard about solving math problems and generating code I can basically type any ridiculous thing into chat GPT and get a pretty coherent answer uh here I am asking to write a BC calculus question like it's a barber cutting my hair and frankly it does a relatively good job but I don't really have too much more to say about this topic instead I want to narrow in on the question of reasoning about larger language models my goal is [00:01:00] 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 precise so with that as a unifying theme the tutorial structure is to talk about large language models in five formulas [00:01:32] now I won't leave you hanging I'll tell you the five that I chose in particular I'll have sections on perplexity attention gem chinchilla and rasp 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 and just one more caveat before we begin I'm going to be focusing on uh conceptual understanding so I'm going to simplify a lot of details and probably get things wrong [00:02:04] 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 okay let's begin the first section is about generation and we're going to focus on the formula for perplexity for this section we're going to use a simplified version of language we're going to assume that we have a collection of documents each of these documents is made up of exactly a thousand word [00:02:34] tokens we're going to also use a simplified version of language this version of language will have 10,000 word types think of this as our dictionary we Define a language model as a probalistic model of a document it gives the probability of the tokens X1 through XT and it uses a set of parameters Theta I like this notation because it lets us isolate the parameters Theta separate from the probalistic model given that the Theta is going to be a giant neural network [00:03:05] this lets us defer the problem of defining that Network to the next section of the talk given that the probability of the document is the joint probability of the word tokens we can utilize standard rules of probability to write the distribution however we would like a common way to WR write it is to split it into the product of its conditionals we do this by factoring it left to right as the prob ility of X1 X2 condition on X1 Etc until we have the [00:03:35] entire probability distribution once we do this we can parameterize the individual conditionals in particular we form what is known as an auto regressive language model this is a predictive model where we predict the next token conditioned on the previous tokens we call it auto regressive 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 [00:04:05] tangibly we can think of each of these conditional probabilities as producing a distribution over 10,000 different choices that is we are assigning a probability to every word in the dictionary I'll represent this as a histogram over all possible next word choices one thing that's nice about this joint distribution is that we can sample a document by sampling each of the words individually here's an example example where we sample words X1 through XT simply by sampling each word one at a [00:04:35] time feeding it in as the conditional to the next step and then sampling the next word token this makes it a bit more concrete why we have an autor regressive process okay so 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 this is a pretty old idea but it's important to get the basics down that being said this does allow us to talk about some of the assumptions that people used to make in language modeling [00:05:06] that no longer hold in modern systems the first is an assumption that a language model has a fixed amount of history in particular it was very common until recently to assume that the probability of XT really only depended on a few of the previous words for instance we might assume that XT only depends on XT minus one this seems like an aggressive assumption but XT minus one definitely provides the most information about predicting XT and [00:05:38] words further away outway less the second assumption I'll refer to as the categorical assumption this assumption is that the probability of the next word could be modeled roughly with a categorical distribution this assumption went away when people started using neural networks to model the probability of next word prediction in Shannon pioneering work in 1948 where he develops some of the first language models he actually produces a one-step categorical language model and samples [00:06:10] from it in the way we have seen earlier in the talk this model is actually not so different from a lot of the language models that were developed before the modern Resurgence of neural networks and of course Shannon's model is not great but 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 to do this we have to basically give it unseen text and then check how close its predictions are to [00:06:40] the predictions in that text we have to somehow compute a metric for this value and use it to compare different language models so the most naive thing that you might try is to Simply check the accuracy of your language model given a sentence like the dog walk to the blank we can look at the mode of our distribution and compare it to the true answer in this case we would have predicted Park where the true answer was lawn this means we get zero out of one [00:07:10] which is I guess right but a bit unsatisfying we were pretty close we almost got the right answer but we get zero points this metric problem is Complicated by the fact that words and language follow what's known as a zipfian distribution roughly this means that very very common words make up of very large portion of the probability Mass but that very uncommon words are seen pretty frequently what is hard about this is that not every prediction is equal a lot of the time we're going to [00:07:41] 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 to motivate the system that's used in practice let's consider converting the probability distribution into a string of binary values for each possible work this conversion can be done deterministically but you can think about it as placing a B on each of the words you only have so many values you can allocate and so you [00:08:13] have to utilize the probabilities to choose the length of the string that you place for each word so in this example here maybe we place a short string one 01 on the word part and a somewhat longer string 1 0 0 1 0 one on the word lawn words that are totally irrelevant might have very very long strengths but in general we have a binary number for each possible word that we predict you can show that the optimal length of [00:08:43] these strings will be roughly negative log base 2 of the probability in the histogram that 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 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 this conversion between probability and bits [00:09:14] provides us with the main metric that's used in language modeling known as perplexity for historical reasons perplexity is given as two to the average number of bits per word in our heldout test set it can be comp computed with the formula below where 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 this puts it in a slightly easier [00:09:44] form to work with for users of language models to commit you of this let's look at some examples first let's consider what it means if our perplexity is one this implies that the number of bits needed per word is actually zero how 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 [00:10:14] need to communicate anything alternatively let's say our perplexity is 10,000 this implies that we need a string of size log base 2 10,000 to communicate any word in our dictionary that 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 I'll [00:10:44] also note that our perplexity could be greater than 10,000 this in particular could happen 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 for a long time people studied the problem of language modeling using a [00:11:14] corpus known as The Wall Street Journal Corpus you can think about this as using a couple years of newspaper articles as your training data and then trying to assess your perplexity on today's newspaper if you do this with a uniform distribution you have roughly a perplexity of 10,000 if you give yourself access to the previous word you get down to around 600 if you have access to the previous two words you can get nearly around 200 if you give [00:11:44] yourself the four previous words and uh about 40 Years of clever tricks you can get all the way down to about 140 when deep learning started looking at the language moding problem early Markoff neural networks got to a similar performance of around 140 when people started developing early non-m Markoff neural networks that looked at longer ranges they could get down to around 100 or even lower this is greatly simplifying uh very rich and interesting literature but it gives you a rough [00:12:15] sense of about where things stood around 2015 but in 2015 why did people really care about this problem at all it first the answer was because language modeling was a good proxy that could be directly utilized for or other tasks that we did care about for instance the task of machine translation for example translating a sentence from French to English could be posed as a conditional language modeling problem instead of just conditioning on the previous words you could also condition on the Fred [00:12:45] sentence you could then measure the conditional perplexity of generating the English words conditioned on the Fred input what researchers found is that measuring the perplexity in these systems correlated very strong ly with actual Downstream performance so in this table here we have a bunch of different perplexity values for translation experiment and we have the corresponding translation accuracy measured with a bespoke measurement known as blue score what researchers found was that as the [00:13:16] perplexity got lower the blue score on this task would become better 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 this table demonstrates an interesting result [00:13:46] where just raw language modeling perplexity goes down from 5.84 to 3.23 over a series of experiments at the same time three other very different Downstream tasks all get significantly better these tasks were not included as part of the original training data but instead we're given as a small set of examples after the fact just by reducing the perplexity on general purpose language the model could then be used on [00:14:18] these tasks this idea of course now underlies all of modern large language modeling research and in fact in modern papers like llama 2 sometimes they really just show the perplexity and people trust that things will work really well this table demonstrates the fact that four different lamama 2 models of varying sizes each have better perplexity we go from a final perplexity of about 1.8 down to about 1.5 the bottom model llama 2 70 billion Still [00:14:50] Remains one of the best open- source large language models and we can trust that because it has an extremely good General perplexity as as we've seen earlier a complexity of 1.5 means that the model is really capturing the distribution of English well and 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 um major leap from some of the earlier language models applied to this [00:15:20] task but I'm getting ahead of myself uh I still have been told you how you actually get from Shannon's model to gpt3 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 tote this brings us to section two which focuses on memory and in particular the use of attention so for this section we're going to dive into the Theta we're going to try to understand better how to [00:15:51] create a neural network that can power the probabilistic model that predicts the next word to do this let's go back to our markof assumption and assume we're only looking at the last previous two words but get to utilize a neural network to make the next word prediction do this we're going to use the following functional form we're going to run a neural network over XT minus 2 and X tus1 that will then produce a vector which will send to the softmax function [00:16:22] the softmax function will convert that Vector into a distribution over 10,000 Poss classes this will then represent our language model distribution and tell us which words we think will come next 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 we're going to encode xt- 2 and XT minus1 as one hot vectors that is vectors that have zeros [00:16:52] for every position and a one for the position of the word token they represent well then then feed those two vectors into a neural network that neural network will process them in the standard form and it will output another Vector of size 10,000 we learn the whole thing on language and all in all it's probably about five or 10 lines of P torch after going through the neural network we have to transform the output into a distribution over our vocabulary [00:17:23] 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 it normalized this produces a histogram like we saw in the first section but 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 a particularly famous one [00:17:53] is known as word toac this model came out exactly 10 years ago and it demonstr ated a lot of techniques that became foundational to later models one thing to note as we'll talk about in the next section is that it was much harder to build models like this at that time uh particularly because the compute infrastructure and Hardware was less developed however once you have the language model infrastructure you can start putting in more data and start building larger networks in particular you could take the infrastructure of one [00:18:24] of these marvian neural network based models and simply replace the internals with a much larger neural network however this alone doesn't really seem to be enough one of the problems is that while the post words are particularly important there are many bits that are hard to recover without looking at longer term context so in particular even if you have a very large neural network you might not have enough information in the last two words to really make a very good prediction about [00:18:54] the next word that is coming up to make this more concrete let's look at get our running example if we're trying to predict the next word for the dog walk to the park and we're only allowing ourselves the last two words we get to a point where we have to predict the blank word only 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 sance itself we don't know who is going or what the verb was and that [00:19:25] information can really help us get some of the harder bits in this prediction problem one famous type of language modeling problem are examples where there is a proper noun that is going to fill a slot but that proper noun was mentioned much earlier in the document consider for example reading a newspaper article where you mention a person's name in the intro paragraph later you might bring up the same person and need to recall what their last name is in theory this is relatively rare and [00:19:55] mostly you can get the easy bits just by saying some proper proun however to get the hard bits of exactly who that person was requires a very long-term memory these and similar examples really motivate the use of fully autor regressive models these are models that have the ability to utilize all previous tokens this obviously is a pretty simple idea and there have been many different models that have tried this we're going to focus though on one particular use of an approach known as attenion that allows us to build fully autor regressive models of relatively long [00:20:26] range you might first ask why a simple neural network couldn't just do this can't we just take the Markoff neural network language modle that we saw earlier 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 [00:20:57] structure that get constructed in the document we're going to focus in on one particular Solution that's at the center of all modern large language models this is an idea known as attention 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 as an example let's return to our simple sentence we're going to have the words the dog walk to the blank [00:21:28] and we're going to want to use that history to predict the next word to do that we're going to need three different pieces of information we'll have one vector known as the query which has looked at the whole sentence so far additionally we'll have a lookup table which has a key and a value for each previous position based on the query we will match the key that we think will be most relevant to our next word prediction from that key we'll then extract the corresponding value that value will then be passed to a neural [00:21:59] network which we can utilize to predict the next word in our sequence recourse steps the query matches the key the best match is selected and then we return the value of that key 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 [00:22:31] useful derivative if we write it as a one-dimensional function we can see that we get a flat structure with a derivative of zero instead we need another function that lets us softly select which key we would like to use the common approach to this is to Simply replace the artmax function with the softmax function that we saw before this softmax instead of producing a distribution over work word types produces a distribution over previous [00:23:01] token positions this distribution is computed softly and has a non-trivial derivative at every location in 2D this softmax function can be drawn with a sigmoid shape which has a nice derivative at every location so here's our new process in step one we use the query to score our key instead of picking the highest scoring value we instead use a softmax to normalize the scores in step three instead of producing a [00:23:31] single value we utilize the softmax to average over the different values waiting them by how well their key match the query here's what this looks like in practice we start with the same query key and value we then match the query to the keys to get a score for each location and we compute a softmax which gives us a histogram over these key locations we then use that histogram to average together the values this produces a new Vector that is some [00:24:04] average of the different values weighted by the histogram the final step is exactly the same instead of using a single value we 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 this attention operation was Central to the key work in language modeling known as attention is all you need this paper written in 2017 [00:24:35] introduces a neural net architecture known as a Transformer that relies heavily on this attention step 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 this parameterizes our language model and produces a nice neural network for predicting the next word distribution the diagram of the Transformer [00:25:05] 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 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 and I've told you [00:25:35] what attention is but not why it's the best way to do this sort of long formed 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 paralyzable it's a nice combination of long range dependency and something that runs fast [00:26:05] on Modern Hardware to make this a bit more clear let me note that in practice we're actually going to be Computing several queries simultaneously at once these queries can be grouped together in a matrix in the same way the memory that has the key and Value Store can also be written as two separate matrices when we combine the query and the key and and compute the soft mags we end up with many different histograms representing all the combinations of the queries and [00:26:35] the keys this 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 this produces es each of the value outputs which are then used to predict the next word the punchline of this process is that the entire attention step can be [00:27:06] written as two Matrix multiplies with a soft Max around the inner one going back to our three steps we can simply read off this mathematical formula them to see that 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 soft Maxes and Matrix multipli but then why is that [00:27:38] actually efficient so far we've now described a full autor regressive 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 I've been bouncing around this question of why language models suddenly got so much better and [00:28:08] if I'm honest the answer is pretty simple it's mostly because of gpus 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 together 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 [00:28:39] models were possible to be built as we've seen earlier the softmax function is central for predicting the next word in the sequence as well as for its use in attention in order to compute the softmax function we need to normalize the distribution this means taking a a sum over every word type in our 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 [00:29:10] 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 a similar example is the calculation of [00:29:40] matrix multiplication 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 I'm going to use this as a running example to teach you a little bit about how gpus work 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 [00:30:11] 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 for memory within a GPU each one of these threads are grouped together into a block in this picture here we can see one GPU block and for this example it corresponds to 12 different individual threads each of these threads again has [00:30:42] to run the same code but they can also read and write from a bit of memory that's seen by the entire block reading and writing to this memory is quite efficient finally the whole GPU consists of a grid this grid has all of the blocks that we've previously seen and each thread in the grid can additionally read from a set of global memory This Global memory is shared by the entire grid but it's quite inefficient to read and write from the main rules of GPU [00:31:14] 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 I think at this point you've got the main idea but let's go through some examples it would be very bad if each of the individual threads was reading and [00:31:45] writing to Global memory by themselves you'd get parallelism but it would be really slow to calculate things in the ideal world we first load from the global memory into to 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 [00:32:15] it can be a little bit counterintuitive and seeing it applied in practice for the first time can be quite challenging to make things more tangible let's run through an example of 3x3 matrix multiplication for this example we're going to have two square matrices A and B both of these will be 3x3 matrices and we'll compute each element of a b by multiplying columns in a with rows in B and then summing up the results here's an example of computing the second row [00:32:46] First Column here we multiply the First Column of a with the second row of B and then sum up the results if we were to do this naively we would have one thread compute each of the outputs in AB in order to compute this value we would do six Global reads the column of a and the row of B do each of the multiplications sum them up with the thread and then write it back out to Global memory note [00:33:16] that doing this requires six Global reads for each thread a better method is going to be to first read from Global memory into the blocks memory we can then calculate important intermediate results within the block itself and then finally write back out to the final value let's look at how this works in practice in step one we read from Global memory into the block memory we'll do this by reading the whole Matrix into our block memory yielding 2 * 9 Global [00:33:48] 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 read 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 we're done we can simply copy the shared AB calculated Matrix back out to Global memory for use in the next [00:34:19] operation 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 Base method all of the reads happen in the first stack which yields 9 * 2 reads for a total of 18 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 [00:34:51] to arbitrarily large matrices you might ask then how you would do a 6x6 Matrix multiplic with blocks of the same size and the answer is that you end up having to do it in multiple steps for step one instead of reading in the entire Matrix we read in a 3X3 block of each of the A and B matrices again this yields 2x9 reads into our 3x3 blocks once we've done this we can calculate a part of the final value by multiplying together the [00:35:23] three values of the column of the top Matrix and three values of the row of the bottom Matrix we then use a single thread reading from block memory to compute the cell 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 [00:35:53] sum them into the final value between the first and the 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 * B Matrix once this is done we can take our 3x3 block and write it back out to Global memory this gives us a 3X3 part of the full a by B output Matrix while this is happening other [00:36:24] blocks are completing the rest of the a * B Matrix each of these are again only doing 2 * 9 reads each time in this case we end up doing 36 total Global reads per block in order to compute the full final Matrix and that's the main operation for this section I probably should have just called it matrix multiplication in practice this operation is often called gem when applied on gpus the GPU operation lets [00:36:54] us to a generalized version of this Matrix multiplier 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 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 attension as well as [00:37:26] the core neural network blocks that are utilized in a Transformer 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 but you might wonder why we're focusing so much on speed isn't enough enough you can run a language bottle on your laptop isn't that good enough why do we have to [00:37:56] 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 so we've talked about generative models and we've talked about Transformers in this section we're going to focus on the p pre-training in particular we're going to focus on the question of scaling these sorts of language models to be trained on lots and lots of data with very large neural [00:38:26] network unlike some of the other sections of this talk the key decisions in scaling seem quite simple we have to decide how big of a neural network to use that means how many parameters should we try to fit in our neural network layers and we have to decide how much training data to use 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 what's [00:38:57] 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 pre-training scales as the product of their two sizes you can think of this horis by thinking that each token that goes to the N Network needs to touch every one of the neural network parameters this is what forms the multiplicative relationship and form the total compute of the system we can [00:39:29] see this relationship in three famous models in the Bert Bas model released in 2018 there were about 109 million parameters and the system was trained on 250 billion tokens this yielded a compute about 1.6 e to the 20 when we jump to a model like pal which was trained in 2022 it has 540 billion parameters and was trained on 780 billion tokens this yielded a total compute of about 2.5 * 10 24th so models [00:40:02] 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 FS so you might ask which of these variables we want to change in a paper on scaling laws researchers at open AI 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 [00:40:34] 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 given this relationship 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 the the problem with this argument is that even if you're extremely GPU Rich you still have a compute budget and you have to determine how to best utilize the [00:41:05] compute that you have available for instance 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 this is not really a theoretical problem for example the pal model which I mentioned earlier [00:41:35] 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 approach we'll use to study this problem is to write down the formula for a power wall 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 the [00:42:07] formula looks a little bit complicated so instead let's draw it as a picture the formula tells us that our perplexity can be predicted as a function of some value a over the number of parameters we use to the alpha term plus a second term which is B over the amount of data we use to a beta term we then add in an additional e term which acts as a bias and corresponds to the best possible [00:42:37] perplexity you could get for the language 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 this formula and its fit are explored in a paper known as chinchillo the main results of chinchillo 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 [00:43:08] 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 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 pal model which used a scaling that favorite 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 [00:43:38] 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 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 Gilla model which we've seen uh those equal scaling ends up being on this line where [00:44:08] other models such as gpt3 or approaches like Megatron touring which is 530 billing parameters end up overparameterized for the amount of tokens that they use there of course lots of uh 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 chinula scaling and how to take advantage of various constraints there are some really important caveats to this scaling [00:44:39] 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 every inference of the actual large language model so every time you actually call Cat GPT 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 [00:45:09] 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 the most notable example of this was the original L 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 [00:45:39] 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 so just to summarize the key way we produce large language models is through this initial pre-training stage particularly when working with generative pre-trained 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 [00:46:09] 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 yet this is still a bit unsatisfactory we know that these models learn a ton from this pre-training 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 [00:46:39] hard tasks part five algorithms 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 it structure but also able to do complex algorithmic tasks such as remember a previous name [00:47:10] 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 at the time 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 one interesting one is [00:47:40] the problem of associative memory here the task is to look at the less than symbol in the context and generate whatever symbol came to the last of fact when you next produce it so in this case here we would generate the letter A which came before the first less than symbol this is obviously an abstraction of the original task um but it is a good representation of what makes the problem hard 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 [00:48:12] for so many different domains and honestly I actually have no idea uh there are 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 version of this question that look into how might language models 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 there are many different formal [00:48:42] 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 just as in the first section of this talk we looked at the probalistic properties of these models and abstracted Away the Transformer part in this section we'll look at the Transformer structure but abstract away the probalistic or stochastic part of [00:49:13] the model basically we're going to be building these little finite State autom like machines that represent Transformer properties 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 if you recall attention has three parts a query key and a value those three parts are then combined to produce an output which then gets [00:49:43] utilized to predict the next word or fed into another layer of attention in our WRA visualization we're going to have a square like the one shown on this Slide the key will be at the top the query will be at the left and the value will be at 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 [00:50:15] 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 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 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 [00:50:45] each of the words in our sentence 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 [00:51:15] 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 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 okay here's another program for this program [00:51:47] what we'd like to do is we'd like to take our input and shift everything one word to the right 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 [00:52:19] 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 the rasp language allows you to combine logic with these attention operations for instance in this match we were trying to [00:52:49] match previous symbols that have the same token value as our current token we can you can see that the only thing that matches is the greater than 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 with these three basic operators we can already start building interesting programs so this Ras program consists of [00:53:19] two layers of attention the first finds the greater than 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 greater than operator arbitrarily far in the past with just two layers of these neural networks operating together we can start to build up pretty interesting programs here's that same program applied to a longer example here [00:53:51] we're able to read this whole input and figure out that Q was the token that was before at the previous greater than symbol again this is a simple example but with the 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 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 [00:54:22] long decimal digits and produce the correct answer 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 uh this particular example it's particularly neat they were able to take the adding circuit from the previous slide and show that they could add a back door to that circuit it's a relatively complex WRA program but they're able to show that for particular inputs they can cause the model to not [00:54:52] add but instead output a mean message I do feel like this section 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 [00:55:22] 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 so end with a short conclusion so we started with this wild example Chad GPT is able to take a pretty ridiculous question 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 handil on how each part of this [00:55:52] 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 while I'm pretty comfortable each of the five formulas I've listed will continue to be pretty important there's already starting to be all sorts of different properties that people are exploring for the first section on perplexity we're seeing all sorts of work into incorporating human [00:56:22] feedback into language models which Chang is the objective and the goals when we talk about 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 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 when we talk about scaling there's this assumption that things will [00:56:52] just keep on getting bigger and bigger but for a lot of Lang languages that's not possible it would be really amazing to produce models that could work with more human scale amounts of data finally I touched a bit on the rasp language and trying to understand how certain algorithms may be coded in Transformers but 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 [00:57:22] to come 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 or how to think about them uh even as we use them for many different tasks so and [00:57:52] there uh thanks so much for listening and uh feel free to ask questions in the comments about further resources or other ways to get involved in any of these areas thanks so much