17 January 2023
Andrej Karpathy · watch on YouTube ↗ · click any timestamp to jump the video
Machine-generated by an AI from the transcript and the top comments. Not my writing, and it may contain errors.
Just under two hours of live coding, building a small GPT on Shakespeare from an empty file. The first forty minutes are data loading and a bigram baseline and can be skimmed if you have written a training loop before. It becomes essential at 47:11, where the matrix-multiply trick is set up, and the section he labels the crux of the video starts at 62:00. The six short notes from 71:38 are the densest part. This transcript was produced locally with faster-whisper; the video has no author-supplied caption track.
The definition the whole video builds toward. Every token emits two vectors: a query, which is roughly what am I looking for, and a key, roughly what do I contain. The affinity between two tokens is just the dot product of one's query with the other's key, so tokens whose interests align end up listening to each other. His worked example is a vowel looking back for consonants.
Before any of the machinery appears, he spends fifteen minutes on a toy demonstration that multiplying by a lower-triangular matrix of ones computes a running average — and that varying the numbers in that matrix turns it into an arbitrary weighted aggregation. Attention is then just the question of where those weights come from. Teaching the trick before the thing it implements is what makes the reveal land.
The first version of token communication is a plain average of everything preceding — literally a bag of words, written with Python for loops. He immediately volunteers that this is an extremely weak form of interaction which throws away all information about arrangement. Naming the inadequacy up front is what makes each later addition feel earned rather than arbitrary.
The reframing that generalises everything: nodes holding vectors, each aggregating from the nodes that point at it, weighted by data. Language modelling just happens to use one particular graph — token one pointed to by itself, token eight by all its predecessors. The triangular mask is a property of the task, not of attention.
It operates over a set of vectors, so position simply does not exist unless you add it — which is what positional encoding is for. The contrast with convolution is the clarifying bit: a convolutional filter is inherently about spatial layout, whereas attention has to be told, separately, that order matters at all.
Not cosmetic. With unit-variance inputs the raw scores have variance on the order of the head size, and those scores feed a softmax. Large values make softmax converge toward one-hot — each token would aggregate from exactly one other token, which is useless at initialisation. The scaling keeps the distribution diffuse so learning can start.
His preferred way to picture it: a residual pathway running straight from input to loss, which computation forks off from and rejoins by addition. Because addition distributes gradient equally to both branches, the signal reaches the input unimpeded. The second half is the part usually left out — the blocks are initialised to contribute almost nothing, so the network starts out nearly a straight wire and the blocks come online during training.
Not an assistant — a document completer that babbles internet instead of Shakespeare. Ask it a question and it may well reply with more questions, because that is what a document containing a question tends to do next. He calls it totally unaligned, and the fine-tuning stage exists entirely to fix that. Clearest explanation of the distinction I have seen in a tutorial.
He pastes in the batch norm from an earlier video and changes a zero to a one. The consequence matters more than the edit: because normalisation no longer spans examples, the running buffers disappear and so does any distinction between training and test behaviour.
Attention lets the tokens talk; the feed-forward layer then processes each token independently. Interleaving the two is the entire structure of the transformer, repeated.
A detail that prevents a common confusion: with a batch of four and a block size of eight you do not have thirty-two communicating nodes, you have four separate pools of eight. The batched matrix multiply keeps the examples completely independent.
The closing claim, and it is the useful one: what he just wrote is architecturally near-identical to the large models, differing by a factor of somewhere between ten thousand and a million in size, plus an infrastructure problem involving thousands of GPUs. The ten-million-parameter Shakespeare model against 175 billion makes the comparison concrete.
The top comment by a wide margin marvels at someone between Tesla and OpenAI casually posting what many consider the best introduction to the subject, for free.
A college lecturer says they return to it not only for the content but to study how to explain a hard topic — which matches the structure of the thing, where each concept is introduced as the solution to a problem you have just watched fail.
A recurring note: people who had looked at the transformer architecture diagram for years without it meaning anything report that this is where it finally resolved.
One of the more cheering replies — a near-seventy-year-old who recently started with Python reports coming away actually understanding how the model works.
Auto-generated captions. Click any line or timestamp to seek; the current line highlights as the video plays.
0:00Hi everyone. So by now you have probably heard of chatGPT. It has taken the world and the AI
0:05community by storm and it is a system that allows you to interact with an AI and give it text-based
0:12tasks. So for example, we can ask chatGPT to write us a small haiku about how important it is that
0:17people understand AI and then they can use it to improve the world and make it more prosperous.
0:21So when we run this, AI knowledge brings prosperity for all to see, embrace its power.
0:28Okay, not bad. And so you could see that chatGPT went from left to right and generated all these words
0:34sequentially. Now I asked it already the exact same prompt a little bit earlier and it generated
0:40a slightly different outcome. AI's power to grow, ignorance holds us back, learn, prosperity,
0:46waits. So pretty good in both cases and slightly different. So you can see that chatGPT is a
0:52probabilistic system and for any one prompt it can give us multiple answers sort of replying to it.
0:58Now this is just one example of a prompt. People have come up with many, many examples
1:02and there are entire websites that index interactions with chatGPT. And so many of
1:08them are quite humorous, explain HTML to me like I'm a dog, write release notes for chess2,
1:14write a note about Elon Musk buying a Twitter and so on. So as an example, please write a
1:21breaking news article about a leaf falling from a tree and a shocking turn of events. A leaf has
1:27fallen from a tree in the local park. Witnesses report that the leaf which was previously attached
1:31to a branch of a tree detached itself and fell to the ground. Very dramatic. So you can see that
1:37this is a pretty remarkable system and it is what we call a language model because it models
1:44the sequence of words or characters or tokens more generally and it knows how sort of words
1:50follow each other in English language. And so from its perspective what it is doing is it is
1:55completing the sequence. So I give it the start of a sequence and it completes the sequence with
2:01the outcome and so it's a language model in that sense. Now I would like to focus on the
2:06under the hood of under the hood components of what makes chatGPT work. So what is the neural
2:12network under the hood that models the sequence of these words? And that comes from this paper
2:18called Attention is All You Need. In 2017 a landmark paper, landmark paper in AI that
2:26proposed the transformer architecture. So GPT is short for Generatively Pre-trained Transformer.
2:34So transformer is the neural net that actually does all the heavy lifting under the hood.
2:38It comes from this paper in 2017. Now if you read this paper, this reads like a pretty
2:44random machine translation paper and that's because I think the authors didn't fully
2:47anticipate the impact that the transformer would have on the field. And this architecture
2:52that they produced in the context of machine translation in their case actually ended up
2:56taking over the rest of AI in the next five years after. And so this architecture with minor
3:03changes was copy pasted into a huge amount of applications in AI in more recent years.
3:09And that includes at the core of chatGPT. Now we are not going to, what I'd like to do now is
3:15I'd like to build out something like chatGPT. But we're not going to be able to of course reproduce
3:21chatGPT. This is a very serious production grade system. It is trained on a good chunk of internet
3:28and then there's a lot of pre-training and fine-tuning stages to it. And so it's very
3:32complicated. What I'd like to focus on is just to train a transformer-based language model.
3:38And in our case it's going to be a character level language model. I still think that is a
3:44very educational with respect to how these systems work. So I don't want to train on the
3:48chunk of internet. We need a smaller dataset. In this case I propose that we work with my favorite
3:53toy dataset. It's called Tiny Shakespeare. And what it is is basically it's a concatenation of all
3:59of the works of Shakespeare in my understanding. And so this is all of Shakespeare in a single
4:04file. This file is about one megabyte and it's just all of Shakespeare. And what we are going
4:10to do now is we're going to basically model how these characters follow each other. So for
4:16example given a chunk of these characters like this, given some context of characters in the past,
4:22the transformer neural network will look at the characters that I've highlighted and is going to
4:26predict that G is likely to come next in the sequence. And it's going to do that because we're
4:31going to train that transformer on Shakespeare. And it's just going to try to produce character
4:37sequences that look like this. And in that process is going to model all the patterns inside this data.
4:43So once we've trained the system, I'd just like to give you a preview,
4:47we can generate infinite Shakespeare. And of course it's a fake thing that looks kind of like
4:53Shakespeare. Apologies for, there's some jank that I'm not able to resolve in here. But you can
5:04see how this is going character by character and it's kind of like predicting Shakespeare-like
5:09language. So verily my lord, the sites have left the again the king coming with my curses with
5:17precious pale. And then Traniosa something else etc. And this is just coming out of the
5:22transformer in a very similar manner as it would come out in chat GPT. In our case character by
5:28character in chat GPT, it's coming out on the token by token level. And tokens are these sort
5:34of like little subword pieces. So they're not word level, they're kind of like word chunk level.
5:40And now I've already written this entire code to train these transformers. And it is in a GitHub
5:49repository that you can find and it's called a nano GPT. So nano GPT is a repository that you can
5:54find on my GitHub. And it's a repository for training transformers on any given text. And
6:01what I think is interesting about it, because there's many ways to train transformers,
6:05but this is a very simple implementation. So it's just two files of 300 lines of code each.
6:10One file defines the GPT model, the transformer, and one file trains it on some given text data set.
6:16And here I'm showing that if you train it on a open web text data set, which is a fairly large
6:20data set of web pages, then I reproduce the performance of GPT-2. So GPT-2 is an early
6:28version of OpenAI's GPT from 2017, if I recall correctly. And I've only so far reproduced the
6:35smallest 124 million parameter model. But basically, this is just proving that the codebase is
6:40correctly arranged. And I'm able to load the neural network weights that OpenAI has released later.
6:47So you can take a look at the finished code here in nano GPT. But what I would like to do in this
6:52lecture is I would like to basically write this repository from scratch. So we're going to
6:58begin with an empty file. And we're going to define a transformer piece by piece. We're going
7:04to train it on the tiny Shakespeare data set. And we'll see how we can then generate infinite
7:09Shakespeare. And of course, this can copy paste to any arbitrary text data set that you like.
7:15But my goal really here is to just make you understand and appreciate how under the hood
7:20chat GPT works. And really all that's required is a proficiency in Python, and some basic
7:27understanding of copulus and statistics. And it would help if you also see my previous videos
7:33on the same YouTube channel, in particular, my Make More series, where I define smaller and
7:40simpler neural network language models. So multi-layer perceptrons and so on, it really
7:46introduces the language modeling framework. And then here in this video, we're going to focus
7:50on the transformer neural network itself. Okay, so I created a new Google Colab Jupyric Notebook
7:57here. And this will allow me to later easily share this code that we're going to develop
8:00together with you so you can follow along. So this will be in a video description later.
8:06Now, here, I've just done some preliminaries, I downloaded the data set, the tiny Shakespeare
8:10data set at this URL. And you can see that it's about a one megabyte file.
8:15Then here, I open the input dot txt file and just read in all the text of the string.
8:19And we see that we are working with 1 million characters roughly. And the first 1000
8:24characters, if we just print them out, are basically what you would expect. This is the
8:28first 1000 characters of the tiny Shakespeare data set, roughly up to here. So, so far, so good.
8:35Next, we're going to take this text. And the text is a sequence of characters in Python.
8:40So when I call the set constructor on it, I'm just going to get the set of all the
8:45characters that occur in this text. And then I call list on that, to create a list of those
8:52characters instead of just a set, so that I have an ordering an arbitrary ordering.
8:56And then I sort that. So basically, we get just all the characters that occur in the entire data
9:01set, and they're sorted. Now, the number of them is going to be our vocabulary size, these are the
9:06possible elements of our sequences. And we see that when I print here the characters,
9:12there's 65 of them in total, there's a space character, and then all kinds of special
9:17characters. And then capitals and lowercase letters. So that's our vocabulary. And that's the
9:23sort of like possible characters that the model can see or emit. Okay, so next, we would like to
9:30develop some strategy to tokenize the input text. Now, when people say tokenize, they mean convert
9:37the raw text as a string to some sequence of integers, according to some note, according to some
9:43vocabulary of possible elements. So as an example, here, we are going to be building a character level
9:49language model. So we're simply going to be translating individual characters into integers.
9:53So let me show you a chunk of code that sort of does that for us. So we're building both the
9:58encoder and the decoder. And let me just talk through what's happening here. When we encode
10:04an arbitrary text, like Hi there, we're going to receive a list of integers that represents that
10:11string. So for example, 46, 47, etc. And then we also have the reverse mapping. So we can take
10:18this list and decode it to get back the exact same string. So it's really just like a translation
10:24to integers and back for arbitrary string. And for us, it is done on a character level.
10:30Now, the way this was achieved is we just iterate over all the characters here and create a lookup
10:35table from the character to the integer and vice versa. And then to encode some string,
10:40we simply translate all the characters individually. And to decode it back, we use the reverse mapping
10:46and concatenate all of it. Now, this is only one of many possible encodings or many possible sort
10:51of tokenizers. And it's a very simple one. But there's many other schemas that people have come
10:56up with in practice. So for example, Google uses a sentence piece. So sentence piece will also
11:02encode text into integers, but in a different schema, and using a different vocabulary. And
11:10sentence piece is a subword sort of tokenizer. And what that means is that you're not encoding
11:16entire words, but you're not also encoding individual characters. It's a subword unit level.
11:22And that's usually what's adopted in practice. For example, also open AI has this library called
11:26tick token that uses a byte pair encoding tokenizer. And that's what GPT uses. And you
11:34can also just encode words into like hell world into a list of integers. So as an example,
11:40I'm using the tick token library here, I'm getting the encoding from GPT two or that was used for
11:45GPT two, instead of just having 65 possible characters, or tokens, they have 50,000 tokens.
11:53And so when they encode the exact same string, hi there, we only get a list of three integers.
11:59But those integers are not between zero and 64. They are between zero and 5000,
12:045,000 256. So basically, you can trade off the codebook size and the sequence lengths.
12:12So you can have very long sequences of integers with very small vocabularies, or we can have a
12:17short sequences of integers with very large vocabularies. And so typically people use in
12:25practice, these subword encodings, but I'd like to keep our tokenizer very simple. So we're using
12:31character level tokenizer. And that means that we have very small codebooks, we have very simple
12:36encode and decode functions. But we do get very long sequences as a result. But that's the level
12:43at which we're going to stick with this lecture, because it's the simplest thing. Okay, so now
12:46that we have an encoder and a decoder, effectively tokenizer, we can tokenize the entire training set
12:52of Shakespeare. So here's a chunk of code that does that. And I'm going to start to use the PyTorch
12:57library, and specifically the torch dot tensor from the PyTorch library. So we're going to take
13:02all of the text in tiny Shakespeare, encode it, and then wrap it into a torch dot tensor,
13:07to get the data tensor. So here's what the data tensor looks like when I look at just the
13:12first 1000 characters, or the 1000 elements of it. So we see that we have a massive sequence of
13:17integers. And this sequence of integers here is basically an identical translation of the first 1000
13:23characters here. So I believe, for example, that zero is a new line character, and maybe one is
13:29a space, not 100% sure. But from now on, the entire data set of text is rerepresented as just
13:35it's just stretched out as a single very large sequence of integers. Let me do one more thing
13:41before we move on here. I'd like to separate out our data set into a train and a validation split.
13:47So in particular, we're going to take the first 90% of the data set, and consider that to be the
13:52training data for the transformer. And we're going to withhold the last 10% at the end of it,
13:58to be the validation data. And this will help us understand to what extent our model is overfitting.
14:03So we're going to basically hide and keep the validation data on the side.
14:06Because we don't want just a perfect memorization of this exact Shakespeare. We want a neural network
14:12that sort of creates Shakespeare like text. And so it should be fairly likely for it to produce
14:18the actual like, stowed away, true Shakespeare text. And so we're going to use this to get a
14:26sense of the overfitting. Okay, so now we would like to start plugging these text sequences or
14:31integer sequences into the transformer so that it can train and learn those patterns.
14:36Now, the important thing to realize is we're never going to actually feed entire text into
14:40transformer all at once, that would be computationally very expensive and prohibitive.
14:45So when we actually train a transformer on a lot of these data sets, we only work with chunks of
14:50the data set. And when we train the transformer, we basically sample random little chunks out of
14:54the training set and train them just chunks at a time. And these chunks have basically some kind of
15:00length and some maximum length. Now, the maximum length, typically, at least in the code I usually
15:06write is called block size. You can you can find it under different names like context length or
15:12something like that. Let's start with the block size of just eight. And let me look at the first
15:16train data characters, the first block size plus one characters, I'll explain why plus one in a second.
15:22And so this is the first nine characters in the sequence in the training set. Now,
15:29what I'd like to point out is that when you sample a chunk of data like this,
15:33so say these nine characters out of the training set, this actually has multiple examples packed
15:39into it. And that's because all of these characters follow each other. And so what this
15:46thing is going to say when we plug it into a transformer, is we're going to actually
15:50simultaneously train it to make prediction at every one of these positions. Now, in the in a chunk
15:56of nine characters, there's actually eight individual examples packed in there. So there's
16:01the example that when 18, when in the context of 1847, luckily comes next, in a context of 18 and
16:0947, 56 comes next, in a context of 1847, 56, 57 can come next, and so on. So that's the eight
16:18individual examples. Let me actually spell it out with code. So here's a chunk of code to illustrate.
16:25X are the inputs to the transformer, it will just be the first block size characters. Y will be the
16:33next block size characters, so it's offset by one. And that's because Y are the targets
16:38for each position in the input. And then here I'm iterating over all the block size of eight.
16:45And the context is always all the characters in X up to T and including T. And the target is always
16:52the teeth character, but in the targets array Y. So let me just run this. And basically, it spells
16:59out what I said in words. These are the eight examples hidden in a chunk of nine characters
17:05that we sampled from the training set. I want to mention one more thing. We train on all the
17:13eight examples here with context between one all the way up to context of block size.
17:19And we train on that not just for computational reasons, because we happen to have the sequence
17:22already or something like that. It's not just done for efficiency. It's also done to make the
17:28transformer network be used to seeing contexts all the way from as little as one all the way to
17:34block size. And we'd like the transformer to be used to seeing everything in between.
17:39And that's going to be useful later during inference, because while we're sampling,
17:43we can start to set a sampling generation with as little as one character of context.
17:47And then transformer knows how to predict next character with all the way up to just
17:51context of one. And so then it can predict everything up to block size. And after block
17:56size, we have to start truncating, because the transformer will never receive more than
18:01block size inputs when it's predicting the next character. Okay, so we've looked at the
18:06time dimension of the tensors that are going to be feeding into the transformer.
18:09There's one more dimension to care about, and that is the batch dimension. And so as we're
18:14sampling these chunks of text, we're going to be actually, every time we're going to feed them
18:18into a transformer, we're going to have many batches of multiple chunks of text that are
18:22always like stacked up in a single tensor. And that's just done for efficiency, just so that we
18:26can keep the GPUs busy, because they are very good at parallel processing of data. And so we
18:34just want to process multiple chunks all at the same time. But those chunks are processed
18:38completely independently, they don't talk to each other, and so on. So let me basically just
18:42generalize this and introduce a batch dimension. Here's a chunk of code. Let me just run it,
18:48and then I'm going to explain what it does. So here, because we're going to start sampling
18:54random locations in the data sets to pull chunks from, I am setting the seed so that
19:00in the random number generator, so that the numbers I see here are going to be the same
19:03numbers you see later, if you try to reproduce this. Now the batch size here is how many independent
19:08sequences we are processing every forward backward pass of the transformer. The block size, as I
19:14explained, is the maximum context length to make those predictions. So let's say batch size 4,
19:20block size 8, and then here's how we get batch for any arbitrary split. If the split is a training
19:26split, then we're going to look at train data, otherwise, and val data. That gives us the data
19:32array. And then when I generate random positions to grab a chunk out of, I actually
19:39generate batch size number of random offsets. So because this is 4, Ix is going to be
19:474 numbers that are randomly generated between 0 and len of data minus block size. So it's just
19:53random offsets into the training set. And then x's, as I explained, are the first block size
20:00characters starting at i. The y's are the offset by 1 of that, so just add plus 1.
20:08And then we're going to get those chunks for every one of integers i in Ix, and use a torch.stack
20:15to take all those one-dimensional tensors, as we saw here, and we're going to stack them up
20:23at rows. And so they all become a row in a 4 by 8 tensor. So here's what I'm printing then.
20:32When I sample a batch xb and yb, the inputs of the transformer now are,
20:38the input x is the 4 by 8 tensor, 4 rows of 8 columns, and each one of these is a chunk of
20:47the training set. And then the targets here are in the associated array y, and they will come in
20:54to the transformer all the way at the end to create the loss function. So they will give us
21:01the correct answer for every single position inside x. And then these are the four independent rows.
21:09So spelled out, as we did before, this 4 by 8 array contains a total of 32 examples.
21:17And they're completely independent as far as the transformer is concerned.
21:22So when the input is 24, the target is 43, or rather 43 here in the y array. When the input is 2443,
21:31the target is 58. When the input is 2443, 58, the target is 5, etc. Or like when it is a
21:385258, 1, the target is 58. Right, so you can sort of see this spelled out. These are the 32
21:45independent examples packed in to a single batch of the input x, and then the desired targets are
21:52in y. And so now this integer tensor of x is going to feed into the transformer.
22:01And that transformer is going to simultaneously process all these examples, and then look up the
22:06correct integers to predict in every one of these positions in the tensor y. Okay, so now
22:12that we have our batch of input that we'd like to feed into a transformer, let's start basically
22:17feeding this into neural networks. Now we're going to start off with the simplest possible neural
22:21network, which in the case of language modeling, in my opinion, is the bigram language model.
22:25And we've covered the bigram language model in my Make More series in a lot of depth.
22:30And so here I'm going to sort of go faster, and let's just implement the PyTorch module
22:34directly that implements the bigram language model. So I'm importing the PyTorch nn module.
22:41For reproducibility. And then here I'm constructing a bigram language model, which is a subclass of
22:46nn module. And then I'm calling it and I'm passing in the inputs and the targets.
22:53And I'm just printing. Now when the inputs and targets come here, you see that I'm just taking
22:57the index, the inputs x here, which I renamed to IDX, and I'm just passing them into this token
23:04embedding table. So what's going on here is that here in the constructor, we are creating a
23:10token embedding table, and it is of size vocab size by vocab size. And we're using an endot
23:17embedding, which is a very thin wrapper around basically a tensor of shape vocab size by vocab
23:23size. And what's happening here is that when we pass IDX here, every single integer in our input
23:30is going to refer to this embedding table, and is going to pluck out a row of that embedding
23:34table corresponding to its index. So 24 here, we'll go to the embedding table, and we'll pluck
23:40out the 24th row. And then 43 will go here and pluck out the 43rd row, etc. And then PyTorch is
23:47going to arrange all of this into a batch by time by channel tensor. In this case batch is
23:54four, time is eight, and C, which is the channels, is vocab size or 65. And so we're just going to
24:02pluck out all those rows, arrange them in a B by T by C. And now we're going to interpret this as
24:07the logits, which are basically the scores for the next character in a sequence. And so what's
24:13happening here is we are predicting what comes next based on just the individual identity of a
24:19single token. And you can do that because I mean, currently the tokens are not talking to each
24:24other, and they're not seeing any context except for they're just seeing themselves. So I'm a
24:29token number five, and then I can actually make pretty decent predictions about what comes next,
24:35just by knowing that I'm token five, because some characters follow other characters in typical
24:42scenarios. So we saw a lot of this in a lot more depth in the Make More series. And here,
24:47if I just run this, then we currently get the predictions, the scores, the logits for every
24:53one of the four by eight positions. Now that we've made predictions about what comes next,
24:57we'd like to evaluate the loss function. And so in Make More series, we saw that a good way to measure
25:03a loss or like a quality of the predictions is to use the negative log likelihood loss, which is
25:08also implemented in PyTorch under the name cross entropy. So what we'd like to do here is loss is
25:15the cross entropy on the predictions and the targets. And so this measures the quality of the
25:20logits with respect to the targets. In other words, we have the identity of the next character.
25:26So how well are we predicting the next character based on logits? And intuitively, the correct
25:33dimension of logits, depending on whatever the target is, should have a very high number.
25:39And all the other dimensions should be very low number, right? Now, the issue is that this won't
25:44actually, this is what we want, we want to basically output the logits and the loss.
25:51This is what we want. But unfortunately, this won't actually run. We get an error message.
25:57But intuitively, we want to measure this. Now, when we go to the PyTorch cross entropy
26:05documentation here, we're trying to call the cross entropy and its functional form.
26:11So that means we don't have to create like a module for it. But here, when we go to the documentation,
26:17you have to look into the details of how PyTorch expects these inputs. And basically,
26:21the issue here is PyTorch expects, if you have multi dimensional input, which we do,
26:26because we have a b by t by c tensor, then it actually really wants the channels to be
26:32the second dimension here. So if you so basically, it wants a b by c by t, instead of a b by t by c.
26:42And so it's just the details of how PyTorch treats these kinds of inputs. And so we don't
26:50actually want to deal with that. So we're going to do instead is we need to basically reshape our
26:53logits. So here's what I like to do. I like to take basically give names to the dimensions.
26:58So logits dot shape is b by t by c and unpack those numbers. And then let's say that logits
27:05equals logits dot view. And we want it to be a b times c, b times t by c. So just a two dimensional
27:12array. Right, so we're going to take all the we're going to take all of these positions here,
27:20and we're going to stretch them out in a one dimensional sequence and preserve the channel
27:25dimension as the second dimension. So we're just kind of like stretching out the array. So it's
27:30two dimensional. And in that case, it's going to better conform to what PyTorch sort of expects
27:35in its dimensions. Now we have to do the same two targets, because currently targets are
27:41of shape b by t. And we want it to be just b times t. So one dimensional. Now, alternatively,
27:50you could always still just do minus one, because PyTorch will guess what this should be if you
27:54want to lay it out. But let me just be explicit and save you times t. Once we reshape this,
28:00it will match the cross entropy case. And then we should be able to evaluate our loss.
28:07Okay, so that right now, and we can do loss. And so currently, we see that the loss is 4.87.
28:15Now, because our we have 65 possible vocabulary elements, we can actually guess at what the
28:20loss should be. And in particular, we covered negative log likelihood in a lot of detail,
28:26we are expecting log or lon of one over 65 and negative of that. So we're expecting the
28:35loss to be about 4.17, but we're getting 4.87. And so that's telling us that the initial
28:41predictions are not super diffuse, they've got a little bit of entropy. And so we're guessing
28:45wrong. So yes, but actually, we are able to evaluate the loss. Okay, so now that we can
28:53evaluate the quality of the model on some data, we'd like to also be able to generate from the
28:59model. So let's do the generation. Now I'm going to go again a little bit faster here, because I
29:03covered all this already in previous videos. So here's a generate function for the model.
29:12So we take some, we take the same kind of input idx here. And basically, this is the current
29:21context of some characters in a batch in some batch. So it's also b by t. And the job of
29:28generate is to basically take this b by t and extend it to b by t plus one plus two plus three.
29:33And so it's just basically it continues the generation in all the batch dimensions in the
29:38time dimension. So that's its job. And it will do that for max new tokens. So you can see here on
29:44the bottom, there's going to be some stuff here. But on the bottom, whatever is predicted
29:48is concatenated on top of the previous idx along the first dimension, which is the time dimension,
29:54to create a b by t plus one. So that becomes a new idx. So the job of generate is to take
30:00a b by t and make it a b by t plus one plus two plus three, as many as we want maximum tokens. So
30:06this is the generation from the model. Now inside the generation, what are we doing? We're taking
30:11the current indices, we're getting the predictions. So we get those are in the logits. And then the
30:19loss here is going to be ignored, because we're not we're not using that and we have no targets
30:24that are sort of ground truth targets that we're going to be comparing with. Then once we get
30:29the logits, we are only focusing on the last step. So instead of a b by t by c, we're going to
30:35pluck out the negative one, the last element in the time dimension, because those are the
30:41predictions for what comes next. So that gives us the logits, which we then convert to probabilities
30:46via softmax. And then we use Torch that multinomial to sample from those probabilities.
30:51And we ask PyTorch to give us one sample. And so idx next will become a b by one,
30:57because in each one of the batch dimensions, we're going to have a single prediction for
31:01what comes next. So this num samples equals one will make this be a one. And then we're going
31:08to take those integers that come from the sampling process according to the probability distribution
31:12given here. And those integers got just concatenated on top of the current sort of like running stream
31:18of integers. And this gives us a b by t plus one. And then we can return that. Now one thing here
31:25is you see how I'm calling self of idx, which will end up going to the forward function,
31:31I'm not providing any targets. So currently, this would give an error because targets is
31:37sort of like not given. So targets has to be optional. So targets is none by default.
31:42And then if targets is none, then there's no loss to create. So it's just losses none. But else,
31:50all of this happens and we can create a loss. So this will make it so if we have the targets,
31:56we provide them and get a loss. If we have no targets, we'll just get the logits.
32:01So this here will generate from the model. And let's take that for a ride now.
32:11So I have another code chunk here, which will generate for the model from the model. And okay,
32:15this is kind of crazy. So maybe let me let me break this down. So these are the idx, right?
32:22I'm creating a batch will be just one time will be just one. So I'm creating a little one by one
32:31tensor, and it's holding a zero. And the D type the data type is integer. So zero is going to
32:38be how we kick off the generation. And remember that zero is, is the element standing for a
32:45newline character. So it's kind of like a reasonable thing to to feed in as the very
32:48first character in a sequence to be the newline. So it's going to be idx, which we're going to feed
32:55in here, then we're going to ask for 100 tokens, and then n.generate will continue that. Now
33:02because generate works on the level of batches, we then have to index into the zero throw
33:09to basically unblock the single batch dimension that exists. And then that gives us a time steps,
33:20just a one dimensional array of all the indices, which we will convert to simple Python list from
33:26PyTorch tensor, so that that can feed into our decode function and convert those integers into
33:33text. So let me bring this back. And we're generating 100 tokens, let's run. And here's
33:40the generation that we achieved. So obviously, it's garbage. And the reason it's garbage is
33:44because this is totally random model. So next up, we're going to want to train this model.
33:49Now, one more thing I wanted to point out here is this function is written to be general.
33:54But it's kind of like ridiculous right now, because we're feeding in all this,
33:59we're building out this context, and we're concatenating it all. And we're always feeding it
34:04all into the model. But that's kind of ridiculous, because this is just a simple bigram model. So to
34:10make, for example, this prediction about K, we only needed this W. But actually, what we fed into
34:15the model is we fed the entire sequence. And then we only looked at the very last piece,
34:20and predicted K. So the only reason I'm writing it in this way is because right now, this is a
34:26bigram model. But I'd like to keep this function fixed. And I'd like it to work later, when our
34:33characters actually, basically look further in the history. And so right now, the history is not used.
34:39So this looks silly. But eventually, the history will be used. And so that's why we want to do it
34:45this way. So just a quick comment on that. So now, we see that this is random. So let's train
34:51the model. So it becomes a bit less random. Okay, let's now train the model. So first,
34:56what I'm going to do is I'm going to create a PyTorch optimization object. So here we are using
35:01the optimizer Adam W. Now, in the Make More series, we've only ever used stochastic gradient descent,
35:07the simplest possible optimizer, which you can get using the SGD instead. But I want to use Adam,
35:12which is a much more advanced and popular optimizer. And it works extremely well. For
35:18typical good setting for the learning rate is roughly three, negative four. But for very,
35:23very small networks, like is the case here, you can get away with much, much higher learning rates,
35:27running negative three or even higher probably. But let me create the optimizer object,
35:32which will basically take the gradients and update the parameters using the gradients.
35:38And then here, our batch size up above was only four. So let me actually use something bigger,
35:43let's say 32. And then for some number of steps, we are sampling a new batch of data,
35:49we're evaluating the loss, we're zeroing out all the gradients from the previous step,
35:54getting the gradients for all the parameters, and then using those gradients to update our
35:58parameters. So typical training loop, as we saw in the Make More series. So let me now run this
36:05for say 100 iterations, and let's see what kind of losses we're gonna get.
36:11So we started around 4.7. And now we're getting down to like 4.6, 4.5, etc. So the optimization
36:18is definitely happening. But let's sort of try to increase the number of iterations and only print at
36:25the end, because we probably will not train for longer. Okay, so we're down to 3.6 roughly,
36:36roughly down to three. This is the most janky optimization. Okay, it's working. Let's just
36:48do 10,000. And then from here, we want to copy this. And hopefully, we're going to get something
36:57reasonable. And of course, it's not going to be Shakespeare from a background model. But at
37:01least we see that the loss is improving. And hopefully, we're expecting something a bit more
37:06reasonable. Okay, so we're down at about 2.5 ish. Let's see what we get. Okay, dramatic improvement
37:14certainly on what we had here. So let me just increase the number of tokens. Okay, so we see
37:20that we're starting to get something at least like reasonable ish. Certainly not Shakespeare.
37:29But the model is making progress. So that is the simplest possible model. So now what I'd like to do
37:35is obviously that this is a very simple model, because the tokens are not talking to each other.
37:42So given the previous context of whatever was generated, we're only looking at the very last
37:46character to make the predictions about what comes next. So now these now these tokens have to
37:51start talking to each other and figuring out what is in the context so that they can make better
37:56predictions for what comes next. And this is how we're going to kick off the transformer.
38:01Okay, so next, I took the code that we developed in this Jupyter Notebook, and I converted it to be a
38:05script. And I'm doing this because I just want to simplify our intermediate work, which is just the
38:10final product that we have at this point. So in the top here, I put all the hyper parameters
38:16that we find, I introduced a few, and I'm going to speak to that in a little bit. Otherwise,
38:21a lot of this should be recognizable. Reproducibility, read data, get the encoder and the decoder,
38:27create the train and test splits, use the data loader that gets a batch of the inputs and targets.
38:36This is new, and I'll talk about it in a second. Now, this is the background language model that
38:40we developed, and it can forward and give us a logits and loss and it can generate.
38:45And then here, we are creating the optimizer and this is the training loop. So everything here
38:52should look pretty familiar. Now, some of the small things that I added, number one,
38:56I added the ability to run on a GPU if you have it. So if you have a GPU, then you can,
39:02this will use CUDA instead of just CPU, and everything will be a lot more faster.
39:06Now, when device becomes CUDA, then we need to make sure that when we load the data,
39:12we move it to device. When we create the model, we want to move the model parameters to device.
39:19So as an example, here we have the nn embedding table and it's got a dot weight inside it,
39:25which stores the sort of lookup table. So that would be moved to the GPU so that all
39:30the calculations here happen on the GPU and they can be a lot faster. And then finally,
39:34here when I'm creating the context that feeds into generate, I have to make sure that I create
39:39on the device. Number two, what I introduced is the fact that here in the training loop,
39:48here I was just printing the loss that item inside the training loop. But this is a very noisy
39:54measurement of the current loss because every batch will be more or less lucky. And so what
39:59I want to do usually is I have an estimate loss function. And the estimate loss basically then
40:07goes up here. And it averages up the loss over multiple batches. So in particular,
40:15we're going to iterate eval iter times. And we're going to basically get our loss. And then we're
40:20going to get the average loss for both splits. And so this will be a lot less noisy. So here
40:25when we call the estimate loss, we're going to report the pretty accurate train and validation loss.
40:32Now, when we come back up, you'll notice a few things here. I'm setting the model to a value
40:36evaluation phase. And down here, I'm resetting it back to training phase. Now right now for our model
40:42as is, this doesn't actually do anything. Because the only thing inside this model is this
40:47nn.embedding. And this network would behave both would behave the same in both evaluation
40:55mode and training mode. We have no dropout layers, we have no batch norm layers, etc.
41:00But it is a good practice to think through what mode your neural network is in. Because some layers
41:05will have different behavior at inference time or training time. And there's also this context
41:12manager, torch.nograd. And this is just telling PyTorch that everything that happens inside this
41:17function, we will not call dot backward on. And so PyTorch can be a lot more efficient with its
41:22memory use, because it doesn't have to store all the intermediate variables, because we're never
41:27going to call backward. And so it can it can be a lot more efficient in that way. So also a
41:32good practice to tell PyTorch when we don't intend to do back propagation. So right now,
41:39this script is about 120 lines of code of and that's kind of our starter code. I'm calling it
41:46by gram.py. And I'm going to release it later. Now running this script gives us output in the
41:52terminal. And it looks something like this. It basically, as I ran this code, it was giving
41:58me the train loss and val loss. And we see that we convert to somewhere around 2.5 with the
42:03bigram model. And then here's the sample that we produced at the end. And so we have everything
42:09packaged up in the script. And we're in a good position now to iterate on this. Okay, so we
42:14are almost ready to start writing our very first self attention block for processing these tokens.
42:21Now, before we actually get there, I want to get you used to a mathematical trick that is used
42:26in the self attention inside a transformer. And it's really just like at the heart of an
42:32efficient implementation of self attention. And so I want to work with this toy example to just
42:37get you used to this operation. And then it's going to make it much more clear once we actually
42:41get to it in the script again. So let's create a b by t by c, where b, t and c are just four,
42:49eight and two in this toy example. And these are basically channels. And we have batches. And
42:55we have the time component. And we have some information at each point in the sequence. So c.
43:02Now, what we would like to do is we would like these tokens. So we have up to eight tokens here
43:07in a batch. And these eight tokens are currently not talking to each other, and we would like them
43:12to talk to each other, we'd like to couple them. And in particular, we don't we we want to couple
43:18them in a very specific way. So the token, for example, at the fifth location, it should not
43:24communicate with tokens in the sixth, seventh and eighth location. Because those are future
43:29tokens in the sequence. The token on the fifth location should only talk to the one in the
43:34fourth, third, second and first. So it's only so information only flows from previous context to the
43:40current timestamp. And we cannot get any information from the future because we are about to try to
43:44predict the future. So what is the easiest way for tokens to communicate? Okay, the easiest
43:52way I would say is, okay, if we're up to if we're a fifth token, and I'd like to communicate with my
43:57past, the simplest way we can do that is to just do a week is to just do an average of all the of
44:04all the preceding elements. So for example, if I'm the fifth token, I would like to take the channels
44:10that make up that are information at my step, but then also the channels from the fourth step,
44:16third step, second step in the first step, I'd like to average those up. And then that would
44:20become sort of like a feature vector that summarizes me in the context of my history.
44:26Now, of course, just doing a sum or like an average is an extremely weak form of interaction,
44:30like this communication is extremely lossy, we've lost a ton of information about spatial
44:35arrangements of all those tokens. But that's okay. For now, we'll see how we can bring that
44:39information back later. For now, what we would like to do is for every single batch element
44:44independently, for every teeth token, in that sequence, we'd like to now calculate the average
44:52of all the vectors in all the previous tokens, and also at this token. So let's write that out.
45:00I have a small snippet here. And instead of just fumbling around,
45:03let me just copy paste it and talk to it. So in other words, we're going to create x
45:10and the BOW is short for bag of words, because bag of words is, is kind of like a term that
45:16people use when you are just averaging up things. So it's just a bag of words. Basically, there's a
45:21word stored on every one of these eight locations, and we're doing a bag of words, which is averaging.
45:27So in the beginning, we're going to say that it's just initialized at zero. And then I'm doing
45:31a for loop here. So we're not being efficient yet, that's coming. But for now, we're just
45:35iterating over all the batch dimensions independently, iterating over time. And then the previous tokens
45:43are at this batch dimension, and then everything up to and including the teeth token. Okay.
45:51So when we slice out x in this way, x prev becomes of shape, how many t elements there were in the
45:59past, and then of course, see, so all the two dimensional information from these little tokens.
46:05So that's the previous sort of chunk of tokens from my current sequence. And then I'm just doing
46:13the average or the mean over the zero dimensions. So I'm averaging out the time here. And I'm just
46:19going to get a little C one dimensional vector, which I'm going to store in x bag of words. So I
46:25can run this. And this is not going to be very informative, because let's see. So this is x of
46:32zero. So this is the zero batch element, and then expo at zero. Now, you see how the at the first
46:40location here, you see that the two are equal. And that's because it's we're just doing an average
46:46of this one token. But here, this one is now an average of these two. And now this one is an
46:53average of these three, and so on. So, and this last one is the average of all of these elements.
47:02So vertical average, just averaging up all the tokens now gives this outcome here. So this is
47:09all well and good. But this is very inefficient. Now, the trick is that we can be very, very
47:14efficient about doing this using matrix multiplication. So that's the mathematical trick. And let me
47:19show you what I mean. Let's work with the toy example here. Let me run it and I'll explain.
47:25I have a simple matrix here that is three by three of all ones. A matrix B of just random numbers,
47:32and it's a three by two, and a matrix C, which will be three by three, multiply three by two,
47:37which will give out a three by two. So here, we're just using matrix multiplication. So A
47:44multiply B gives us C. So how are these numbers in C achieved, right? So this number in the top left
47:55is the first row of A dot product with the first column of B. And since all the row of A right now
48:02is all just ones, then the dot product here with this column of B is just going to do a sum of
48:09this column. So two plus six plus six is 14. The element here in the output of C is also the first
48:16column here, the first row of A multiplied now with the second column of B. So seven plus four plus
48:23plus five is 16. Now, you see that there's repeating elements here. So this 14 again
48:28is because this row is again all ones, and it's multiplying the first column of B. So we get 14.
48:33And this one is, and so on. So this last number here is the last row dot product last column.
48:40Now the trick here is the following. This is just a boring number of,
48:46it's just a boring array of all ones. But Torch has this function called trill,
48:52which is short for a triangular, something like that. And you can wrap it in Torch. Once,
48:58and it will just return the lower triangular portion of this. So now it will basically zero out
49:07these guys here. So we just get the lower triangular part. Well, what happens if we do
49:11that? So now we'll have A like this and B like this. And now what are we getting here in C?
49:20Well, what is this number? Well, this is the first row times the first column. And because this
49:26is zeros, these elements here are now ignored. So we just get a two. And then this number here
49:34is the first row times the second column. And because these are zeros, they get ignored,
49:38and it's just seven. The seven multiplies this one. But look what happened here, because this
49:44is one and then zeros, we, what ended up happening is we're just plucking out the row,
49:49this row of B, and that's what we got. Now here, we have 110. So here, 110 dot product with these
49:58two columns will now give us two plus six, which is eight, and seven plus four, which is 11. And
50:03because this is 111, we ended up with the addition of all of them. And so basically, depending on how
50:10many ones and zeros we have here, we are basically doing a sum currently of the variable number of
50:18these rows, and that gets deposited into C. So currently, we're doing sums because these are ones,
50:24but we can also do average, right? And you can start to see how we can do average of the rows of B,
50:31sort of in an incremental fashion, because we don't have to, we can basically normalize
50:36these rows so that they sum to one, and then we're going to get an average. So if we took A,
50:42and then we did A equals A divide, a torch dot sum in the, of A, in the one dimension,
50:54and then let's keep them as true. So therefore, the broadcasting will work out. So if I rerun this,
51:00you see now that these rows now sum to one. So this row is one, this row is 0.5, 0.5 is zero,
51:07and here we get one-thirds. And now when we do A multiply B, what are we getting?
51:12Here we are just getting the first row, first row. Here now we are getting the average of the first
51:18two rows. Okay, so two and six average is four, and four and seven average is 5.5.
51:26And on the bottom here, we're now getting the average of these three rows. So the average of
51:33all of elements of B are now deposited here. And so you can see that by manipulating these
51:40elements of this multiplying matrix, and then multiplying it with any given matrix,
51:45we can do these averages in this incremental fashion, because we just get, and we can
51:52manipulate that based on the elements of A. Okay, so that's very convenient. So let's swing back up
51:57here and see how we can vectorize this and make it much more efficient using what we've learned.
52:02So in particular, we are going to produce an array A, but here I'm going to call it way
52:08short for weights. But this is our A, and this is how much of every row we want to average up.
52:17And it's going to be an average because you can see that these rows sum to one.
52:21So this is our A, and then our B in this example, of course, is x. So what's going to happen here
52:29now is that we are going to have an expo two. And this expo two is going to be way multiplying
52:38rx. So let's think this through way is t by t. And this is matrix multiplying in PyTorch,
52:45a B by t by c. And it's giving us the what shape. So PyTorch will come here and it will see that
52:53these shapes are not the same. So it will create a batch dimension here. And this is a batch matrix
52:59multiply. And so it will apply this matrix multiplication and all the batch elements in
53:05parallel and individually. And then for each batch element, there will be a t by t multiplying
53:11t by c exactly as we had below. So this will now create B by t by c. And expo two will now
53:23become identical to expo. So we can see that torch.allclose of expo and expo two should be true.
53:35Now, so this kind of like, convinces us that these are in fact the same. So expo and expo two,
53:45if I just print them. Okay, we're not going to be able to, okay, we're not going to be able to just
53:52stare it down. But well, let me try expo basically just at the zero element and expo two at the
53:58zero element. So just the first batch. And we should see that this and that should be identical,
54:04which they are. Right? So what happened here, the trick is we were able to use batch matrix multiply
54:11to do this aggregation, really. And it's a weighted aggregation. And the weights are specified in
54:19this t by t array. And we're basically doing weighted sums. And these weighted sums are
54:26according to the weights inside here, they take on sort of this triangular form. And so that means
54:33that a token at the t dimension will only get sort of information from the tokens preceding it.
54:41So that's exactly what we want. And finally, I would like to rewrite it in one more way.
54:45And we're going to see why that's useful. So this is the third version. And it's also identical to the
54:51first and second. But let me talk through it, it uses softmax. So trill here is this matrix,
55:00lower triangular ones, way begins as all zero. Okay, so if I just print way in the beginning,
55:09it's all zero, then I use masked fill. So what this is doing is weight that masked fill,
55:17it's all zeros. And I'm saying for all the elements where trill is equal to equal zero,
55:23make them be negative infinity. So all the elements where trill is zero will become
55:28negative infinity now. So this is what we get. And then the final line here is softmax.
55:37So if I take a softmax along every single sub div is negative one, so along every single row,
55:43if I do a softmax, what is that going to do? Well, softmax is, it's also like a normalization
55:53operation, right? And so spoiler alert, you get the exact same matrix. Let me bring back the softmax.
56:01And recall that in softmax, we're going to exponentiate every single one of these.
56:06And then we're going to divide by the sum. And so if we exponentiate every single element here,
56:11we're going to get a one. And here, we're going to get basically zero, zero, zero, zero,
56:16everywhere else. And then when we normalize, we just get one, here, we're going to get one,
56:21one, and then zeros, and then softmax will again divide, and this will give us point five,
56:26point five, and so on. And so this is also the same way to produce this mask. Now,
56:33the reason that this is a bit more interesting, and the reason we're going to end up using it and
56:37self attention is that these weights here begin with zero. And you can think of this as like an
56:45interaction strength, or like an affinity. So basically, it's telling us how much of each
56:52token from the past do we want to aggregate an average up. And then this line is saying,
56:59tokens from the past cannot communicate by setting them to negative infinity,
57:04we're saying that we will not aggregate anything from those tokens. And so basically,
57:09this then goes through softmax and through the weighted, and this is the aggregation through
57:13matrix multiplication. And so what this is now is, you can think of these as these zeros
57:20are currently just set by us to be zero. But a quick preview is that these affinities
57:26between the tokens are not going to be just constant at zero, they're going to be data dependent,
57:31these tokens are going to start looking at each other. And some tokens will find other tokens
57:36more or less interesting. And depending on what their values are, they're going to find each other
57:41interesting to different amounts, and I'm going to call those affinities, I think. And then here
57:46we are saying the future cannot communicate with the past, we're going to clamp them.
57:51And then when we normalize and sum, we're going to aggregate sort of their values depending on how
57:57interestingly they find each other. And so that's the preview for self attention. And basically,
58:03long story short from this entire section is that you can do weighted aggregations of your past
58:08elements by having by using matrix multiplication of a lower triangular fashion. And then the
58:16elements here in the lower triangular part are telling you how much of each element fuses into
58:22this position. So we're going to use this trick now to develop the self attention block.
58:27So first, let's get some quick preliminaries out of the way.
58:30First, the thing I'm kind of bothered by is that you see how we're passing and vocab size into
58:34the constructor, there's no need to do that because vocab size is already defined up top
58:38as a global variable. So there's no need to pass this stuff around. Next, what I want to do is,
58:44I don't want to actually create, I want to create like a level of indirection here where we don't
58:48directly go to the embedding for the logits. But instead, we go through this intermediate phase,
58:54because we're going to start making that bigger. So let me introduce a new variable and embed.
59:01It's short for number of embedding dimensions. So an embed here will be say 32. That was a
59:09suggestion from GitHub co-pilot, by the way, it also shows 32, which is a good number.
59:15So this is an embedding table and only 32 dimensional embeddings. So then here,
59:21this is not going to give us logits directly. Instead, this is going to give us token embeddings.
59:26That's what I'm going to call it. And then to go from the token embeddings to the logits,
59:30we're going to need a linear layer. So self dot lmhead, let's call it short for
59:34language modeling head, is n and linear from n embed up to vocab size. And then when we
59:40swing over here, we're actually going to get the logits by exactly what the co-pilot says.
59:46Now we have to be careful here because this C and this C are not equal. This is an embed C
59:53and this is vocab size. So let's just say that n embed is equal to C. And then this just creates
1:00:00one spurious layer of interaction through a linear layer. But this should basically run.
1:00:12So we see that this runs and this currently looks kind of spurious, but we're going to build on top
1:00:18of this. Now next up, so far, we've taken these indices and we've encoded them based on the
1:00:24identity of the tokens inside IDX. The next thing that people very often do is that we're not just
1:00:31encoding the identity of these tokens, but also their position. So we're going to have a second
1:00:36position embedding table here. So solve that position embedding table is an embedding of block
1:00:42size by n embed. And so each position from zero to block size minus one will also get its own
1:00:48embedding vector. And then here, first let me decode b by t from IDX.shape. And then here,
1:00:56we're also going to have a plus embedding, which is the positional embedding. And these are,
1:01:00this is tortoise range. So this will be basically just integers from zero to t minus one.
1:01:06And all of those integers from zero to t minus one get embedded through the table to create a t by c.
1:01:12And then here, this gets renamed to just say x. And x will be the addition of the token embeddings
1:01:19with the positional embeddings. And here, the broadcasting note will work out. So b by t by c
1:01:25plus t by c, this gets right aligned, a new dimension of one gets added, and it gets
1:01:30broadcasted across batch. So at this point, x holds not just the token identities, but the positions
1:01:38at which these tokens occur. And this is currently not that useful, because of course, we just have
1:01:42a simple bigram model. So it doesn't matter if you're in the fifth position, the second
1:01:46position or wherever, it's all translation invariant at this stage. So this information
1:01:50currently wouldn't help. But as we work on the self attention block, we'll see that this
1:01:55starts to matter. Okay, so now we get the crux of self attention. So this is probably the
1:02:03most important part of this video to understand. We're going to implement a small self attention
1:02:08for a single individual head as they're called. So we start off with where we were. So all of
1:02:14this code is familiar. So right now, I'm working with an example where I changed number of channels
1:02:19from two to 32. So we have a four by eight arrangement of tokens. And each token and the
1:02:25information at each token is currently 32 dimensional, but we just are working with random numbers.
1:02:31Now we saw here that the code as we had it before, does a simple weight, simple average of all the
1:02:39past tokens, and the current token. So it's just the previous information and current information
1:02:44is just being mixed together in an average. And that's what this code currently achieves.
1:02:49And it does so by creating this lower triangular structure, which allows us to mask out this
1:02:55weight matrix that we create. So we mask it out, and then we normalize it. And currently,
1:03:02when we initialize the affinities between all the different sort of tokens or nodes, I'm going to
1:03:08use those terms interchangeably. So when we initialize the affinities between all the
1:03:13different tokens to be zero, then we see that way gives us this structure where every single row
1:03:19has these uniform numbers. And so that's what then in this matrix multiply makes it so that
1:03:27we're doing a simple average. Now, we don't actually want this to be all uniform, because
1:03:35different tokens will find different other tokens more or less interesting, and we want that to be
1:03:41data dependent. So for example, if I'm a vowel, then maybe I'm looking for consonants in my past,
1:03:46and maybe I want to know what those consonants are, and I want that information to flow to me.
1:03:51And so I want to now gather information from the past, but I want to do it in a data dependent way.
1:03:56And this is the problem that self attention solves. Now, the way self attention solves this
1:04:01is the following. Every single node or every single token at each position will emit two
1:04:08vectors, it will emit a query, and it will emit a key. Now, the query vector, roughly speaking,
1:04:16is what am I looking for? And the key vector, roughly speaking, is what do I contain?
1:04:23And then the way we get the affinities between these tokens now in a sequence,
1:04:28is we basically just do a dot product between the keys and the queries.
1:04:33So my query dot products with all the keys of all the other tokens,
1:04:37and that dot product now becomes way. And so if the key and the query are sort of aligned,
1:04:47they will interact to a very high amount. And then I will get to learn more about that specific token,
1:04:54as opposed to any other token in the sequence. So let's implement this now. We're going to implement
1:05:02a single what's called head of self attention. So this is just one head, there's a hyper parameter
1:05:10involved with these heads, which is the head size. And then here, I'm initializing linear modules,
1:05:17and I'm using bias equals false. So these are just going to apply matrix multiply with some fixed weights.
1:05:23And now let me produce a key and q, k and q, by forwarding these modules on x.
1:05:31So the size of this will now become b by t by 16, because that is the head size.
1:05:38And the same here, b by t by 16, that's being that size. So you see here that when I forward
1:05:49this linear on top of my x, all the tokens in all the positions in the b by t arrangement,
1:05:55all of them in parallel and independently produce a key and a query. So no communication has happened
1:06:01yet. But the communication comes now, all the queries will dot product with all the keys.
1:06:08So basically, what we want is we want way now, or the affinities between these to be query
1:06:14multiplying key. But we have to be careful with, we can't matrix multiply this, we actually need to
1:06:20transpose k. But we have to be also careful because these are when you have the batch dimension.
1:06:27So in particular, we want to transpose the last two dimensions, dimension negative one and
1:06:32dimension negative two. So negative two, negative one. And so this matrix multiply now will
1:06:40basically do the following b by t by 16 matrix multiplies b by 16 by t to give us b by t by t.
1:06:56So for every row of b, we're now going to have a t squared matrix giving us the affinities.
1:07:02And these are now the way. So they're not zeros, they are now coming from this dot product between
1:07:08the keys and the queries. So this can now run, I can run this. And the weighted aggregation now
1:07:15is a function in a data abandoned manner between the keys and queries of these nodes.
1:07:20So just inspecting what happened here, the way takes on this form. And you see that before way
1:07:28was just a constant. So it was applied in the same way to all the batch elements. But now
1:07:33every single batch elements will have different sort of way, because every single batch element
1:07:39contains different tokens at different positions. And so this is now data dependent. So when we look
1:07:45at just the zero row, for example, in the input, these are the weights that came out. And so you
1:07:51can see now that they're not just exactly uniform. And in particular, as an example here for the last
1:07:57row, this was the eighth token. And the eighth token knows what content it has, and it knows at
1:08:02what position it's in. And now the eighth token, based on that, creates a query, hey, I'm looking
1:08:09for this kind of stuff. I'm a vowel, I'm on the eighth position, I'm looking for any consonants at
1:08:14positions up to four. And then all the nodes get to emit keys. And maybe one of the channels could
1:08:21be I am a consonant, and I am in a position up to four. And that key would have a high number in
1:08:28that specific channel. And that's how the query and the key when they dot product, they can find
1:08:32each other and create a high affinity. And when they have a high affinity, like say,
1:08:37this token was pretty interesting to to this eighth token. When they have a high affinity,
1:08:44then through the softmax, I will end up aggregating a lot of its information into
1:08:48my position. And so I'll get to learn a lot about it. Now, just this, we're looking at way
1:08:56after this has already happened. Let me erase this operation as well. So let me erase the masking and
1:09:02the softmax, just to show you the under the hood internals and how that works. So without the
1:09:07masking and the softmax way comes out like this, right? This is the outputs of the dot products.
1:09:14And these are the raw outputs, and they take on values from negative, you know, to to positive
1:09:18two, etc. So that's the raw interactions and raw affinities between all the nodes. But now,
1:09:25if I'm a fifth node, I will not want to aggregate anything from the sixth node, seventh node,
1:09:30and the eighth node. So actually, we use the upper triangular masking. So those are not allowed to
1:09:36communicate. And now, we actually want to have a nice distribution. So we don't want to aggregate
1:09:44negative point one one of this node, that's crazy. So instead, we exponentiate and normalize.
1:09:49And now we get a nice distribution that sums to one. And this is telling us now in the data
1:09:53dependent manner, how much information to aggregate from any of these tokens in the past. So that's
1:10:00way, and it's not zeros anymore. But it's calculated in this way. Now, there's one more
1:10:07part to a single self attention head. And that is that when we do the aggregation,
1:10:12we don't actually aggregate the tokens exactly, we aggregate, we produce one more value here.
1:10:17And we call that the value. So in the same way that we produce key inquiry, we're also going
1:10:24to create a value. And then here, we don't aggregate x, we calculate a v, which is just
1:10:34achieved by propagating this linear on top of x again. And then we output way multiplied by v.
1:10:43So v is the elements that we aggregate, or the vector that we aggregate instead of the raw x.
1:10:48And now, of course, this will make it so that the output here of the single head will be 16
1:10:54dimensional, because that is the head size. So you can think of x as kind of like private
1:11:00information to this token, if you if you think about it that way. So x is kind of private to
1:11:05this token. So I'm a fifth token at some, and I have some identity, and my information is kept
1:11:12in vector x. And now, for the purposes of the single head, here's what I'm interested in.
1:11:17Here's what I have. And if you find me interesting, here's what I will communicate to you. And that's
1:11:24stored in v. And so v is the thing that gets aggregated for the purposes of this single head
1:11:29between the different nodes. And that's basically the self attention mechanism. This is,
1:11:36this is what it does. There are a few notes that I would make like to make about attention.
1:11:42Number one, attention is a communication mechanism. You can really think about it as a
1:11:46communication mechanism, where you have a number of nodes in a directed graph, where basically you
1:11:52have edges pointing between those like this. And what happens is every node has some vector of
1:11:57information, and it gets to aggregate information via a weighted sum from all the nodes that point to it.
1:12:05And this is done in a data dependent manner. So depending on whatever data is actually stored at
1:12:09each node at any point in time. Now, our graph doesn't look like this, our graph has a different
1:12:15structure, we have eight nodes, because the block size is eight, and there's always eight tokens.
1:12:21And the first node is only pointed to by itself. The second node is pointed to by the first node
1:12:27and itself, all the way up to the eighth node, which is pointed to by all the previous nodes
1:12:32and itself. And so that's the structure that our directed graph has, or happens to have in an
1:12:38autoregressive sort of scenario like language modeling. But in principle, attention can be
1:12:42applied to any arbitrary directed graph, and it's just a communication mechanism between the nodes.
1:12:47The second note is that notice that there's no notion of space. So attention simply acts over
1:12:53like a set of vectors in this graph. And so by default, these nodes have no idea where they
1:12:58are positioned in the space. And that's why we need to encode them positionally, and sort of
1:13:02give them some information that is anchored to a specific position, so that they sort of know where
1:13:08they are. And this is different than for example, from convolution, because if you run, for
1:13:12example, a convolution operation over some input, there's a very specific sort of layout of the
1:13:17information in space, and the convolutional filters sort of act in space. And so it's not like an
1:13:24attention. An attention is just a set of vectors out there in space, they communicate. And if you
1:13:29want them to have a notion of space, you need to specifically add it, which is what we've done
1:13:34when we calculated the repositioning code encodings and added that information to the vectors.
1:13:40The next thing that I hope is very clear is that the elements across the batch dimension,
1:13:44which are independent examples, never talk to each other, they're always processed independently.
1:13:48And this is a batch matrix multiply that applies basically a matrix multiplication,
1:13:53kind of imperiled across the batch dimension. So maybe it would be more accurate to say that
1:13:57in this analogy of a directed graph, we really have, because the batch size is four,
1:14:02we really have four separate pools of eight nodes, and those eight nodes only talk to each other.
1:14:07But in total, there's like 32 nodes that are being processed. But there's
1:14:12sort of four separate pools of eight, you can look at it that way.
1:14:15The next note is that here in the case of language modeling, we have this specific
1:14:21structure of directed graph where the future tokens will not communicate to the past tokens.
1:14:27But this doesn't necessarily have to be the constraint in the general case. And in fact,
1:14:31in many cases, you may want to have all of the nodes talk to each other fully. So as an
1:14:37example, if you're doing sentiment analysis, or something like that with a transformer,
1:14:40you might have a number of tokens, and you may want to have them all talk to each other fully.
1:14:45Because later you are predicting, for example, the sentiment of the sentence.
1:14:49And so it's okay for these nodes to talk to each other. And so in those cases,
1:14:54you will use an encoder block of self attention. And all it means that it's an encoder block,
1:15:00is that you will delete this line of code, allowing all the nodes to completely talk to
1:15:04each other. What we're implementing here is sometimes called a decoder block. And it's
1:15:08called a decoder, because it is sort of like decoding language. And it's got this autoregressive
1:15:16format, where you have to mask with the triangulant matrix, so that nodes from the future never talk
1:15:22to the past, because they would give away the answer. And so basically, an encoder blocks,
1:15:27you would delete this, allow all the nodes to talk in decoder blocks, this will always be
1:15:32present, so that you have this triangular structure. But both are allowed and attention
1:15:36doesn't care attention supports arbitrary connectivity between nodes. The next thing
1:15:40I wanted to comment on is you keep me you keep hearing me say attention, self attention, etc.
1:15:45There's actually also something called cross attention, what is the difference?
1:15:49So basically, the reason this attention is self attention is because the keys,
1:15:56queries and the values are all coming from the same source from x. So the same source x produces
1:16:03keys, queries and values. So these nodes are self attending. But in principle, attention is
1:16:08much more general than that. So for example, an encoder, decoder, transformers, you can have a case
1:16:14where the queries are produced from x, but the keys and the values come from a whole separate
1:16:19external source, and sometimes from encoder blocks that encode some context that we'd like to
1:16:24condition on. And so the keys and the values will actually come from a whole separate source,
1:16:29those are nodes on the side. And here we're just producing queries, and we're reading off
1:16:33information from the side. So cross attention is used when there's a separate source of nodes,
1:16:40we'd like to pull information from into our nodes. And it's self attention if we just have
1:16:45nodes that would like to look at each other and talk to each other. So this attention here
1:16:50happens to be self attention. But in principle, attention is a lot more general. Okay, and the
1:16:57last note at this stage is, if we come to the attention is all you need paper here,
1:17:01we've already implemented attention. So given query key and value, we've multiplied the query
1:17:07and the key, we've softmaxed it, and then we are aggregating the values. There's one more thing
1:17:12that we're missing here, which is the dividing by one over square root of the head size,
1:17:16the decay here is the head size. Why aren't they doing this one is important. So they call it a
1:17:22scaled attention. And it's kind of like an important normalization to basically have.
1:17:28The problem is if you have unit Gaussian inputs, so zero mean unit variance, k and q are unit
1:17:33Gaussian. And if you just do way naively, then you see that your way actually will be the
1:17:38variance will be on the order of head size, which in our case is 16. But if you multiply by
1:17:44one over head size square root, so this is square root, and this is one over,
1:17:48then the variance of way will be one. So it will be preserved. Now, why is this important?
1:17:54You'll notice that way here will feed into softmax. And so it's really important,
1:18:01especially at initialization, that way be fairly diffuse. So in our case here,
1:18:06we sort of locked out here and way had a fairly diffuse numbers here. So like this. Now the
1:18:15problem is that because of softmax, if way takes on very positive and very negative numbers inside
1:18:20it, softmax will actually converge towards one hot vectors. And so I can illustrate that here.
1:18:28Say we are applying softmax to a tensor of values that are very close to zero,
1:18:32then we're going to get a diffuse thing out of softmax. But the moment I take the exact same thing
1:18:37and I start sharpening it, making it bigger by multiplying these numbers by eight, for example,
1:18:42you'll see that the softmax will start to sharpen. And in fact, it will sharpen towards the max.
1:18:47So we'll sharpen towards whatever number here is the highest. And so, basically, we don't want
1:18:52these values to be too extreme, especially at initialization. Otherwise, softmax will be
1:18:56way too peaky. And you're basically aggregating information from like a single node, every node
1:19:02just aggregates information from a single other node. That's not what we want, especially at
1:19:06initialization. And so the scaling is used just to control the variance at initialization.
1:19:12Okay, so having said all that, let's now take our self attention knowledge and let's take it
1:19:16for a spin. So here in the code, I created this head module and implements a single head of
1:19:23self attention. So you give it a head size. And then here it creates the key query and the value
1:19:28linear layers. Typically people don't use biases in these. So those are the linear projections
1:19:34that we're going to apply to all of our nodes. Now here, I'm creating this trill variable.
1:19:39Trill is not a parameter of the module. So in sort of PyTorch naming conventions,
1:19:44this is called a buffer, it's not a parameter. And you have to call it you have to assign it
1:19:48to the module using the register buffer. So that creates the trill, the lower triangular matrix.
1:19:54And we're given the input x, this should look very familiar. Now, we calculate the keys,
1:19:58the queries, we calculate the attention scores inside way. We normalize it. So we're using scaled
1:20:05attention here, then we make sure that sure doesn't communicate with the past. So this makes
1:20:10it a decoder block, and then softmax and then aggregate the value and output.
1:20:15Then here in the language model, I'm creating a head in the constructor and I'm calling it self
1:20:20attention head. And the head size I'm going to keep as the same and embed just for now.
1:20:27And then here, once we've encoded the information with the token embeddings and the position
1:20:32embeddings, we're simply going to feed it into the self attention head. And then the output of
1:20:37that is going to go into the decoder language modeling head and create the logits. So this
1:20:45is sort of the simplest way to plug in a self attention component into our network right now.
1:20:51I had to make one more change, which is that here in the generate, we have to make sure that our
1:20:58IDX that we feed into the model, because now we're using positional embeddings, we can never
1:21:04have more than block size coming in. Because if IDX is more than block size, then our position
1:21:09embedding table is going to run out of scope because it only has embeddings for up to block
1:21:13size. And so therefore I added some code here to crop the context that we're going to feed into self
1:21:22so that we never pass more than block size elements. So those are the changes and let's
1:21:27now train the network. Okay, so I also came up to the script here and I decreased the learning
1:21:31rate because the self attention can't tolerate very, very high learning rates. And then I also
1:21:37increased number of iterations because the learning rate is lower. And then I trained it
1:21:40and previously we were only able to get to up to 2.5. And now we are down to 2.4. So we
1:21:45definitely see a little bit of an improvement from 2.5 to 2.4 roughly. But the text is still not
1:21:51amazing. So clearly the self attention hat is doing some useful communication. But we still
1:21:58have a long way to go. Okay, so now we've implemented the scale dot product attention.
1:22:02Now next up and the attention is all you need paper. There's something called multi head
1:22:06attention. And what is multi head attention? It's just applying multiple attentions in parallel
1:22:12and concatenating the results. So they have a little bit of diagram here. I don't know if
1:22:17this is super clear. It's really just multiple attentions in parallel. So let's implement that
1:22:23fairly straightforward. If we want a multi head attention, then we want multiple heads of self
1:22:28attention running in parallel. So in PyTorch, we can do this by simply creating multiple heads.
1:22:36So however many heads you want, and then what is the head size of each. And then we run all of them
1:22:43in parallel into a list and simply concatenate all of the outputs. And we're concatenating over
1:22:50the channel dimension. So the way this looks now is we don't have just a single attention
1:22:56that has a head size of 32. Because remember an embed is 32. Instead of having one
1:23:03communication channel, we now have four communication channels in parallel. And each one of these
1:23:09communication channels typically will be smaller correspondingly. So because we have four
1:23:15communication channels, we want eight dimensional self attention. And so from each communication
1:23:20channel, we're getting together eight dimensional vectors. And then we have four of them. And
1:23:25that concatenates to give us 32, which is the original and embed. And so this is kind of
1:23:30similar to if you're familiar with convolutions, this is kind of like a group convolution.
1:23:34Because basically, instead of having one large convolution, we do convolution in groups.
1:23:39And that's multi headed self attention. And so then here, we just use SA heads self attention
1:23:46heads instead. Now I actually ran it and scrolling down, I ran the same thing. And then we now get
1:23:54this down to 2.28, roughly. And the operation is still the generation is still not amazing.
1:24:00But clearly, the validation loss is improving because we were at 2.4 just now. And so it
1:24:05helps to have multiple communication channels, because obviously, these tokens have a lot to talk
1:24:10about. They want to find the constants, the vowels, they want to find the vowels just from
1:24:14certain positions, they want to find any kinds of different things. And so it helps to create
1:24:20multiple independent channels of communication, gather lots of different types of data, and then
1:24:25decode the output. Now going back to the paper for a second, of course, I didn't explain this figure
1:24:29in full detail, but we are starting to see some components of what we've already implemented. We
1:24:33have the positional encodings, token encodings that add, we have the masked multi headed
1:24:38attention implemented. Now, here's another multi headed attention, which is a cross attention
1:24:43to an encoder, which we haven't, we're not going to implement in this case, I'm going to
1:24:47come back to that later. But I want you to notice that there's a feed forward part here,
1:24:52and then this is grouped into a block that gets repeated again and again.
1:24:56Now the feed forward part here is just a simple multi layer perceptron.
1:25:02So the multi headed, so here position wise feed forward networks, is just a simple little MLP.
1:25:08So I want to start basically in a similar fashion, also adding computation into the network.
1:25:13And this computation is on a per node level. So I've already implemented it. And you can see the diff
1:25:20highlighted on the left here when I've added or changed things. Now before we had the self
1:25:24multi headed self attention that did the communication, but we went way too fast to calculate the logits.
1:25:31So the tokens looked at each other, but didn't really have a lot of time to think on what they
1:25:35found from the other tokens. And so what I've implemented here is a little feed forward single
1:25:41layer. And this little layer is just a linear followed by a relulon linearity. And that's that's
1:25:47it. So it's just a little layer. And then I call it feed forward and embed. And then this feed forward
1:25:56is just called sequentially right after the self attention. So we self attend, then we feed forward.
1:26:02And you'll notice that the feed forward here when it's applying linear, this is on a per token
1:26:06level, all the tokens do this independently. So the self attention is the communication. And then
1:26:12once they've gathered all the data, now they need to think on that data individually. And so that's
1:26:16what feed forward is doing. And that's why I've added it here. Now when I train this, the
1:26:22validation loss actually continues to go down now to 2.24, which is down from 2.28. The output
1:26:28still look kind of terrible, but at least we've improved the situation. And so as a preview,
1:26:34we're going to now start to intersperse the communication with the computation. And that's
1:26:40also what the transformer does when it has blocks that communicate and then compute, and it groups
1:26:46them and replicates them. Okay, so let me show you what we'd like to do. We'd like to do something
1:26:52like this, we have a block. And this block is basically this part here, except for the cross
1:26:57attention. Now the block basically intersperse communication and the computation, the
1:27:03computation, the communication is done using multi-headed self attention. And then the
1:27:07computation is done using a feed forward network on all the tokens independently.
1:27:12Now, what I've added here also is you'll notice this takes the number of embeddings in the
1:27:19embedding dimension and number of heads that we would like, which is kind of like group sizing
1:27:23group convolution. And I'm saying that number of heads we like is four. And so because this is
1:27:2832, we calculate that because this 32, the number of heads should be four,
1:27:34the head size should be eight, so that everything sort of works out channel wise.
1:27:39So this is how the transformer structures sort of the sizes, typically. So the head size will
1:27:45become eight, and then this is how we want to intersperse them. And then here, I'm trying to
1:27:49create blocks, which is just a sequential application of block, block, block, so that we're
1:27:54interspersing communication feed forward many, many times. And then finally, we decode. Now,
1:28:00actually tried to run this. And the problem is this doesn't actually give a very good answer,
1:28:05and very good result. And the reason for that is we're starting to actually get like a pretty deep
1:28:10neural net. And deep neural nets suffer from optimization issues. And I think that's what we're
1:28:14kind of like slightly starting to run into. So we need one more idea that we can borrow from the
1:28:20transformer paper to resolve those difficulties. Now there are two optimizations that dramatically
1:28:25help with the depth of these networks, and make sure that the networks remain optimizable. Let's
1:28:30talk about the first one. The first one in this diagram is you see this arrow here. And then this
1:28:36arrow and this arrow, those are skip connections, or sometimes called residual connections. They come
1:28:42from this paper, the procedural learning for image recognition from about 2015, that introduced the
1:28:49concept. Now, these are basically what it means is you transform data, but then you have a skip
1:28:55connection with addition from the previous features. Now the way I like to visualize it,
1:29:02that I prefer, is the following. Here the computation happens from the top to bottom.
1:29:07And basically you have this residual pathway, and you are free to fork off from the residual
1:29:13pathway, perform some computation, and then project back to the residual pathway via addition.
1:29:19So you go from the inputs to the targets, only the plus and plus and plus. And the reason this
1:29:26is useful is because during backpropagation, remember from our micro-grad video earlier,
1:29:31addition distributes gradients equally to both of its branches that fed as the input. And so
1:29:38the supervision or the gradients from the loss basically hop through every addition node all
1:29:45the way to the input, and then also fork off into the residual blocks. But basically you have
1:29:52this gradient superhighway that goes directly from the supervision all the way to the input,
1:29:57unimpeded. And then these original blocks are usually initialized in the beginning,
1:30:01so they contribute very, very little, if anything, to the residual pathway. They are
1:30:05initialized that way. So in the beginning, they are almost kind of like not there. But then during
1:30:10the optimization, they come online over time, and they start to contribute. But at least at the
1:30:17initialization, you can go from directly supervision to the input gradient this unimpeded and just flows,
1:30:23and then the blocks over time kick in. And so that dramatically helps with the optimization.
1:30:29So let's implement this. So coming back to our block here, basically what we want to do is
1:30:33we want to do x equals x plus self-attention, and x equals x plus self-depth feedforward.
1:30:41So this is x, and then we fork off and do some communication and come back. And we fork off and
1:30:46we do some computation and come back. So those are residual connections. And then swinging back up
1:30:52here, we also have to introduce this projection. So nn.linear. And this is going to be from
1:31:02after we concatenate this. This is the size unimped. So this is the output of the
1:31:06self-attention itself. But then we actually want to apply the projection, and that's the result.
1:31:14So the projection is just a linear transformation of the outcome of this layer.
1:31:19So that's the projection back into the residual pathway. And then here in the feedforward,
1:31:23it's going to be the same thing. I could have a self-depth projection here as well,
1:31:28but let me just simplify it and let me couple it inside the same sequential container.
1:31:35And so this is the projection layer going back into the residual pathway.
1:31:39And so that's it. So now we can train this. So I implemented one more small change. When you
1:31:46look into the paper again, you see that the dimensionality of input and output is 512 for
1:31:52them. And they're saying that the inner layer here in the feedforward has dimensionality of 2048.
1:31:57So there's a multiplier of 4. And so the inner layer of the feedforward network
1:32:03should be multiplied by 4 in terms of channel sizes. So I came here and I multiplied 4 times
1:32:07embed here for the feedforward, and then from 4 times unimbed coming back down to unimbed
1:32:13when we go back to the projection. So adding a bit of computation here and growing that
1:32:18layer that is in the residual block on the side of the residual pathway. And then I train this,
1:32:24and we actually get down all the way to 2.08 validation loss. And we also see that the
1:32:29network is starting to get big enough that our train loss is getting ahead of validation loss.
1:32:33So we started to see a little bit of overfitting. And our generations here are still not amazing,
1:32:42but at least you see that we can see is here, this now, grief, sank. This starts to almost look
1:32:48like English. So yeah, we're starting to really get there. Okay, and the second innovation that
1:32:53is very helpful for optimizing very deep neural networks is right here. So we have this addition
1:32:58now that's the residual part, but this norm is referring to something called layer norm.
1:33:02So layer norm is implemented in PyTorch. It's a paper that came out a while back here.
1:33:10And layer norm is very, very similar to bash norm. So remember back to our Make More Series
1:33:15Part 3, we implemented bash normalization. And bash normalization basically just makes sure
1:33:21that across the bash dimension, any individual neuron had unit Gaussian distribution. So it was
1:33:31zero mean and unit standard deviation, one standard deviation output. So what I did here is I'm
1:33:37copy pasting the bash norm 1D that we developed in our Make More Series. And see here, we can
1:33:42initialize, for example, this module, and we can have a batch of 32 100 dimensional vectors
1:33:48feeding through the bash norm layer. So what this does is it guarantees that when we look at
1:33:54just the zero column, it's a zero mean one standard deviation. So it's normalizing every single column
1:34:02of this input. Now the rows are not going to be normalized by default, because we're just
1:34:08normalizing columns. So let's not implement layer norm. It's very complicated. Look,
1:34:14we come here, we change this from zero to one. So we don't normalize the columns, we normalize
1:34:20the rows. And now we've implemented layer norm. So now the columns are not going to be normalized.
1:34:30But the rows are going to be normalized for every individual example, it's 100 dimensional vector is
1:34:35normalized in this way. And because our computation now does not span across examples, we can delete
1:34:43all of this buffer stuff, because we can always apply this operation, and don't need to maintain
1:34:49any running buffers. So we don't need the buffers. We don't, there's no distinction between training
1:34:57and test time. And we don't need these running buffers, we do keep gamma and beta, we don't need
1:35:04the momentum, we don't care if it's training or not. And this is now a layer norm. And it
1:35:11normalizes the rows instead of the columns. And this here is identical to basically this here.
1:35:19So let's now implement layer norm in our transformer. Before I incorporate the layer
1:35:24norm, I just wanted to note that, as I said, very few details about the transformer have changed in
1:35:28the last five years. But this is actually something that slightly departs from the original
1:35:32paper. You see that the add and norm is applied after the transformation. But now it is a bit
1:35:40more basically common to apply the layer norm before the transformation. So there's a reshuffling
1:35:45of the layer norms. So this is called the pre-norm formulation, and that's the one that
1:35:49we're going to implement as well. So slight deviation from the original paper.
1:35:53Basically, we need two new layer norms. Layer norm one is an n dot layer norm,
1:35:59and we tell it how many, what is the embedding dimension. And we need the second layer norm.
1:36:05And then here, the layer norms are applied immediately on x. So self dot layer norm one
1:36:11applied on x, and self dot layer norm two applied on x before it goes into self attention and feed
1:36:17forward. And the size of the layer norm here is an embed, so 32. So when the layer norm is
1:36:24normalizing our features, it is the normalization here happens, the mean and the variance are
1:36:32taken over 32 numbers. So the batch and the time act as batch dimensions, both of them.
1:36:38So this is kind of like a per token transformation that just normalizes the features and makes them
1:36:45unit mean, unit Gaussian at initialization. But of course, because these layer norms inside it
1:36:52have these gamma and beta trainable parameters, the layer normal eventually create outputs that
1:36:59might not be unit Gaussian, but the optimization will determine that. So for now, this is the,
1:37:05this is incorporating the layer norms and let's train them up. Okay, so I let it run. And we see
1:37:10that we get down to 2.06, which is better than the previous 2.08. So a slight improvement by adding
1:37:16the layer norms. And I'd expect that they help even more if we had bigger and deeper network.
1:37:21One more thing I forgot to add is that there should be a layer norm here also typically,
1:37:26as at the end of the transformer and right before the final linear layer that decodes into
1:37:31vocabulary. So I added that as well. So at this stage, we actually have a pretty complete
1:37:37transformer coming to the original paper. And it's a decoder only transformer. I'll talk about
1:37:42that in a second. But at this stage, the major pieces are in place so we can try to scale this
1:37:47up and see how well we can push this number. Now in order to scale up the model, I had to
1:37:51perform some cosmetic changes here to make it nicer. So I introduced this variable called inlayer,
1:37:57which just specifies how many layers of the blocks we're going to have. I create a bunch of blocks
1:38:02and we have a new variable number of heads as well. I pulled out the layer norm here. And so
1:38:08this is identical. Now one thing that I did briefly change is I added a dropout. So dropout
1:38:14is something that you can add right before the residual connection back, right before the
1:38:19connection back into the residual pathway. So we can drop out that as the last layer here.
1:38:25We can drop out here at the end of the multi-headed extension as well. And we can also drop out here
1:38:32when we calculate the basically affinities. And after the softmax, we can drop out some of those.
1:38:38So we can randomly prevent some of the nodes from communicating. And so dropout comes from
1:38:44this paper from 2014 or so. And basically it takes your neural net. And it randomly,
1:38:53every forward backward pass shuts off some subset of neurons. So randomly drops them to zero
1:39:00and trains without them. And what this does effectively is because the mask of what's being
1:39:06dropped out has changed every single forward backward pass, it ends up kind of training an
1:39:11ensemble of subnetworks. And then at test time, everything is fully enabled and kind of all of
1:39:16those subnetworks are merged into a single ensemble if you can, if you want to think about it that way.
1:39:21So I would read the paper to get the full detail. For now, we're just going to stay on the level of
1:39:25this is a regularization technique. And I added it because I'm about to scale up the model quite
1:39:30a bit. And I was concerned about overfitting. So now when we scroll up to the top,
1:39:36we'll see that I changed a number of hyper parameters here about our neural net.
1:39:40So I made the batch size be much larger now 64. I changed the block size to be 256. So previously
1:39:46was just eight, eight characters of context. Now it is 256 characters of context to predict the 257th.
1:39:54I brought down the learning rate a little bit because the neural net is now much bigger. So
1:39:58I brought down the learning rate. The embedding dimension is now 384. And there are six heads.
1:40:04So 384, divide six means that every head is 64 dimensional as a standard. And then there are
1:40:12going to be six layers of that. And the dropout will be a point two. So every four backward past
1:40:1720% of all of these intermediate calculations are disabled and dropped to zero. And then I already
1:40:25trained this and I ran it. So drum roll, how well does it perform? So let me just scroll up here.
1:40:31We get a validation loss of 1.48, which is actually quite a bit of an improvement on what
1:40:38we had before, which I think was 2.07. So we went from 2.07 all the way down to 1.48,
1:40:43just by scaling up this neural net with the code that we have. And this, of course, ran for a
1:40:47lot longer. This may be trained for, I want to say about 15 minutes on my A100 GPU. So that's a
1:40:53pretty good GPU. And if you don't have a GPU, you're not going to be able to reproduce this.
1:40:58On a CPU, this would be, I would not run this on a CPU or a MacBook or something like that.
1:41:03You'll have to bring down the number of layers and the embedding dimension and so on.
1:41:08But in about 15 minutes, we can get this kind of a result. And I'm printing some of the Shakespeare
1:41:15here. But what I did also is I printed 10,000 characters, so a lot more, and I wrote them to
1:41:19a file. And so here we see some of the outputs. So it's a lot more recognizable as the input text
1:41:27file. So the input text file just for reference look like this. So there's always like someone
1:41:33speaking in this manner. And our predictions now take on that form. Except of course, they're
1:41:41they're nonsensical when you actually read them. So it is every crimpy be house. Oh,
1:41:48those preparation. We give heed. Oh, sent me you mighty Lord.
1:42:01Anyway, so you can read through this. It's nonsensical, of course, but this is just a
1:42:05transformer trained on the character level for 1 million characters that come from Shakespeare.
1:42:10So it's sort of like blabbers on in Shakespeare like manner. But it doesn't, of course,
1:42:15make sense at this scale. But I think I think still a pretty good demonstration of what's possible.
1:42:22So now I think that kind of like concludes the programming section of this video.
1:42:28We basically kind of did a pretty good job and implementing this transformer. But the picture
1:42:35doesn't exactly match up to what we've done. So what's going on with all these additional parts
1:42:39here? So let me finish explaining this architecture and why it looks so funky. Basically,
1:42:44what's happening here is what we implemented here is a decoder only transformer. So there's no
1:42:50component here. This part is called the encoder. And there's no cross attention block here. Our
1:42:56block only has a self attention and the feed forward. So it is missing this third in between
1:43:02piece here. This piece does cross attention. So we don't have it and we don't have the encoder.
1:43:07We just have the decoder. And the reason we have a decoder only is because we are just
1:43:12generating text and it's unconditioned on anything. We're just blabbering on according
1:43:17to a given data set. What makes it a decoder is that we are using the triangular mask
1:43:22in our transformer. So it has this autoregressive property where we can just go and sample from
1:43:27it. So the fact that it's using the triangular mask to mask out the attention makes it a decoder
1:43:34and it can be used for language modeling. Now, the reason that the original paper had an encoder
1:43:39decoder architecture is because it is a machine translation paper. So it is concerned with a
1:43:44different setting in particular. It expects some tokens that encode say for example French
1:43:52and then it is expected to decode the translation in English. So typically these here are special
1:43:58tokens. So you are expected to read in this and condition on it. And then you start off the
1:44:04generation with a special token called start. So this is a special new token that you introduce
1:44:10and always place in the beginning. And then the network is expected to output. Neural networks
1:44:16are awesome and then a special end token to finish the generation. So this part here
1:44:23will be decoded exactly as we've done it. Neural networks are awesome, will be identical to what we
1:44:28did. But unlike what we did, they want to condition the generation on some additional information. And
1:44:36in that case, this additional information is the French sentence that they should be translating.
1:44:41So what they do now is they bring the encoder. Now the encoder reads this part here. So we're
1:44:48only going to take the part of French and we're going to create tokens from it exactly as we've
1:44:53seen in our video. And we're going to put a transformer on it. But there's going to be
1:44:58no triangular mask. And so all the tokens are allowed to talk to each other as much as they want.
1:45:03And they're just encoding whatever's the content of this French sentence. Once they've encoded it,
1:45:11they basically come out in the top here. And then what happens here is in our decoder,
1:45:16which does the language modeling, there's an additional connection here to the outputs of
1:45:22the encoder. And that is brought in through cross-attention. So the queries are still generated
1:45:28from X, but now the keys and the values are coming from the side. The keys and the values
1:45:33are coming from the top generated by the nodes that came outside of the encoder. And those tops,
1:45:41the keys and the values there, the top of it feed in on the side into every single block of
1:45:46the decoder. And so that's why there's an additional cross-attention. And really what it's
1:45:51doing is it's conditioning the decoding, not just on the past of this current decoding,
1:45:57but also on having seen the full, fully encoded French prompt sort of. And so it's an encoder
1:46:06decoder model, which is why we have those two transformers, an additional block, and so on.
1:46:10So we did not do this because we have nothing to encode. There's no conditioning. We just have
1:46:15a text file and we just want to imitate it. And that's why we are using a decoder only
1:46:19transformer exactly as done in GPT. Okay. So now I wanted to do a very brief walkthrough of
1:46:25Nanogpt, which you can find on my GitHub. And Nanogpt is basically two files of interest.
1:46:31There's train.py and model.py. train.py is all the boilerplate code for training the network.
1:46:37It is basically all the stuff that we had here as the training loop. It's just that it's a
1:46:43lot more complicated because we're saving and loading checkpoints and pre-trained weights.
1:46:47And we are decaying the learning rate and compiling the model and using distributed
1:46:51training across multiple nodes or GPUs. So the training.py gets a little bit more hairy,
1:46:56complicated. There's more options, et cetera. But the model.py should look very, very
1:47:02similar to what we've done here. In fact, the model is almost identical.
1:47:07So first, here we have the causal self-attention block. And all of this should look very,
1:47:12very recognizable to you. We're producing queries, keys, values. We're doing dot products. We're
1:47:18masking, applying softmax, optionally dropping out. And here we are pooling the values.
1:47:24What is different here is that in our code, I have separated out the multi-headed attention
1:47:31into just a single individual head. And then here I have multiple heads and I explicitly
1:47:36concatenate them. Whereas here, all of it is implemented in a batched manner inside a single
1:47:42causal self-attention. And so we don't just have a B and a T and a C dimension. We also end up with
1:47:47a fourth dimension, which is the heads. And so it just gets a lot more sort of hairy because we
1:47:53have four-dimensional array tensors now, but it is equivalent mathematically. So the exact same
1:47:59thing is happening as what we have. It's just a bit more efficient because all the heads are
1:48:03now treated as a batch dimension as well. Then we have the multilayer perceptron. It's using
1:48:09the Galun nonlinearity, which is defined here, except instead of Raulu. And this is done just
1:48:14because OpenAI used it and I want to be able to load their checkpoints. The blocks of the
1:48:19transformer are identical, the communicate and the compute phase, as we saw. And then the
1:48:24GPT will be identical. We have the position encodings, token encodings, the blocks, the
1:48:28layer norm at the end, the final linear layer. And this should look all very recognizable.
1:48:35And there's a bit more here because I'm loading checkpoints and stuff like that.
1:48:39I'm separating out the parameters into those that should be weight decayed and those that shouldn't.
1:48:44But the generate function should also be very, very similar. So a few details are different,
1:48:48but you should definitely be able to look at this file and be able to understand a lot of
1:48:53the pieces now. So let's now bring things back to chat GPT. What would it look like if we wanted
1:48:58to train chat GPT ourselves and how does it relate to what we learned today? Well, to train
1:49:03the chat GPT, there are roughly two stages. First is the pre-training stage and then the
1:49:07fine-tuning stage. In the pre-training stage, we are training on a large chunk of internet
1:49:14and just trying to get a first decoder-only transformer to babble text. So it's very,
1:49:19very similar to what we've done ourselves, except we've done like a tiny little baby
1:49:25pre-training step. And so in our case, this is how you print a number of parameters.
1:49:32I printed it and it's about 10 million. So this transformer that I created here to create
1:49:36little Shakespeare transformer was about 10 million parameters. Our dataset is roughly 1
1:49:43million characters, so roughly 1 million tokens. But you have to remember that OpenAI uses
1:49:48different vocabulary. They're not on the character level. They use these subword chunks of words.
1:49:54And so they have a vocabulary of 50,000 roughly elements. And so their sequences are a bit more
1:50:00condensed. So our dataset, the Shakespeare dataset would be probably around 300,000
1:50:05tokens in the OpenAI vocabulary, roughly. So we trained about 10 million parameter model and
1:50:12roughly 300,000 tokens. Now, when you go to the GPT-3 paper and you look at the
1:50:19transformers that they trained, they trained a number of transformers of different sizes,
1:50:24but the biggest transformer here has 175 billion parameters. So ours is again 10 million.
1:50:31They used this number of layers in the transformer. This is the n-embed. This is the number of heads
1:50:37and this is the head size. And then this is the batch size. So ours was 65.
1:50:44And the learning rate is similar. Now, when they train this transformer, they trained on 300
1:50:49billion tokens. So again, remember ours is about 300,000. So this is about a million fold increase.
1:50:57And this number would not be even that large by today's standards. You'd be going up
1:51:01one trillion above. So they are training a significantly larger model on a good chunk of
1:51:09the internet. And that is the pre-training stage. But otherwise, these hyperparameters should be
1:51:14fairly recognizable to you. And the architecture is actually nearly identical to what we implemented
1:51:18ourselves. But of course, it's a massive infrastructure challenge to train this.
1:51:23You're talking about typically thousands of GPUs having to talk to each other to train
1:51:28models of this size. So that's just the pre-training stage. Now, after you complete the
1:51:33pre-training stage, you don't get something that responds to your questions with answers,
1:51:38and it's not helpful, and etc. You get a document completer. So it babbles, but it doesn't
1:51:45babble Shakespeare, it babbles internet. It will create arbitrary news articles and documents,
1:51:49and it will try to complete documents because that's what it's trained for. It's trying to
1:51:53complete the sequence. So when you give it a question, it would potentially just give you
1:51:58more questions. It would follow with more questions. It will do whatever it looks like
1:52:02that some closed document would do in the training data on the internet. And so who knows,
1:52:07you're getting kind of like undefined behavior. It might basically answer to questions with other
1:52:12questions. It might ignore your question. It might just try to complete some news article.
1:52:17It's totally unaligned, as we say. So the second fine-tuning stage is to actually align it to be an
1:52:23assistant. And this is the second stage. And so this chat GPT blog post from OpenAI talks a little
1:52:30bit about how this stage is achieved. We basically, there's roughly three steps to this stage.
1:52:38So what they do here is they start to collect training data that looks specifically like what
1:52:43an assistant would do. So they have documents that have the format where the question is on top,
1:52:47and then an answer is below. And they have a large number of these, but probably not on the order of
1:52:52the internet. This is probably on the order of maybe thousands of examples. And so they then
1:52:59fine-tune the model to basically only focus on documents that look like that. And so you're
1:53:05starting to slowly align it. So it's going to expect a question at the top, and it's going to
1:53:09expect to complete the answer. And these very, very large models are very sample efficient during
1:53:15their fine-tuning. So this actually somehow works. But that's just step one. That's just fine-tuning.
1:53:20So then they actually have more steps where, okay, the second step is you let the model respond,
1:53:25and then different raters look at the different responses and rank them for their preferences to
1:53:30which one is better than the other. They use that to train a reward model. So they can predict,
1:53:35basically using a different network, how much of any candidate response would be desirable. And then
1:53:43once they have a reward model, they run PPO, which is a form of policy gradient reinforcement learning
1:53:49optimizer, to fine-tune this sampling policy so that the answers that GPT now generates
1:53:58are expected to score a high reward according to the reward model. And so basically, there's a whole
1:54:05lining stage here, or fine-tuning stage. It's got multiple steps in between there as well. And it takes
1:54:11the model from being a document completer to a question-answer. And that's like a whole separate
1:54:17stage. A lot of this data is not available publicly. It is internal to OpenAI, and it's
1:54:23much harder to replicate this stage. And so that's roughly what would give you a chat GPT.
1:54:29And NanoGPT focuses on the pre-training stage. Okay, and that's everything that I wanted to
1:54:34cover today. So we trained, to summarize, a decoder-only transformer following this famous
1:54:42paper, Attention is All You Need from 2017. And so that's basically a GPT. We trained it on
1:54:49tiny Shakespeare and got sensible results. All of the training code is roughly 200 lines of code.
1:54:57I will be releasing this code base. So also, it comes with all the Git log commits along the way
1:55:04as we built it up. In addition to this code, I'm going to release the notebook, of course,
1:55:11the Google Colab. And I hope that gave you a sense for how you can train these models,
1:55:17like say GPT-3, that will be architecturally basically identical to what we have,
1:55:22but they are somewhere between 10,000 and 1 million times bigger, depending on how you count.
1:55:27And so that's all I have for now. We did not talk about any of the fine-tuning stages that
1:55:33typically go on top of this. So if you're interested in something that's not just language
1:55:37modeling, but you actually want to, you know, say perform tasks, or you want them to be aligned in
1:55:42a specific way, or you want to detect sentiment or anything like that, basically, anytime you don't
1:55:47want something that's just a document completer, you have to complete further stages of fine-tuning,
1:55:52which we did not cover. And that could be simple, supervised fine-tuning, or it can be something
1:55:57more fancy, like we see in chat.gpt, where we actually train a reward model and then do
1:56:01runs of PPO to align it with respect to the reward model. So there's a lot more that can
1:56:06be done on top of it. I think for now we're starting to get to about two hours mark,
1:56:10so I'm going to kind of finish here. I hope you enjoyed the lecture, and yeah,
1:56:17go forth and transform. See you later.