1 00:00:00,799 --> 00:00:05,399 Ah, hello everyone. (sighs) This is a test recording, 2 00:00:05,539 --> 00:00:10,500 uh, for the talk I'm going to give in the Ooty Retreat, uh, 3 00:00:10,559 --> 00:00:13,779 in June 2025. My name is Samuel Shadrack. 4 00:00:14,420 --> 00:00:16,979 It's 3:00 PM right now, so the 5 00:00:17,840 --> 00:00:22,639 talk will go till 5:00 PM, you know, assuming this test recording works out. 6 00:00:24,500 --> 00:00:29,479 I'm going to introduce machine learning and introduce large language models at 7 00:00:29,500 --> 00:00:34,199 a technical level. Uh, (door thuds) I'm going to be covering a lot of material, 8 00:00:34,959 --> 00:00:36,800 so if you're completely new to this, 9 00:00:37,479 --> 00:00:40,920 it's possible you will not (person coughs) understand all of it, but you'll understand 10 00:00:40,979 --> 00:00:41,759 most of it. 11 00:00:42,439 --> 00:00:47,080 The point is that you get an overview, not that you understand every single detail, you 12 00:00:47,099 --> 00:00:51,700 know, just within two hours. That is going to require, you know, more homework exercises, 13 00:00:51,739 --> 00:00:54,519 more practice, more, you know, writing your own code, 14 00:00:55,439 --> 00:00:59,039 you know, deriving the formulas yourself. Like, that will take a bit more time. 15 00:00:59,139 --> 00:01:00,599 It can't be done in two hours. 16 00:01:01,439 --> 00:01:05,239 But once you have this overview, then you will be better placed to, you know, do more of 17 00:01:05,259 --> 00:01:08,439 that learning even on your own or following in an online course. 18 00:01:10,899 --> 00:01:13,819 I'm assuming a background of 19 00:01:15,259 --> 00:01:17,419 basically, like, class 12, 20 00:01:18,739 --> 00:01:22,799 you know, science background, so I'm assuming you know calculus, I'm assuming you know 21 00:01:22,799 --> 00:01:24,159 differential equations. 22 00:01:25,699 --> 00:01:29,739 Uh, not much of differential equation, but yeah, you should understand basic calculus, 23 00:01:29,739 --> 00:01:31,659 you should understand matrix multiplication. 24 00:01:32,639 --> 00:01:37,459 Uh, we are going to use concepts of calculus and matrices in today's 25 00:01:37,679 --> 00:01:38,159 talk. 26 00:01:41,639 --> 00:01:45,799 Okay, so let's begin. First, let's start with the problem. 27 00:01:46,639 --> 00:01:50,899 So this is a very famous problem in machine learning called MNIST. 28 00:01:51,179 --> 00:01:51,519 Uh, 29 00:01:52,219 --> 00:01:56,999 MNIST stands for Modified NIST. So NIST is the standards body in the 30 00:01:57,079 --> 00:02:00,899 US, and also many countries have their own st- standards bodies, 31 00:02:02,959 --> 00:02:06,659 uh, that define, you know, various units, 32 00:02:07,479 --> 00:02:12,419 like, you know, what is a meter, what's a kilometer? You know, how big should a map be? 33 00:02:12,439 --> 00:02:15,479 What is one degree on a map? You know, these types of things. 34 00:02:16,279 --> 00:02:17,899 So, they have defined a standard 35 00:02:18,539 --> 00:02:23,279 dataset, which is for optical character recognition. 36 00:02:24,079 --> 00:02:28,839 So basically, you have these photos of people who have written various letters... 37 00:02:28,839 --> 00:02:33,379 not letters, sorry, numbers. So people have written, you know, three or two or seven or 38 00:02:33,399 --> 00:02:36,039 one. Like, any number they've written on a piece of paper. 39 00:02:36,059 --> 00:02:37,839 We have taken a photo of that piece of paper. 40 00:02:38,779 --> 00:02:39,219 Uh, 41 00:02:40,219 --> 00:02:44,499 we have converted it more specifically into a 28 by 28 pixel image. 42 00:02:44,559 --> 00:02:48,359 It's grayscale, so it's not, you know, RGB. It's not red, green and blue. 43 00:02:48,359 --> 00:02:48,759 It's just 44 00:02:50,099 --> 00:02:54,359 either black or white or gray. So the gray value can be between zero and one. 45 00:02:54,819 --> 00:02:57,939 So, you know, 0.8 value means, you know, that pixel is 46 00:02:58,979 --> 00:03:02,419 very close to black, and maybe a 0.2 means it's very close to white. 47 00:03:03,359 --> 00:03:05,939 Now we have, you know, 60,000 such photographs. 48 00:03:07,299 --> 00:03:09,439 We already know which classes they're put into, 49 00:03:10,459 --> 00:03:15,259 and now we want to write a program that will take an image as input and tell you, you 50 00:03:15,319 --> 00:03:15,539 know, 51 00:03:16,159 --> 00:03:18,839 what, uh, number is this. "Is this zero? Is this one? 52 00:03:18,879 --> 00:03:21,679 Is this two?" And, like, we already know the answers, but 53 00:03:22,499 --> 00:03:24,899 we don't want the program to memorize the answers. 54 00:03:26,079 --> 00:03:30,039 Uh, you know, o- one way to solve this problem would be just, you know, we already know 55 00:03:30,039 --> 00:03:34,259 the answers, why not just have like a dictionary in the (laughs) program that says, you 56 00:03:34,279 --> 00:03:38,719 know, "If you notice the first image, say this. If you notice the second image, do this. 57 00:03:38,959 --> 00:03:41,299 If else, if else, if else," you know, 60,000 times. 58 00:03:41,999 --> 00:03:45,779 But we are not going to do that, and there's a reason why we are not going to do that. 59 00:03:46,599 --> 00:03:51,199 The reason we are not going to do that is because we also want this program to work 60 00:03:51,419 --> 00:03:53,759 on, uh, images which we have not seen. 61 00:03:54,159 --> 00:03:59,039 So apart from these 60,000 photos, there's another 10,000 photos which we 62 00:03:59,059 --> 00:04:01,179 have not seen. Like, we do not know what they are. 63 00:04:02,639 --> 00:04:05,779 I mean, they're actually there on the internet, but right now we are not allowed to look 64 00:04:05,799 --> 00:04:06,879 at them. That's the rule. 65 00:04:07,899 --> 00:04:11,699 And we want our program to also perform well on these 10,000 images, 66 00:04:12,399 --> 00:04:14,979 and those are also, again, you know, photos of zero, one, two, three, four. 67 00:04:15,959 --> 00:04:16,359 So, 68 00:04:17,879 --> 00:04:22,479 we can't hard code it. We need some actual, you know, way of solving this problem so that 69 00:04:22,519 --> 00:04:24,299 it will work on the photos we have never seen. 70 00:04:26,119 --> 00:04:27,619 I'm going to give everyone, like, 71 00:04:28,619 --> 00:04:30,919 maybe two minutes. Just try thinking 72 00:04:32,019 --> 00:04:35,139 how would you solve this problem and maybe you can 73 00:04:36,339 --> 00:04:37,919 share your approach. Like, 74 00:04:38,799 --> 00:04:42,639 I hope the problem is clear to everyone. Has everyone understood what the problem is? 75 00:04:48,739 --> 00:04:49,759 Okay. So, 76 00:04:50,619 --> 00:04:51,579 here is a 77 00:04:53,819 --> 00:04:58,459 neural network based solution. So, in machine learning, we have this concept called 78 00:04:58,519 --> 00:04:59,859 neural network, which is 79 00:05:00,719 --> 00:05:04,099 basically a special type of program. It's ultimately a program. 80 00:05:04,099 --> 00:05:06,019 It takes input and it gives you output, 81 00:05:08,259 --> 00:05:11,699 but it is a bit different from regular programs (door creaks) or the programs you write 82 00:05:11,719 --> 00:05:15,439 in, you know, C or Python or whatever. Like, we will also write this in Python, but we... 83 00:05:15,479 --> 00:05:17,479 the way we write this program is a little bit different. 84 00:05:18,439 --> 00:05:18,819 So, 85 00:05:19,679 --> 00:05:20,659 if you see these, 86 00:05:21,579 --> 00:05:23,919 this line here, this is basically our program. 87 00:05:24,859 --> 00:05:27,139 Can you see? Y equals ReLU of, 88 00:05:28,239 --> 00:05:28,619 uh, 89 00:05:29,479 --> 00:05:32,779 this into this. So this at symbol is matrix multiplication. 90 00:05:33,039 --> 00:05:33,399 Uh, 91 00:05:34,339 --> 00:05:35,059 it's just, you know, 92 00:05:35,759 --> 00:05:39,519 like A at B just means, you know, matrix A multiplied by matrix B. 93 00:05:40,299 --> 00:05:42,159 (door creaks) So we take this matrix, 94 00:05:42,819 --> 00:05:46,539 B mul- matrix multiplied with W2 and we apply this ReLU function on it. 95 00:05:48,359 --> 00:05:50,039 So x here is the input. 96 00:05:50,939 --> 00:05:55,719 So the input is this, you know, 60,000 photographs. We have converted this into a matrix. 97 00:05:57,299 --> 00:05:57,799 Uh, 98 00:05:58,439 --> 00:06:01,779 we will multiply this with a weight matrix called W1. 99 00:06:01,779 --> 00:06:06,359 I will talk about quickly what all of these are, but...Yeah, x is our 100 00:06:06,379 --> 00:06:06,959 input. 101 00:06:07,640 --> 00:06:11,099 We will multiply the weight matrix. We will apply this ReLu function. 102 00:06:11,319 --> 00:06:15,059 We will multiply that with w2. Then we apply this ReLu function again. 103 00:06:15,839 --> 00:06:20,639 And y is our final answer. Like for each of these 60,000 images, you know, which, 104 00:06:20,820 --> 00:06:22,299 uh, class are they in? 105 00:06:23,819 --> 00:06:24,939 Okay. And 106 00:06:26,280 --> 00:06:29,759 now I will talk about, you know, what is ReLu, what is w1, what's w2, what are all these 107 00:06:29,780 --> 00:06:30,380 things. So 108 00:06:31,319 --> 00:06:33,659 may, actually, no, maybe I'll first even just start with the input. 109 00:06:33,659 --> 00:06:35,540 So x is the input, right? 110 00:06:37,179 --> 00:06:40,059 (cat meowing) So x has dimensions n cross d. 111 00:06:40,539 --> 00:06:42,539 If you've studied matrices, you know, right? 112 00:06:42,539 --> 00:06:42,879 Like a 113 00:06:43,739 --> 00:06:47,179 matrix has two dimensions, like, you know, its length and its breadth. 114 00:06:48,019 --> 00:06:52,039 So n is the number of images. In this case, you know, n is 60,000. 115 00:06:52,859 --> 00:06:57,679 So x is a 60,000 times 784 matrix. So there are 60,000 116 00:06:57,699 --> 00:06:58,699 images and- 117 00:06:58,699 --> 00:06:58,780 (sniffs) 118 00:06:58,780 --> 00:07:00,619 ... each image has 784 pixels. 119 00:07:02,099 --> 00:07:03,699 Uh, yeah, one important thing here, 120 00:07:04,819 --> 00:07:07,999 we are not storing the images in a 28 cross 28. 121 00:07:07,999 --> 00:07:11,119 We're just storing it as, you know, single line of 784 numbers. 122 00:07:12,799 --> 00:07:15,939 Uh, this is a very common thing that's done in machine learning. 123 00:07:16,899 --> 00:07:18,459 I'm not going to talk about why. 124 00:07:18,459 --> 00:07:18,499 (clicks tongue) 125 00:07:18,499 --> 00:07:20,599 But yeah, that's just a thing that's done. 126 00:07:20,679 --> 00:07:20,839 So 127 00:07:22,079 --> 00:07:25,379 we... Instead of having, you know, dun-dun-dun-dun-dun, dun-dun-dun-dun-dun, 128 00:07:25,419 --> 00:07:28,419 dun-dun-dun-dun-dun, like we have, you know, 28 cross 28, we're just going to 129 00:07:29,319 --> 00:07:33,739 put everything in this one line. So 28 numbers, then the next 28 numbers, the next 28 130 00:07:33,759 --> 00:07:38,059 numbers, next 28 numbers and so on. You know, you have 70, 784 of them in line. 131 00:07:38,939 --> 00:07:43,699 And in the same way you have, you know, 60,000 such rows. So each row is one image. 132 00:07:43,699 --> 00:07:46,859 There are, you know, 60,000 of them. So this is our input. 133 00:07:48,639 --> 00:07:51,579 Then we multiply with this weight matrix called w1. 134 00:07:51,619 --> 00:07:56,199 So w1 is this 784 cross and 800 matrix. 135 00:07:57,459 --> 00:08:02,099 Uh, and what is w1? We will find out what w1 is. Right now, just assume it's a matrix. 136 00:08:02,139 --> 00:08:04,239 We have somehow found out the value of w1. 137 00:08:05,159 --> 00:08:08,439 Then we multiply it. Then we do this thing called ReLu. 138 00:08:09,339 --> 00:08:13,319 Now, what is a ReLu? Uh, so here is the definition of ReLu. 139 00:08:13,899 --> 00:08:18,639 ReLu basically means if you have a matrix for each cell, if it's a positive cell, just 140 00:08:18,699 --> 00:08:20,899 keep it as it is. If it's a negative cell, 141 00:08:21,499 --> 00:08:23,699 you remove that value and you put zero instead. 142 00:08:24,499 --> 00:08:28,659 So if you had, you know, let's say, three, zero, minus two, four. 143 00:08:29,799 --> 00:08:32,919 (metal clanks) So now this will become three, zero, zero, four. 144 00:08:33,039 --> 00:08:35,939 So that negative value just got removed and we put zero there. 145 00:08:36,419 --> 00:08:40,799 And all the positive values, we just kept them as it is. So ReLu is nothing special. 146 00:08:40,799 --> 00:08:44,919 It just means remove all the negative values in that matrix and replace them with zero. 147 00:08:47,799 --> 00:08:49,059 So, yeah, we 148 00:08:49,979 --> 00:08:54,859 multiplied x with this w1 matrix. We made all the negative values 149 00:08:54,939 --> 00:08:55,459 zero. 150 00:08:56,419 --> 00:08:59,139 Then we multiply this with the w2 matrix. 151 00:08:59,199 --> 00:09:02,059 Now w2's dimensions are 800 cross 4. 152 00:09:03,619 --> 00:09:08,339 Wait, fuck. That's a mistake. It should be 800 cross 10. N- just give me a minute. 153 00:09:08,339 --> 00:09:09,859 I'll just quickly edit that 154 00:09:11,539 --> 00:09:15,619 because there are 10 classes, you know, 0 to 10. So this has to be 10. 155 00:09:15,619 --> 00:09:16,119 Uh, 156 00:09:24,679 --> 00:09:25,399 hmm. 157 00:09:27,899 --> 00:09:30,659 Yeah, cool. Okay, so we have 800 cross 10. 158 00:09:31,319 --> 00:09:34,579 And finally we get y, which is right now 60,000 cross 10. 159 00:09:34,599 --> 00:09:38,299 Like if you do, you know, the dimensional analysis of this, you will get the answer, you 160 00:09:38,299 --> 00:09:40,939 know, 60,000 cross 10. Like 60,000 161 00:09:41,879 --> 00:09:46,639 times 784 cross 784 times 800 means 60,000 times 8, 60,000 162 00:09:46,659 --> 00:09:47,419 cross 800. 163 00:09:48,239 --> 00:09:51,859 Then you have 60,000 cross 800 times, you know, 800 cross 10. 164 00:09:52,619 --> 00:09:55,519 That leave you 60,000 cross 10. 165 00:09:56,279 --> 00:09:59,999 So now we have 60,000 images. And now why are there 10 dimensions here? 166 00:10:01,059 --> 00:10:04,299 Like if... Think about it, there are 60,000 rows. In each row there are 10 numbers. 167 00:10:04,299 --> 00:10:05,479 Why do we need 10 numbers? 168 00:10:06,119 --> 00:10:10,079 So what you'll get in these 10 numbers is actually 10 probabilities 169 00:10:12,019 --> 00:10:16,279 telling you which number is this. So let's say we had, you know, number called three. 170 00:10:16,379 --> 00:10:18,739 We convert into, you know, 784 kind of thing. 171 00:10:19,959 --> 00:10:22,619 And like this there are many images. At the end we will get... 172 00:10:23,119 --> 00:10:25,199 For each image we're getting ten probabilities. 173 00:10:26,619 --> 00:10:30,599 Now these are actually log probabilities. They're not even probabilities. 174 00:10:31,299 --> 00:10:33,779 So here is an example of what we might get at the end. 175 00:10:34,459 --> 00:10:36,119 And this is just for one image right now. 176 00:10:36,119 --> 00:10:39,499 Actually, there'll be 60,000 of these, but for one image... 177 00:10:41,659 --> 00:10:41,859 and 178 00:10:43,039 --> 00:10:47,099 actually there will be ten of this. Uh, fuck, this is annoying. 179 00:10:48,179 --> 00:10:50,059 I should edit that as well. Yeah. 180 00:10:53,779 --> 00:10:58,099 (clicks tongue) I'm just going to make things easy for now and 181 00:10:59,679 --> 00:11:01,939 go with LN0, you know, six times. 182 00:11:05,899 --> 00:11:08,099 (keyboard clicking) Two, three, four, five, six. 183 00:11:09,599 --> 00:11:10,839 Minus inf. 184 00:11:19,599 --> 00:11:22,119 (keyboard clicking) One, two, three, four, five, six. 185 00:11:23,679 --> 00:11:26,919 Zero, zero, zero, zero, zero, zero, zero, zero. 186 00:11:28,659 --> 00:11:31,399 Yeah, I think that's about right. 187 00:11:33,779 --> 00:11:37,139 Up- (door clanks) load this and here we go. 188 00:11:40,879 --> 00:11:41,479 Ooh. 189 00:11:44,039 --> 00:11:44,779 Yeah, cool. 190 00:11:45,419 --> 00:11:50,119 So our y is actually... Here is what y looks like. So we have some values in here. 191 00:11:51,079 --> 00:11:55,579 These are all run... numbers we get, you know, after doing the two matrix multiplications 192 00:11:55,579 --> 00:11:57,979 and the two ReLus. This is what we, you know, get at the end. 193 00:12:00,839 --> 00:12:05,419 I can close all this. Yeah. So we got minus 1.89, minus 0.28, minus 194 00:12:05,479 --> 00:12:06,819 1.89, minus three, 195 00:12:07,459 --> 00:12:11,223 minus infinity.And infinity just means, you know, very large numbers. 196 00:12:11,243 --> 00:12:14,703 This could be in a minus, you know, 10 to the power of 32 or 197 00:12:15,343 --> 00:12:19,083 whatever, you know, register we are using here for float numbers, so this is, like, the 198 00:12:19,103 --> 00:12:22,384 largest number we can put in that register. So it's as good as infinity. 199 00:12:24,463 --> 00:12:27,264 And y dash is what our actual answer is. 200 00:12:27,264 --> 00:12:30,323 So remember, we know what the actual answer is, so here is y dash. 201 00:12:32,603 --> 00:12:36,383 So now we're going to take y, which is our predicted answer, 202 00:12:37,603 --> 00:12:40,763 and by the way, this y is actually log probability. 203 00:12:40,763 --> 00:12:42,763 So what this thing actually means is 204 00:12:43,743 --> 00:12:46,543 we are predicting, you know, 15% probability it's a zero, 205 00:12:47,203 --> 00:12:51,983 75% probability it's a one, a 15% probability it's a two, 206 00:12:52,243 --> 00:12:57,103 a 5% probability it's a three, and you know, 0% probability it's anything else. 207 00:12:57,263 --> 00:12:59,843 Like, this is what this actually means, like, you know, 208 00:13:00,543 --> 00:13:04,783 ln of 0.15 is, you know, -1.89. 209 00:13:06,763 --> 00:13:07,763 So yeah, 210 00:13:08,543 --> 00:13:12,103 this is the meaning of the thing but actually when we run this program, what we're going 211 00:13:12,103 --> 00:13:13,563 to get is we're going to get this thing. 212 00:13:15,103 --> 00:13:18,703 Uh, assuming we're using good weight matrices, W1, W2. 213 00:13:18,763 --> 00:13:23,223 We have not yet told you where did we get W1 and W2 from, but if you're using a good 214 00:13:23,243 --> 00:13:27,203 value for W1 and a W2, we're, we might get something like this. 215 00:13:28,843 --> 00:13:30,163 Or we might get something different. 216 00:13:30,163 --> 00:13:31,483 It's just this thing is saying, you know, 217 00:13:32,243 --> 00:13:37,043 this character is probably a one because, you know, it's saying there's, there's a 75% 218 00:13:37,083 --> 00:13:38,003 probability it's a one. 219 00:13:39,523 --> 00:13:40,123 And now, 220 00:13:41,463 --> 00:13:43,803 yes, here's the predicted, here's the actual thing. 221 00:13:44,723 --> 00:13:49,283 Now there is a question of, yeah, how do we find good W1 and W2? 222 00:13:50,503 --> 00:13:52,223 So now here is an important question. 223 00:13:52,923 --> 00:13:53,623 And again, I'll, 224 00:13:54,343 --> 00:13:54,763 I know I've 225 00:13:55,383 --> 00:13:58,963 given you a lot of information. I want you to just think for five minutes, you know. 226 00:14:00,763 --> 00:14:04,343 Here's the program. We want to find the W1 and the W2. 227 00:14:04,683 --> 00:14:06,863 These are the dimensions of W1 and W2, 228 00:14:07,803 --> 00:14:08,183 and 229 00:14:09,003 --> 00:14:11,423 we want to ensure at the end, 230 00:14:12,923 --> 00:14:17,363 these predicted values should match close with the actual 231 00:14:17,383 --> 00:14:21,983 values. So like if we manage to get exactly, you know, the exact actual values, you know, 232 00:14:22,923 --> 00:14:24,443 here, ln0, ln1. 233 00:14:25,043 --> 00:14:27,743 What, what is ln1? Ln1 is zero, right? Yeah. 234 00:14:29,503 --> 00:14:34,063 So if we could get, you know, minus infinity, zero, minus infinity, minus infinity, minus 235 00:14:34,103 --> 00:14:36,543 infinity, minus infinity, minus infinity and so on, like, 236 00:14:37,183 --> 00:14:38,863 that is our ideal answer. 237 00:14:42,163 --> 00:14:45,163 And you know, otherwise anything else is a bad answer. 238 00:14:53,403 --> 00:14:56,183 Also, we have this measure of how bad is our answer, like 239 00:14:56,943 --> 00:15:00,003 our prediction is not a hundred percent accurate. There is some inaccuracy here. 240 00:15:00,003 --> 00:15:03,363 So how inaccurate is it? So that thing is called loss. 241 00:15:04,383 --> 00:15:04,723 So thi- 242 00:15:05,763 --> 00:15:09,263 this which tells you how far away is your prediction from the truth. 243 00:15:09,303 --> 00:15:13,483 So this is defined typically, there are multiple ways to define it, but common way of 244 00:15:13,483 --> 00:15:14,103 knowing it is 245 00:15:15,083 --> 00:15:17,663 negative sum of y dot y dash. 246 00:15:18,523 --> 00:15:21,143 So these are our predicted log probabilities. 247 00:15:21,583 --> 00:15:23,783 This is, you know, our actual ground truth. 248 00:15:25,163 --> 00:15:28,783 We are going to dot product this. So dot product means what? Just this one we're taking. 249 00:15:28,783 --> 00:15:29,303 So this, 250 00:15:30,423 --> 00:15:31,443 we just took this. 251 00:15:34,363 --> 00:15:35,043 And 252 00:15:36,063 --> 00:15:37,483 yeah, that's our loss. So 253 00:15:38,163 --> 00:15:40,743 what probability did we give, you know, the correct answer. 254 00:15:40,863 --> 00:15:44,283 In this case, we are just keeping it simple. We're just keeping that as a loss. 255 00:15:44,703 --> 00:15:47,143 In practice here, there are other loss functions also. 256 00:15:47,343 --> 00:15:47,643 But 257 00:15:48,783 --> 00:15:51,883 yeah, for now it's just the log probability of the correct answer. 258 00:15:53,063 --> 00:15:54,623 That is our loss. 259 00:15:55,683 --> 00:15:59,203 So we want this to be less. And remember this, we have now for one image. 260 00:15:59,563 --> 00:16:02,283 We will add this value up for now 60,000 images. 261 00:16:02,283 --> 00:16:04,003 We will get, you know, 60,000 such values. 262 00:16:04,003 --> 00:16:04,923 We will add them all up 263 00:16:05,583 --> 00:16:08,103 and that tells us, you know, how inaccurate our thing is. 264 00:16:08,903 --> 00:16:12,723 So yeah, I, again, I hope this whole problem setup is 265 00:16:13,603 --> 00:16:17,823 clear. Now we need to figure out a good W1 and a good W2. 266 00:16:17,863 --> 00:16:20,703 Like how will we search for a good W1 and a good W2? 267 00:16:21,823 --> 00:16:23,043 That's the question here. 268 00:16:28,063 --> 00:16:29,363 Okay. So 269 00:16:33,603 --> 00:16:36,723 how, uh, how many options do we have for W2? 270 00:16:36,743 --> 00:16:40,243 Like, we need to search for a good W2, but like what's the search space? 271 00:16:40,263 --> 00:16:41,643 The search space is this, you know, 272 00:16:43,183 --> 00:16:47,803 784 cross 800 search space. 784 cross 273 00:16:47,823 --> 00:16:49,003 800. That's 274 00:16:50,063 --> 00:16:52,463 627,000 values, right? 275 00:16:54,243 --> 00:16:55,923 That is also wrong. 276 00:16:58,123 --> 00:17:02,643 627,000 values and W2 has 277 00:17:03,963 --> 00:17:08,563 8,000 values. 278 00:17:13,243 --> 00:17:14,063 2 to the 279 00:17:14,063 --> 00:17:21,763 power... 280 00:17:24,483 --> 00:17:24,883 Yeah. 281 00:17:27,003 --> 00:17:29,003 Huh. Lot of mistakes. Yeah. 282 00:17:32,723 --> 00:17:35,523 Yeah. So how many values do we have? We have... 283 00:17:36,783 --> 00:17:38,283 Here we have around... 284 00:17:41,223 --> 00:17:43,703 This is not changing. Cool. 285 00:17:44,403 --> 00:17:48,883 784 times 800. We have around six to seven... 286 00:17:49,643 --> 00:17:52,323 Around 600,000 values basically. And 287 00:17:56,283 --> 00:17:59,063 W2 has, you know, around 8,000 288 00:17:59,063 --> 00:18:11,939 values.That's 289 00:18:12,079 --> 00:18:12,799 60 s- 290 00:18:13,840 --> 00:18:14,239 sorry, 291 00:18:15,659 --> 00:18:19,819 784 into 800 into 800 into 10. 292 00:18:21,540 --> 00:18:24,620 (keyboard clicking) This is... 293 00:18:25,419 --> 00:18:28,179 That was in... Yeah, this is around 5 billion, I guess. 294 00:18:29,620 --> 00:18:34,260 2 to the power 5... That means 75 into 10 to the power of 9. (keyboard clicking) Cool. 295 00:18:34,279 --> 00:18:35,380 So we have 2 to the power of 627,000 into 8,000, that's around 5 billion values. 296 00:18:35,380 --> 00:18:35,380 Uh, so, we have... Even if in each cell we can only set, you know, 0 or 1, and actually 297 00:18:35,380 --> 00:18:35,380 we can set values in between also. We can set, you know, 0.3, we can set 0.5. 298 00:18:35,380 --> 00:18:35,380 We can set, you know, 50,000. We can set, you know, -87.5. You know, we can set anything. 299 00:18:35,380 --> 00:18:35,380 But even if we could set only 0 or 1, we have, you know, 2 to the power of 5 billion 300 00:18:35,380 --> 00:18:35,380 possibilities. And in practice because we can set many values, we have actually a lot 301 00:18:35,380 --> 00:18:35,380 more possibilities than that. So, we definitely cannot do any sort of brute force. 302 00:18:35,380 --> 00:18:35,380 We cannot even do, like, a very smart sort of brute force. 303 00:18:35,380 --> 00:18:35,380 Like, even if you reduce this, you know, 2 to the power of 5 billion into, you know, 2 to 304 00:18:35,380 --> 00:18:35,380 the power of 1 million possibilities, like, maybe you find some symmetry and you say, 305 00:18:35,380 --> 00:18:35,380 "Okay, these cells are similar to these cells, so we don't have to actually search them 306 00:18:35,380 --> 00:18:35,380 all together. Let's just search these ones and we'll copy the values in there." Like, 307 00:18:35,380 --> 00:18:38,620 often when we have brute force problems, this is how people, you know, make it efficient, 308 00:18:38,620 --> 00:18:42,339 right? They try to reduce the search space. 309 00:18:42,439 --> 00:18:43,779 But even that's not really 310 00:18:43,800 --> 00:18:48,980 going 311 00:18:49,019 --> 00:18:52,719 to work here because our search space is just so crazily large. 312 00:18:52,719 --> 00:18:55,320 So we can't really do it. We have to find something better. 313 00:18:55,339 --> 00:18:59,699 And also, 314 00:18:59,699 --> 00:19:03,819 actually, if you look at this 315 00:19:04,199 --> 00:19:07,379 problem more, you will understand there's also not, like, an exact answer. 316 00:19:07,379 --> 00:19:11,899 Like, you can't 317 00:19:11,899 --> 00:19:14,339 just look at this formula and calculate what they look. 318 00:19:14,719 --> 00:19:15,759 Here is the formula for calculating, 319 00:19:16,219 --> 00:19:25,999 you 320 00:19:25,999 --> 00:19:27,459 know, the optimal 321 00:19:28,139 --> 00:19:30,380 W1 and W2. Like, to optimize loss, you know, here is the optimal W1, here's the optimal 322 00:19:30,380 --> 00:19:30,380 W2. There is no direct formula to do this either. 323 00:19:30,380 --> 00:19:30,659 Like, just by looking at it you won't understand this. 324 00:19:30,680 --> 00:19:34,419 If you spend maybe a month (laughs) just 325 00:19:34,419 --> 00:19:37,319 thinking about this problem, you'll realize, like, yeah, there's no way you're getting an 326 00:19:37,319 --> 00:19:37,360 exact formula for this. So what does that leave us with? 327 00:19:37,360 --> 00:19:37,360 Remember, we need to optimize loss given a W1 and W2. 328 00:19:37,360 --> 00:19:37,360 So, we are going to bring in a concept from calculus. 329 00:19:37,360 --> 00:19:37,360 Uh, in calculus, when you have to optimize one variable with respect to other variables, 330 00:19:37,360 --> 00:19:37,360 what do we do, typically? Yeah. So we take, you know, d loss by dw1 equals to 0. 331 00:19:37,360 --> 00:19:37,419 Like you remember, right? Maxima, minima. Yeah. 332 00:19:37,419 --> 00:19:37,960 So, that's kind of what we're going to do. Uh, there are a couple of complications. 333 00:19:37,960 --> 00:19:37,960 The first complication is w1 is not a number, it's a matrix. 334 00:19:37,960 --> 00:19:37,960 Because it's a matrix, you know, that complicates things a little bit. The... 335 00:19:37,960 --> 00:19:37,960 Because, you know, if you just do dL by dw1 and d... 336 00:19:37,960 --> 00:19:37,960 If w1 was a scale of value, then, you know, we can just do dL by dw1. 337 00:19:37,960 --> 00:19:37,960 W1 is a matrix, so if you do dL by dw1, now that is going to be another matrix. 338 00:19:37,960 --> 00:19:37,960 And you will say, "Okay, well, wh- why can't we just put that matrix equal to 0?" Like, 339 00:19:37,960 --> 00:19:37,960 that mat-... dL by dw1 is a matrix, we'll calculate it. 340 00:19:37,960 --> 00:19:37,960 We'll say, you know, that matrix equal to 0 matrix. 341 00:19:37,960 --> 00:19:37,960 And similarly, 342 00:19:39,499 --> 00:19:44,319 you know, dL by dw2 will be 343 00:19:45,759 --> 00:19:50,279 one more matrix. We will get that matrix, we'll equate that to 0 344 00:19:50,279 --> 00:19:51,939 matrix. We will solve this entire system and we'll get the answer. 345 00:19:51,939 --> 00:19:53,119 And, like, first of all you can try to do that. 346 00:19:53,119 --> 00:19:56,599 It's still going to be a very large system to solve. 347 00:19:56,599 --> 00:19:58,119 Uh, if someone wants, they can try this out. I have not tried it. You can try it out. 348 00:19:58,119 --> 00:19:58,159 But, y- like, you can imagine, right, you're going to be getting... 349 00:19:58,159 --> 00:20:03,079 Like, again, you're going to get, you know, a system of five billion equations. 350 00:20:04,359 --> 00:20:04,639 Like, 351 00:20:05,279 --> 00:20:07,399 five 352 00:20:07,399 --> 00:22:59,379 billion 353 00:23:01,199 --> 00:23:05,279 linear equations, and you're going to try to solve them, I think, and... 354 00:23:05,279 --> 00:23:08,879 Well, okay, you can put it, you know, in a linear algebra form, maybe you'll get it 355 00:23:08,879 --> 00:23:11,159 solved, but, yeah, that's again hard. 356 00:23:12,419 --> 00:23:14,959 And that's even more not going to work when you make... 357 00:23:15,599 --> 00:23:18,159 like, when you add more matrices and you make this thing bigger. 358 00:23:18,239 --> 00:23:19,819 Like, instead of, you know, 359 00:23:20,839 --> 00:23:25,499 w1 relu, w2 relu, what if I then did, you know, w3 relu, w4 360 00:23:25,519 --> 00:23:30,099 relu, w5 relu? Like, you know, multiply this with another matrix, take a relu, 361 00:23:30,439 --> 00:23:34,159 multiply that with one more matrix, take a relu, multiply that with one more matrix, take 362 00:23:34,179 --> 00:23:34,659 a relu. 363 00:23:35,519 --> 00:23:39,479 And in practice, actually, for more difficult problems, we need to do that. 364 00:23:40,679 --> 00:23:44,459 Like, right now we have picked an easy problem, so this gets solved pretty well just 365 00:23:44,479 --> 00:23:48,919 with, you know, two layers, like two matrix multiplications and relus. 366 00:23:48,919 --> 00:23:49,219 But 367 00:23:49,919 --> 00:23:50,239 for 368 00:23:50,899 --> 00:23:52,939 harder problems we might take more. So, 369 00:23:53,799 --> 00:23:58,619 you know, for transformers, like, you know, your latest AI use, like, GPT 3.5, 370 00:23:58,659 --> 00:24:02,099 GPT 4 or whatever, you might take, you know, like, 80 layers or even more. 371 00:24:02,819 --> 00:24:06,419 So there will be, you know, 80 matrix multiplications running one after the other. 372 00:24:07,159 --> 00:24:09,999 And then you'll be like, "Well, how do we find the optimal 373 00:24:10,919 --> 00:24:15,839 matrix?" Like, you're not going to be able to solve, you know, d loss by dw1 374 00:24:15,879 --> 00:24:20,747 equal 0, like, analytically.What you can do... 375 00:24:23,327 --> 00:24:26,227 Is this thing called gradient descent, which is 376 00:24:27,848 --> 00:24:28,967 basically... 377 00:24:31,468 --> 00:24:32,447 Yeah. Basically, 378 00:24:33,107 --> 00:24:34,768 move along the direction 379 00:24:35,528 --> 00:24:39,608 decided by the slope. So, wherever you are right now, so I mean, I can make... 380 00:24:39,608 --> 00:24:41,707 This might be easier to draw with, you know, a figure. 381 00:24:41,827 --> 00:24:42,108 But 382 00:24:43,268 --> 00:24:44,887 whatever your loss value is right now, 383 00:24:45,928 --> 00:24:48,247 find out dL by dW in there 384 00:24:49,088 --> 00:24:51,947 and whichever direction it's pointing to, just move in that direction 385 00:24:52,747 --> 00:24:56,268 because that is... If you move along that direction, that's where the loss will reduce 386 00:24:56,347 --> 00:24:57,127 the fastest. 387 00:25:03,367 --> 00:25:06,548 So, this is the important trick in gradient descent. 388 00:25:09,627 --> 00:25:14,167 So, we are going to find dL by dW1 and dL by dW2. 389 00:25:14,787 --> 00:25:19,467 We'll move in that direction and we'll find ano- new W1 and W2 that are slightly 390 00:25:19,527 --> 00:25:20,847 better than the previous ones. 391 00:25:22,167 --> 00:25:25,587 And then we can repeat this process. You know, once we have W1 and W2, 392 00:25:26,267 --> 00:25:30,227 we'll find the gradient, we'll move a little bit to find a slightly better one. 393 00:25:30,687 --> 00:25:34,707 We'll find the gradient again. We'll find another W1 and W2 that are even better than 394 00:25:34,727 --> 00:25:37,587 that. We'll repeat this, repeat this, repeat this. 395 00:25:37,587 --> 00:25:40,167 You know, after enough times repeating this process, 396 00:25:41,447 --> 00:25:44,167 we'll get a really good W1 and W2. 397 00:25:45,107 --> 00:25:49,247 So for this, we need to find dL by dW1 and dL by dW2. 398 00:25:50,587 --> 00:25:55,447 This is... Now actually finding these gradients is kind of straightforward. 399 00:25:55,447 --> 00:25:57,107 It's also kind of boring. Uh, 400 00:25:58,287 --> 00:26:01,407 if you have, you know, half an hour to 401 00:26:02,667 --> 00:26:07,107 sit with a pen and paper and compute these gradients, it can totally be done. 402 00:26:08,787 --> 00:26:11,127 So I mean, dL by dW2 is basically 403 00:26:11,727 --> 00:26:14,967 dL by dW2 IJ for each of these IJ. 404 00:26:16,227 --> 00:26:21,067 So it's like, you know, we had matrix times matrix ReLU matrix ReLU. 405 00:26:22,007 --> 00:26:22,987 So now D 406 00:26:23,607 --> 00:26:27,187 and O... So this is Y. And then we did, you know, dot product with the actual answers, we 407 00:26:27,207 --> 00:26:27,907 took a sum. 408 00:26:28,647 --> 00:26:30,627 So now this D loss by... 409 00:26:31,767 --> 00:26:33,247 We have to do with each of these values. 410 00:26:33,247 --> 00:26:36,127 If this value changes, you know, how much does the loss change? 411 00:26:36,127 --> 00:26:38,347 If this value changes, how much does the loss change? 412 00:26:38,947 --> 00:26:40,947 And so on for all the values in this matrix. 413 00:26:41,587 --> 00:26:44,827 That's dL by dW1. Same way we look at this matrix. 414 00:26:45,667 --> 00:26:48,107 We'll be like, if this value changed, how much should the loss change? 415 00:26:48,107 --> 00:26:50,047 If this value changed, how much should the loss change? 416 00:26:50,687 --> 00:26:51,367 And so on. 417 00:26:52,027 --> 00:26:56,267 We'll get one more matrix of those changes. That's dL by dW2. 418 00:26:56,767 --> 00:26:59,547 And yeah, those are just your gradients. 419 00:26:59,547 --> 00:26:59,907 Like, 420 00:27:01,607 --> 00:27:03,467 if you change each, if you change this... 421 00:27:03,467 --> 00:27:08,127 Like for each of these values in W1, if you changed it by some epsilon, it will change 422 00:27:08,167 --> 00:27:09,847 by, you know, this much times epsilon. 423 00:27:09,847 --> 00:27:14,367 Like, you change it by epsilon, you change it by maybe 0.03 epsilon. 424 00:27:14,367 --> 00:27:16,967 So 0.03 is, you know, your gradient. 425 00:27:28,767 --> 00:27:33,007 Yeah. I mean, I can maybe find a nicer way of writing this. 426 00:27:33,067 --> 00:27:34,127 Just give me a minute. I will 427 00:27:35,027 --> 00:27:37,107 write this in a kind of nice way. 428 00:27:38,467 --> 00:27:39,907 Like, uh... 429 00:27:52,227 --> 00:27:54,867 The loss is equal to... 430 00:28:02,427 --> 00:28:03,707 Into, you know, 431 00:28:05,727 --> 00:28:06,947 something. 432 00:28:19,667 --> 00:28:22,347 Then we took, you know, a ReLU of the entire thing. 433 00:28:33,907 --> 00:28:38,727 Let me just use a... This bracket to differentiate, to 434 00:28:38,767 --> 00:28:40,387 keep it, you know, looking separate. 435 00:28:41,727 --> 00:28:44,967 And then we're gonna multiply that with yet another matrix. 436 00:28:52,607 --> 00:28:53,127 And 437 00:28:54,787 --> 00:28:56,907 then we took, you know, yet another ReLU. 438 00:29:00,987 --> 00:29:03,927 So I'm using another bracket for that. 439 00:29:11,247 --> 00:29:15,067 Then what did we do? We multiplied this with the actual answer, correct? 440 00:29:18,747 --> 00:29:19,187 Yeah. 441 00:29:26,007 --> 00:29:28,767 So this we multiplied with 442 00:29:30,487 --> 00:29:33,207 w dot y dash. 443 00:29:35,567 --> 00:29:39,847 And then what did we do? We took a negative sum of the entire thing, right? Yeah. 444 00:29:41,827 --> 00:29:45,247 So sum we can just, we can just say, you know, magnitude of this thing. 445 00:29:46,847 --> 00:29:49,787 Sum, sum... Yeah, I'll, I'll just use the magnitude 446 00:29:51,047 --> 00:29:52,127 notation. 447 00:30:03,687 --> 00:30:04,927 So here we have, you know, 448 00:30:06,867 --> 00:30:08,367 X00. 449 00:30:11,407 --> 00:30:14,747 This goes all the way to, you know, XND, 450 00:30:16,847 --> 00:30:24,023 and-Here 451 00:30:24,083 --> 00:30:25,443 we have w1, 452 00:30:27,763 --> 00:30:30,383 0, 0. 453 00:30:32,043 --> 00:30:34,124 (background noise) 454 00:30:36,303 --> 00:30:37,144 And this 455 00:30:38,983 --> 00:30:39,784 goes to 456 00:30:41,064 --> 00:30:43,463 w1 or w1 dimensions. 457 00:30:48,324 --> 00:30:49,263 W1, 458 00:30:51,223 --> 00:30:52,203 D, D1. 459 00:30:53,823 --> 00:30:57,023 Instead of calling this D1, let's just call it h. 460 00:30:57,243 --> 00:30:59,163 I think h is more convenient. 461 00:31:00,563 --> 00:31:00,944 Yeah. 462 00:31:03,563 --> 00:31:06,244 So we have w1 Dh. 463 00:31:15,183 --> 00:31:19,163 Then similarly we have, you know, w2, 0, 0. 464 00:31:21,903 --> 00:31:23,683 That goes still, 465 00:31:25,863 --> 00:31:26,503 yeah, 466 00:31:27,783 --> 00:31:30,723 w2 Hc. 467 00:31:34,224 --> 00:31:36,023 Yeah. I think that looks good. 468 00:31:46,044 --> 00:31:46,723 Uh, 469 00:31:48,283 --> 00:31:50,583 yeah, just give me a minute. I need to 470 00:31:52,824 --> 00:31:54,223 get that looking. 471 00:31:56,684 --> 00:32:03,563 (background noise) 472 00:32:08,523 --> 00:32:09,284 Yeah. 473 00:32:22,743 --> 00:33:02,123 (background noise) 474 00:33:02,123 --> 00:33:04,824 Where are we? This is not yet updated. 475 00:33:13,864 --> 00:33:16,643 (background noise) Yeah, this looks good. 476 00:33:22,023 --> 00:33:25,343 Yeah. So we just want to find new dL by dx. 477 00:33:26,303 --> 00:33:28,104 No, sorry, Dw1 0, 0. 478 00:33:28,923 --> 00:33:32,223 Yeah. And remember like these are partial derivatives. 479 00:33:33,043 --> 00:33:33,523 Uh, 480 00:33:34,224 --> 00:33:37,823 I hope everyone remembers, you know, partial derivatives versus total derivatives. 481 00:33:37,983 --> 00:33:38,324 Uh, 482 00:33:39,403 --> 00:33:43,783 partial derivative basically means when we are doing dL by Dw1 0, 0, we are assuming 483 00:33:43,783 --> 00:33:47,923 everything else is constant. So now all these values are constant, you know, everything 484 00:33:47,944 --> 00:33:48,563 in x, 485 00:33:49,383 --> 00:33:53,983 all the other values in w1 itself are constant, all the values in, you know, 486 00:33:54,183 --> 00:33:57,543 w2 are assumed constant. You know, I true is assumed constant. 487 00:33:57,663 --> 00:33:57,864 So, 488 00:33:58,483 --> 00:34:02,303 you know, x0, 0 is the... Sorry, w1 0, 0 is the only thing that's changing. 489 00:34:02,783 --> 00:34:04,423 We change this by, you know, an Epsilon. 490 00:34:04,423 --> 00:34:06,543 (bell ringing) We're just checking by how much will L change, 491 00:34:07,603 --> 00:34:08,983 by what multiple of Epsilon. 492 00:34:10,923 --> 00:34:11,183 So 493 00:34:12,463 --> 00:34:14,943 now this we can actually solve, you know, using chain rule. 494 00:34:15,743 --> 00:34:18,464 So we can write it, 495 00:34:22,583 --> 00:34:23,563 you know, this way like... 496 00:34:24,323 --> 00:34:28,023 Wait, that's bad. I can't use H here and H here as well. 497 00:34:28,203 --> 00:34:30,203 I will have to change this to something else. 498 00:34:31,703 --> 00:34:33,683 Uh, (metal clanking) 499 00:34:33,683 --> 00:34:47,323 (background noise) 500 00:34:51,983 --> 00:34:55,863 let's just make this small h. I think that looks good. 501 00:34:57,744 --> 00:35:00,303 And yeah, this will become a small h. 502 00:35:05,024 --> 00:35:09,623 You know the notation that we have stepped is below as well, so I'd rather not change 503 00:35:09,684 --> 00:35:10,363 this. 504 00:35:12,843 --> 00:35:13,703 Mm. 505 00:35:16,343 --> 00:35:19,303 Yeah, I think I'll just have to make this small h then. Yeah. 506 00:35:21,043 --> 00:35:25,383 Yeah, but that's kind of an annoying way of doing this. 507 00:35:30,183 --> 00:35:31,723 (background noise) Just give me some other letter. 508 00:35:31,783 --> 00:35:34,363 Uh, 509 00:35:38,783 --> 00:35:41,084 e? Okay. 510 00:35:42,003 --> 00:35:46,024 I, I don't know. This is probably not what's used in, you know, the standard textbooks or 511 00:35:46,024 --> 00:35:49,244 whatever, but I'm just going to make this thing e, you know, fuck it. 512 00:35:50,883 --> 00:35:51,423 Uh, 513 00:35:54,084 --> 00:35:56,823 yeah, let me just check, is h sitting anywhere? 514 00:35:58,723 --> 00:36:00,904 This case-sensitive. 515 00:36:05,683 --> 00:36:07,143 Yeah. 516 00:36:10,703 --> 00:36:14,523 (background noise) Cool. Yeah, cool. This is good. 517 00:36:18,783 --> 00:36:23,763 (background noise) 518 00:36:27,060 --> 00:36:30,260 Yeah. So we just made that number of dimensions E. 519 00:36:33,139 --> 00:36:34,659 So E is, you know, 800. 520 00:36:37,399 --> 00:36:39,679 This is why they called, you know, 800 hidden layers. 521 00:36:39,679 --> 00:36:41,959 Like if you ever go see, you know, the... 522 00:36:42,979 --> 00:36:47,379 Some like standard lecture on this, this is a very standard problem, you know, two layer 523 00:36:47,379 --> 00:36:49,559 fully connected network trained on MNIST. 524 00:36:50,619 --> 00:36:50,939 Uh, 525 00:36:52,119 --> 00:36:54,039 this thing will be called this number of hidden layers2. 526 00:36:54,039 --> 00:36:55,639 So there are 800 hidden layers 527 00:36:56,260 --> 00:36:57,839 in our neural network. 528 00:36:59,579 --> 00:37:02,320 So you can solve this. If you solve this, you will get, 529 00:37:03,399 --> 00:37:05,179 uh, this answer. So 530 00:37:06,420 --> 00:37:10,280 this is gradient W1, you know, dL by dW1 is this answer, and 531 00:37:12,039 --> 00:37:14,579 uh, dL by dW2 is this answer. 532 00:37:17,079 --> 00:37:18,619 So basically, 533 00:37:21,280 --> 00:37:25,119 this x dot T is x transpose. I'm just clarifying, you know, the notation. 534 00:37:25,739 --> 00:37:30,279 And maybe I'll just start here. h dot T is, you know, h transpose, uh, B hat. 535 00:37:32,799 --> 00:37:37,299 Uh, here we are taking minus Y true, so minus Y true times mask. 536 00:37:37,299 --> 00:37:38,599 Now mask is 537 00:37:39,599 --> 00:37:41,679 all the positive values of a2. ... 538 00:37:42,339 --> 00:37:45,379 a2. Okay. a2 is h times W2. 539 00:37:47,259 --> 00:37:49,579 Can you see? It is a little bit complicated, but yeah. 540 00:37:51,899 --> 00:37:53,939 Like let's just go back here. 541 00:37:55,059 --> 00:37:57,499 So h times W2. So basically they took this thing, 542 00:37:58,139 --> 00:38:00,479 like this, this much is h times W2 down here. 543 00:38:02,879 --> 00:38:05,519 They checked which of those values are greater than zero, 544 00:38:07,799 --> 00:38:10,539 and s- some of those are one and some of those are zero. 545 00:38:11,739 --> 00:38:14,139 So we made that, you know, a mask, mask2. 546 00:38:15,239 --> 00:38:17,759 And by the... Why did we end up with something weird like this, 547 00:38:18,379 --> 00:38:21,739 where put... the positive values are made one, the negative values are made zero? 548 00:38:21,739 --> 00:38:22,439 It's basically the 549 00:38:23,279 --> 00:38:26,559 positive values, their gradients are one, and the negative values, their gradients are 550 00:38:26,579 --> 00:38:29,139 zero. Like, this is the gradient of ReLU. 551 00:38:31,339 --> 00:38:35,659 Like, if you have a function that goes, you know, ting, ting, so now the gradient of 552 00:38:35,679 --> 00:38:38,259 these values is zero, and, you know, the gradient of these values is one. 553 00:38:38,819 --> 00:38:41,179 So that's why you ended up with this thing called a mask. 554 00:38:43,139 --> 00:38:45,479 If you would actually try deriving this, you'll understand, I think. 555 00:38:45,479 --> 00:38:48,379 I don't have time to, you know, derive the entire thing right now. 556 00:38:50,319 --> 00:38:51,019 The... 557 00:38:51,639 --> 00:38:51,939 So 558 00:38:52,659 --> 00:38:55,999 you'll, you'll get a mask, you'll do minus Y true into this mask. 559 00:38:56,759 --> 00:39:00,939 Then you'll do h transpose matrix multiplied with this entire thing, that gives you your 560 00:39:00,979 --> 00:39:01,699 gradient. 561 00:39:02,639 --> 00:39:04,859 Then now that you have this gradient... 562 00:39:08,119 --> 00:39:12,399 okay. I think that you keep separate, and separately you also need to do this calculation 563 00:39:13,259 --> 00:39:15,299 for, you know, gradient of W1. 564 00:39:16,599 --> 00:39:20,059 So this derivation, I'm going to leave it as homework. 565 00:39:21,119 --> 00:39:25,139 Uh, those of who, you who are good at linear algebra, you can probably get this 566 00:39:25,159 --> 00:39:27,379 derivation done. It's basically, you know, 567 00:39:28,019 --> 00:39:29,919 here is the original thing we started with. 568 00:39:30,439 --> 00:39:34,759 Assume you know, except W100, every, just keep everything constant. 569 00:39:34,839 --> 00:39:37,719 Then except W101, keep everything constant. 570 00:39:37,799 --> 00:39:42,699 Except W102, keep everything constant, and just find dL by, you know, that value. 571 00:39:45,999 --> 00:39:49,859 So you'll get these derivations. You know, we have found the matrices. 572 00:39:50,779 --> 00:39:53,119 So, but this works at one particular point. 573 00:39:53,119 --> 00:39:53,359 So, 574 00:39:53,999 --> 00:39:57,959 you know, access constraint we have assumed in our input, provide you through the actual 575 00:39:57,999 --> 00:39:59,399 answers is constraint. 576 00:40:00,139 --> 00:40:02,219 W1 and W2 are the things we need to find. 577 00:40:02,219 --> 00:40:06,859 So if we are given a current W1 and W2, we can find another gradients at that 578 00:40:06,919 --> 00:40:07,759 particular position. 579 00:40:09,179 --> 00:40:09,679 Then 580 00:40:10,359 --> 00:40:12,039 we will update W1 and W2, 581 00:40:12,999 --> 00:40:16,959 like we multiply these ingredients with a small epsilon, just a tiny amount. 582 00:40:16,979 --> 00:40:19,719 The actual number is something we can decide later, 583 00:40:20,379 --> 00:40:23,799 how big we want epsilon to be. That is a whole separate topic. 584 00:40:25,099 --> 00:40:28,839 Uh, there is this thing called an atom optimizer that is used nowadays. 585 00:40:29,019 --> 00:40:30,619 I will get to that later. Not right now. 586 00:40:30,799 --> 00:40:33,699 But right now, just assume E is some very small constraint value. 587 00:40:33,699 --> 00:40:36,019 It could be, you know, 0.0001 or whatever. 588 00:40:38,219 --> 00:40:41,719 So we multiply that with our gradient and we update W1 and W2. 589 00:40:41,999 --> 00:40:46,039 So we get as, uh, a new W1 and W2, which has slightly less loss. 590 00:40:47,899 --> 00:40:51,079 The L now for this new W1 and W2 is a little bit less, 591 00:40:52,079 --> 00:40:54,079 or well, sorry, I think L is negative, so 592 00:40:55,119 --> 00:40:56,899 L will go up actually, yeah. 593 00:40:58,439 --> 00:41:01,599 And then we repeat this process millions of times. 594 00:41:02,979 --> 00:41:03,959 Like, we will have to do 595 00:41:04,919 --> 00:41:09,479 these calculations, we'll have to do them millions of times and keep updating W1 and W2, 596 00:41:10,479 --> 00:41:11,999 and at the end we'll get a good answer. 597 00:41:16,699 --> 00:41:17,239 Uh, 598 00:41:18,399 --> 00:41:18,699 yeah. 599 00:41:19,399 --> 00:41:20,479 This is what 600 00:41:21,259 --> 00:41:23,439 a two layer fully connected network 601 00:41:24,059 --> 00:41:24,639 is. 602 00:41:27,199 --> 00:41:32,059 And how did we manage to run this? Uh, we will probably need a GPU to run 603 00:41:32,099 --> 00:41:32,439 this. 604 00:41:33,279 --> 00:41:34,759 So the GPU is basically 605 00:41:36,079 --> 00:41:36,259 a 606 00:41:37,359 --> 00:41:40,979 chip that does matrix multiplications very fast. 607 00:41:41,279 --> 00:41:42,479 Like, that's the whole thing. 608 00:41:43,179 --> 00:41:47,179 The GPUs were actually way back, you know, before deep learning was very popular, GPUs 609 00:41:47,179 --> 00:41:48,519 were actually designed for gaming. 610 00:41:49,379 --> 00:41:54,339 And for gaming you need to render like, you know, vectors into 611 00:41:54,379 --> 00:41:55,099 graphics. 612 00:41:56,079 --> 00:41:57,539 So your vector might be, you know, 613 00:41:58,399 --> 00:42:01,619 draw a line from 0, 1 to 4, 3. 614 00:42:02,679 --> 00:42:05,619 And we need an actual, you know, pixel matrix. 615 00:42:05,639 --> 00:42:08,859 Like we need all, each of those pixels, you know, set to 1, 1, 1, 1, 1, 1. 616 00:42:10,259 --> 00:42:12,659 Like, even right now, if you look at this entire thing, 617 00:42:13,699 --> 00:42:16,259 like what's showing on your monitor, actually, 618 00:42:16,999 --> 00:42:20,519 initially the instructions are like, okay, you know, draw a rectangle here for this tab. 619 00:42:20,519 --> 00:42:22,999 You know, put some text inside here. Now what the hell is text? 620 00:42:23,019 --> 00:42:25,239 Now go and get this, you know, font sheet from here. 621 00:42:25,899 --> 00:42:29,543 There is some, you know, bunch of images for each of these letters.... 622 00:42:29,644 --> 00:42:30,224 you know, draw 623 00:42:31,443 --> 00:42:32,343 lines here, 624 00:42:33,263 --> 00:42:35,703 draw, you know, a rounded rectangle here and, you know, 625 00:42:36,743 --> 00:42:40,403 this gets a lot more complicated when you are playing a game, for example. 626 00:42:40,443 --> 00:42:42,783 So when you have a game, you have, you know, a 3D game engine, 627 00:42:43,704 --> 00:42:47,523 so you need to convert all those lines into, you know, this ultimate 2D view. 628 00:42:48,864 --> 00:42:53,083 So doing all those conversions also in practice you end up doing using a bunch of matrix 629 00:42:53,124 --> 00:42:54,583 multiplications, so 630 00:42:55,243 --> 00:42:57,443 that's what GPUs were originally designed for. 631 00:42:58,003 --> 00:43:02,424 Today, you know, GPUs are just being used for AI and people are even 632 00:43:02,463 --> 00:43:05,563 redesigning GPUs to become better at AI. So, yeah. 633 00:43:07,663 --> 00:43:11,383 So that's why, you know, GPUs were originally good for matrix multiplication. 634 00:43:13,143 --> 00:43:16,323 And if you look at these, you know, the only things we have mainly are only 635 00:43:16,783 --> 00:43:20,103 computationally expensive operations here are matrix multiplication. 636 00:43:21,443 --> 00:43:24,924 We also have this, you know, zero check, is this positive or negative like across the 637 00:43:24,943 --> 00:43:29,003 matrix? This can also be parallelized and GPUs can parallelize this as well. 638 00:43:29,023 --> 00:43:33,603 Like, again, remember these are like millions of values there. 639 00:43:35,183 --> 00:43:36,423 A2 is, 640 00:43:37,023 --> 00:43:38,203 uh, how big is A2? 641 00:43:41,983 --> 00:43:44,183 60,000 cross 10, right? Yeah. 642 00:43:45,103 --> 00:43:48,903 Yeah, so this is 600,000 values. We would rather, you know, check all of this in 643 00:43:48,943 --> 00:43:50,743 parallel, not check them, you know, one by one. 644 00:43:51,023 --> 00:43:51,323 So, 645 00:43:52,463 --> 00:43:54,203 yeah, all this can be parallelized. 646 00:43:54,843 --> 00:43:58,083 We are not doing any copy-paste anywhere and that is important. 647 00:43:58,223 --> 00:44:01,643 If we had to, you know, copy-paste, you know, one million values from here to there, that 648 00:44:01,663 --> 00:44:03,223 would be a problem, it will become slow. 649 00:44:04,383 --> 00:44:07,943 And because we're doing this on every iteration, so like, we have to copy one million 650 00:44:07,983 --> 00:44:09,803 values, you know, one million times. 651 00:44:09,843 --> 00:44:13,163 So that's like a trillion copy-paste operations and that's bad. 652 00:44:13,263 --> 00:44:16,003 So, if we had copy-paste that would be a problem. 653 00:44:17,023 --> 00:44:18,183 We don't have that either. 654 00:44:18,803 --> 00:44:21,803 We just need to in place, we need to update, you know, w1 and w2. 655 00:44:24,583 --> 00:44:27,243 So, yeah, this will run quickly on a GPU. 656 00:44:29,203 --> 00:44:29,743 And, 657 00:44:30,783 --> 00:44:33,683 yeah, okay, there's a note, I have simplified some steps here. 658 00:44:34,603 --> 00:44:38,563 So one thing which I've simplified is, in practice we are not going to be... 659 00:44:39,023 --> 00:44:42,843 typically, we are not going to be processing our entire 60,000 images all at once. 660 00:44:43,443 --> 00:44:45,683 We will typically split it into batches. 661 00:44:46,223 --> 00:44:50,403 We will do, like the gradient descent we'll be doing, we'll be taking different batches 662 00:44:50,423 --> 00:44:51,683 in different iterations. 663 00:44:52,303 --> 00:44:57,243 And this you can actually maybe go read online how batch SGD works but, basically we su- 664 00:44:57,283 --> 00:45:00,543 spread the 60,000 into batches of let's say 100 images each. 665 00:45:00,983 --> 00:45:04,783 We compute the gradient not over the entire 60,000 images, but over 100 images. 666 00:45:05,463 --> 00:45:06,483 So our x, 667 00:45:07,523 --> 00:45:11,563 when we do this anti-gradient calculation, we will take x as, you know, just 668 00:45:12,003 --> 00:45:16,923 100 cross, uh, 784 instead of, you know, 60,000 cross 784. 669 00:45:17,563 --> 00:45:20,183 We'll compute the gradients, we'll update w1 and w2. 670 00:45:20,863 --> 00:45:23,963 Now we will take a different batch, and th- here's where things get weird, 671 00:45:24,783 --> 00:45:27,643 uh, and then we will compute the gradient on that. 672 00:45:28,623 --> 00:45:31,723 But now you will ask, "Uh, isn't that a whole separate problem?" You know? 673 00:45:32,643 --> 00:45:36,543 Like the optimal w1 and w2 for this batch and optimal w1 and w2 for that batch are 674 00:45:36,543 --> 00:45:37,123 different like, 675 00:45:37,743 --> 00:45:40,903 this is an entire separate gradient descent problem, that's now a separate gradient 676 00:45:40,923 --> 00:45:44,363 descent problem but, we'll take one gradient from here, we'll take one gradient from 677 00:45:44,403 --> 00:45:47,743 there, we'll make a third batch, we'll take another gradient from there. 678 00:45:48,583 --> 00:45:51,483 Then we will use all of these gradients, we'll add up all of them, you know. 679 00:45:52,063 --> 00:45:55,543 Epsilon plus this, epsilon plus this, epsilon plus this, and we'll keep doing all of 680 00:45:55,603 --> 00:45:55,903 this. 681 00:45:57,263 --> 00:46:01,403 So this is called, you know, batch SGD. SD is stochastic gradient descent. 682 00:46:02,043 --> 00:46:02,963 And here there's a 683 00:46:04,283 --> 00:46:08,343 bit of randomness that will come in because we are not taking all the images at once, 684 00:46:08,343 --> 00:46:09,763 we're like just randomly selecting 685 00:46:10,403 --> 00:46:13,503 different batches and taking gradients from all of them and adding it up. 686 00:46:14,183 --> 00:46:17,963 There's a bit of randomness that comes in, but this will also give a good enough 687 00:46:17,963 --> 00:46:18,543 solution. 688 00:46:19,423 --> 00:46:24,163 Yeah, and by the way, the only reason we didn't take all of them is, again, 60,000 images 689 00:46:24,163 --> 00:46:27,543 might not fit in your GPU advance, and that's the only reason, like, you... 690 00:46:28,023 --> 00:46:32,923 If you can fit in 60,000 G- like images into the GPU you might be honestly better off 691 00:46:32,923 --> 00:46:37,183 just doing the normal gradient descent, like, the regular way of doing it what I've 692 00:46:37,183 --> 00:46:37,963 described, but 693 00:46:38,863 --> 00:46:39,763 if you had even more 694 00:46:41,203 --> 00:46:45,263 data like instead of 60,000 images you had, you know, like a billion images, 695 00:46:46,183 --> 00:46:50,743 and in modern datasets you have literally, you know, billions of images, billions of, you 696 00:46:50,763 --> 00:46:52,863 know, text documents, so, 697 00:46:53,783 --> 00:46:56,723 you can't fit all of that into GPU, so, yeah. 698 00:46:57,923 --> 00:47:00,203 Then we use an optimizer called ADAM. 699 00:47:01,483 --> 00:47:06,143 ADAM is basically just a way of changing these epsilon values, so, we 700 00:47:06,203 --> 00:47:10,883 start with big epsilon values like we first descend f- fast and then we reduce the 701 00:47:10,883 --> 00:47:14,363 epsilon values, so Adam is just, how do you do this? 702 00:47:14,363 --> 00:47:14,583 And 703 00:47:15,723 --> 00:47:19,603 the main reason to do this is just to again save cost like 704 00:47:20,483 --> 00:47:23,363 how do we do this anti-gradient descent more quickly? 705 00:47:24,323 --> 00:47:27,203 And last what? Cross-entropy loss, 706 00:47:28,083 --> 00:47:31,243 yeah I think I've already defined this using cross-entropy loss only. 707 00:47:32,803 --> 00:47:37,403 But yeah cross-entropy loss is basically this thing about the log probabilities like we 708 00:47:37,403 --> 00:47:40,563 didn't use actual probabilities we used log probabilities like that, 709 00:47:41,263 --> 00:47:44,903 the loss that comes out at the end of this is called cross-entropy loss. 710 00:47:47,263 --> 00:47:47,643 Yeah. 711 00:47:48,383 --> 00:47:49,123 So now, 712 00:47:49,783 --> 00:47:51,143 here's some homework. 713 00:47:52,123 --> 00:47:55,043 So the first homework is I have already told you, you know you need to 714 00:47:56,943 --> 00:48:00,083 derive this backprop formula yourself. 715 00:48:01,783 --> 00:48:06,283 This thing just actually derive this or at least go online search for a 716 00:48:06,303 --> 00:48:10,803 derivation you know look at it like you can ask O3 to do this stuff for you like 717 00:48:11,163 --> 00:48:15,503 actually O3 computed this you know you can ask O3 give me more steps you can be like you 718 00:48:15,503 --> 00:48:18,283 know O3 I do not understand this you know explain this or whatever. 719 00:48:19,423 --> 00:48:22,283 Like I can literally show you what I'm talking about. Uh... 720 00:48:26,623 --> 00:48:28,223 Explain this 721 00:48:29,883 --> 00:48:30,583 in about 722 00:48:31,283 --> 00:48:31,743 10, 10 723 00:48:33,003 --> 00:48:37,455 to 15 minutes.Like I ment- 724 00:48:37,775 --> 00:48:39,695 mentioned a point earlier, you know, like... 725 00:48:41,076 --> 00:48:41,235 Wait, 726 00:48:41,855 --> 00:48:43,636 let's, uh, copy this. 727 00:48:46,036 --> 00:48:46,855 (keyboard clicking) And put 728 00:48:48,416 --> 00:48:49,735 this in here. 729 00:48:55,076 --> 00:48:57,395 (keyboard clicking) Why are we doing 730 00:48:59,595 --> 00:49:01,755 A2 greater than zero? 731 00:49:04,395 --> 00:49:08,176 (keyboard clicking) Like can you see how we did A2 greater than zero? 732 00:49:08,195 --> 00:49:10,235 I'm just asking O3, you know, why are we doing that? 733 00:49:16,795 --> 00:49:17,715 (breathing) 734 00:49:19,135 --> 00:49:22,536 (chair creaking) So, you can ask for re- these types of, you know, specific questions. 735 00:49:24,815 --> 00:49:29,516 Yes. What does it say? Okay, it just... yeah, this is just defining the problem. 736 00:49:32,175 --> 00:49:34,075 Okay, there are... chain rule. 737 00:49:36,015 --> 00:49:38,775 Yeah, that's an lengthy-ass derivation. 738 00:49:39,095 --> 00:49:39,496 Uh, 739 00:49:40,376 --> 00:49:41,715 yeah, here is our explanation. 740 00:49:42,595 --> 00:49:44,476 Why is A2 greater than zero? 741 00:49:45,296 --> 00:49:50,135 Yes, a ReLu is this function, you know, zero if negative, you know, Z is 742 00:49:50,155 --> 00:49:51,155 positive. And 743 00:49:51,916 --> 00:49:53,336 its derivative is zero and one. 744 00:49:54,355 --> 00:49:56,155 That's why we did this. 745 00:49:57,955 --> 00:50:01,515 Yeah, I mean, you can sit and read all this if you, for instance want to understand, you 746 00:50:01,515 --> 00:50:02,995 know, how to do the gradient of ReLu. 747 00:50:04,795 --> 00:50:06,816 So you can... yeah, that's your first homework. 748 00:50:07,436 --> 00:50:10,956 The second homework which is more advanced, but again, you will find this online. 749 00:50:11,315 --> 00:50:13,976 You don't have f-... You can do this from scratch if 750 00:50:14,995 --> 00:50:19,015 you really want to understand this or, you know, you're a bit of a masochist, uh, but 751 00:50:19,635 --> 00:50:24,475 you will find this assignment existing online too, which is, you know, actually train 752 00:50:24,475 --> 00:50:27,095 this fully connected network with, you know, 800 hidden, 753 00:50:28,135 --> 00:50:29,595 uh, dimensions. 754 00:50:31,195 --> 00:50:32,856 Use Pytorch. Uh, 755 00:50:33,715 --> 00:50:37,955 you will get a 1.6% error rate, which means in out of the 60,000, 756 00:50:38,895 --> 00:50:43,276 uh, examples, not... (keyboard clicking) Sorry, on the remaining 10,000, right? 757 00:50:43,276 --> 00:50:44,015 (keyboard clicking) So on the remaining 758 00:50:44,736 --> 00:50:46,435 10,000 you'll get, 759 00:50:48,635 --> 00:50:52,276 (keyboard clicking) uh, what am I doing? Yeah, 98.4% correct. 760 00:50:52,276 --> 00:50:52,495 So 761 00:50:53,475 --> 00:50:58,136 9,840 examples you'll get correct. And about 160 you'll get wrong 762 00:50:58,375 --> 00:51:01,315 on the remaining 10,000 examples which you've not seen till now. 763 00:51:02,915 --> 00:51:05,536 (chair creaking) So that's pretty good. 764 00:51:05,595 --> 00:51:08,595 (child talking in background) So yeah, that's what 765 00:51:09,355 --> 00:51:13,835 you will do. (paper rustling) Yeah, and also I want, if you're doing this 766 00:51:13,835 --> 00:51:16,655 assignment, like you're actually sitting and writing this in Pytorch, 767 00:51:17,355 --> 00:51:20,515 uh, I would recommend you actually write down these gradients. 768 00:51:20,515 --> 00:51:22,835 Otherwise Pytorch has this feature called autograd, 769 00:51:23,636 --> 00:51:27,496 which means Pytorch can compute these gradient formulas on its own and use them. 770 00:51:28,395 --> 00:51:28,755 Like, 771 00:51:29,455 --> 00:51:34,355 this is part of the reason why NVIDIA is on track to become a trillion 772 00:51:34,395 --> 00:51:35,195 dollar company, 773 00:51:36,155 --> 00:51:37,496 uh, is, you know, 774 00:51:39,416 --> 00:51:42,075 Pytorch does these types of calculations on their own. 775 00:51:42,075 --> 00:51:45,275 So ML engineers don't need to sit and, you know, break their head on how to do 776 00:51:46,575 --> 00:51:47,455 gradients, 777 00:51:48,496 --> 00:51:50,335 like how to calculate this gradient formulas. 778 00:51:50,336 --> 00:51:54,115 This is not the only reason NVIDIA is so rich, but, uh, you know, it's part of the 779 00:51:54,115 --> 00:51:54,475 reason. 780 00:51:55,095 --> 00:51:55,996 (paper rustling) Uh, so yeah, 781 00:51:56,615 --> 00:51:58,755 I would recommend for the homework assignment, you know, 782 00:51:58,795 --> 00:52:03,756 (child playing in background) uh, yeah, use these actual formulas which are written 783 00:52:03,775 --> 00:52:04,936 down here. Don't, you know, 784 00:52:05,555 --> 00:52:08,315 ask Pytorch to i- autograd to, you know, just do it. 785 00:52:11,655 --> 00:52:12,035 Cool. 786 00:52:13,575 --> 00:52:15,135 Any questions? I think 787 00:52:16,315 --> 00:52:19,695 this- (bell ringing) what we've done here, this is actually the most important thing in 788 00:52:19,735 --> 00:52:20,535 the lesson today. 789 00:52:21,235 --> 00:52:24,295 The remaining one hour is a bit less important (laughs) than this one hour. 790 00:52:25,355 --> 00:52:29,595 If you have just understood this very well, you have understood, uh, like a lot about how 791 00:52:29,615 --> 00:52:30,695 deep learning works. 792 00:52:34,935 --> 00:52:35,455 Cool. 793 00:52:36,975 --> 00:52:38,975 Now (coughs) comes the next section. 794 00:52:39,595 --> 00:52:39,895 Yeah. 795 00:52:40,655 --> 00:52:41,895 So here is a question. 796 00:52:43,115 --> 00:52:47,235 Why did we do all this, you know? Why did we use two layers? Why did we use a ReLu? 797 00:52:47,775 --> 00:52:49,955 Why did we define our loss function this way? 798 00:52:50,155 --> 00:52:53,095 Like we could have defined the loss function different ways, you know, actually if we 799 00:52:53,115 --> 00:52:53,735 wanted to. 800 00:52:55,055 --> 00:52:55,335 Like 801 00:52:56,095 --> 00:52:58,975 here I took, you know, this value and ignored all the other values. 802 00:52:59,115 --> 00:53:01,675 And also I said this should be a log probability. 803 00:53:01,695 --> 00:53:04,295 Like why did we decide all this, you know? 804 00:53:05,795 --> 00:53:10,635 And in practice, when you go into proper deep learning, you will get even more questions. 805 00:53:11,935 --> 00:53:13,835 Like you will learn things, you know... 806 00:53:14,575 --> 00:53:16,275 There will be a lot of these types of tricks. 807 00:53:16,275 --> 00:53:19,415 Like okay, use RMS norm here, use Softmax here, 808 00:53:20,355 --> 00:53:25,215 define attention layers in this way, use these many layers, like we use two layers, you 809 00:53:25,215 --> 00:53:26,535 know. Why two layers? Like... 810 00:53:27,135 --> 00:53:29,315 And this in practice becomes a lot more layers. 811 00:53:29,975 --> 00:53:34,675 You know, what's a residual layer? Like, we are going to encounter a lot of tricks 812 00:53:35,295 --> 00:53:38,955 and whenever a question comes in your mind, you know, why are we doing this trick? 813 00:53:39,415 --> 00:53:41,595 Always by default, here is the answer. 814 00:53:43,035 --> 00:53:45,315 Um, I'm going to just open the answer right now for you. 815 00:53:53,915 --> 00:53:54,015 (person coughing) 816 00:53:55,655 --> 00:53:55,715 (water running) 817 00:53:55,715 --> 00:53:56,595 Should we bow? 818 00:53:57,235 --> 00:53:59,795 Yeah, he's a king. Seems like I'm always thanking you for something. 819 00:54:00,335 --> 00:54:02,015 (clears throat) What are you doing? 820 00:54:02,595 --> 00:54:04,655 Uh, we, we don't do that here. 821 00:54:06,015 --> 00:54:06,635 Yeah. 822 00:54:07,295 --> 00:54:08,875 So that's the answer. 823 00:54:09,755 --> 00:54:11,795 We do not ask these types of questions here. 824 00:54:12,675 --> 00:54:15,575 Do not ever ask me, "Why did we do this?" and "Why did we not do that?" 825 00:54:17,335 --> 00:54:18,295 The answer is 826 00:54:19,075 --> 00:54:21,515 some guy tried it a few years back, it worked out. 827 00:54:21,855 --> 00:54:25,155 Some guy tried a little bunch of similar things, those did not work and this worked 828 00:54:25,275 --> 00:54:28,975 because who knows? And that's the answer. That's always the answer. 829 00:54:29,115 --> 00:54:31,215 You're not going to get no- a better answer than that. 830 00:54:32,135 --> 00:54:35,435 If you think you have a better answer than that, often you're actually mistaken. 831 00:54:37,355 --> 00:54:41,075 And that comes here. Like lots of people will give intuitions for why, you know, this 832 00:54:41,095 --> 00:54:45,735 trick works and that does not work.But often, these intuitions are created 833 00:54:45,736 --> 00:54:47,695 after we have actually tried the experiment. 834 00:54:47,696 --> 00:54:47,955 Like, 835 00:54:49,216 --> 00:54:50,395 we tried an experiment, 836 00:54:51,015 --> 00:54:52,876 some particular trick worked, you know. 837 00:54:53,536 --> 00:54:56,296 Let's say this tool ReLU worked very well. 838 00:54:57,335 --> 00:55:00,376 Instead of ReLU, we could have refined a different, you know, activation function. 839 00:55:01,175 --> 00:55:03,555 But ReLU worked very well and now, let's... 840 00:55:03,555 --> 00:55:08,135 We will now invent some, you know, intuition for why ReLU is a good idea and why, I don't 841 00:55:08,135 --> 00:55:09,735 know, why ReLU is a bad idea or whatever. 842 00:55:11,995 --> 00:55:15,975 Uh, yeah. ReLU is the real thing. Uh, I can show you. 843 00:55:16,995 --> 00:55:18,815 Like, there are other loss functions. 844 00:55:25,055 --> 00:55:28,235 Yeah. So this is what ReLU looks like. Here's the formula for ReLU. 845 00:55:29,075 --> 00:55:33,595 ReLU is just, you know, negative X, make it zero, positive X, make it, you know, X. 846 00:55:35,175 --> 00:55:37,335 And here is a bit more ... Yeah. 847 00:55:40,636 --> 00:55:42,135 So, yeah, people just 848 00:55:44,516 --> 00:55:46,575 create their intuitions afterwards. 849 00:55:47,535 --> 00:55:51,855 Also, yeah, especially modern day ML, because there's only a few labs that are doing, you 850 00:55:51,855 --> 00:55:53,215 know, the latest research, 851 00:55:54,055 --> 00:55:57,135 uh, only people in those labs are going to have these intuitions anyway. 852 00:55:57,395 --> 00:55:57,815 Like, 853 00:55:58,536 --> 00:56:00,355 they are the ones who are running hundreds of runs. 854 00:56:00,376 --> 00:56:04,355 Like, to get an intuition, you need to try some things and see that some of them work and 855 00:56:04,355 --> 00:56:07,195 others did not work. Like, it's a very empirical, 856 00:56:08,315 --> 00:56:11,575 uh, field. It's not a theoretical or research kind of field. 857 00:56:11,915 --> 00:56:15,135 It's just pure, tried this, it worked, tried that, it did not work. 858 00:56:16,335 --> 00:56:19,996 And right now, the people trying things are only, like only if you have, you know, a 859 00:56:20,015 --> 00:56:24,055 company that has raised, you know, millions of dollars in funding, you can try things 860 00:56:24,075 --> 00:56:24,915 anymore. So, 861 00:56:25,695 --> 00:56:27,595 they are the only ones who have these intuitions anymore. 862 00:56:28,555 --> 00:56:31,655 Yeah. So if you ask me, you know, why all these older things? 863 00:56:31,655 --> 00:56:31,855 So, 864 00:56:32,515 --> 00:56:36,755 here are, you know, older architectures, here are, you know, older activation functions, 865 00:56:36,776 --> 00:56:41,775 older loss functions, older optimizers, you know, older regularization 866 00:56:41,835 --> 00:56:46,075 tactics. You know, a lot of these things, you will read about these in, you know, five 867 00:56:46,115 --> 00:56:49,515 year old tutorials. You will not read them, about them in today's tutorials. 868 00:56:49,616 --> 00:56:50,015 And 869 00:56:50,795 --> 00:56:53,756 the only reason is just, you know, we found something else that works better. 870 00:56:53,795 --> 00:56:55,056 That's the only reason. 871 00:56:56,636 --> 00:56:57,035 Cool. 872 00:56:58,036 --> 00:56:59,095 Now, we are coming 873 00:56:59,716 --> 00:57:00,056 to 874 00:57:01,376 --> 00:57:02,015 modern day 875 00:57:02,835 --> 00:57:04,615 machine learning, deep learning, you know. 876 00:57:05,335 --> 00:57:06,895 We are coming to transformers. 877 00:57:09,395 --> 00:57:13,755 So we have, you know, ChatGPT or Grok-3, or you know, Claude. 878 00:57:13,755 --> 00:57:16,636 Any model you are using today is basically a transformer. 879 00:57:19,455 --> 00:57:23,115 Uh, here is the lecture from which I first studied transformers. 880 00:57:23,115 --> 00:57:24,895 I still think it's an amazing lecture. 881 00:57:30,895 --> 00:57:33,895 It's by this guy called Justin Johnson. 882 00:57:34,175 --> 00:57:39,055 It's four years old, but I think it's still relevant, even despite being a four year old 883 00:57:39,815 --> 00:57:43,455 video. You know, four year old videos would otherwise be considered ancient in 884 00:57:44,315 --> 00:57:47,815 deep learning speak. So if you have a video that's four years old and it's still relevant 885 00:57:47,815 --> 00:57:48,475 today, you know, 886 00:57:49,275 --> 00:57:50,395 that means it's a good 887 00:57:51,035 --> 00:57:51,515 video. 888 00:57:53,715 --> 00:57:57,135 So yeah, here's lecture 13. This will make more sense if you actually see some of the 889 00:57:57,155 --> 00:58:00,995 previous lectures. It is going to make less sense if you just jump into it directly. 890 00:58:01,035 --> 00:58:04,955 But yeah, if you want to understand transformers more in depth, that's the lecture for 891 00:58:04,975 --> 00:58:05,195 you. 892 00:58:06,375 --> 00:58:09,975 And here is an actual transformer we are going to look at. 893 00:58:14,175 --> 00:58:19,115 This is LLaMA-4. It was trained by Facebook, uh, as the model 894 00:58:19,175 --> 00:58:19,615 card. 895 00:58:21,755 --> 00:58:22,175 Yeah. 896 00:58:23,115 --> 00:58:26,135 Llama-4 Scout, LLaMA-4 Maverick. 897 00:58:28,235 --> 00:58:29,535 Hmm. So, 898 00:58:31,875 --> 00:58:36,795 yeah. Scout has 109 billion parameters, Maverick has 400 899 00:58:36,815 --> 00:58:40,895 billion parameters. I'll explain later on, you know, what is parameters. 900 00:58:40,915 --> 00:58:41,855 Right now, I'm just saying, 901 00:58:42,635 --> 00:58:45,015 like, this is basically the size of the model. 902 00:58:45,095 --> 00:58:46,915 Like, it's the size of the weight matrices. 903 00:58:46,915 --> 00:58:49,315 You know, we talked about, right, like W1, W2, so. 904 00:58:49,615 --> 00:58:51,375 Transformers also have weight matrices. 905 00:58:51,375 --> 00:58:51,875 This is the 906 00:58:52,555 --> 00:58:56,275 sum of, you know, all of the weight matrices, in total, how big are these matrices. 907 00:58:57,295 --> 00:59:01,315 So there are a total of 400 billion numbers inside all these weight tr- matrices, 908 00:59:01,535 --> 00:59:02,175 basically. 909 00:59:05,055 --> 00:59:09,495 And yeah, this was released just two months back, in April 2025. 910 00:59:14,075 --> 00:59:15,515 I think I'm going to have water. 911 00:59:29,335 --> 00:59:34,015 Okay, cool. Uh, this was trained on five million GPU hours. 912 00:59:34,815 --> 00:59:39,175 It basically means you took a GPU, you ran it for, like, 500 years. 913 00:59:39,715 --> 00:59:40,735 500 years, right? 914 00:59:41,675 --> 00:59:42,375 Five 915 00:59:43,895 --> 00:59:45,395 million hours 916 00:59:46,835 --> 00:59:49,495 is... Okay, yeah, 570 hours. So, 917 00:59:50,435 --> 00:59:54,095 you take, like, 570 GPUs and you run them together for a year 918 00:59:55,975 --> 01:00:00,875 or, or, you know, you can take more. You can take, like, uh, like, what, 7000 GPUs and 919 01:00:00,895 --> 01:00:02,135 you can run them for a month. 920 01:00:04,815 --> 01:00:05,735 Uh, like, 921 01:00:06,655 --> 01:00:08,655 a GPU looks like this, by the way. 922 01:00:13,915 --> 01:00:15,855 Yeah. Here is what it... 923 01:00:18,075 --> 01:00:19,855 GPU looks like. 924 01:00:20,795 --> 01:00:22,455 It's just a card like this. 925 01:00:23,455 --> 01:00:26,615 And the cooling is, I think not... Yeah, this includes the cooling. 926 01:00:27,815 --> 01:00:31,035 Like, separately you need, you know, these are the GPUs inside and you need the cooling 927 01:00:31,035 --> 01:00:32,015 unit separately. 928 01:00:36,435 --> 01:00:37,635 So, yeah. 929 01:00:39,955 --> 01:00:41,975 So now I'm going to actually look at the code. 930 01:00:42,035 --> 01:00:44,375 The code is not that too difficult to understand. 931 01:00:45,575 --> 01:00:49,315 If you already have studied transformers, you might actually understand a fair amount of 932 01:00:49,315 --> 01:00:52,791 the code.Yeah. 933 01:00:54,091 --> 01:00:58,591 So again, I don't have time to go through the entire thing, but yeah. 934 01:00:59,351 --> 01:01:03,152 See here, this is called attention. This thing called attention is the most important 935 01:01:03,192 --> 01:01:05,311 part of a transformer. 936 01:01:06,531 --> 01:01:08,832 Uh, here we defined forward. 937 01:01:09,752 --> 01:01:11,311 This is a forward pass, so 938 01:01:12,812 --> 01:01:16,711 you know, what we do to our input to get output, that is called a forward pass. 939 01:01:21,211 --> 01:01:23,831 Here is forward for the transformer block. 940 01:01:24,111 --> 01:01:26,432 Again, I will explain all this soon. 941 01:01:30,251 --> 01:01:32,291 Yeah, I think this is where our actual... 942 01:01:36,031 --> 01:01:37,791 Yeah. Okay, so this is LLaMA 4. 943 01:01:39,291 --> 01:01:41,131 And where is LLaMA 4 forward? 944 01:01:42,652 --> 01:01:46,492 Okay, LLaMA 4 inference mode. This is where we generate... Ah. Yeah, transformer input. 945 01:01:46,492 --> 01:01:46,492 Okay. Yeah. Where is this line called? 93. 946 01:01:46,492 --> 01:01:46,711 Yeah, okay, model equals transformer model args. 947 01:01:46,711 --> 01:01:47,071 So here's the input we give to the transformer and we... No, no, no. 948 01:01:47,091 --> 01:01:47,131 Here's the 949 01:01:47,131 --> 01:01:55,772 hyperparameters 950 01:01:55,791 --> 01:01:56,452 and I 951 01:02:20,052 --> 01:02:21,351 won't explain that right now. 952 01:02:24,532 --> 01:02:28,451 Yeah, I think this script is less useful. It's this one that's useful. 953 01:02:31,212 --> 01:02:34,131 Yeah, return transformer output. 954 01:02:35,831 --> 01:02:36,751 Trans... 955 01:02:37,372 --> 01:02:40,391 Oh, yeah, yeah, yeah. This is useful. Very useful. 956 01:02:40,532 --> 01:02:40,811 (sniffs) 957 01:02:41,431 --> 01:02:44,591 So this is where we start. We start with some tokens. 958 01:02:44,591 --> 01:02:48,631 Okay, so we got self model input. 959 01:02:49,331 --> 01:02:53,071 Tokens equals model input or tokens. So here is already a first step that has happened. 960 01:02:54,491 --> 01:02:55,171 Ah. 961 01:02:56,231 --> 01:02:56,771 Where are we? 962 01:02:59,031 --> 01:03:03,131 Yep, so here is a bunch of steps that happen in a transformer. 963 01:03:03,711 --> 01:03:07,131 The first step is called tokenize. We use a tokenizer. 964 01:03:08,231 --> 01:03:12,531 And tokenizer is pretty straightforward. Basically means we... 965 01:03:14,931 --> 01:03:18,911 Actually first, sorry, I should maybe start with, what the hell does a transformer do? 966 01:03:18,931 --> 01:03:23,431 So a transformer takes a sequence of tokens as input and it gives you log 967 01:03:23,451 --> 01:03:25,331 probabilities about what the next token should be. 968 01:03:26,351 --> 01:03:26,651 So 969 01:03:27,571 --> 01:03:29,651 tokens are just words, you know, like 970 01:03:31,511 --> 01:03:36,351 I'll say, "My name is Samuel. What is your name?" So my is a token, name is a 971 01:03:36,391 --> 01:03:38,911 token, is is a token, Samuel is a token. 972 01:03:39,751 --> 01:03:41,531 There's a sequence of tokens and 973 01:03:42,291 --> 01:03:45,451 a transformer will be given, you know, a sequence of tokens, you know. 974 01:03:45,451 --> 01:03:48,031 "My name is Samuel. What is?" 975 01:03:49,211 --> 01:03:52,951 Predict the next token. And it will give you, you know, log probabilities of like 976 01:03:54,391 --> 01:03:59,331 what is a probability, what is your probability, what is my probability, 977 01:03:59,351 --> 01:04:03,611 what is their probability. Like you will get, you know, a sequence of numbers giving the 978 01:04:03,651 --> 01:04:06,151 probabilities for every possible next token. 979 01:04:07,411 --> 01:04:10,351 So input is some tokens, output is some tokens. 980 01:04:12,631 --> 01:04:12,891 So 981 01:04:14,611 --> 01:04:17,791 yeah, first we convert our words into tokens. 982 01:04:18,271 --> 01:04:23,151 The tokenizer just says instead of my there is a dictionary of maybe some 100,000 983 01:04:23,211 --> 01:04:26,631 popular words in English and they're all numbered, you know, zero, one, two, three, four 984 01:04:26,691 --> 01:04:30,331 till 100,000. Instead of my, maybe my is, you know, word number 985 01:04:31,451 --> 01:04:35,871 4,058. So instead of my we'll put, you know, 4058. 986 01:04:37,131 --> 01:04:39,291 Okay, my name, maybe name is, I don't know, 987 01:04:39,911 --> 01:04:43,711 50732 word or it's the 50,732nd word. 988 01:04:43,711 --> 01:04:44,311 So instead of 989 01:04:45,091 --> 01:04:47,511 this we'll put, you know, 50732 or whatever. 990 01:04:48,531 --> 01:04:51,071 Yeah, actually we won't put 50732. We will 991 01:04:52,371 --> 01:04:57,191 fill 100,000 zeros and on the 50732 index instead of zero we'll put a 992 01:04:57,251 --> 01:04:59,611 one. That's actually how we will fill the matrix. 993 01:05:00,811 --> 01:05:03,431 But yeah, remember we need matrices. Our entire... 994 01:05:03,471 --> 01:05:07,451 All this work has to wor- do- get done on matrices. We can't do work with words. 995 01:05:07,451 --> 01:05:10,851 We have to start with some X. X has to be some matrix input. 996 01:05:11,531 --> 01:05:14,131 So we'll convert our words into, you know, a matrix like this. 997 01:05:14,131 --> 01:05:16,351 So that's just what tokenization is. 998 01:05:17,891 --> 01:05:21,711 Then next step, we'll have this thing called a positional embedding. 999 01:05:22,211 --> 01:05:23,551 Where's the positional embedding? 1000 01:05:25,571 --> 01:05:27,691 Yeah, start position. 1001 01:05:29,391 --> 01:05:32,771 Okay, so we got a position. What did we do with the position? 1002 01:05:35,151 --> 01:05:38,111 We converted that into frequency. 1003 01:05:42,231 --> 01:05:46,011 What did we do with the frequencies? Okay, we took the frequency here and 1004 01:05:48,391 --> 01:05:50,151 we defined layer. 1005 01:05:51,291 --> 01:05:54,331 Like, you know, if you have to sit down and understand this code, yes, it will take time. 1006 01:05:54,511 --> 01:05:59,491 I'm going to do... Take a shortcut. I'm going to copy this entire code and I'm going 1007 01:05:59,551 --> 01:06:01,391 to ask our AI to tell us 1008 01:06:02,211 --> 01:06:03,731 what does everything here mean. 1009 01:06:05,911 --> 01:06:06,991 Explain 1010 01:06:07,831 --> 01:06:11,571 all the steps in simple English. 1011 01:06:17,151 --> 01:06:17,611 Bam. 1012 01:06:20,691 --> 01:06:25,431 Yeah, uh, O3 is a great tool for learning anything by the way. 1013 01:06:25,531 --> 01:06:28,891 Like O3 sometimes makes mistakes, it can understand it. 1014 01:06:29,091 --> 01:06:31,471 And that can be a problem if you're a beginner. 1015 01:06:32,551 --> 01:06:35,371 But apart from that, O3 is a great way of learning anything. 1016 01:06:41,691 --> 01:06:42,871 Cool. So we 1017 01:06:43,611 --> 01:06:45,711 have a huge number of steps, you know. 1018 01:06:48,851 --> 01:06:51,851 Utility pieces. Let's skip the utility pieces right 1019 01:06:51,871 --> 01:06:56,295 now.Yeah, top level transformer. 1020 01:06:56,995 --> 01:07:01,035 It stores hyper parameters, I'm going to ignore that. It... 1021 01:07:03,596 --> 01:07:08,055 Yeah. Yeah, this is the thing we actually want, this is the forward pass. 1022 01:07:09,675 --> 01:07:11,955 It converts tokens to embeddings. 1023 01:07:12,815 --> 01:07:15,855 If there's an image, we will also convert that into embeddings. 1024 01:07:15,915 --> 01:07:19,335 And right now I'm not going to talk about images, let's just consider text for now. 1025 01:07:19,335 --> 01:07:24,216 But we can also convert an image into embeddings and then we do that right at the 1026 01:07:24,255 --> 01:07:27,956 start, and then, you know, rest all the steps there on the same for text or images. 1027 01:07:29,395 --> 01:07:29,796 Yeah, we 1028 01:07:30,735 --> 01:07:35,096 put these positional embeddings, which is basically along with the token we also kind of 1029 01:07:35,136 --> 01:07:35,535 store 1030 01:07:36,255 --> 01:07:41,155 which position it starts instead of, you know, my name is Samuel, it becomes like my 1031 01:07:41,255 --> 01:07:43,835 one name two is three, Samuel four. 1032 01:07:44,616 --> 01:07:47,855 So even if you're in the middle of the sentence somewhere you are at like Samuel four, 1033 01:07:47,915 --> 01:07:49,635 you already know you're at the fourth word. 1034 01:07:50,015 --> 01:07:52,495 Otherwise, when you're somewhere in between, you 1035 01:07:53,096 --> 01:07:55,075 may have forgotten where exactly are you. 1036 01:07:57,595 --> 01:07:59,955 Yeah, then this, this attention mask, 1037 01:08:00,755 --> 01:08:02,255 I'm going to skip that. 1038 01:08:06,935 --> 01:08:10,855 Yeah, I will, I will have to ask this a bunch, bunch of more questions to actually 1039 01:08:10,875 --> 01:08:15,095 simplify this. Like, you can ask or refer more questions, you know, what the hell is 1040 01:08:15,115 --> 01:08:18,735 this? You know, what is this? You can keep asking, it will eventually simplify it, but it 1041 01:08:18,735 --> 01:08:19,635 will take a lot of time. 1042 01:08:20,315 --> 01:08:23,015 I have actually already simplified this stuff, and 1043 01:08:24,755 --> 01:08:25,955 I'm going to follow that. 1044 01:08:26,775 --> 01:08:28,295 So here is the simplified version. 1045 01:08:29,035 --> 01:08:32,955 See, we tokenized, we put positional embeddings, we got image embeddings. 1046 01:08:33,595 --> 01:08:38,375 Then we have N of these, so this entire thing actually repeats, you know, maybe 32 1047 01:08:38,415 --> 01:08:43,115 times or 64 times or 80 times. So let's say we repeat this 80 times. 1048 01:08:43,855 --> 01:08:48,575 We do something called an RMS norm, we do something called multi-headed attention, which 1049 01:08:48,575 --> 01:08:50,755 is the most important part of this entire model. 1050 01:08:51,635 --> 01:08:54,755 This is basically with our input, we, you know, computed something, we computed 1051 01:08:54,795 --> 01:08:56,735 something, we computed something, we computed something. 1052 01:08:57,375 --> 01:09:01,215 We did this entire thing 80 times, then we did a projection and we got an answer. 1053 01:09:01,875 --> 01:09:05,835 So we did something called a multi-headed attention, we did add residual, we did RMS 1054 01:09:05,855 --> 01:09:06,315 norm, 1055 01:09:06,935 --> 01:09:09,955 we did a feed-forward layer, we did add residual. 1056 01:09:10,795 --> 01:09:12,215 So this is the entire thing we did, 1057 01:09:13,055 --> 01:09:17,715 and I don't have time to explain all of these in detail, so I'm just going to explain the 1058 01:09:17,715 --> 01:09:19,155 steps which are really important. 1059 01:09:20,575 --> 01:09:22,795 Steps which are important and the steps which are simple. 1060 01:09:23,795 --> 01:09:28,615 So the steps which are simple are like tokenizing is simple, embedding is simple, RMS 1061 01:09:28,675 --> 01:09:29,535 norm is simple, 1062 01:09:30,355 --> 01:09:33,395 residual is simple. So I'll just explain some of these which are, like, very simple, 1063 01:09:34,275 --> 01:09:37,615 and the thing which is important. So the important thing is this projection 1064 01:09:38,315 --> 01:09:40,535 and this, uh, dot product attention. 1065 01:09:42,675 --> 01:09:46,095 So I'm, I'm going to explain some of these more and some of these less. 1066 01:09:48,455 --> 01:09:52,795 So anyways, keeping the big picture, we started with the sequence of tokens, 1067 01:09:53,675 --> 01:09:56,575 we did this giant amount of calculation on it, 1068 01:09:57,435 --> 01:10:02,115 we got output, you know, log probabilities for what the next token should be, 1069 01:10:03,195 --> 01:10:07,175 and this entire thing is described in this, you know, 500 lines of code here. 1070 01:10:07,175 --> 01:10:11,335 If you go through this, like most of these steps are all, all the code is just sitting 1071 01:10:11,335 --> 01:10:15,255 here. Some of it, a little bit might be setting in ffn.pie for example. 1072 01:10:15,995 --> 01:10:18,195 So here this, you know, another 50 lines of code but, 1073 01:10:18,835 --> 01:10:22,535 yeah, the vast majority of the code for this entire process is just sitting in this one 1074 01:10:22,595 --> 01:10:23,895 script. If you actually 1075 01:10:24,755 --> 01:10:29,315 sit and read this properly you will understand everything I've talked about here is 1076 01:10:29,315 --> 01:10:30,235 actually sitting in there. 1077 01:10:31,155 --> 01:10:35,075 And okay, a bunch of these steps are all actually optional, like rotary embedding, 1078 01:10:35,095 --> 01:10:39,835 optional, headwise RMS norm, optional, temperature tuning for long context, this is 1079 01:10:39,895 --> 01:10:43,115 optional. So you can actually, you can even kind of ignore these for now. 1080 01:10:43,175 --> 01:10:44,155 These are not important. 1081 01:10:46,975 --> 01:10:51,035 Cool. So we started with the sequence of tokens, we did all sorts of 1082 01:10:51,375 --> 01:10:55,115 computation on it, we got an output, now what do we do with this output? 1083 01:10:57,215 --> 01:10:58,575 With these log probabilities? 1084 01:11:05,215 --> 01:11:08,975 Remember right now, actually these log probabilities are completely inaccurate, they're 1085 01:11:08,995 --> 01:11:09,715 like garbage, 1086 01:11:11,115 --> 01:11:13,035 because we don't know what our weight matrices are. 1087 01:11:13,555 --> 01:11:16,495 So like when we do these calculations we will need weight matrices. 1088 01:11:16,495 --> 01:11:16,795 So like 1089 01:11:17,495 --> 01:11:21,615 we will need a weight matrix here for embeddings, we will need a weight matrix for the 1090 01:11:21,635 --> 01:11:22,455 projection, 1091 01:11:23,295 --> 01:11:25,975 we will need a weight matrix for this projection, 1092 01:11:27,855 --> 01:11:31,435 we will need... Actually we'll need three weight matrices here, 1093 01:11:32,855 --> 01:11:35,395 we will need weight matrices for the feed-forward layer. 1094 01:11:35,415 --> 01:11:38,235 So there's a lot of weight matrices here which right now we don't know, 1095 01:11:38,855 --> 01:11:41,715 and we're just going to initialize them with some junk values. 1096 01:11:42,795 --> 01:11:44,515 And yep, we are 1097 01:11:45,235 --> 01:11:47,695 going to do gradient descent remember? Yeah. 1098 01:11:48,295 --> 01:11:52,395 We are going to define some loss at the end and we are going to do dL by dW. 1099 01:11:53,695 --> 01:11:55,355 So first of all, what is loss? 1100 01:11:57,435 --> 01:12:01,775 Loss is how much our predicted answers differ from the actual answers. 1101 01:12:02,415 --> 01:12:05,135 So we will get some log probabilities of the next token, 1102 01:12:05,955 --> 01:12:09,975 and we know the actual next token. So right now with that data we are assuming we 1103 01:12:09,975 --> 01:12:11,315 actually know the next token. 1104 01:12:12,175 --> 01:12:16,395 So now this, given this and this, we can mul- dot product these two, we will get a 1105 01:12:16,415 --> 01:12:20,635 cross-entropy loss. This is just a value telling us, you know, how bad is our prediction, 1106 01:12:20,655 --> 01:12:23,955 which right now when we start off it's pretty, pretty bad. 1107 01:12:25,855 --> 01:12:27,835 Then we compute, you know, gradients, 1108 01:12:29,635 --> 01:12:32,575 you know, dL by dW, and there are a lot of Ws here. 1109 01:12:32,595 --> 01:12:37,075 There's, you know, this projection W, there's, you know, the dot product W, there is a... 1110 01:12:37,815 --> 01:12:39,475 No, sorry, here there's no W, there's, 1111 01:12:40,455 --> 01:12:43,115 yeah, this, this W, there's, you know, tokenizing W. 1112 01:12:43,115 --> 01:12:46,215 But yeah, the most important weight matrices are actually this one, the one in the 1113 01:12:46,215 --> 01:12:48,875 projection step, and multi-headed attention. 1114 01:12:49,035 --> 01:12:53,155 When I describe multi-headed attention I'll talk to you about what this projection step 1115 01:12:53,175 --> 01:12:58,024 actually is.So, 1116 01:12:58,104 --> 01:13:00,983 anyways, there is a weight matrix that is required for this. 1117 01:13:00,983 --> 01:13:01,383 We will do 1118 01:13:02,344 --> 01:13:07,223 d loss by that W. We will then get, we will then update that 1119 01:13:07,344 --> 01:13:07,703 W. 1120 01:13:08,963 --> 01:13:11,803 We will set a good epsilon. Remember epsilon? 1121 01:13:12,043 --> 01:13:12,603 Uh, 1122 01:13:13,723 --> 01:13:17,764 epsilon is basically once we know the gradient, how much do we nudge our weight matrices, 1123 01:13:17,764 --> 01:13:20,403 right? Do we nudge our weights a little bit or do we nudge them a lot? 1124 01:13:21,024 --> 01:13:24,023 So, we will use Adam with weight decay to 1125 01:13:26,023 --> 01:13:28,003 decide how much to update our weights. 1126 01:13:29,243 --> 01:13:33,623 And then, we will spend hundreds of billions of dollars, uh, 1127 01:13:34,424 --> 01:13:39,243 running this gradient descent, just updating W, update W, update W, and you spend 1128 01:13:39,264 --> 01:13:40,923 hundreds of billions of dollars this way. 1129 01:13:42,083 --> 01:13:42,663 Uh, 1130 01:13:43,743 --> 01:13:47,703 yeah, this might distract if I bring this up right now. But actually, I'll just... 1131 01:13:55,683 --> 01:13:58,103 Is my note not working? What? Hmm. 1132 01:13:59,563 --> 01:14:02,643 Yeah, my note is fine. It's slow. 1133 01:14:03,283 --> 01:14:03,703 Okay. 1134 01:14:04,323 --> 01:14:07,503 (instrumental music plays) 1135 01:14:08,743 --> 01:14:13,323 Stargate, put that name down in your books, 'cause I think you're gonna hear a lot about 1136 01:14:13,343 --> 01:14:18,163 it in the future. A new American company that will invest 500 billion dollars 1137 01:14:18,263 --> 01:14:20,423 at least in AI infrastructure. 1138 01:14:21,143 --> 01:14:23,163 The data centers are actually under construction. 1139 01:14:23,363 --> 01:14:25,623 The first of them are under construction in Texas. 1140 01:14:26,163 --> 01:14:28,463 The Abilene location, which is our first location. 1141 01:14:29,283 --> 01:14:33,003 Abilene is a town on the western central plains of Texas. 1142 01:14:33,363 --> 01:14:36,323 We kinda scratch and fight for everything good that comes our way. 1143 01:14:36,363 --> 01:14:40,603 I had people reaching out and they said, "When is the President going to come see you?" 1144 01:14:40,923 --> 01:14:42,863 And I said, "Well, he hasn't texted me yet, so." 1145 01:14:43,003 --> 01:14:44,263 (laughs) 1146 01:14:44,263 --> 01:14:46,623 I'm thrilled we get to do this in the United States of America. 1147 01:14:46,903 --> 01:14:49,523 I think this will be the most important project of this era. 1148 01:14:51,963 --> 01:14:56,883 Welcome to Abilene, Texas. This is the site of Project 1149 01:14:56,963 --> 01:15:01,883 Stargate, a very mysterious, much talked about, much speculated about project, 1150 01:15:01,983 --> 01:15:02,783 that brings together- 1151 01:15:03,223 --> 01:15:06,483 Yeah, so by the way, this is a documentary on the 1152 01:15:07,963 --> 01:15:12,463 largest computer data center being built in the world, as of, you know, June 1153 01:15:12,523 --> 01:15:14,223 2025. Uh, 1154 01:15:16,623 --> 01:15:19,643 you know, one year later this data might be outdated, but as of today it's the biggest 1155 01:15:19,643 --> 01:15:21,343 one, uh, under construction. 1156 01:15:21,983 --> 01:15:26,583 For OpenAI they've committed like $500 billion over, I think, four years. 1157 01:15:26,923 --> 01:15:29,303 Like, the first year investment has gone through for sure. 1158 01:15:29,483 --> 01:15:31,103 The remaining years seems uncertain. 1159 01:15:32,543 --> 01:15:35,803 Uh, I only bring this up because the point is, you know, 1160 01:15:36,863 --> 01:15:41,723 we're spending all this money to ultimately update the weight matrix here, 1161 01:15:41,883 --> 01:15:44,103 basically. And remember, there are 80 of them. 1162 01:15:44,203 --> 01:15:46,463 Like, they could be 80, they could be more than 100. 1163 01:15:46,483 --> 01:15:46,783 Like, 1164 01:15:47,883 --> 01:15:51,483 we do this, you know, 100 times so there are 100 weight matrices. 1165 01:15:51,543 --> 01:15:54,903 We are going to keep updating those weight matrices 100 times. 1166 01:15:56,543 --> 01:15:58,283 That's what the data center is for. 1167 01:16:00,003 --> 01:16:04,703 So yeah, now let me talk about what are these weight matrices that we are going to 1168 01:16:04,723 --> 01:16:05,243 update 1169 01:16:06,403 --> 01:16:08,503 by spending billions of dollars on them. 1170 01:16:10,143 --> 01:16:10,583 Yeah. 1171 01:16:12,723 --> 01:16:13,663 So, this is 1172 01:16:14,323 --> 01:16:15,543 the attention block. 1173 01:16:16,763 --> 01:16:19,983 So, the most important part is that x is our input. You remember, right? 1174 01:16:19,983 --> 01:16:23,483 We started with images here, we're starting with text. So x is just an input. 1175 01:16:24,563 --> 01:16:28,943 It has a huge number of examples, and each example has, you know, an array of 1176 01:16:28,983 --> 01:16:29,603 numbers. 1177 01:16:30,823 --> 01:16:32,563 So, these are the tokens essentially. 1178 01:16:34,063 --> 01:16:37,663 We're going, our weight matrices are WQ, WK and WV. 1179 01:16:38,503 --> 01:16:42,663 We multiply x with each of these, we'll get, you know, three matrices, QKV. 1180 01:16:43,683 --> 01:16:46,923 Then we will do this thing and we will get y, which is our output. 1181 01:16:48,803 --> 01:16:51,983 And this is the most important step in a transformer. 1182 01:16:52,123 --> 01:16:56,103 That's why I'm going to explain it a lot in detail, because it's the most important. 1183 01:16:57,523 --> 01:17:01,283 So, I guess this is straightforward, you understood what's a prediction. 1184 01:17:02,323 --> 01:17:05,523 Now we are doing Q multiplied by K transpose. 1185 01:17:05,643 --> 01:17:08,223 So K transpose, we added a mask. 1186 01:17:09,623 --> 01:17:10,763 A mask 1187 01:17:11,403 --> 01:17:15,543 is a matrix that has some values which are hardcoded to zero, and 1188 01:17:17,243 --> 01:17:18,903 some values are hardcoded to one. 1189 01:17:19,923 --> 01:17:23,303 Actually, no, ma- mask has, a mask has some value zero, some values one. 1190 01:17:23,323 --> 01:17:24,823 So when you add a zero, nothing happens. 1191 01:17:24,823 --> 01:17:28,823 But when you add a one, this entire thing gets, like, locked into one. 1192 01:17:30,543 --> 01:17:30,903 I'll, 1193 01:17:31,523 --> 01:17:34,343 I'll explain to you what softmax is, then it will make sense. 1194 01:17:35,023 --> 01:17:36,903 So softmax is basically, 1195 01:17:37,983 --> 01:17:41,803 if you have a matrix, you want to do softmax on this matrix, it's like, for the first, 1196 01:17:42,223 --> 01:17:44,023 you will first do e to the power all the values. 1197 01:17:44,023 --> 01:17:46,943 So just all the values that are there, just replace them with e to the power of that 1198 01:17:46,943 --> 01:17:49,823 value. So you started with, you know, one, two, three, four. 1199 01:17:51,043 --> 01:17:55,743 Okay no, one, two, three, minus one. So e to the power one, e to the power two, e to the 1200 01:17:55,763 --> 01:17:57,243 power three, e to the power minus one. 1201 01:17:57,283 --> 01:18:00,103 So we just replaced all the values with e to the power of the same value. 1202 01:18:01,983 --> 01:18:05,203 Then we will divide each value with the sum of those rows. 1203 01:18:05,203 --> 01:18:05,443 So, 1204 01:18:07,003 --> 01:18:08,503 instead of one we now have 1205 01:18:09,103 --> 01:18:13,303 e to the power one upon e to the, divided by e to the power one plus e to the power two. 1206 01:18:14,203 --> 01:18:16,323 Like let's say our values in that row were one and two. 1207 01:18:17,663 --> 01:18:21,063 Now that two will become, you know, e to the power two divided by e to the power one plus 1208 01:18:21,063 --> 01:18:21,763 e to the power two. 1209 01:18:22,523 --> 01:18:24,683 On the next row we had what, three and minus one, right? 1210 01:18:24,683 --> 01:18:28,483 So three will become e to the power three divided by e to the power three plus e to the 1211 01:18:28,483 --> 01:18:29,203 power minus one. 1212 01:18:29,923 --> 01:18:31,283 Minus one will become, you know... 1213 01:18:32,423 --> 01:18:37,123 Yeah, maybe I can just give an example here that seems pretty easy to do, I should maybe 1214 01:18:37,243 --> 01:18:39,063 just actually give an example here 1215 01:18:40,103 --> 01:18:41,303 of softmax. 1216 01:18:49,023 --> 01:18:57,663 Softmax... 1217 01:18:58,443 --> 01:19:00,223 One, two, three, four, five, six, seven, 1218 01:19:00,223 --> 01:19:04,592 eight.This 1219 01:19:04,611 --> 01:19:05,231 one... 1220 01:19:08,931 --> 01:19:10,371 Equals... 1221 01:19:23,711 --> 01:19:24,872 Square. 1222 01:19:27,071 --> 01:19:28,651 By equaling square. 1223 01:19:29,531 --> 01:19:32,311 B square by B plus 1224 01:19:33,331 --> 01:19:34,491 A square. 1225 01:19:35,571 --> 01:19:37,471 And cube by... 1226 01:19:39,952 --> 01:19:42,331 Cube plus, you know, one by eight. 1227 01:19:44,092 --> 01:19:46,311 And cube by... 1228 01:19:48,952 --> 01:19:51,931 And cube plus one by, one by eight. (typing sound) Yeah. 1229 01:19:52,291 --> 01:19:56,092 (hums) (typing sound) Yeah. So here is an example. 1230 01:19:56,092 --> 01:19:56,092 Here, let me just call this an example. (hums) (typing sound) Yeah. Nevermind. 1231 01:19:56,092 --> 01:19:56,092 So here's an, here's an example. This is what softmax actually does. 1232 01:19:56,092 --> 01:19:59,472 If you want an intuition for what softmax 1233 01:20:00,231 --> 01:20:01,011 does, 1234 01:21:06,071 --> 01:21:10,751 basically the 1235 01:21:10,791 --> 01:21:14,091 large values in a row, they become close to one, the smaller values in a row, they become 1236 01:21:14,131 --> 01:21:14,851 close to zero. 1237 01:21:15,791 --> 01:21:18,611 If we had done row normalization, that's what would have happened anyway. 1238 01:21:18,711 --> 01:21:18,911 Like 1239 01:21:20,231 --> 01:21:22,531 if we had just done, you know, two divided by one plus two 1240 01:21:23,631 --> 01:21:26,871 (laughs) , basically the large values like two divided by one plus two is two by three. 1241 01:21:27,031 --> 01:21:30,371 So that's close to one and one divided by one plus two is, you know, 1242 01:21:31,311 --> 01:21:35,311 one by three that's close to zero. But we, because we did this E to the power, it's just 1243 01:21:35,351 --> 01:21:39,951 going to be that thing on steroids. So two will become even more close to one and this 1244 01:21:39,971 --> 01:21:41,571 one will become even more close to zero. 1245 01:21:42,211 --> 01:21:45,331 This three will go even more close to one, and this will go even more close to zero. 1246 01:21:46,371 --> 01:21:49,091 Like this thing is almost one basically, and 1247 01:21:49,811 --> 01:21:51,531 this thing is almost zero basically. 1248 01:21:52,691 --> 01:21:53,151 I mean I, 1249 01:21:53,871 --> 01:21:57,391 yeah maybe I'll just compute these values and leave it like why not? 1250 01:21:58,331 --> 01:22:00,411 I guess that's 0.268. 1251 01:22:09,071 --> 01:22:11,991 Square, 0.731. 1252 01:22:13,711 --> 01:22:15,171 0.731. 1253 01:22:24,771 --> 01:22:26,711 0.982. 1254 01:22:39,271 --> 01:22:40,131 Ah. 1255 01:22:45,591 --> 01:22:55,191 (hums) 1256 01:22:57,591 --> 01:22:59,071 0.208. 1257 01:23:01,371 --> 01:23:02,731 208. 1258 01:23:14,631 --> 01:23:15,651 Yeah, so... 1259 01:23:17,311 --> 01:23:18,551 Motherfucker. 1260 01:23:24,631 --> 01:23:29,611 Yes, so the one was already like a small value went close to zero, two 1261 01:23:29,631 --> 01:23:32,411 was big, it went close to one. Like even within the row. 1262 01:23:33,091 --> 01:23:35,951 Three went close to one, one minus one went really close to zero. 1263 01:23:37,011 --> 01:23:38,511 So softmax does this. 1264 01:23:39,411 --> 01:23:42,631 And remember before we do the softmax, we're going to add a mask. 1265 01:23:43,231 --> 01:23:44,931 So the mask basically says 1266 01:23:46,191 --> 01:23:48,231 some of these values we're just going to add 1267 01:23:49,131 --> 01:23:53,711 something so large that, well, only that gets paid attention to or something so 1268 01:23:53,791 --> 01:23:56,671 small, you know, that is not going to get paid attention to at all. 1269 01:23:57,831 --> 01:24:00,911 Basically, we are going to hard set some of these values to one and some of them to zero, 1270 01:24:01,031 --> 01:24:04,491 depending on what the mask says. I'll talk about what the mask says. 1271 01:24:04,551 --> 01:24:04,731 But 1272 01:24:06,011 --> 01:24:06,291 yeah, 1273 01:24:07,251 --> 01:24:09,711 so we're going to do soft mask of this entire thing. 1274 01:24:10,191 --> 01:24:12,311 Oh here we're just dividing a constant value. 1275 01:24:13,111 --> 01:24:17,171 The dimension of q, so if q is, you know, 1000 dimensions we're just going to do square 1276 01:24:17,211 --> 01:24:20,871 root 1000 which is 30, we're just going to divide all the values by 30. 1277 01:24:22,151 --> 01:24:25,451 It's not that important. But anyways we did softmax something 1278 01:24:26,151 --> 01:24:29,231 and then the softmax thing we are going to multiply with V. 1279 01:24:29,351 --> 01:24:31,251 So V is again X at WV. 1280 01:24:33,871 --> 01:24:36,051 So softmax something times V. 1281 01:24:37,731 --> 01:24:42,431 So this is basically telling us, like in this matrix V, we want to pay attention to some 1282 01:24:42,491 --> 01:24:46,331 parts. Like some parts of V are important and some parts of V are not important. 1283 01:24:47,091 --> 01:24:51,651 And this entire softmax thing is going to tell us what parts of V are important and what 1284 01:24:51,671 --> 01:24:52,931 parts of V are not important. 1285 01:24:54,171 --> 01:24:57,971 Oh, and by the way, you should always, always, always remember 1286 01:24:59,191 --> 01:24:59,691 this. 1287 01:25:00,631 --> 01:25:03,431 Deep learning is alchemy. Deep learning is black box. 1288 01:25:03,451 --> 01:25:08,435 We don't have mathematical proofs for anything.When I'm telling you, okay, you know, 1289 01:25:08,455 --> 01:25:12,355 here is all the reasons, these are all intuitions. We don't have proof for all this. 1290 01:25:13,516 --> 01:25:16,856 There might be some other secret reason why this thing is working well, which we have no 1291 01:25:16,856 --> 01:25:17,516 idea about. 1292 01:25:18,795 --> 01:25:22,835 Uh, I'm just giving some intuitions of, you know, why does this particular formula happen 1293 01:25:22,855 --> 01:25:23,695 to look so good. 1294 01:25:25,036 --> 01:25:26,175 You know, this, this thing. 1295 01:25:30,795 --> 01:25:31,876 Yeah. So, 1296 01:25:33,595 --> 01:25:35,915 this basically says... So, 1297 01:25:36,915 --> 01:25:38,116 we are going to 1298 01:25:38,875 --> 01:25:42,695 give a lot of weight to some parts of v and we're going to ignore other parts of v. 1299 01:25:44,396 --> 01:25:47,156 That's what softmax times means. So now putting this all together. 1300 01:25:47,895 --> 01:25:50,055 First, we, what, we multiplied with... 1301 01:25:50,955 --> 01:25:54,195 We multiplied Q and K transpose, we added a mask. 1302 01:25:54,716 --> 01:25:57,155 Mask said, "Okay, let's hide something." Some things are 1303 01:25:58,235 --> 01:26:01,535 as good as zero. We're just not going to consider information. 1304 01:26:02,295 --> 01:26:04,276 We are going to softmax it, which means 1305 01:26:05,295 --> 01:26:08,595 only some information we pay attention to, some information we don't pay attention to. 1306 01:26:08,615 --> 01:26:12,455 That's step two. Step three, we are going to do this on v. 1307 01:26:13,395 --> 01:26:16,795 So, at least some parts of v get paid attention to and some parts of v don't get paid 1308 01:26:16,815 --> 01:26:17,595 attention to. 1309 01:26:18,215 --> 01:26:20,515 So, when you put this all together, it basically says 1310 01:26:21,875 --> 01:26:24,355 we are hiding some parts of x, which is our input. 1311 01:26:24,575 --> 01:26:27,155 Like, rem- remember, QKV are not some magic things. 1312 01:26:27,175 --> 01:26:30,055 QKV are just X multiplied with different weight matrices. 1313 01:26:33,155 --> 01:26:35,655 So, we are hiding some part of X from ourselves. 1314 01:26:36,915 --> 01:26:38,255 We hide some part of X. 1315 01:26:39,455 --> 01:26:39,855 We 1316 01:26:42,295 --> 01:26:45,855 pay more attention to another part of X. Remember, this is also X. 1317 01:26:46,375 --> 01:26:48,935 Like, it's kind of like, you know, three incarnations of X. 1318 01:26:49,735 --> 01:26:50,195 Like... 1319 01:26:52,555 --> 01:26:56,195 Yeah, like this is, you know, first incarnation, this is the second incarnation, and the 1320 01:26:56,235 --> 01:26:57,375 third incarnation. So the 1321 01:26:58,755 --> 01:27:02,295 third... The first two incarnations of X, along with the mask, 1322 01:27:03,635 --> 01:27:07,175 or first two incarnations of X with some values masked out tell you which 1323 01:27:07,795 --> 01:27:11,195 parts of third incarnation of X to, you know, keep and which to discard. 1324 01:27:11,475 --> 01:27:12,475 It's kind of like this. 1325 01:27:16,615 --> 01:27:17,355 Yeah. And 1326 01:27:19,835 --> 01:27:22,235 then, we get an output. 1327 01:27:24,195 --> 01:27:27,435 And now remember, we don't actually know these weight matrices either. 1328 01:27:27,435 --> 01:27:29,515 We're going to find them using gradient descent. 1329 01:27:30,995 --> 01:27:35,575 We're going to start with some really bad values of WQ, WK and 1330 01:27:35,675 --> 01:27:38,915 WV. We are going to compute a loss at the end after, you know, 1331 01:27:40,035 --> 01:27:41,635 80 rounds of all this crap. 1332 01:27:42,675 --> 01:27:44,775 Uh, we're going to get a loss at the end. 1333 01:27:45,495 --> 01:27:50,155 We're going to compute the dL by dW. Like for all these weight matrices, we're going to 1334 01:27:50,175 --> 01:27:55,135 get a dL by dWq, dL by dWk, dL by dV. We're going to update 1335 01:27:55,135 --> 01:27:58,295 these weight matrices so they're like, are a little bit less bad. 1336 01:27:58,755 --> 01:28:02,775 We're going to do this, you know, trillions of times until these are actually pretty 1337 01:28:02,835 --> 01:28:03,135 good. 1338 01:28:04,675 --> 01:28:07,295 So ultimately, we want to find really good matrices here 1339 01:28:08,195 --> 01:28:12,455 that tell us basically in like exactly which situation, what part of X to pay attention 1340 01:28:12,475 --> 01:28:14,335 to, what part of X to not pay attention to. 1341 01:28:15,535 --> 01:28:17,635 And yeah, remember there are going to be 80 of these. 1342 01:28:17,635 --> 01:28:17,915 So, 1343 01:28:18,995 --> 01:28:23,895 we are going to do like 80 rounds of this thing with 80 different sets of, you know, 1344 01:28:24,055 --> 01:28:25,115 weight matrices. 1345 01:28:27,595 --> 01:28:31,115 Sorry, I think that's a lot of things, moving pieces, but yeah. 1346 01:28:31,115 --> 01:28:33,075 So, in one round we have 1347 01:28:34,455 --> 01:28:36,415 these three weight matrices, which we don't know. 1348 01:28:37,335 --> 01:28:41,255 We did this game about paying attention to some parts of our input. 1349 01:28:41,875 --> 01:28:45,975 We repeat this with another new set of three matrices with some other in- with, like, 1350 01:28:46,015 --> 01:28:49,375 that... The first round's output goes as the second thing's input. 1351 01:28:50,255 --> 01:28:53,215 And there we have another three weight matrices. We do again this game. 1352 01:28:53,855 --> 01:28:55,855 Then we do it again with another set of weight matrices. 1353 01:28:55,855 --> 01:28:58,035 Like this, we have, you know, 80 sets of matrices. 1354 01:28:58,535 --> 01:29:00,515 Finally, we get an output right at the end. 1355 01:29:00,855 --> 01:29:01,715 Then we compute 1356 01:29:02,395 --> 01:29:06,635 the gradients on all these 80 sets of matrices. We update all the 80 sets of matrices. 1357 01:29:06,635 --> 01:29:08,955 Then we come back, we redo this. 1358 01:29:09,915 --> 01:29:11,935 So, that's how this entire thing is working. 1359 01:29:26,355 --> 01:29:26,835 Yeah. 1360 01:29:30,375 --> 01:29:32,915 Yeah. Now, here are some intuitions on 1361 01:29:33,615 --> 01:29:35,755 why does this technique work in the first place? 1362 01:29:35,815 --> 01:29:36,375 Remember, 1363 01:29:37,175 --> 01:29:41,115 we don't actually know. This is guesswork. But our guesswork is first of all... 1364 01:29:42,195 --> 01:29:44,295 This entire thing can be done parallel. 1365 01:29:45,835 --> 01:29:50,555 Like otherwise, if, if you did not know anything about Transformers, you just knew, 1366 01:29:52,055 --> 01:29:56,035 uh, this like, okay, look, here is our input, our input is, you know, billions of tokens 1367 01:29:56,055 --> 01:29:57,515 for which we know the next token. 1368 01:29:58,355 --> 01:30:01,475 And now for some tokens we don't know, we have to predict the next token. 1369 01:30:02,375 --> 01:30:05,975 Like, you will typically formulate this as like a sequential problem, you know. 1370 01:30:07,075 --> 01:30:10,635 Like, you know, with weather data is kind of sequential or, you know, stock market data 1371 01:30:10,635 --> 01:30:12,075 is, you know, sequential. We'll put, you know, 1372 01:30:12,935 --> 01:30:17,775 X1, X2, X3, X4, X5, X6. And now, okay, given each of these, you know, pred X7s you'll do 1373 01:30:18,175 --> 01:30:20,895 some calculation on X1, you'll do some calculation on X2. 1374 01:30:21,375 --> 01:30:24,475 X is the first token, then something on the second token, you'll do some calculation on 1375 01:30:24,495 --> 01:30:28,175 the third token. You will go till the Nth token and then you'll try to predict, you know, 1376 01:30:28,195 --> 01:30:29,875 what the N plus one token, like... 1377 01:30:30,535 --> 01:30:33,955 Most algorithms, if you'll try on this type of problem, are sequential. 1378 01:30:37,075 --> 01:30:38,475 So, this could 1379 01:30:39,895 --> 01:30:44,015 be, you know, n-grams. Well, n-grams you can kind of parallelize too. 1380 01:30:44,095 --> 01:30:45,775 Okay, le- leave aside n-grams. But, 1381 01:30:46,815 --> 01:30:47,595 uh... 1382 01:30:49,855 --> 01:30:50,955 Where were we? Yeah. 1383 01:30:51,935 --> 01:30:55,655 Yeah. So RNNs were a very popular method called recurrent neural networks. 1384 01:30:56,095 --> 01:31:00,475 Now nobody uses RNNs, but once upon a time, RNNs were used, and 1385 01:31:01,555 --> 01:31:04,055 they do this type of sequential computation. 1386 01:31:04,075 --> 01:31:04,375 Like, 1387 01:31:05,175 --> 01:31:08,575 you have a weight matrix that does something on the first metoken, then weight matrix 1388 01:31:08,595 --> 01:31:11,375 does something on the second token, then something on the third token and so 1389 01:31:11,435 --> 01:31:15,463 on.Ah. 1390 01:31:16,583 --> 01:31:21,443 And the problem with training an RNN, especially a really big RNN, 1391 01:31:23,583 --> 01:31:24,163 is 1392 01:31:25,243 --> 01:31:27,063 you have to do, 1393 01:31:28,064 --> 01:31:31,043 like you have to update a lot of gradients in sequence. 1394 01:31:32,084 --> 01:31:36,484 First of all, that's slower, like GPUs can do computational parallel stuff faster 1395 01:31:37,144 --> 01:31:41,223 compared to doing a lot of sequential stuff and then do gradients across all of it. 1396 01:31:42,703 --> 01:31:47,463 And there is this vanishing and exploding gradients problem, which is a common problem if 1397 01:31:47,504 --> 01:31:49,243 you do a lot of machine learning, which is 1398 01:31:49,963 --> 01:31:53,724 sometimes these gradients become very s- close to zero and sometimes they become very 1399 01:31:53,763 --> 01:31:54,744 close to infinity. 1400 01:31:55,603 --> 01:31:59,604 And when that happens, you need to have some tricks up your sleeve to deal with it. 1401 01:32:00,063 --> 01:32:03,103 I mean ideally it should not happen, but if it happens you need to have some tricks to 1402 01:32:03,123 --> 01:32:03,843 deal with it. 1403 01:32:05,183 --> 01:32:05,823 And, 1404 01:32:07,103 --> 01:32:08,103 yeah, one trick 1405 01:32:08,883 --> 01:32:12,383 to make sure this happens less is just don't do a lot of sequential steps. 1406 01:32:12,383 --> 01:32:15,983 If you do a lot of sequential steps you're more likely to get gradients that go to zero 1407 01:32:16,003 --> 01:32:16,963 or go to infinity. 1408 01:32:23,543 --> 01:32:24,523 Yeah, so, 1409 01:32:26,323 --> 01:32:27,083 that's why 1410 01:32:28,403 --> 01:32:32,523 this attention we are doing it is good, because we use this concept of an attention mask 1411 01:32:32,543 --> 01:32:33,223 which just says 1412 01:32:33,823 --> 01:32:36,823 some parts of this we are going to hide now. Which parts of this we are going to hide? 1413 01:32:37,663 --> 01:32:41,183 That actually we can choose. Like we can choose what to put in this attention mask. 1414 01:32:41,203 --> 01:32:43,723 Like we can hide the last token and say like predict this. 1415 01:32:44,083 --> 01:32:46,483 We can hide the first token and say, you know, go predict this. 1416 01:32:47,083 --> 01:32:50,423 We can hide the fourth token and say, you know, now let's predict the fourth token. 1417 01:32:51,343 --> 01:32:55,323 And this, when we are training the model we can keep changing the attention mask as we 1418 01:32:55,343 --> 01:32:55,843 choose. 1419 01:32:57,523 --> 01:33:00,803 But this means now our entire problem is like parallelized, like 1420 01:33:01,483 --> 01:33:04,063 we had you know ten tokens, we hid the fifth token. 1421 01:33:04,063 --> 01:33:07,343 We know the first four and the last five and now we are predicting the fifth token. 1422 01:33:08,323 --> 01:33:11,643 And we will do another gradient descent to predict good values of the fifth token. 1423 01:33:13,723 --> 01:33:15,323 So this is parallelized and, 1424 01:33:16,963 --> 01:33:21,403 yeah, another good thing about this is we can dynamically decide what to pay attention 1425 01:33:21,463 --> 01:33:22,043 to, like 1426 01:33:22,643 --> 01:33:25,263 here we have lots of tokens, you know, which of these are important? 1427 01:33:25,263 --> 01:33:25,903 Like if I say 1428 01:33:26,703 --> 01:33:29,263 my name is Samuel what is 1429 01:33:30,343 --> 01:33:31,203 question mark. 1430 01:33:32,083 --> 01:33:36,863 Now what is might be more important there and maybe Samuel is not important there. 1431 01:33:36,863 --> 01:33:37,163 Like, 1432 01:33:38,203 --> 01:33:41,143 if it was my name is Alice, your name, 1433 01:33:42,083 --> 01:33:42,423 sorry, 1434 01:33:43,323 --> 01:33:47,023 what is question mark, you know, the answer is probably the same, you know, that being 1435 01:33:47,043 --> 01:33:48,743 Samuel or Alice makes no difference. 1436 01:33:49,783 --> 01:33:53,943 But what is, that is is very important because it says okay here is a verb, now after 1437 01:33:53,963 --> 01:33:56,963 this verb you can maybe put a preposition, you can put a noun. 1438 01:33:57,243 --> 01:33:59,063 You can't put another verb, that's for sure. 1439 01:33:59,923 --> 01:34:01,083 These types of rules 1440 01:34:01,983 --> 01:34:06,363 you can reason about. So you need to know which parts of this sentence (coughs) , 1441 01:34:06,823 --> 01:34:09,063 and by the way right now I gave you just six words, right? 1442 01:34:09,103 --> 01:34:11,403 In practice they could be in the thousands of words. 1443 01:34:12,503 --> 01:34:12,923 Like, 1444 01:34:14,343 --> 01:34:18,603 like remember O3s here are giving you thousands of words of output, so it needs to be 1445 01:34:18,623 --> 01:34:20,563 able to do this reasoning on thousands of words. 1446 01:34:20,563 --> 01:34:20,783 So 1447 01:34:21,683 --> 01:34:24,943 it needs to be able to dynamically pay attention to what is important here. 1448 01:34:25,603 --> 01:34:29,763 Now some older methods like CNNs, convolutional neural networks were like a really 1449 01:34:29,783 --> 01:34:31,163 popular method. Again, 1450 01:34:32,043 --> 01:34:35,063 extinct now. Nobody uses CNNs anymore, but 1451 01:34:36,403 --> 01:34:40,843 like at most today they, people do like one CNN layer to convert your image into an 1452 01:34:40,863 --> 01:34:44,423 embedding and then they fit that into your transformer. 1453 01:34:44,543 --> 01:34:48,123 But yeah, nobody trains like end-to-end, you know, CNNs anymore. 1454 01:34:48,323 --> 01:34:48,563 (coughs) 1455 01:34:50,523 --> 01:34:50,843 So 1456 01:34:52,323 --> 01:34:56,443 CNNs basically maintain these sliding windows and like window, window, window, window. 1457 01:34:57,163 --> 01:34:59,263 If you have multiple layers you have multiple windows 1458 01:35:00,243 --> 01:35:04,443 that decide, okay, now you have to pay attention to this part of the sequence, then we 1459 01:35:04,443 --> 01:35:07,163 will pay attention to this, then we will take a smaller window there. 1460 01:35:07,923 --> 01:35:10,423 Like if you study CNNs you will understand what I am talking about. 1461 01:35:10,923 --> 01:35:14,163 But CNNs kind of fi- hard code this thing what to pay attention to. 1462 01:35:15,063 --> 01:35:17,703 LSTMs and RNNs, uh, 1463 01:35:18,843 --> 01:35:22,483 maintain one context across the entire thing which is kind of bad. 1464 01:35:22,703 --> 01:35:23,183 Uh, 1465 01:35:25,263 --> 01:35:26,643 so it's like if you had to do 1466 01:35:27,823 --> 01:35:30,743 my name is Samuel what is question mark. 1467 01:35:31,523 --> 01:35:33,563 Let us say the context said okay 1468 01:35:34,883 --> 01:35:38,303 what is, is paid more attention to and 1469 01:35:39,023 --> 01:35:41,383 name is paid more attention to. Cool. 1470 01:35:41,843 --> 01:35:46,523 Now, okay we predicted the next token, we said next token is your. 1471 01:35:46,923 --> 01:35:51,163 Now again we want to predict, okay, my name is Samuel, what is your 1472 01:35:51,643 --> 01:35:56,283 question mark. Now again attention we will say okay name important, what is 1473 01:35:56,323 --> 01:35:57,723 important, but actually now 1474 01:35:58,383 --> 01:36:01,003 what is is not important. Name is very important because 1475 01:36:01,843 --> 01:36:05,963 my name is Samuel, what is your question mark. It's probably going to be name, right? 1476 01:36:06,803 --> 01:36:11,323 So we need to be able to now dynamically change, okay, now we don't want to pay attention 1477 01:36:11,363 --> 01:36:14,003 to those things, we want to pay attention to a different thing like 1478 01:36:15,563 --> 01:36:20,043 we need to be able to change this context like even as we go thr- through each token we 1479 01:36:20,063 --> 01:36:23,463 need to be changing our context of what is important to pay attention to. 1480 01:36:24,143 --> 01:36:28,783 And LSTMs and RNNs can't do that as well as a transformer can. 1481 01:36:30,083 --> 01:36:32,863 So yeah, that's another reason transformers are good. 1482 01:36:39,003 --> 01:36:42,223 Yeah. Here is some more, you know, miscellaneous stuff. 1483 01:36:42,323 --> 01:36:42,643 (coughs) 1484 01:36:43,563 --> 01:36:44,843 In the transformer 1485 01:36:46,083 --> 01:36:49,483 I guess tokenization again already briefly talked about 1486 01:36:50,463 --> 01:36:55,143 you convert words into numbers. Positional embeddings is basically instead of 1487 01:36:55,183 --> 01:36:55,463 saying 1488 01:36:56,123 --> 01:37:01,003 my name is Samuel it becomes my one name two is three Samuel four, and 1489 01:37:01,023 --> 01:37:03,043 actually we don't store one two three four we 1490 01:37:03,743 --> 01:37:07,783 do like cosine one cosine two cosine three cosine four or something which is 1491 01:37:09,123 --> 01:37:13,583 yeah a bit unusual but that's the thing again we tried it, it works you know here's 1492 01:37:14,403 --> 01:37:16,403 one more in your bag of tricks to 1493 01:37:16,583 --> 01:37:20,995 use.The root mean square norm is very 1494 01:37:21,095 --> 01:37:23,116 straightforward. It basically says, 1495 01:37:25,055 --> 01:37:27,095 like we have done a lot of normalization, right? 1496 01:37:27,095 --> 01:37:31,756 We did the normal row normalization, which is divide each cell by 1497 01:37:32,275 --> 01:37:34,975 some of all cells in the row. We did soft max which is, you know, 1498 01:37:35,955 --> 01:37:40,155 do the same thing but E to the power, you know, E to the power cell divided by some of E 1499 01:37:40,155 --> 01:37:41,896 to the power all cells in the row. 1500 01:37:43,095 --> 01:37:46,255 So RMS is yet another... RMS is a divide 1501 01:37:47,495 --> 01:37:51,835 each cell by sum of squares of all the cells in the row. 1502 01:37:51,895 --> 01:37:54,056 This is just yet another way of normalizing. 1503 01:37:55,715 --> 01:37:57,695 Yeah, we do a lot of normalizing 1504 01:37:58,615 --> 01:38:02,055 partly because of this, you know, vanishing, exploding gradient problem. 1505 01:38:02,055 --> 01:38:06,575 Like, we don't want gradients to ever go to zero or infinity, therefore we need to keep 1506 01:38:06,575 --> 01:38:08,015 normalizing stuff all the time. 1507 01:38:11,855 --> 01:38:15,076 Yeah, and this will ensure our values always stay between zero and one. 1508 01:38:16,275 --> 01:38:18,475 Yeah, then there is this thing called residual layer. 1509 01:38:18,476 --> 01:38:20,595 Again, maybe I will go back to the big picture. 1510 01:38:21,495 --> 01:38:25,715 The big picture is, yeah, we tokenized, we put the positions, we put our images into 1511 01:38:25,755 --> 01:38:27,455 something that's not images. So from, 1512 01:38:28,435 --> 01:38:31,175 once you have reached here you can forget the fact that they're images. 1513 01:38:31,375 --> 01:38:32,095 We just now have 1514 01:38:32,895 --> 01:38:36,075 numbers, you know. We have all these vectors coming in, that's it. 1515 01:38:36,575 --> 01:38:38,075 So we did, you know, normalizing, 1516 01:38:38,775 --> 01:38:43,455 we did our, you know, projection into X, we made into QKV and that 1517 01:38:43,535 --> 01:38:45,655 QKV we made back into an output. 1518 01:38:46,275 --> 01:38:47,835 We got some output at the end. 1519 01:38:49,075 --> 01:38:50,995 Uh, this projection is straightforward, we just... 1520 01:38:51,415 --> 01:38:52,955 Yeah, this slide I don't think we need to explain, it's 1521 01:38:53,915 --> 01:38:55,635 just whatever output we got from here, 1522 01:38:56,875 --> 01:39:00,055 here meaning our soft max times V, we got some output. 1523 01:39:00,335 --> 01:39:03,975 We got a Y, we multiplied that by, you know, yet another ma- matrix and we got the 1524 01:39:03,995 --> 01:39:08,075 output. So from here we will multiply one more weight matrix, we will get the output. 1525 01:39:08,075 --> 01:39:10,035 So here is our multi-headed attention done. 1526 01:39:10,995 --> 01:39:12,575 Now we do a thing called add residual. 1527 01:39:12,575 --> 01:39:17,195 Now add residual basically means imagine from here we had X, we did multi-headed 1528 01:39:17,215 --> 01:39:21,095 attention, we got Y. Now instead of taking Y, we will take Y plus X. 1529 01:39:22,515 --> 01:39:25,295 Right, we will take our input, we will add it to the output and we will keep that. 1530 01:39:25,295 --> 01:39:28,535 And it is a little bit of a weird trick but it works, we do this in machine learning all 1531 01:39:28,575 --> 01:39:28,995 the time. 1532 01:39:30,035 --> 01:39:32,775 Again just, it works therefore we do it. 1533 01:39:32,775 --> 01:39:33,215 Uh, 1534 01:39:38,775 --> 01:39:39,575 yeah, so 1535 01:39:40,335 --> 01:39:42,695 Y equals FX was done in the previous step. 1536 01:39:43,275 --> 01:39:46,235 So now we will take FX plus X and we will keep this. 1537 01:39:47,835 --> 01:39:50,335 And why? Our intuition is basically, 1538 01:39:51,235 --> 01:39:55,275 remember we are going to be a- doing this entire thing on like 80 layers, we are going to 1539 01:39:55,315 --> 01:39:58,235 do this entire 80 times with 80 sets of eight matrices. 1540 01:39:59,295 --> 01:40:00,795 So if you had just done, you know, 1541 01:40:02,215 --> 01:40:06,515 F of F of F of F of F of F of F of X like 80 times and then you would do, run 1542 01:40:07,035 --> 01:40:08,995 gradient descent on this to find you know the 1543 01:40:10,275 --> 01:40:12,695 best weight matrices for each of those Fs. 1544 01:40:12,815 --> 01:40:15,055 Like these are all 80 different Fs by the way, like the 1545 01:40:17,215 --> 01:40:19,915 weight matrices on, inside each of these Fs are different. 1546 01:40:20,915 --> 01:40:23,555 So if you had to do this you are more like again... 1547 01:40:23,595 --> 01:40:27,555 Remember sequential is bad for exploding and vanishing gradients. 1548 01:40:28,335 --> 01:40:29,775 That is a heuristic. 1549 01:40:31,415 --> 01:40:31,915 Uh, 1550 01:40:33,335 --> 01:40:37,895 therefore we are more likely to get exploding or vanishing gradients however if 1551 01:40:37,935 --> 01:40:38,895 sometimes we 1552 01:40:40,375 --> 01:40:40,715 do, 1553 01:40:41,435 --> 01:40:41,695 like, 1554 01:40:42,315 --> 01:40:47,015 we did F of X plus X, F of this thing plus 3000 thing, F of this thing plus 3000 thing, F 1555 01:40:47,035 --> 01:40:48,375 of this thing plus 3000 thing, 1556 01:40:48,975 --> 01:40:50,135 and then now we... 1557 01:40:50,815 --> 01:40:53,375 And each of these Fs have their own weight matrices. 1558 01:40:53,395 --> 01:40:56,115 Now we run unigradian descent across this entire thing 1559 01:40:57,295 --> 01:40:58,535 then, uh, 1560 01:41:03,815 --> 01:41:04,095 yeah, 1561 01:41:06,115 --> 01:41:08,575 then we are more likely to get, uh, 1562 01:41:11,855 --> 01:41:13,715 good, uh, weight matrices 1563 01:41:14,935 --> 01:41:18,575 at the end. Like we are going to get gradients that are not going to zero or infinity, we 1564 01:41:18,575 --> 01:41:19,575 are going to get... 1565 01:41:23,955 --> 01:41:27,975 Yeah, I guess you got what I said. Okay, then there is thing called mixture of experts 1566 01:41:27,975 --> 01:41:32,815 which I am not describing now. It's kind of a recent trick like two years ago nobody did 1567 01:41:32,855 --> 01:41:34,575 this, now everyone is doing it. 1568 01:41:35,475 --> 01:41:39,195 And because it's a recent trick also there's a possibility it will not get used in the 1569 01:41:39,235 --> 01:41:40,175 future, like 1570 01:41:40,815 --> 01:41:44,295 maybe it will and maybe it will not, it's just I don't know how to predict that. 1571 01:41:45,095 --> 01:41:48,215 But yeah this is a common trick that has come up in the last couple of years. 1572 01:41:49,835 --> 01:41:53,335 Uh, feed forward is pretty straightforward, you can search what feed forward is. 1573 01:41:53,335 --> 01:41:55,655 Feed forward typically nowadays is just, 1574 01:41:56,335 --> 01:41:59,095 just fully connected network in our two layer. Remember this thing? 1575 01:42:00,315 --> 01:42:04,515 Yeah, so this entire thing is like one tiny tiny piece in our entire big, giant big 1576 01:42:04,555 --> 01:42:05,235 transformer. 1577 01:42:07,015 --> 01:42:11,755 Yeah, so we put here, you know, we put two weight matrices, our W1 and W2 they go in here 1578 01:42:12,515 --> 01:42:13,895 then we add a residual again. 1579 01:42:15,075 --> 01:42:18,415 We do this entire thing 80 times, we do one last projection step. 1580 01:42:18,415 --> 01:42:22,255 Projection just means we multiply it with yet another matrix and we do a ReLU. 1581 01:42:23,195 --> 01:42:27,095 Oh yeah, I think I forgot to mention we do a lot of ReLUs like, 1582 01:42:30,235 --> 01:42:34,255 like when we do residual we will probably do a ReLU, when we do projection we will do a 1583 01:42:34,255 --> 01:42:37,035 ReLU, when we do this projection we will do a ReLU like, 1584 01:42:38,375 --> 01:42:43,315 like a ReLU just happens so often it's not kind of worth even mentioning like ReLU is 1585 01:42:43,315 --> 01:42:46,155 that common. We just do ReLUs everywhere. 1586 01:42:51,195 --> 01:42:51,435 Yeah, 1587 01:42:53,175 --> 01:42:58,095 so we have done this lots and lots of computation and then we are going to do 1588 01:42:58,975 --> 01:43:00,275 gradient descent, 1589 01:43:04,255 --> 01:43:08,795 i- and we are going to find our weight matrices. I think that's basically the overview. 1590 01:43:10,995 --> 01:43:15,935 There's one final topic that I need to cover or I want to cover which is the 1591 01:43:15,935 --> 01:43:17,735 hyper-parameters which is like 1592 01:43:18,475 --> 01:43:20,575 how big should this entire thing be, you know. 1593 01:43:21,195 --> 01:43:26,179 How many layers do we pick? How big our weight matrices do we want to pick?And these 1594 01:43:26,199 --> 01:43:28,039 are just typical values in the past. 1595 01:43:29,900 --> 01:43:33,459 So number of layers is usually between 32 and 128. 1596 01:43:34,199 --> 01:43:36,619 How big are weight matrices is, you know, so 1597 01:43:37,239 --> 01:43:42,020 GPT-2 was the first good transformer that was trained back in 2019. 1598 01:43:43,579 --> 01:43:45,819 It had 1.5 billion 1599 01:43:46,539 --> 01:43:48,559 parameters. So parameters is just 1600 01:43:49,479 --> 01:43:51,639 how many numbers are in your weight matrix. 1601 01:43:51,760 --> 01:43:52,079 So 1602 01:43:52,920 --> 01:43:54,619 let's go way, way, way back. 1603 01:43:56,139 --> 01:43:56,779 So this 1604 01:43:57,620 --> 01:44:00,200 two-layer network we trained, this has, 1605 01:44:01,239 --> 01:44:01,679 uh, 1606 01:44:02,819 --> 01:44:04,799 so the total number of parameters in this 1607 01:44:06,399 --> 01:44:10,559 network are 784 times 800, that's W1, plus, 1608 01:44:11,280 --> 01:44:13,059 you know, 800 times 1609 01:44:13,920 --> 01:44:17,639 10. Yeah, so we, this is around 600,000 parameters, 1610 01:44:19,299 --> 01:44:24,059 635,000 parameter. So this is a 635,000 parameter network. 1611 01:44:24,119 --> 01:44:29,019 So GPT-2 is a 1.5 billion parameter network. 1612 01:44:30,279 --> 01:44:34,719 There are going to be 1.5 billion numbers inside the network, 1613 01:44:35,399 --> 01:44:39,019 and we are going to do gradient descent to find, you know, gradients for each of these 1614 01:44:39,039 --> 01:44:40,639 1.5 billion numbers, 1615 01:44:41,459 --> 01:44:44,739 and update them a little bit, and update them a little bit more, update them a little bit 1616 01:44:44,779 --> 01:44:46,159 more. We'll keep doing that. 1617 01:44:47,939 --> 01:44:48,379 So, 1618 01:44:49,899 --> 01:44:50,159 yeah, 1619 01:44:51,379 --> 01:44:56,259 GPT-2 is 1.5 billion parameters. GPT-3 we made a big jump to 175 1620 01:44:56,279 --> 01:44:57,499 billion parameters. 1621 01:44:58,159 --> 01:45:00,919 GPT-4 is around two trillion parameters. 1622 01:45:00,939 --> 01:45:02,459 You know, our latest models are now, 1623 01:45:03,079 --> 01:45:05,679 they've crossed 10 trillion parameters. 1624 01:45:08,139 --> 01:45:08,399 Yeah. 1625 01:45:09,479 --> 01:45:12,639 Uh, this is a Mixture of Experts thing. I will skip that for now. 1626 01:45:14,699 --> 01:45:18,319 So if you're using O3, it's most likely based on GPT-4. 1627 01:45:18,959 --> 01:45:20,999 So O3 is, you know, O3 is this thing. 1628 01:45:20,999 --> 01:45:25,719 How 1629 01:45:25,979 --> 01:45:27,719 to bake chocolate 1630 01:45:28,959 --> 01:45:29,339 cake, 1631 01:45:31,699 --> 01:45:32,799 no egg. 1632 01:45:44,999 --> 01:45:45,879 Yeah. So 1633 01:45:47,359 --> 01:45:50,759 this is, uh, O3, this is based on GPT-4. 1634 01:45:52,119 --> 01:45:54,379 I mean, it's not just straightforward GPT-4. 1635 01:45:54,379 --> 01:45:55,019 We have done some 1636 01:45:55,759 --> 01:45:58,979 shenanigans on top at the end with somebody else. 1637 01:45:58,979 --> 01:46:00,639 We'll have to talk about in a different lecture. 1638 01:46:00,839 --> 01:46:01,099 But 1639 01:46:01,799 --> 01:46:03,079 yeah, it's based on GPT-4. 1640 01:46:04,899 --> 01:46:09,179 It has two trillion parameters which we found using this gradient descent. 1641 01:46:09,259 --> 01:46:09,559 Uh, 1642 01:46:11,739 --> 01:46:14,379 each parameter is typically stored as four bytes. 1643 01:46:15,579 --> 01:46:19,279 So when we say two trillion parameters, what we actually mean is eight terabytes. 1644 01:46:20,679 --> 01:46:25,419 You know, each byte, each parameter is four bytes so 2T into 4 is 8T. 1645 01:46:26,679 --> 01:46:28,219 8 terabytes, uh, 1646 01:46:29,679 --> 01:46:33,479 nowadays there are also tricks being used, like mixed precision training was used in 1647 01:46:33,519 --> 01:46:37,259 DeepSeek, which is instead of four-byte weights, we will use, you know, two-byte weights 1648 01:46:37,379 --> 01:46:40,139 somewhere and somewhere we will use four bytes. 1649 01:46:42,079 --> 01:46:43,159 Like for some parts of the 1650 01:46:43,859 --> 01:46:47,619 network we will do gradients on four-byte weights and some parts of the networks we will 1651 01:46:47,659 --> 01:46:51,659 do gradient descent on two-byte weights, and apparently that also works. 1652 01:46:53,359 --> 01:46:57,779 So the basic thing is if you do gradient descent on smaller, like use 1653 01:46:58,359 --> 01:46:59,799 less bytes for the weights, 1654 01:47:00,559 --> 01:47:04,499 you will, you can calculate it faster but also you'll get less accurate answers. 1655 01:47:05,699 --> 01:47:09,359 So for training, I mean typically people still use the entire full four bytes. 1656 01:47:09,759 --> 01:47:12,619 For inference, inference is, you know, once you've trained the model. 1657 01:47:13,119 --> 01:47:17,159 Train means once you've got really good weight matrices, now you just want to do this on 1658 01:47:17,159 --> 01:47:17,999 some new answer. 1659 01:47:19,719 --> 01:47:22,939 Like you have trained this on our existing data, and on some new data you need to try it 1660 01:47:22,979 --> 01:47:27,579 out. When you're trying it out on new data, then people do, you know, two bytes, they 1661 01:47:27,579 --> 01:47:29,719 take one byte weight, you know, four-bit weights. 1662 01:47:29,719 --> 01:47:33,739 We have, even have now like literally less than two-bit weights. 1663 01:47:33,759 --> 01:47:34,039 Like 1664 01:47:34,859 --> 01:47:37,859 there's this thing called ternary bits which is like, uh, 1665 01:47:39,819 --> 01:47:40,859 you round a number, 1666 01:47:42,799 --> 01:47:45,439 uh, I forget it, I'm not explaining ternary bits right now. 1667 01:47:45,439 --> 01:47:48,539 But yeah, we can r- make each weight really, really small. 1668 01:47:49,719 --> 01:47:54,219 We can make it even like as small as two bits per weight or even less than two bits, so 1669 01:47:54,219 --> 01:47:55,719 it's 1.58 bits. 1670 01:47:58,759 --> 01:47:59,299 Uh, 1671 01:48:00,379 --> 01:48:01,799 so this is our model size. 1672 01:48:02,639 --> 01:48:04,579 Yeah, and how did we pick these numbers? 1673 01:48:04,579 --> 01:48:09,099 Like who, why did we take the decision of GPT-2 should be this big, GPT-3 should be this 1674 01:48:09,119 --> 01:48:09,479 big? 1675 01:48:10,259 --> 01:48:12,919 It depends on how much data and how much compute we have. 1676 01:48:13,519 --> 01:48:14,859 So data is our input, 1677 01:48:15,799 --> 01:48:19,559 like we will train this on some, like initially, you know, we started with 60,000 images. 1678 01:48:19,659 --> 01:48:23,439 So here for text we'll start, you know, some few million sentences or some trillion 1679 01:48:23,459 --> 01:48:24,219 sentences. 1680 01:48:25,859 --> 01:48:29,499 So how much of that do we have and how much compute do we have, which is just how many 1681 01:48:29,559 --> 01:48:30,519 GPUs do we have? 1682 01:48:32,299 --> 01:48:37,179 And you can calculate the requirements using this law called the Chinchilla Scaling 1683 01:48:37,219 --> 01:48:37,499 Law. 1684 01:48:38,559 --> 01:48:42,679 Uh, I'm not explaining Chinchilla Scaling Law right now, but you can read it on 1685 01:48:42,699 --> 01:48:44,579 Wikipedia, it's pretty straightforward. 1686 01:48:46,399 --> 01:48:46,759 Like, 1687 01:48:48,059 --> 01:48:52,579 C is how much computation you need, n is the number of parameters, 1688 01:48:53,199 --> 01:48:57,219 b is the size of your dataset, and there's like a straightforward relationship between 1689 01:48:57,219 --> 01:48:57,779 all of them. 1690 01:48:59,339 --> 01:49:02,919 So given some amount of compute, it just tells you, look here's how much. 1691 01:49:04,179 --> 01:49:04,499 Like, 1692 01:49:05,919 --> 01:49:09,339 given some amount of compute, given some amount of data, here's the loss you're going to 1693 01:49:09,379 --> 01:49:10,099 get at the end. 1694 01:49:10,819 --> 01:49:14,539 And using those you can reverse calculate, okay, well if you want to optimally use that 1695 01:49:14,559 --> 01:49:17,459 compute, you know, how much compute and how much data should be used at once. 1696 01:49:18,479 --> 01:49:23,059 You will calculate that. And yeah, an important thing is that usually compute is the 1697 01:49:23,079 --> 01:49:28,027 bottleneck, not data.Like usually we have 1698 01:49:28,207 --> 01:49:29,547 more data than we actually need 1699 01:49:30,347 --> 01:49:32,587 and less GPUs than we actually need 1700 01:49:33,307 --> 01:49:35,147 (laughs) . So more GPUs, bigger model. 1701 01:49:35,988 --> 01:49:39,607 So you know, this first train on a smaller data center, maybe for an hour. 1702 01:49:39,647 --> 01:49:44,047 This maybe was trained in hundreds of hours. Not hundreds, sorry, thousands of hours. 1703 01:49:44,967 --> 01:49:48,588 Here, by the time we are reaching here, we are running into you know, millions of GPU 1704 01:49:48,588 --> 01:49:49,127 hours. 1705 01:49:52,387 --> 01:49:52,727 So 1706 01:49:54,147 --> 01:49:58,987 yeah, it's just to train a bigger model, in short, you need more compute and more data, 1707 01:49:59,188 --> 01:49:59,528 but 1708 01:50:00,348 --> 01:50:01,988 we have data and we don't have compute. 1709 01:50:01,988 --> 01:50:05,347 So the more compute you have, the bigger model you will get to train, basically. 1710 01:50:09,147 --> 01:50:09,627 And 1711 01:50:10,688 --> 01:50:12,547 this is just if you want to train it well. 1712 01:50:12,627 --> 01:50:12,967 Like 1713 01:50:13,807 --> 01:50:18,447 remember when we were here, when we defined, we never said how many times to do a 1714 01:50:18,447 --> 01:50:19,627 gradient update. Like 1715 01:50:20,427 --> 01:50:23,848 we could just update the weight matrices, you know, 1000 times and be like, look here's 1716 01:50:23,848 --> 01:50:27,347 our performance. It's just this performance won't be that good. 1717 01:50:27,847 --> 01:50:30,607 If you want the performance to be really good, you will have to update it, you know, 1718 01:50:30,607 --> 01:50:32,147 millions and millions of times. And 1719 01:50:33,527 --> 01:50:36,407 then there's a question of, you know, how many GPUs do you have? 1720 01:50:36,407 --> 01:50:39,187 How many years can you run these GPUs for and so on. 1721 01:50:41,147 --> 01:50:45,867 And yeah, Epoch AI is this nonprofit... 1722 01:50:46,707 --> 01:50:48,827 Are they nonprofit or for profit? I'm not sure. 1723 01:50:48,907 --> 01:50:49,247 Uh, 1724 01:50:50,827 --> 01:50:53,967 let's say it's once in 2028, we'll run out of data. 1725 01:50:54,047 --> 01:50:54,327 Like 1726 01:50:54,947 --> 01:50:58,127 literally every... all the data on the internet, you know, 1727 01:50:59,247 --> 01:51:02,867 will have been fed into a model and there's no more data out there 1728 01:51:04,107 --> 01:51:05,467 and then we'll run out of data. 1729 01:51:06,147 --> 01:51:08,427 And there are some tricks that might bypass that too. 1730 01:51:09,127 --> 01:51:12,527 I- I- My personal guess is some of those tricks will work, but, 1731 01:51:13,587 --> 01:51:16,667 uh, yeah, we are going to run out of data in 2028. 1732 01:51:16,707 --> 01:51:17,027 And 1733 01:51:17,807 --> 01:51:20,667 compute, we're, we're never going to run out of compute. 1734 01:51:24,047 --> 01:51:27,987 That's just a question of, you know, like what fraction of the world economy do you want 1735 01:51:28,007 --> 01:51:30,187 to invest into, you know, GPUs? And 1736 01:51:31,627 --> 01:51:34,967 GPUs keep getting faster. So you know, every 18 months the- 1737 01:51:35,767 --> 01:51:36,467 you can get 1738 01:51:37,347 --> 01:51:38,627 twice the number of 1739 01:51:39,607 --> 01:51:41,647 matrix multiplications per dollar. 1740 01:51:42,947 --> 01:51:47,067 So the number of dollars being spent is going up and how many matrix multiplications you 1741 01:51:47,087 --> 01:51:48,827 can get per dollar is also going up. 1742 01:51:49,767 --> 01:51:53,987 And that seems the trajectory at least for, you know, the next five years, ten years. 1743 01:51:54,107 --> 01:51:54,347 So 1744 01:51:55,087 --> 01:51:57,087 yeah, you can see more on the forecasts 1745 01:51:58,807 --> 01:52:01,467 for how much data and how much compute we will use 1746 01:52:02,687 --> 01:52:06,907 in the next, you know, five years here. Like Epoch has forecast. 1747 01:52:07,927 --> 01:52:09,767 Ah, my net is fucking slow. 1748 01:52:12,867 --> 01:52:17,467 Mm, yeah, so our compute spending is- we're doing 5x every year. 1749 01:52:17,467 --> 01:52:19,367 So it's like exponentially going up. 1750 01:52:20,487 --> 01:52:22,347 Data runs out in 2028. 1751 01:52:25,687 --> 01:52:27,947 We are spending three times as many dollars 1752 01:52:28,907 --> 01:52:30,147 on training every year. 1753 01:52:31,107 --> 01:52:33,267 So it's first we start with, you know, like 1754 01:52:33,927 --> 01:52:38,447 $10 million, then $30 million, $100 million, $30 million. 1755 01:52:38,647 --> 01:52:40,807 Sorry, $300 million, a billion dollars. 1756 01:52:41,467 --> 01:52:44,667 Right now I think Grok is trained on seven billion, right? Yeah. 1757 01:52:45,527 --> 01:52:46,307 So Grok, 1758 01:52:47,647 --> 01:52:51,387 trained by Elon Musk, required a data center that cost $7 1759 01:52:51,507 --> 01:52:52,667 billion. 1760 01:52:54,527 --> 01:52:57,007 Half a billion dollars was spent on Grok itself. 1761 01:52:58,547 --> 01:53:02,727 Uh, but just to acquire that ha-hardware, like once you have acquired it, what else are 1762 01:53:02,747 --> 01:53:06,107 you going to use it for? It- you're going to just use it either for training or for 1763 01:53:06,147 --> 01:53:06,947 inference. So 1764 01:53:08,627 --> 01:53:13,107 this is the cost to acquire all that infrastructure, including the GPUs, including the 1765 01:53:13,147 --> 01:53:15,147 networking, the cooling and so on. 1766 01:53:17,807 --> 01:53:18,427 Uh, 1767 01:53:20,087 --> 01:53:20,487 yeah. 1768 01:53:23,687 --> 01:53:27,107 I think that's pretty much it for my presentation. 1769 01:53:28,867 --> 01:53:29,387 Uh, 1770 01:53:31,507 --> 01:53:34,967 yeah here are some rules of thumb. If you don't want to actually sit and solve the 1771 01:53:34,987 --> 01:53:38,507 Chinchilla scaling thing directly, here are some rules of thumb that say 1772 01:53:39,767 --> 01:53:42,487 for these many parameters, here's how much data you need. 1773 01:53:42,507 --> 01:53:46,687 And for these many parameters and this much data, here's how much compute you need. 1774 01:53:48,287 --> 01:53:50,927 So yeah, I think that's basically it for my presentation.