Let's reproduce GPT-2 (124M) — Transcript
Full transcript
- 0:00hi everyone so today we are going to be
- 0:02continuing our Zero to Hero series and
- 0:04in particular today we are going to
- 0:06reproduce the gpt2 model the 124 million
- 0:09version of it so when openi released
- 0:13gpt2 this was 2019 and they released it
- 0:16with this blog post on top of that they
- 0:19released this paper and on top of that
- 0:21they released this code on GitHub so
- 0:23open a/
- 0:24gpt2 now when we talk about reproducing
- 0:27gpt2 we have to be careful because in
- 0:29particular in this video we're going to
- 0:30be reproducing the 124 million parameter
- 0:33model so the thing to realize is that
- 0:35there's always a miniseries when these
- 0:37are releases are made so there are the
- 0:40gpt2 miniseries made up of models at
- 0:42different sizes and usually the biggest
- 0:45model is called the
- 0:46gpt2 but basically the reason we do that
- 0:49is because you can put the model sizes
- 0:51on the x-axis of plots like this and on
- 0:53the Y AIS you put a lot of uh Downstream
- 0:55metrics that you're interested in like
- 0:57translation summarization question
- 0:58answering and so on and you can chart
- 1:00out these scaling laws so basically as
- 1:03the model size increases you're getting
- 1:05better and better at Downstream metrics
- 1:07and so in particular for
- 1:09gpt2 if we scroll down in paper there
- 1:12are four models in the gpt2 miniseries
- 1:15starting at 124 million all the way up
- 1:18to 1558 million now the reason my
- 1:22numbers the way I say them disagree with
- 1:23this table is that this table is wrong
- 1:25if you actually go to the uh gpt2 uh
- 1:29GitHub repo they sort of say that um
- 1:32there was an error in how they added up
- 1:33the parameters but basically this is the
- 1:35124 million parameter model Etc so the
- 1:38124 million parameter had 12 layers in
- 1:40the Transformer and it had 768 channels
- 1:44in the Transformer 768 dimensions and
- 1:47I'm going to be assuming some
- 1:48familiarity with what these terms mean
- 1:50because I covered all of this in my
- 1:51previous video let's build gpt2 uh let's
- 1:54build GPT from scratch so I covered that
- 1:56in the previous video in this playlist
- 1:59now if we do everything correctly and
- 2:01everything works out well by the end of
- 2:03this video we're going to see something
- 2:04like this where we're looking at the
- 2:06validation loss which basically um
- 2:10measures how good we are at predicting
- 2:11the next token in a sequence on some
- 2:13validation data that the model has not
- 2:15seen during training and we see that we
- 2:17go from doing that task not very well
- 2:20because we're initializing from scratch
- 2:22all the way to doing that task quite
- 2:23well um by the end of the training and
- 2:26hopefully we're going to beat the gpt2
- 2:28uh 124 M model
- 2:30now previously when they were working on
- 2:32this this is already 5 years ago so this
- 2:35was probably a fairly complicated
- 2:36optimization at the time and the gpus
- 2:38and the compute was a lot smaller today
- 2:41you can reproduce this model in roughly
- 2:42an hour or probably less even and it
- 2:45will cost you about 10 bucks if you want
- 2:47to do this on the cloud uh Cloud Compu a
- 2:49sort of computer that you can all rent
- 2:52and if you pay $10 for that computer you
- 2:54wait about an hour or less you can
- 2:56actually achieve a model that is as good
- 2:58as this model that open ey released and
- 3:02uh one more thing to mention is unlike
- 3:04many other models open ey did release
- 3:06the weights for gpt2 so those weights
- 3:08are all available in this repository but
- 3:11the gpt2 paper is not always as good
- 3:14with all of the details of training so
- 3:16in addition to the gpt2 paper we're
- 3:18going to be referencing the gpt3 paper
- 3:20which is a lot more Concrete in a lot of
- 3:22the hyp parameters and optimization
- 3:24settings and so on um and it's not a
- 3:27huge departure in the architecture from
- 3:29the GPT 2 uh version of the model so
- 3:31we're going to be referencing both gpt2
- 3:33and gpt3 as we try to reproduce gpt2 124
- 3:36M uh so let's go so the first thing I
- 3:40would like to do is actually start at
- 3:41the end or at the Target so in other
- 3:43words let's load the GPT to 124 M model
- 3:47as it was released by openi and maybe
- 3:48take it for a spin let's sample some
- 3:50tokens from it now the issue with that
- 3:52is when you go into the code base of
- 3:54gpt2 and you go into the source and you
- 3:56click in on the model. pi you'll realize
- 3:58that actually this is using tensorflow
- 4:01so the original gpt2 code here was
- 4:03written in tensor flow which is
- 4:06um you know not let's just say not used
- 4:09as much anymore um so we'd like to use
- 4:12pytorch uh because it's a lot friendlier
- 4:14easier and I just personally like a lot
- 4:16more the problem with that is the
- 4:17initial code is intenser flow we'd like
- 4:19to use pytorch so instead uh to get the
- 4:21target we're going to use the hugging
- 4:23face Transformers um code which I like a
- 4:27lot more so when you go into the
- 4:28Transformers source Transformers models
- 4:30gpt2 modeling gpt2 Pi you will see that
- 4:33they have the gpt2 implementation of
- 4:35that Transformer here in this
- 4:37file um and it's like medium readable
- 4:42but not fully readable um but what it
- 4:45does is it did all the work of
- 4:47converting all those weights uh from
- 4:50tensor flow to pytorch Friendly and so
- 4:52it's much easier to load and work with
- 4:54so in particular we can look at the
- 4:56gpt2 um model here and we can load it
- 4:59using hugging face Transformers so
- 5:01swinging over this is what that looks
- 5:03like from Transformers import the DP GT2
- 5:07LM head model and then from pre-train
- 5:12gpt2 uh now one awkward thing about this
- 5:15is that when you do gpt2 as the model
- 5:17that we're loading this actually is the
- 5:19124 million parameter model if you want
- 5:22the actual the gpt2 the 1.5 billion then
- 5:25you actually want to do- XL so this is
- 5:28the 12 4 M our Target now what we're
- 5:32doing is when we actually get this we're
- 5:33initializing the uh pytorch NN module as
- 5:37defined here in this
- 5:38class from it I want to get just the
- 5:41state dict which is just a raw tensors
- 5:44so we just have um the tensors of that
- 5:46file and by the way here this is a
- 5:49jupyter notebook uh but this is jupyter
- 5:51notebook running inside vs code uh so I
- 5:54like to work with it all in a single
- 5:56sort of interface so I like to use vs
- 5:57code so this is the jupyter notebook
- 6:00extension inside the es
- 6:03code so when we get the state dick this
- 6:06is just a dict so we can print the key
- 6:09and the value which is the tensor and
- 6:11let's just look at the shapes so these
- 6:13are sort of
- 6:14the uh different parameters inside the
- 6:17gbt2 model and their shape so the W
- 6:22weight for token
- 6:25embedding is of size
- 6:2750257 by 768 where this is coming from
- 6:31is that we have
- 6:3250257 tokens in the gpt2 vocabulary um
- 6:37and the tokens by the way these are
- 6:39exactly the tokens that we spoken about
- 6:40in the previous video on my tokenization
- 6:43Series so the previous videos just
- 6:45before this I go into a ton of detail on
- 6:47tokenization gpt2 tokenizer happens to
- 6:49have this many tokens for each
- 6:53token we have a 768 dimensional
- 6:56embedding that is the distributed
- 6:58representation that stands in for that
- 7:01token so each token is a little string
- 7:03piece and then the 768 numbers are the
- 7:06vector that represents that
- 7:08token and so this is just our lookup
- 7:10table for tokens and then here we have
- 7:13the lookup table for the positions so
- 7:16because gbt2 has a maximum sequence
- 7:18length of
- 7:191024 we have up to 1,24 positions that
- 7:23each token can be attending to in the
- 7:25past and every one of those positions in
- 7:28gpd2 has a fixed Vector of
- 7:31768 that is learned by
- 7:33optimization um and so this is the
- 7:36position embedding and the token
- 7:38embedding um and then everything here is
- 7:41just the other weights and biases and
- 7:43everything else of this
- 7:45Transformer so when you just take for
- 7:47example the positional embeddings and
- 7:49flatten it out and take just the 20
- 7:51elements you can see that these are just
- 7:52the parameters these are weights floats
- 7:56just we can take and we can plot them so
- 7:59these are the position embeddings and we
- 8:01get something like this and you can see
- 8:03that this has structure and it has
- 8:04structure because what we what we have
- 8:07here really is every Row in this
- 8:10visualization is a different position a
- 8:12fixed absolute position in um the range
- 8:16from 0 to
- 8:171024 and each row here is the
- 8:19representation of that position and so
- 8:23it has structure because these
- 8:24positional embeddings end up learning
- 8:26these sinusoids and cosiness um that
- 8:29sort of like represent each of these
- 8:31positions and uh each row here stands in
- 8:35for that position and is processed by
- 8:36the Transformer to recover all the
- 8:38relative positions and uh sort of
- 8:41realize which token is where and um
- 8:44attend to them depending on their
- 8:45position not just their
- 8:47content so when we actually just look
- 8:49into an individual column inside these
- 8:53and I just grabbed three random columns
- 8:55you'll see that for example here we are
- 8:57focusing on every every single um
- 9:01Channel and we're looking
- 9:03at what that channel is doing as a
- 9:07function of uh position from one from Z
- 9:11to
- 9:121223
- 9:14really and we can see that some of these
- 9:15channels basically like respond more or
- 9:17less to different parts of the position
- 9:19Spectrum so this green channel uh really
- 9:22likes to fire for everything after 200
- 9:26uh up to 800 but not less a lot less and
- 9:30has a sharp drop off here near zero so
- 9:33who knows what these embeddings are
- 9:34doing and why they are the way they are
- 9:36you can tell for example that because
- 9:37they're a bit more Jagged and they're
- 9:38kind of noisy you can tell that this
- 9:40model was not fully trained and the more
- 9:43trained this model was the more you
- 9:45would expect to smooth this out and so
- 9:47this is telling you that this is a
- 9:48little bit of an undertrained model um
- 9:51but in principle actually these curves
- 9:53don't even have to be smooth this should
- 9:55just be totally random noise and in fact
- 9:57in the beginning of the optimization it
- 9:58is complete random noise because this
- 10:01position embedding table is initialized
- 10:03completely at random so in the beginning
- 10:05you have jaggedness and the fact that
- 10:07you end up with something smooth is
- 10:09already kind of impressive um that that
- 10:11just falls out of the optimization
- 10:13because in principle you shouldn't even
- 10:14be able to get any single graph out of
- 10:16this that makes sense but we actually
- 10:18get something that looks a little bit
- 10:19noisy but for the most part looks
- 10:21sinusoidal like um in the original
- 10:24Transformer um in the original
- 10:26Transformer paper the attention is all
- 10:28you need paper the positional embeddings
- 10:30are actually initialized and fixed if I
- 10:32remember correctly to sinusoids and
- 10:34cosiness of uh different frequencies and
- 10:37that's the positional coding and it's
- 10:38fixed but in gpt2 these are just
- 10:40parameters and they're trained from
- 10:41scratch just like any other parameter uh
- 10:44and that seems to work about as well and
- 10:46so what they do is they kind of like
- 10:47recover these sinusoidal like features
- 10:50during the
- 10:52optimization we can also look at any of
- 10:54the other matrices here so here I took
- 10:57the first layer of the
- 11:00Transformer and looking at like one of
- 11:02its weights and just the first block of
- 11:05300 by 300 and you see some structure
- 11:08but like again like who knows what any
- 11:10of this is if you're into mechanistic
- 11:12interpretability you might get a real
- 11:14kick out of trying to figure out like
- 11:16what is going on what is this structure
- 11:18and what does this all mean but we're
- 11:19not going to be doing that in this video
- 11:21but we definitely see that there's some
- 11:22interesting structure and that's kind of
- 11:24cool what we're mostly interested in is
- 11:26we've loaded the weights of this model
- 11:28that was released by open Ai and now
- 11:30using the hogging face Transformers we
- 11:33can not just get all the raw weights but
- 11:35we can also get the um what they call
- 11:39Pipeline and sample from it so this is
- 11:42the prefix hello I'm a language model
- 11:44comma and then we're sampling uh 30
- 11:47tokens and we getting five sequences and
- 11:50I ran this and this is what it produced
- 11:53um hell language
- 11:55model but what I'm really doing is
- 11:57making a human readable document there
- 11:59are other languages but those are dot
- 12:01dot dot so you can read through these if
- 12:03you like but basically these are five
- 12:05different completions of the same prefix
- 12:07from this uh gbt
- 12:092124m now uh if I go here I took this
- 12:13example from here and sadly even though
- 12:16we are fixing the seed we are getting
- 12:18different Generations from the snippet
- 12:21than what they got so presumably the
- 12:24code changed um but what we see though
- 12:28at this stage that's important is that
- 12:29we are getting coherent text so we've
- 12:32loaded the model successfully we can
- 12:34look at all its parameters and the keys
- 12:36tell us where in the model these come
- 12:39from and we want to actually write our
- 12:41own gpt2 class so that we have full
- 12:43understanding of what's happening there
- 12:44we don't want to be working with
- 12:46something like uh the modeling gpt2 Pi
- 12:49because it's just too complicated we
- 12:50want to write this from scratch
- 12:51ourselves so we're going to be
- 12:53implementing the GPT model here in
- 12:54parallel and as our first task let's
- 12:57load the gpt2 124 M into the class that
- 13:01we're going to develop here from scratch
- 13:04that's going to give us confidence that
- 13:06we can load the open ey model and
- 13:08therefore there's a setting of Weights
- 13:10that exactly is the 124 model but then
- 13:13of course what we're going to do is
- 13:14we're going to initialize the model from
- 13:15scratch instead and try try to train it
- 13:18ourselves um on a bunch of documents
- 13:20that we're going to get and we're going
- 13:22to try to surpass that model so we're
- 13:24going to get different weights and
- 13:25everything's going to look different
- 13:27hopefully better even um
- 13:29but uh we're going to have a lot of
- 13:31confidence that because we can load the
- 13:32openi model we are in the same model
- 13:34family and model class and we just have
- 13:36to ReDiscover a good setting of the
- 13:37weights uh but from scratch so let's now
- 13:41write the gbt2 model and let's load the
- 13:43weights and make sure that we can also
- 13:45generate text that looks coherent okay
- 13:48so let's now swing over to the attention
- 13:49is all un need paper that started
- 13:51everything and let's scroll over to the
- 13:53model architecture the original
- 13:55Transformer now remember that gpt2 is
- 13:57slightly modified from the or or
- 13:59Transformer in particular we do not have
- 14:02uh the encoder gpt2 is a decoder only
- 14:05Transformer as we call it so this entire
- 14:07encoder here is missing in addition to
- 14:09that this cross attention here that was
- 14:12using that encoder is also missing so we
- 14:14delete this entire part everything else
- 14:18stays almost the same but there are some
- 14:20differences that we're going to uh sort
- 14:21of look at here so there are two main
- 14:26differences when we go to the gb2 page
- 14:29under 2.3 model we notice that first
- 14:32there's a reshuffling of the layer Norms
- 14:34so they change place and second an
- 14:38additional layer normalization was added
- 14:40here to the final self detention block
- 14:43so basically all the layer Norms here
- 14:46instead of being after the MLP or after
- 14:48the attention they SN before it and an
- 14:50additional layer Norm gets added here
- 14:52right before the final
- 14:54classifier so now let's Implement some
- 14:56of the first sort of skeleton NN module
- 14:59modules here in our GPT NN module and in
- 15:02particular we're going to try to match
- 15:04up this schema here that is used by
- 15:06hugging face Transformers because that
- 15:08will make it much easier to load these
- 15:10weights from this state dict so we want
- 15:12something that reflects uh this schema
- 15:15here so here's what I came up with
- 15:19um basically we see that the main
- 15:22container here that has all the modules
- 15:24is called Transformer so I'm reflecting
- 15:26that with an NN module dict and this is
- 15:29basically a module that allows you to
- 15:30index into the subm modules using keys
- 15:34just like a dictionary uh
- 15:36strings within it we have the weights of
- 15:39the token embeddings WT and that's an N
- 15:41embedding and the weights of the
- 15:44position embeddings which is also just
- 15:45an N embedding and if you remember n
- 15:47embedding is really just a fancy little
- 15:49wrapper module around just a single um
- 15:53single array of numbers a single uh
- 15:56block of numbers just like this it's a
- 15:58single tensor and an embedding is a
- 16:02glorified um wrapper around a tensor
- 16:04that allows you to access its elements
- 16:07uh by indexing into the
- 16:08rows now in addition to that we see here
- 16:11that we have a h and then there's a this
- 16:14is index using numbers instead of
- 16:16indexed using strings so there's a h. 0
- 16:191 2 Etc all the way up till h. 11 and
- 16:23that's because there are 12 layers here
- 16:26in this Transformer so to reflect that
- 16:28I'm creating also an H I think that
- 16:31probably stands for hidden and instead
- 16:33of a module dict this is a model list so
- 16:35we can index it using integers exactly
- 16:37as we see here 01 2 Etc and the modular
- 16:42list has a n layer blocks and the blocks
- 16:46are yet to be defined in a module in a
- 16:48bit in addition to that following the
- 16:50gpt2 paper we have we need an additional
- 16:53final layer Norm that we're going to put
- 16:56in there and then we have the final
- 16:58classifier uh the language model head
- 17:01which um projects from 768 the number of
- 17:05embedding dimensions in this GPT all the
- 17:08way to the vocab size which is
- 17:1050257 and gpt2 uses no bias for this
- 17:13final uh sort of projection so this is
- 17:16the skeleton and you can see that it
- 17:19reflects this so the wte is the token
- 17:22embeddings here it's called output
- 17:24embedding but it's really the token
- 17:26embeddings the PE is the positional
- 17:29codings uh those two pieces of
- 17:31information as we saw previously are
- 17:32going to add and then go into the
- 17:34Transformer the H is the all the blocks
- 17:37in Gray and the LNF is this new layer
- 17:40that gets added here by the gpt2 model
- 17:43and LM head is this linear part here so
- 17:47that's the skeleton of the gpt2 we now
- 17:50have to implement the block okay so
- 17:53let's now recurse to the block itself so
- 17:55we want to define the block um so I'll
- 17:59start putting them here so the block I
- 18:02like to write out like
- 18:04this uh these are some of the
- 18:06initializations and then this is the
- 18:07actual forward pass of what this block
- 18:09computes and notice here that there's a
- 18:12change from the Transformer again that
- 18:14is mentioned in the gpt2 paper so here
- 18:17the layer normalizations are after the
- 18:20application of attention or feed forward
- 18:22in addition to that note that the
- 18:24normalizations are inside the residual
- 18:26stream you see how feed forward is
- 18:28applied and this arrow goes through and
- 18:30through the normalization so that means
- 18:33that your residual pathway has
- 18:35normalizations inside them and this is
- 18:37not very good or desirable uh you
- 18:39actually prefer to have a single uh
- 18:42clean residual stream all the way from
- 18:44supervision all the way down to the
- 18:45inputs the tokens and this is very
- 18:48desirable and nice because the gradients
- 18:51that flow from the top if you remember
- 18:54from your microad addition just
- 18:56distributes gradients during the
- 18:58backwards state to both of its branches
- 19:00equally so addition is a branch in the
- 19:04gradients and so that means that the
- 19:06gradients from the top flows straight to
- 19:08the inputs the tokens through the
- 19:10residual Pathways unchanged but then in
- 19:13addition to that the gradient also flows
- 19:14through the blocks and the blocks you
- 19:17know contribute their own contribution
- 19:18over time and kick in and change the
- 19:20optimization over time but basically
- 19:22clean residual pathway is desirable from
- 19:25an optimization perspective and then the
- 19:28this is the pre-normalization version
- 19:30where you see that RX first goes through
- 19:32the layer normalization and then the
- 19:34attention and then goes uh back out to
- 19:38go to the L ration number two and the
- 19:40multia perceptron sometimes also
- 19:43referred to as a feed forward Network or
- 19:44an FFN and then that goes into the
- 19:47residual stream again and the one more
- 19:50thing that is kind of interesting to
- 19:51note is that recall that attention is a
- 19:53communication operation it is where all
- 19:55the tokens and there's 1,24 tokens lined
- 19:58up in a sequence and this is where the
- 20:00tokens communicate this is where they
- 20:02exchange information so attention is a
- 20:06um aggregation function it's a pooling
- 20:08function it's a weighted sum function it
- 20:12is a reduce operation whereas MLP this
- 20:16uh MLP here happens at every single
- 20:18token individually there's no
- 20:20information being collected or exchanged
- 20:21between the tokens so the attention is
- 20:24the reduce and the MLP is the map and
- 20:27what you end up with is that the
- 20:28Transformer just ends up just being a
- 20:30repeated application of map produce if
- 20:33you want to think about it that way so
- 20:36um this is where they communicate and
- 20:37this is where they think individually
- 20:39about the information that they gathered
- 20:41and every one of these blocks uh
- 20:43iteratively refines the um
- 20:46representation is at the residual stream
- 20:48so this is our block um slightly
- 20:51modified from this picture Okay so let's
- 20:53now move on to the MLP so the MLP block
- 20:57uh I implemented as follows
- 20:59it is relatively straightforward we
- 21:00basically have two linear projections
- 21:02here that are sandwiched in between the
- 21:05G
- 21:06nonlinearity so nn. G approximate is 10h
- 21:11now when we swing on uh swing over to
- 21:13the Pyro documentation this is n.g and
- 21:16it has this format and it has two
- 21:18versions the original version of G which
- 21:20we'll step into into in a bit and the
- 21:22approximate version of Galo which we can
- 21:24request using
- 21:2510 so as you can see just as a preview
- 21:28here G is a basically like a reu except
- 21:32there's no flat exactly Flat Tail here
- 21:35at exactly zero but otherwise it looks
- 21:38very much like a slightly smoother reu
- 21:41it comes from this paper here Gan error
- 21:43linear units and uh you can step through
- 21:46this paper and there's some mathematical
- 21:48calac reasoning that leads to an
- 21:50interpretation that leads to the
- 21:51specific formulation it has to do with
- 21:53stochastic radial risers and the
- 21:56expectation of a modification to
- 21:57Adaptive dropout so you can read through
- 21:59all of that if you'd like here and
- 22:01there's a little bit of history as to
- 22:03why there is an an approximate version
- 22:05of G and that comes from this issue here
- 22:08as far as I can tell and in this issue
- 22:11Daniel Hendrix mentions that at the time
- 22:14when they developed this nonlinearity
- 22:17the Earth function which you need to
- 22:19evaluate the exact G was very slow in
- 22:21tensor flow so they ended up basically
- 22:23developing this approximation and this
- 22:25approximation that then ended up being
- 22:27picked up by Bert and by GP P2 Etc but
- 22:30today there's no real good reason to use
- 22:31the approximate version you'd prefer to
- 22:33just use the exact version um because I
- 22:36my expectation is that there's no big
- 22:38difference anymore and this is kind of
- 22:40like a historical um kind of Quirk um
- 22:43but we are trying to reproduce gpt2
- 22:45exactly and gpt2 used the 10h
- 22:49approximate version so we prefer to
- 22:51stick with
- 22:52that um now one other reason to actually
- 22:55just intuitively use G instead of veru
- 22:57is previously in the in videos in the
- 22:59past we've spoken about the dead reu
- 23:02neuron problem where in this tale of a
- 23:04reu if it's exactly flat at zero any
- 23:07activations that fall there will get
- 23:09exactly zero gradient there's no change
- 23:11there's no adaptation there's no
- 23:13development of the network if any of
- 23:15these activations end in this flat
- 23:17region but the G always contributes a
- 23:20local gradient and so there's always
- 23:22going to be a change always going to be
- 23:23an adaptation and sort of smoothing it
- 23:25out ends up empirically working better
- 23:27in practice as demonstrated in this
- 23:29paper and also as demonstrated by it
- 23:31being picked up by the bird paper gbt2
- 23:33paper and so on so for that reason we
- 23:35adopt this nonlinearity uh here in the
- 23:3810 in the gbt2 reproduction now in more
- 23:41modern networks also like llama 3 and so
- 23:43on this nonlinearity also further
- 23:45changes uh to swiglo and other variants
- 23:48like that uh but for gpt2 they Ed this
- 23:50approximate
- 23:51G okay and finally we have the attention
- 23:54operation so let me paste in my
- 23:57attention
- 24:00so I know this is a lot so I'm going to
- 24:02go through this a bit quickly a bit
- 24:03slowly but not too slowly because we
- 24:05have covered this in the previous video
- 24:07and I would just point you there um so
- 24:10this is the attention operation now in
- 24:12the previous video you will remember
- 24:13this is not just attention this is um
- 24:16multi-headed attention right and so in
- 24:19the previous video we had this
- 24:20multi-headed attention module and this
- 24:23implementation made it obvious that
- 24:25these heads are not actually that
- 24:26complicated uh there's basically
- 24:28in parallel inside every attention block
- 24:32there's multiple heads and they're all
- 24:33functioning in parallel and uh their
- 24:36outputs are just being concatenated and
- 24:38that becomes the output of the
- 24:40multi-headed attention so the heads are
- 24:42just kind of like parallel streams and
- 24:45their outputs get
- 24:46concatenated and so it was very simple
- 24:48and made the head be kind of like U
- 24:51fairly straightforward in terms of its
- 24:54implementation what happens here is that
- 24:56instead of having two separate modules
- 24:58and indeed many more modules that get
- 24:59concatenated all of that is just put
- 25:01into a single uh self attention uh
- 25:04module and instead I'm being very
- 25:07careful and doing a bunch of transpose
- 25:10split um tensor gymnastics to make this
- 25:13very efficient in pych but fundamentally
- 25:15and algorithmically nothing is different
- 25:17from the implementation we saw
- 25:19before um in this uh give
- 25:22repository so to remind you very briefly
- 25:25and I don't want to go in this uh into
- 25:27this in too many in too much time but we
- 25:30have these tokens lined up in a sequence
- 25:32and there's 1,20 of them and then each
- 25:35token at this stage of the attention
- 25:37emits three vectors the query key and
- 25:40the value and first what happens here um
- 25:44is that the queries and the keys have to
- 25:46multiply each other to get sort of the
- 25:49attention um amount like how interesting
- 25:52they find each other so they have to
- 25:54interact multiplicatively so what we're
- 25:56doing here is we're calculating the qkv
- 25:58we splitting it and then there's a bunch
- 26:00of gymnastics as I mentioned here and
- 26:03the way this works is that we're
- 26:04basically making the number of heads and
- 26:06H into a batch Dimension and so it's a
- 26:10batch Dimension just like B so that in
- 26:12these operations that follow pytorch
- 26:14treats B and NH as batches and it
- 26:18applies all the operations on all of
- 26:20them in parallel in both the batch and
- 26:22the
- 26:23heads and the operations that get
- 26:25applied are number one the queries and
- 26:27the keys intera to give us her attention
- 26:30this is the autoaggressive mask that
- 26:32makes sure that the tokens only attend
- 26:35to tokens before them and never to
- 26:37tokens in the
- 26:39future the softmax here normalizes the
- 26:41attention so it sums to one always and
- 26:45then recall from the previous video that
- 26:47doing the attention Matrix multiply with
- 26:48the values is basically a way to do a
- 26:50weighted sum of the values of the tokens
- 26:53that we found interesting at every
- 26:55single token and then the final
- 26:57transpose conf VI and view is just
- 26:59reassembling all of that again and this
- 27:02actually performs the concatenation
- 27:04operation so you can step through this
- 27:06uh slowly if you'd like um but it is
- 27:08equivalent mathematically to our
- 27:10previous implementation is just more
- 27:12efficient in P torch so that's why I
- 27:14chose this implementation
- 27:16instead now in addition to that I'm
- 27:18being careful with how I name my
- 27:19variables so for example cattin is the
- 27:22same as seaten and so actually our keys
- 27:25should basically exactly follow the
- 27:27schema of the hugging face train
- 27:28Transformers code and that will make it
- 27:29very easy for us to now Port over all
- 27:32the weights from exactly this sort of
- 27:34naming conventions because all of our
- 27:36variables are named the same thing but
- 27:39um at this point we have finished the
- 27:41gpt2 implementation and what that allows
- 27:44us to do is we don't have to basically
- 27:46use uh this file from hugging face which
- 27:48is fairly long
- 27:50um this
- 27:52is uh 2,000 lines of code um instead we
- 27:57just have a less than 100 lines of code
- 27:59and this is the complete uh gpd2
- 28:01implementation so at this stage we
- 28:02should just be able to take over all the
- 28:04weights set them and then do generation
- 28:07so let's see what that looks like okay
- 28:09so here I've also changed the GPT config
- 28:11so that the numbers here the H
- 28:13parameters agree with the gpt2 124 M
- 28:15model so the maximum sequence length
- 28:17which I call block size here is 124 the
- 28:21number of tokens is 50250 257 which if
- 28:25you watch my tokenizer video know that
- 28:27this is 50,000 m merges BP merges 256
- 28:31bite tokens the leaves of the BP tree
- 28:35and one special end of text token that
- 28:36delimits different documents and can
- 28:38start generation as well and there are
- 28:4112 layers there are 12 heads in the
- 28:43attention and the dimension of the
- 28:45Transformers was
- 28:46768 so here's how we can now load the
- 28:49parameters from hugging face to uh our
- 28:52code here and initialize the GPT class
- 28:54with those parameters so let me just
- 28:56copy paste a bunch of code
- 28:59here and I'm not going to go through
- 29:00this code too slow too quickly too
- 29:03slowly because um honestly it's not that
- 29:07interesting it's not that exciting we're
- 29:08just loading the weights so it's kind of
- 29:10dry but as I mentioned there are four
- 29:12models in this miniseries of gpt2 this
- 29:15is some of the Jupiter code um code that
- 29:18we had here on the right I'm just pting
- 29:20it over these are the hyper parameters
- 29:22of the gpt2 models uh we're creating the
- 29:24config object and creating our own model
- 29:27and then what's Happening Here is we're
- 29:28creating the state dict both for our
- 29:30model and for the hugging face
- 29:33model um and then what we're doing here
- 29:36is we're going over the hugging face
- 29:38model keys and we're copying over those
- 29:42tensors and in the process we are kind
- 29:45of ignoring a few of the buffers they're
- 29:47not parameters they're buffers so for
- 29:49example attention dobias uh that's just
- 29:51used for the autoaggressive mask and so
- 29:53we are ignoring some of those masks and
- 29:56uh that's it and then then one
- 29:58additional kind of annoyance is that
- 30:00this comes from the tensorflow repo and
- 30:02I'm not sure how this is a little bit
- 30:04annoying but some of the weights are
- 30:05transposed from what pytorch would want
- 30:08and so manually I hardcoded the weights
- 30:10that should be transposed and then we
- 30:12transpose them if that is so and then we
- 30:15return this model so the from
- 30:18pre-trained is a
- 30:20Constructor or class method in Python
- 30:23that Returns the GPT object if we just
- 30:26give it the model type which in our case
- 30:28is gpt2 the smallest model that we're
- 30:30interested in so this is the code and
- 30:33this is how you would use it and um we
- 30:35can pop open the terminal here in vs
- 30:38code and we can python train gbt2 pi and
- 30:44fingers
- 30:46crossed okay so we didn't crash and so
- 30:50we can load the weights and the biases
- 30:52and everything else into our Ann module
- 30:55but now let's also get additional
- 30:57confidence that this is working and
- 30:58let's try to actually generate from this
- 31:00model okay now before we can actually
- 31:01generate from this model we have to be
- 31:03able to forward it we didn't actually
- 31:04write that code yet so here's the
- 31:06forward
- 31:08function so the input to the forward is
- 31:11going to be our indices our tokens uh
- 31:13token indices and they are always of
- 31:16shape B BYT and so we have batch
- 31:19dimension of B and then we have the time
- 31:22dimension of up to T and the T can't be
- 31:26more than the block size the block size
- 31:27is is the maximum sequence length so B
- 31:30BYT indices arranged is sort of like a
- 31:32two-dimensional layout and remember that
- 31:35basically every single row of this is of
- 31:37size up to uh block size and this is T
- 31:41tokens that are in a sequence and then
- 31:43we have B independent sequences stacked
- 31:46up in a batch so that this is
- 31:48efficient now here we are forwarding the
- 31:51position embeddings and the token
- 31:52embeddings and this code should be very
- 31:54recognizable from the previous lecture
- 31:56so um we basically use uh a range which
- 31:59is kind of like a version of range but
- 32:01for pytorch uh and we're iterating from
- 32:04Z to T and creating this uh positions uh
- 32:07sort of uh indices
- 32:10um and then we are making sure that
- 32:12they're in the same device as idx
- 32:14because we're not going to be training
- 32:15on only CPU that's going to be too
- 32:16inefficient we want to be training on
- 32:18GPU and that's going to come in in a
- 32:20bit uh then we have the position
- 32:22embeddings and the token embeddings and
- 32:24the addition operation of those two now
- 32:26notice that the position embed are going
- 32:28to be identical for every single row of
- 32:31uh of input and so there's broadcasting
- 32:33hidden inside this plus where we have to
- 32:36create an additional Dimension here and
- 32:38then these two add up because the same
- 32:40position embeddings apply at every
- 32:41single row of our example stacked up in
- 32:44a batch then we forward the Transformer
- 32:46blocks and finally the last layer norm
- 32:49and the LM head so what comes out after
- 32:52forward is the logits and if the input
- 32:55was B BYT indices then at every single B
- 32:58by T we will calculate the uh logits for
- 33:02what token comes next in the sequence so
- 33:05what is the token B t+1 the one on the
- 33:09right of this token and B app size here
- 33:12is the number of possible tokens and so
- 33:16therefore this is the tensor that we're
- 33:17going to obtain and these low jits are
- 33:19just a softmax away from becoming
- 33:22probabilities so this is the forward
- 33:25pass of the network and now we can get
- 33:27load and so we're going to be able to
- 33:29generate from the model
- 33:30imminently okay so now we're going to
- 33:32try to set up the identical thing on the
- 33:35left here that matches hug and face on
- 33:36the right so here we've sampled from the
- 33:39pipeline and we sampled five times up to
- 33:4230 tokens with the prefix of hello I'm a
- 33:45language model and these are the
- 33:46completions that we achieved so we're
- 33:48going to try to replicate that on the
- 33:49left here so number turn sequences is
- 33:51five max length is 30 so the first thing
- 33:53we do of course is we initialize our
- 33:55model then we put it into evaluation
- 33:57mode now this is a good practice to put
- 33:59the model into eval when you're not
- 34:01going to be training it you're just
- 34:02going to be using it and I don't
- 34:05actually know if this is doing anything
- 34:07right now for the following reason our
- 34:09model up above here contains no modules
- 34:11or layers that actually have a different
- 34:14uh Behavior at training or evaluation
- 34:16time so for example Dropout batch norm
- 34:18and a bunch of other layers have this
- 34:20kind of behavior but all of these layers
- 34:22that we've used here should be identical
- 34:23in both training and evaluation time um
- 34:27so so potentially model that eval does
- 34:29nothing but then I'm not actually sure
- 34:31if this is the case and maybe pytorch
- 34:33internals uh do some clever things
- 34:35depending on the evaluation mode uh
- 34:36inside here the next thing we're doing
- 34:39here is we are moving the entire model
- 34:41to Cuda so we're moving this all of the
- 34:44tensors to GPU so I'm sshed here to a
- 34:47cloud box and I have a bunch of gpus on
- 34:49this box and here I'm moving the entire
- 34:53model and all of its members and all of
- 34:54its tensors and everything like that
- 34:56everything gets shipped off to basically
- 34:59a whole separate computer that is
- 35:01sitting on the GPU and the GPU is
- 35:03connected to the uh CPU and they can
- 35:05communicate but it's basically a whole
- 35:06separate computer with its own computer
- 35:08architecture and it's really well
- 35:09catered to parallel processing tasks
- 35:11like those of running neural networks so
- 35:14I'm doing this so that the model lives
- 35:16on the GPU a whole separate computer and
- 35:19it's just going to make our code a lot
- 35:20more efficient because all of this stuff
- 35:22runs a lot more efficiently on the
- 35:25gpus so that's the model
- 35:29itself now uh the next thing we want to
- 35:31do is we want to start with this as the
- 35:34prefix when we do the generation so
- 35:37let's actually create those prefix
- 35:39tokens so here's the code that I've
- 35:41written we're going to import the tich
- 35:43token library from open Ai and we're
- 35:45going to get the gpt2 encoding so that's
- 35:48the tokenizer for gpt2 and then we're
- 35:51going to encode this string and get a
- 35:54list of integers which are the tokens uh
- 35:57now these integers here should actually
- 35:59be fairly straightforward because we can
- 36:01just copy paste this string and we can
- 36:04sort of inspect what it is in tick
- 36:05tokenizer so just pasting that in these
- 36:08are the tokens that are going to come
- 36:09out so this list of integers is what we
- 36:12expect tokens to become and as you
- 36:15recall if you saw my video of course all
- 36:17the tokens they're just little string
- 36:19chunks right so these are this is the
- 36:21chunc of this string into gpt2
- 36:25tokens so once we have those tokens it's
- 36:27a list of integers we can create a torch
- 36:30tensor out of it in this case it's eight
- 36:32tokens and then we're going to replicate
- 36:34these eight tokens for five times to get
- 36:36five rows of eight tokens and that is
- 36:40our initial um input X as I call it here
- 36:45and it lives on the GPU as well so X now
- 36:48is this idx that we can put into forward
- 36:52to get our logits so that we know what
- 36:55comes as the sixth token
- 36:58uh sorry as the ninth token in every one
- 37:01of these five rows okay and we are now
- 37:04ready to generate so let me paste in one
- 37:05more code block
- 37:07here um so what's happening here in this
- 37:09code block is we have this x which is of
- 37:12size B BYT right so batch by time and
- 37:16we're going to be in every iteration of
- 37:18this loop we're going to be adding a
- 37:19column of new indices into each one of
- 37:22these rows right and so these are the
- 37:24new indices and we're appending them to
- 37:27the the sequence as we're sampling so
- 37:29with each Loop iteration we get one more
- 37:31column into X and all of the operations
- 37:34happen in the context manager of torch.
- 37:36nograd this is just telling pytorch that
- 37:38we're not going to be calling that
- 37:39backward on any of this so it doesn't
- 37:41have to cach all the intermediate
- 37:43tensors it's not going to have to
- 37:44prepare in any way for a potential
- 37:46backward later and this saves a lot of
- 37:48space and also possibly uh some time so
- 37:52we get our low jits we get the loow jits
- 37:54at only the last location we throw away
- 37:57all the other low jits uh we don't need
- 37:59them we only care about the last columns
- 38:01low jits so this is being wasteful uh
- 38:04but uh this is just kind of like an
- 38:06inefficient implementation of
- 38:08sampling um so it's correct but
- 38:10inefficient so we get the last column of
- 38:13loow jits pass it through soft Max to
- 38:14get our probabilities then here I'm
- 38:16doing top case sampling of 50 and I'm
- 38:18doing that because this is the hugging
- 38:20face default so just looking at the
- 38:23hugging face docks here of a pipeline um
- 38:26there's a bunch of
- 38:28quarks that go into hugging face and I
- 38:32mean it's it's kind of a lot honestly
- 38:34but I guess the important one that I
- 38:36noticed is that they're using top K by
- 38:38default which is 50 and what that does
- 38:41is that uh so that's being used here as
- 38:43well and what that does is basically we
- 38:45want to take our probabilities and we
- 38:47only want to keep the top 50
- 38:49probabilities and anything that is lower
- 38:51than the 50th probability uh we just
- 38:54clamp to zero and renormalize and so
- 38:56that way we are never sampling very rare
- 38:59tokens uh the tokens we're going to be
- 39:01sampling are always in the top 50 of
- 39:03most likely tokens and this helps keep
- 39:05the model kind of on track and it
- 39:07doesn't blabber on and it doesn't get
- 39:08lost and doesn't go off the rails as
- 39:10easily uh and it kind of like um sticks
- 39:13in the vicinity of likely tokens a lot
- 39:15better so this is the way to do it in
- 39:17pytorch and you can step through it if
- 39:18you like I don't think it's super
- 39:20insightful so I'll speed through it but
- 39:22roughly speaking we get this new column
- 39:24of of tokens we append them on x and
- 39:27basically The Columns of X grow until
- 39:30this y Loop gets tripped up and then
- 39:33finally we have an entire X of size um 5
- 39:38by 30 in this case in this example and
- 39:41we can just basically print all those
- 39:43individual rows so I'm getting all the
- 39:46rows I'm getting all the tokens that
- 39:48were sampled and I'm using the decode
- 39:50function from Tik tokenizer to get back
- 39:52the string which we can print and so
- 39:55terminal new terminal
- 39:59and let me python train
- 40:08gpt2 okay so these are the generations
- 40:11that we're getting hello I'm a language
- 40:13model not a
- 40:15program um new line new line Etc hello
- 40:19I'm a language model and one of the main
- 40:21things that bothers me when they create
- 40:22languages is how easy it becomes to
- 40:23create something that I me so this will
- 40:26just like blabber on right in all these
- 40:27cases now one thing you will notice is
- 40:29that these Generations are not the
- 40:31generations of hugging face here and I
- 40:35can't find the discrepancy to be honest
- 40:37and I didn't fully go through all these
- 40:39options but probably there's something
- 40:40else hiding in on addition to the top P
- 40:43so I'm not able to match it up but just
- 40:45for correctness um down here Below in
- 40:47the juper notebook and using the hugging
- 40:49face model so this is the hugging face
- 40:52model here I was I replicated the code
- 40:56and if I do this and I run that then I
- 40:59am getting the same results so basically
- 41:03the model internals are not wrong it's
- 41:05just I'm not 100% sure what the pipeline
- 41:08does in hugging face and that's why
- 41:09we're not able to match them up but
- 41:11otherwise the code is correct and we've
- 41:13loaded all the um tensors correctly so
- 41:16we're initializing the model correctly
- 41:18and everything here works so long story
- 41:20short uh We've Port it all the weights
- 41:22we initialize the gpt2 this is the exact
- 41:25opening gpt2 and it can generate
- 41:27sequences and they look sensible and now
- 41:30here of course we're initializing with
- 41:32gbt2 model weights but now we want to
- 41:34initialize from scratch from random
- 41:36numbers and we want to actually train a
- 41:38model that will give us sequences as
- 41:40good as or better than these ones in
- 41:44quality and so that's what we turn to
- 41:46next so it turns out that using the
- 41:48random model is actually fairly
- 41:49straightforward because pytorch already
- 41:51initializes our model randomly and by
- 41:53default so when we create the GPT model
- 41:58and the Constructor this is all um all
- 42:00of these layers and modules have random
- 42:03initializers that are there by default
- 42:05so when these linear layers get created
- 42:07and so on there's default Constructors
- 42:10for example using the Javier
- 42:11initialization that we saw in the past
- 42:13uh to construct the weights of these
- 42:15layers and so creating a random model
- 42:18instead of a gpt2 model is actually
- 42:20fairly straightforward and we would just
- 42:22come here and instead we would create
- 42:24model equals GPT and then we want to use
- 42:28the default config GPT config and the
- 42:31default config uses the 124 M parameters
- 42:33so this is the random model
- 42:35initialization and we can run
- 42:42it and we should be able to get uh
- 42:46results now the results here of course
- 42:48are total garbage carbal and that's
- 42:50because this is random model and so
- 42:51we're just getting all these random
- 42:53token string pieces chunked up totally
- 42:55at random so that's what we have right
- 42:57now uh now one more thing I wanted to
- 42:59point out by the way is in case you do
- 43:01not have Cuda available because you
- 43:03don't have a GPU you can still follow
- 43:04along with uh with what we're doing here
- 43:07uh to some extent uh and probably not to
- 43:10the very end because by the end we're
- 43:11going to be using multiple gpus and
- 43:13actually doing a serious training run uh
- 43:15but for now you can actually follow
- 43:16along decently okay uh so one thing that
- 43:19I like to do in pytorch is I like to
- 43:20autod detect the device that is
- 43:22available to you so in particular you
- 43:24could do that like this
- 43:28so here we are trying to detect a device
- 43:30to run on that has the highest compute
- 43:32capability you can think about it that
- 43:33way so by default we start with CPU
- 43:36which of course is available everywhere
- 43:37because every single computer will have
- 43:39a CPU but then we can try to detect do
- 43:42you have a GPU you so use a Cuda and
- 43:44then if you don't have a Cuda uh do you
- 43:47at least have MPS MPS is the back end
- 43:49for Apple silicon so if you have a
- 43:51Macbook that is fairly new you probably
- 43:53have apple silicon on the inside and
- 43:55then that has a GPU that is actually
- 43:57fairly capable uh depending on which
- 43:59MacBook you have and so you can use MPS
- 44:01which will be potentially faster than
- 44:02CPU and so we can print the device here
- 44:05now once we have the device we can
- 44:07actually use it in place of Puda so we
- 44:11just swap it in and notice that here
- 44:14when we call model on X if this x here
- 44:17is on CPU instead of GPU then it will
- 44:21work fine because here in the forward
- 44:23which is where P to will come when we
- 44:26create a pose we were careful to use the
- 44:28device of idx to create this tensor as
- 44:31well and so there won't be any mismatch
- 44:33where one tensor is on CPU one is on GPU
- 44:36and uh that you can't combine those but
- 44:38here we are um carefully initializing on
- 44:41the correct device as indicated by the
- 44:43input to this model so this will autod
- 44:47detect device for me this will be of
- 44:49course
- 44:50GPU so using device
- 44:54Cuda uh but uh you can also run with um
- 44:58as I mentioned another device and it's
- 45:00not going to be too much slower so if I
- 45:01override device here
- 45:03oops if I override device equals
- 45:07CPU
- 45:08then we'll still print Cuda of course
- 45:11but now we're actually using CPU one 2 3
- 45:164 5 6 okay about 6 seconds and actually
- 45:21we're not using torch compile and stuff
- 45:22like that which will speed up everything
- 45:24a lot faster as well but you can follow
- 45:27even on a CPU I think to a decent extent
- 45:30um so that's note on that okay so I do
- 45:32want to loop around eventually into what
- 45:35it means to have different devices in
- 45:36pytorch and what it is exactly that
- 45:38pytorch does in the background for you
- 45:40when you do something like module. 2
- 45:43device or where you take a torch tensor
- 45:45and do A2 device and what exactly
- 45:48happens and how that works but for now
- 45:49I'd like to get to training and I'd like
- 45:51to start training the model and for now
- 45:53let's just say the device makes code go
- 45:55fast um and let's go into how we can
- 45:58actually train the model so to train the
- 46:00model we're going to need some data set
- 46:02and for me the best debugging simplest
- 46:04data set that I like to use is the tiny
- 46:06Shakespeare data set um and it's
- 46:09available at this URL so you can W get
- 46:11it or you can just search tiny
- 46:12Shakespeare data
- 46:13set and so um I have in my file system
- 46:16as just LS input.txt
- 46:18so I already downloaded it and here I'm
- 46:22reading the data set getting the first
- 46:231,000 characters and printing the first
- 46:26100
- 46:27now remember that gpt2 has uh roughly a
- 46:30compression ratio the tokenizer has a
- 46:32compression ratio of rly 3 to1 so th000
- 46:35characters is roughly 300 tokens here uh
- 46:37that will come out of this in the slice
- 46:39that we're currently getting so this is
- 46:42the first few uh
- 46:44characters and uh if you want to get a
- 46:46few more statistics on this we can do
- 46:48work count on input.txt
- 46:50so we can see that this is uh 40,000
- 46:53lines about 200,000 words in this data
- 46:56set and about 1 million bytes in this
- 46:59file and knowing that this file is only
- 47:01asky characters there's no crazy unic
- 47:03code here as far as I know and so every
- 47:05asky character is encoded with one bite
- 47:08and so this is uh the same number
- 47:10roughly a million characters inside this
- 47:12data set so that's the data set size uh
- 47:15by default very small and minimal data
- 47:17set for debugging to get us off the
- 47:19ground in order to tokenize this data
- 47:21set we're going to get Tik token
- 47:23encoding for gbt2 encode the data uh the
- 47:27first um 1,000 characters and then I'm
- 47:30only going to print the first 24 tokens
- 47:33so these are the tokens as a list of
- 47:36integers and if you can read gpt2 tokens
- 47:38you will see that 198 here you'll
- 47:40recognize that as the slashing character
- 47:42so that is a new line and then here for
- 47:45example we have two new lines so that's
- 47:46198 twice here uh so this is just a
- 47:49tokenization of the first 24 tokens so
- 47:52what we want to do now is we want to
- 47:54actually process these token sequences
- 47:56and feed them into a Transformer and in
- 47:59particular we want them we want to
- 48:01rearrange these tokens into this idx
- 48:05variable that we're going to be feeding
- 48:06into the Transformer so we don't want a
- 48:08single very long onedimensional sequence
- 48:10we want an entire batch where each
- 48:12sequence is up to uh is basically T
- 48:16tokens and T cannot be larger than the
- 48:18maximum sequence length and then we have
- 48:21these t uh tlong uh sequences of tokens
- 48:25and we have B independent examples of
- 48:27sequences so how can we create a b BYT
- 48:30tensor that we can feed into the forward
- 48:32out of these onedimensional
- 48:34sequences so here's my favorite way to
- 48:36to achieve this uh so if we take torch
- 48:39and then we create a tensor object out
- 48:41of this list of integers and just the
- 48:42first 24 tokens my favorite way to do
- 48:45this is basically you do a do view of um
- 48:49of uh for example 4x6 which multiply to
- 48:5224 and so it's just a two-dimensional
- 48:54rearrangement of these tokens and you'll
- 48:56is that when you view this
- 48:57onedimensional sequence as
- 48:58two-dimensional 4x6 here the first six
- 49:03uh tokens uh up to here end up being the
- 49:06first row the next six tokens here end
- 49:09up being the second row and so on and so
- 49:12basically it's just going to stack up
- 49:14this the um every six tokens in this
- 49:18case as independent rows and it creates
- 49:20a batch of tokens in this case and so
- 49:23for example if we are token 25 in the
- 49:26Transformer when we feed this in and
- 49:28this becomes the idx this token is going
- 49:30to see these three tokens and it's going
- 49:33to try to predict that 198 comes
- 49:35next so in this way we are able to
- 49:39create this two-dimensional batch that's
- 49:41that's quite nice now in terms of the
- 49:44label that we're going to need for the
- 49:45Target to calculate the loss function
- 49:47how do we get that well we could write
- 49:49some code inside the forward pass
- 49:51because we know that the next uh token
- 49:53in a sequence which is the label is just
- 49:55to the right of us but you'll notice
- 49:57that actually we for this token at the
- 49:59very end 13 we don't actually have the
- 50:02next correct token because we didn't
- 50:03load it so uh we actually didn't get
- 50:07enough information here so I'll show you
- 50:09my favorite way of basically getting
- 50:11these batches and I like to personally
- 50:14have not just the input to the
- 50:15Transformer which I like to call X but I
- 50:18also like to create the labels uh tensor
- 50:21which is of the exact same size as X but
- 50:24contains the targets at every single
- 50:26position
- 50:27and so here's the way that I like to do
- 50:28that I like to make sure that I fetch
- 50:30plus one uh token because we need the
- 50:32ground Truth for the very last token uh
- 50:35for
- 50:3613 and then when we're creating the
- 50:39input we take everything up to the last
- 50:41token not including and view it as 4x6
- 50:44and when we're creating targets we do
- 50:47the buffer but starting at index one not
- 50:50index zero so we're skipping the first
- 50:52element and we view it in the exact same
- 50:54size and then when I print this
- 50:58here's what happens where we see that
- 51:00basically as an example for this token
- 51:0225 its Target was 198 and that's now
- 51:05just stored at the exact same position
- 51:07in the Target tensor which is 198 and
- 51:10also this last token 13 now has its
- 51:13label which is 198 and that's just
- 51:16because we loaded this plus one here so
- 51:19basically this is the way I like to do
- 51:20it you take long sequences you uh view
- 51:24them in two- dimensional terms so that
- 51:26you get batch of time and then we make
- 51:29sure to load one additional token so we
- 51:31basically load a buffer of tokens of B *
- 51:34t+ one and then we sort of offset things
- 51:37and view them and then we have two
- 51:39tensors one of them is the input to the
- 51:41Transformer and the other exactly is the
- 51:43labels and so let's now reorganize this
- 51:46code and um create a very simple data
- 51:50loader object that tries to basically
- 51:52load these tokens and um feed them to
- 51:55the Transformer and calculate the loss
- 51:57okay so I reshuffled the code here uh
- 51:59accordingly so as you can see here I'm
- 52:01temporarily overwriting U to run a CPU
- 52:05and importing TI token and all of this
- 52:06should look familiar we're loading a
- 52:08th000 characters I'm setting BT to just
- 52:10be 4 and 32 right now just because we're
- 52:13debugging we just want to have a single
- 52:15batch that's very small and all of this
- 52:17should now look familiar and follows
- 52:19what we did on the right and then here
- 52:21we get the we create the model and get
- 52:24the lits and so so here as you see I
- 52:28already ran this only runs in a few
- 52:30seconds but because we have a batch of
- 52:32uh 4X 32 our lits are now of size 4X 32x
- 52:3850257 so those are the lit for what
- 52:40comes next at every position and now we
- 52:43have the labels which are stored in y so
- 52:46now is the time to calculate the loss
- 52:48and then do the backward pass and then
- 52:49the optimization so let's first
- 52:51calculate the
- 52:52loss okay so to calculate the loss we're
- 52:55going to adjust the forward function of
- 52:56this NN module in the model and in
- 52:59particular we're not just going to be
- 53:00returning logits but also we're going to
- 53:02return the loss uh and we're going to
- 53:04not just pass in the input in thees but
- 53:06also the targets uh in y and now we will
- 53:12print not Lo just. shape anymore we're
- 53:14actually going to print the loss
- 53:14function and then c. exit of zero so
- 53:17that we skip some of the sampling logic
- 53:20so now let's swing up to the forward
- 53:21function which gets called there because
- 53:25now we also have these optional
- 53:28targets and when we get the targets we
- 53:30can also calculate uh the loss and
- 53:32remember that we want to basically
- 53:34return uh log just loss and loss by
- 53:36default is none
- 53:39but
- 53:40um let's put this here if uh targets is
- 53:45not none then we want to calculate loss
- 53:49and co-pilot is already getting excited
- 53:51here and calculating the what looks to
- 53:53be correct loss it is using the cross
- 53:55entropy loss as is documented here uh so
- 54:00this is a function in pytorch under the
- 54:03functional now what is actually
- 54:05happening here because it looks a little
- 54:06bit scary uh basically uh the F that
- 54:09cross entropy does not like
- 54:10multi-dimensional inputs it can't take a
- 54:12b BYT by vocap size so what's happening
- 54:15here is that we are flattening out this
- 54:17three-dimensional tensor into just two
- 54:19Dimensions the First Dimension is going
- 54:21to be calculated automatically and it's
- 54:23going to be B * T and then the last
- 54:26Dimension is vocap size so basically
- 54:28this is uh flattening out this
- 54:30three-dimensional tensor of logits to
- 54:32just be two- dimensional B * T all
- 54:35individual examples and vocap size on uh
- 54:39in terms of the length of each row and
- 54:41then it's also flattening out the
- 54:42targets which are also two- dimensional
- 54:44at this stage but we're going to just
- 54:46flatten them out so they're just a
- 54:48single tensor of B * T and this can then
- 54:51pass into cross entropy to calculate a
- 54:52loss which we return so this should
- 54:55basically at this point run because this
- 54:57is not too complicated
- 54:59so let's run it and let's see if we
- 55:03should be printing the
- 55:09loss and here we see that we printed 11
- 55:12uh roughly and so
- 55:16um and notice that this is the tensor of
- 55:18a single element which is this number 11
- 55:21now we also want to be able to calculate
- 55:23a reasonable uh kind of starting point
- 55:25for a random rationalized Network so we
- 55:27covered this in previous videos but our
- 55:29vocabulary size is
- 55:3150257 at initialization of the network
- 55:34you would hope that um every vocab
- 55:37element is getting roughly a uniform
- 55:40probability uh so that we're not
- 55:42favoring at initialization any token way
- 55:45too much we're not confidently wrong at
- 55:47initialization so what we're hoping is
- 55:49that the probability of any arbitrary
- 55:51token is roughly 1 over 50,2 57 and now
- 55:55we can sanity check the loss because
- 55:57remember that the cross entropy loss is
- 55:59just basically the negative um log
- 56:01likelihood so if we now take this
- 56:04probability and we take it through the
- 56:06natural logarithm and then we do the
- 56:08negative that is the loss we expect at
- 56:11initialization and we covered this in
- 56:13previous videos so I would expect
- 56:15something around 10.82 and we're seeing
- 56:17something around 11 so it's not way off
- 56:20this is roughly the probability I expect
- 56:21at initialization so that tells me that
- 56:24the at initialization or probability
- 56:26distribtion is roughly diffused it's a
- 56:27good starting point and we can now uh
- 56:30perform the optimization and tell the
- 56:32network which elements you know should
- 56:34follow correctly in what order so at
- 56:37this point we can do a l step backward
- 56:39calculate the gradients and do an
- 56:40optimization so let's get to that okay
- 56:43so let's do the optimization now um so
- 56:46here we
- 56:47have the loss is this is how we get the
- 56:51loss but now basically we want a load
- 56:53for Loop here so 4 I in range let's do
- 56:5550 steps or something like that uh let's
- 56:58create an Optimizer object in
- 57:00pytorch um and so here we are using the
- 57:04atom um Optimizer which is an
- 57:07alternative to the stochastic radian
- 57:08descent Optimizer SGD that we were using
- 57:11so SGD is a lot simpler atom is a bit
- 57:13more involved and I actually
- 57:14specifically like the atom W variation
- 57:17because in my opinion it kind of just
- 57:19like fixes a bug um so adom w is a bug
- 57:22fix of atom is what I would say when we
- 57:25go to the documentation for atom
- 57:27W oh my
- 57:29gosh we see um that it takes a bunch of
- 57:32hyper parameters and it's a little bit
- 57:34more complicated than the SGD we were
- 57:35looking at before uh because in addition
- 57:37to basically updating the parameters
- 57:39with the gradient uh scaled by the
- 57:41Learning rate it keeps these buffers
- 57:43around and it keeps two buffers the m
- 57:46and the V which it calls the first and
- 57:48the second moment so something that
- 57:49looks a bit like momentum and something
- 57:51that looks a bit like RMS prop if you're
- 57:53familiar with it but you don't have to
- 57:55be it's just kind of a normalization
- 57:57that happens on each gradient element
- 57:59individually and speeds up the
- 58:00optimization especially for language
- 58:02models but I'm not going to go into the
- 58:04detail right here we're going to treat
- 58:06it as a bit of a black box and it just
- 58:08optimizes um the objective faster than
- 58:12SGD which is what we've seen in the
- 58:13previous lectures so let's use it as a
- 58:15black box in our case uh create the
- 58:18optimizer object and
- 58:21then go through the optimization
- 58:28the first thing to always make sure the
- 58:30co-pilot did not forget to zero the
- 58:32gradients so um always remember that you
- 58:35have to start with a zero gradient then
- 58:38when you get your loss and you do a DOT
- 58:39backward dot backward adds to gradients
- 58:42so it deposits gradients it it always
- 58:44does a plus equals on whatever the
- 58:46gradients are which is why you must set
- 58:48them to zero so this accumulates the
- 58:50gradient from this loss and then we call
- 58:52the step function on the optimizer to um
- 58:56update the parameters and to um decrease
- 59:00the
- 59:00loss and then we print a step and the
- 59:03loss do item is used here because loss
- 59:06is a tensor with a single element do
- 59:08item will actually uh convert that to a
- 59:11single float and this float will live
- 59:13not will will live on the CPU so this
- 59:16gets to some of the internals again of
- 59:17the devices but loss is a is a tensor
- 59:20with a single element and it lifts on
- 59:22GPU for me because I'm using gpus when
- 59:25you call item P torch behind the scenes
- 59:28will take that one-dimensional tensor
- 59:30ship it back to the CPU uh memory and
- 59:32convert it into a float that we can just
- 59:35print so this is the optimization and
- 59:38this should probably just
- 59:42work let's see what
- 59:45happens actually sorry let me instead of
- 59:47using CPU override let me delete that so
- 59:50this is a bit faster for me and it runs
- 59:52on Cuda
- 59:58oh expected all tensors to be on the
- 1:00:00same device but found at least two
- 1:00:02devices Cuda zero and CPU so Cuda zero
- 1:00:06is the zeroth GPU because I actually
- 1:00:07have eight gpus on this box uh so the
- 1:00:10zeroth GPU in my box and CPU and model
- 1:00:14we have moved to device but when I was
- 1:00:17writing this code I actually introduced
- 1:00:18a bug because buff we never moved to
- 1:00:21device and you have to be careful
- 1:00:23because you can't just do buff dot two
- 1:00:25of
- 1:00:26device um it's not stateful it doesn't
- 1:00:30convert it to be a device it instead uh
- 1:00:33returns pointer to a new memory which is
- 1:00:35on the device so you see how we can just
- 1:00:37do model that two a device that does not
- 1:00:39apply to tensors you have to do buff
- 1:00:42equals
- 1:00:44um b.2 device and then this should work
- 1:00:49okay so what do we expect to see we
- 1:00:52expect to see a reasonable loss in the
- 1:00:53beginning and then we continue to
- 1:00:55optimize just the single batch and so we
- 1:00:57want to see that we can overfit this
- 1:00:58single batch we can we can crush this
- 1:01:01little batch and we can perfectly
- 1:01:02predict the indices on just this little
- 1:01:04batch and indeed that is roughly what
- 1:01:06we're seeing here
- 1:01:08so um we started off at roughly 10.82 11
- 1:01:12in this case and then as we continue
- 1:01:14optimizing on this single batch without
- 1:01:16loading new examples we are making sure
- 1:01:17that we can overfit a single batch and
- 1:01:20we are getting to very very low loss so
- 1:01:21the Transformer is memorizing this
- 1:01:24single individual batch and one more
- 1:01:26thing I didn't mention is uh the
- 1:01:28learning rate here is 3 E4 which is a
- 1:01:30pretty good default for most uh
- 1:01:33optimizations that you want to run at a
- 1:01:35very early debugging stage so this is
- 1:01:38our simple inter Loop and uh we are
- 1:01:41overfitting a single batch and this
- 1:01:42looks good so now what uh what comes
- 1:01:45next is we don't just want to overfit a
- 1:01:46single batch we actually want to do an
- 1:01:48optimization so we actually need to
- 1:01:50iterate these XY batches and create a
- 1:01:52little data loader uh that makes sure
- 1:01:54that we're always getting a fresh batch
- 1:01:56and that we're actually optimizing a
- 1:01:57reasonable objective so let's do that
- 1:01:59next okay so this is what I came up with
- 1:02:01and I wrote a little data loader
- 1:02:03light um so what this data loader does
- 1:02:06is we're importing the token up here
- 1:02:08we're reading the entire text file from
- 1:02:10this single input.txt
- 1:02:12tokenizing it and then we're just
- 1:02:14printing the number of tokens in total
- 1:02:17and the number of batches in a single
- 1:02:19Epoch of iterating over this data set so
- 1:02:22how many unique batches do we output
- 1:02:24before we loop back around the beginning
- 1:02:26of the document and start reading it
- 1:02:28again so we start off at position zero
- 1:02:31and then we simply walk the document in
- 1:02:33batches of B * T so we take chunks of B
- 1:02:36* T and then always Advance by B * T and
- 1:02:40um it's important to note that we're
- 1:02:42always advancing our position by exactly
- 1:02:44B * T but when we're fetching the tokens
- 1:02:47we're actually fetching from current
- 1:02:49position to B * t + 1 and we need that
- 1:02:52plus one because remember uh we need the
- 1:02:55target token
- 1:02:56um for the last token in the current
- 1:02:58batch and so that way we can do um the
- 1:03:02XY exactly as we did it before and if we
- 1:03:07are to um run out of data we'll just
- 1:03:09loop back around to zero so this is one
- 1:03:12way to write a very very simple data
- 1:03:13loader um that simply just goes through
- 1:03:16the file in chunks and is good enough
- 1:03:19for us uh for current purposes and we're
- 1:03:21going to complexify it later and now
- 1:03:24we'd like to come back around here and
- 1:03:26we'd like to actually use our data
- 1:03:27loader so the import Tik token has moved
- 1:03:29up and actually all of this is now
- 1:03:32useless so instead we just want a train
- 1:03:35loader for the training data and we want
- 1:03:38to use the same hyper parameters for
- 1:03:39four so B size was four and time was
- 1:03:4332 and then here we need to get the XY
- 1:03:47for the current batch so let's see if
- 1:03:49copal gets it because this is simple
- 1:03:51enough uh so we call the next batch and
- 1:03:53then we um make sure that we have to
- 1:03:57move our tensors from CPU to the device
- 1:04:02so here when I converted the tokens
- 1:04:05notice that I didn't actually move these
- 1:04:06tokens to the GPU I left them on CPU
- 1:04:10which is the default um and that's just
- 1:04:12because I'm trying not to waste too much
- 1:04:14memory on the GPU in this case this is a
- 1:04:16tiny data set and it would fit uh but
- 1:04:19it's fine to just uh ship it to GPU
- 1:04:21right now for for our purposes right now
- 1:04:24so we get the next batch we keep the
- 1:04:26data loader simple CPU class and then
- 1:04:29here we actually ship it to the GPU and
- 1:04:31do all the computation and uh let's see
- 1:04:34if this runs so python train gbt2 pi and
- 1:04:39what do we expect to see before this
- 1:04:41actually happens what we expect to see
- 1:04:43is now we're actually getting the next
- 1:04:44batch so we expect to not overfit a
- 1:04:47single batch and so I expect our loss to
- 1:04:50come down but not too much and that's
- 1:04:54because I still expect it to come down
- 1:04:55because in the
- 1:04:5750257 tokens many of those tokens never
- 1:05:00occur in our data set so there are some
- 1:05:02very easy gains to be made here in the
- 1:05:04optimization by for example taking the
- 1:05:06biases of all the loits that never occur
- 1:05:08and driving them to negative infinity
- 1:05:11and that would basically just it's just
- 1:05:12that all of these crazy unic codes or
- 1:05:14different languages those tokens never
- 1:05:16occur so their probability should be
- 1:05:17very low and so the gains that we should
- 1:05:19be seeing are along the lines of
- 1:05:22basically deleting the usage of tokens
- 1:05:24that never occur that's probably most of
- 1:05:26the loss gain that we're going to see at
- 1:05:28this scale right now uh but we shouldn't
- 1:05:30come to a zero uh because um we are only
- 1:05:35doing 50 iterations and I don't think
- 1:05:37that's enough to do an eoch right now so
- 1:05:39let's see what we
- 1:05:40got we um we have 338,000
- 1:05:44tokens which makes sense with our 3:1
- 1:05:47compression ratio because there are 1
- 1:05:48million uh characters so one Epoch with
- 1:05:52the current setting of B and T will take
- 1:05:552, 600 batches and we're only doing 50
- 1:05:58batches of optimization in
- 1:06:01here so we start off in a familiar
- 1:06:03territory as expected and then we seem
- 1:06:05to come down to about 6.6 so basically
- 1:06:09things seem to be working okay right now
- 1:06:11with respect to our expectations so
- 1:06:13that's good okay next I want to actually
- 1:06:16fix a bug that we have in our code um
- 1:06:18it's not a major bug but it is a bug
- 1:06:20with respect to how gpt2 training uh
- 1:06:22should
- 1:06:24happen um
- 1:06:26so the buck is the following we were not
- 1:06:28being careful enough when we were
- 1:06:29loading the weights from hugging face
- 1:06:31and we actually missed a little detail
- 1:06:33so if we come
- 1:06:35here notice that um the shape of these
- 1:06:38two tensors is the same so this one here
- 1:06:42is the token embedding at the bottom of
- 1:06:44the
- 1:06:45Transformer right so and this one here
- 1:06:48is the language modeling head at the top
- 1:06:50of the
- 1:06:51Transformer and both of these are
- 1:06:53basically two-dimensional tensors and
- 1:06:55they shape is identical so here the
- 1:06:59first one is the output embedding the
- 1:07:00token embedding and the second one is
- 1:07:02this linear layer at the very top the
- 1:07:04classifier layer both of them are of
- 1:07:07shape
- 1:07:0850257 X
- 1:07:09768 um this one here is giving us our
- 1:07:13token embeddings at the bottom and this
- 1:07:16one here is taking the 768 channels of
- 1:07:18the Transformer and trying to upscale
- 1:07:21that to 50, 257 to get the Lis for the
- 1:07:24next token so they're both the same
- 1:07:27shape but more than that actually if you
- 1:07:29look at um comparing their elements um
- 1:07:33in pytorch this is an element wise
- 1:07:35equality so then we use do all and we
- 1:07:37see that every single element is
- 1:07:39identical and more than that we see that
- 1:07:42if we actually look at the data pointer
- 1:07:44uh this is what this is a way in pytorch
- 1:07:47to get the actual pointer to the uh data
- 1:07:49and the storage we see that actually the
- 1:07:51pointer is identical so not only are
- 1:07:53these two separate tensors that happen
- 1:07:55to have the same shape and elements
- 1:07:57they're actually pointing to the
- 1:07:58identical tensor so what's happening
- 1:08:02here is that this is a common weight
- 1:08:03tying scheme uh that actually comes from
- 1:08:06the original
- 1:08:08um from the original attention is all
- 1:08:10you need paper and actually even the
- 1:08:12reference before it so if we come
- 1:08:16here
- 1:08:19um eddings and softmax in the attention
- 1:08:22is all you need paper they mentioned
- 1:08:24that in our model we shared the same
- 1:08:26weight Matrix between the two embedding
- 1:08:28layers and the pre softmax linear
- 1:08:30transformation similar to 30 um so this
- 1:08:34is an awkward way to phrase that these
- 1:08:36two are shared and they're tied and
- 1:08:38they're the same Matrix and the 30
- 1:08:40reference is this
- 1:08:42paper um so this came out in
- 1:08:452017 and you can read the full paper but
- 1:08:47basically it argues for this weight
- 1:08:49tying scheme and I think intuitively the
- 1:08:53idea for why you might want to do this
- 1:08:54comes from from this paragraph here and
- 1:08:58basically you you can observe
- 1:09:01that um you actually want these two
- 1:09:04matrices to behave similar in the
- 1:09:07following sense if two tokens are very
- 1:09:10similar semantically like maybe one of
- 1:09:12them is all lowercase and the other one
- 1:09:14is all uppercase or it's the same token
- 1:09:16in a different language or something
- 1:09:17like that if you have similarity between
- 1:09:19two tokens presumably you would expect
- 1:09:21that they are uh nearby in the token
- 1:09:23embedding space but in the exact same
- 1:09:26way you'd expect that if you have two
- 1:09:27tokens that are similar semantically
- 1:09:30you'd expect them to get the same
- 1:09:32probabilities at the output of a
- 1:09:33transformer because they are
- 1:09:35semantically similar and so both
- 1:09:39positions in the Transformer at the very
- 1:09:41bottom and at the top have this property
- 1:09:43that similar tokens should have similar
- 1:09:46embeddings or similar weights and so
- 1:09:49this is what motivates their exploration
- 1:09:51here and they they kind of you know I
- 1:09:53don't want to go through the entire
- 1:09:54paper and and uh you can go through it
- 1:09:57but this is what they observe they also
- 1:09:59observe that if you look at the output
- 1:10:00embeddings they also behave like word
- 1:10:02embeddings um if you um if you just kind
- 1:10:06of try to use those weights as word
- 1:10:08embeddings um so they kind of observe
- 1:10:10this similarity they try to tie them and
- 1:10:13they observe that they can get much
- 1:10:14better performance in that way and so
- 1:10:17this was adopted and the attention is
- 1:10:18all need paper and then it was used
- 1:10:20again in gpt2 as well
- 1:10:24so I couldn't find it in the
- 1:10:26Transformers implementation I'm not sure
- 1:10:28where they tie those embeddings but I
- 1:10:30can find it in the original gpt2 code U
- 1:10:34introduced by open aai so this is um
- 1:10:36openai gpt2 Source model and here where
- 1:10:40they are forwarding this model and this
- 1:10:41is in tensorflow but uh that's okay we
- 1:10:44see that they get the wte token
- 1:10:46embeddings and then here is the incoder
- 1:10:50of the token embeddings and the
- 1:10:52position and then here at the bottom
- 1:10:54they Ed the WT again to do the lits so
- 1:10:58when they get the loits it's a math Mo
- 1:11:00of uh this output from the Transformer
- 1:11:02and the wte tensor is
- 1:11:05reused um and so the wte tensor
- 1:11:08basically is used twice on the bottom of
- 1:11:10the Transformer and on the top of the
- 1:11:12Transformer and in the backward pass
- 1:11:14we'll get gradients contributions from
- 1:11:17both branches right and these gradients
- 1:11:19will add up um on the wte tensor um so
- 1:11:23we'll get a contribution from the
- 1:11:24classifier list
- 1:11:25and then at the very end of the
- 1:11:27Transformer we'll get a contribution at
- 1:11:28the at the bottom of it float floating
- 1:11:31again into the wte uh tensor so we want
- 1:11:35to we are currently not sharing WT and
- 1:11:38our code but we want to do
- 1:11:40that um
- 1:11:44so weight sharing scheme um and one way
- 1:11:48to do this let's see if goil gets it oh
- 1:11:50it does okay uh so this is one way to do
- 1:11:54it
- 1:11:56uh
- 1:11:56basically relatively straightforward
- 1:11:59what we're doing here is we're taking
- 1:12:00the wte do weight and we're simply uh
- 1:12:04redirecting it to point to the LM head
- 1:12:08so um this basically copies the data
- 1:12:11pointer right it copies the reference
- 1:12:14and now the wte weight becomes orphaned
- 1:12:17uh the old value of it and uh pytorch
- 1:12:20will clean it up python will clean it up
- 1:12:23and so we are only left with a single
- 1:12:26tensor and it's going to be used twice
- 1:12:28in the forward pass and uh this is to my
- 1:12:31knowledge all that's required so we
- 1:12:34should be able to use this and this
- 1:12:36should probably train uh we're just
- 1:12:39going to basically be using this exact
- 1:12:40same sensor twice and
- 1:12:44um we weren't being careful with
- 1:12:46tracking the likelihoods but uh
- 1:12:48according to the paper and according to
- 1:12:50the results you'd actually expect
- 1:12:51slightly better results doing this and
- 1:12:53in addition to that one other reason
- 1:12:54that this is very very nice for us is
- 1:12:57that this is a ton of parameters right
- 1:12:59uh what is the size here it's 768 *
- 1:13:0350257 so This Is 40 million parameters
- 1:13:07and this is a 124 million parameter
- 1:13:09model so 40 divide 124 so this is like
- 1:13:1230% of the parameters are being saved
- 1:13:15using this weight time scheme and so
- 1:13:18this might be one of the reasons that
- 1:13:20this is working slightly better if
- 1:13:21you're not training the model long
- 1:13:22enough because of the weight tying uh
- 1:13:25you don't have to train as many
- 1:13:26parameters and so you become more
- 1:13:27efficient um in terms of the training
- 1:13:30process uh because you have fewer
- 1:13:32parameters and you're putting in this
- 1:13:34inductive bias that these two embeddings
- 1:13:36should share similarities between tokens
- 1:13:40so this is the way time scheme and we've
- 1:13:42saved a ton of parameters and we expect
- 1:13:44our model to work slightly better
- 1:13:45because of the scheme okay next I would
- 1:13:47like us to be a bit more careful with
- 1:13:49the initialization and to try to follow
- 1:13:50the way gpt2 initialized their model now
- 1:13:54unfortunately the gpt2 paper and the
- 1:13:55gpt3 paper are not very explicit about
- 1:13:58initialization so we kind of have to
- 1:14:00read between the lines uh and instead of
- 1:14:02going to the paper which is quite vague
- 1:14:04um there's a bit of information in the
- 1:14:07code that open I released so when we go
- 1:14:09to the model.py we see that when they
- 1:14:11initialize their weights they are using
- 1:14:13the standard deviation of
- 1:14:150.02 and that's how they they so this is
- 1:14:19a normal distribution for the weights
- 1:14:21and the standard deviation is
- 1:14:230.02 for the bias they initialize that
- 1:14:25with
- 1:14:26zero and then when we scroll down
- 1:14:30here why is this not scrolling
- 1:14:33um the token embeddings are initialized
- 1:14:36at
- 1:14:370.02 and position embeddings at 0.01 for
- 1:14:40some reason so those are the
- 1:14:42initializations and we'd like to mirror
- 1:14:44that in
- 1:14:45gpt2 uh in our module here so here's a
- 1:14:48snippet of code that I sort of came up
- 1:14:50with very
- 1:14:52quickly so what's happening here is at
- 1:14:55the end of our initializer for the GPT
- 1:14:57module we're calling the apply function
- 1:14:59of NN module and that iterates all the
- 1:15:02sub modules of this module and uh
- 1:15:05applies in it weights function on them
- 1:15:08and so what's happening here is that
- 1:15:11we're in we're iterating all the modules
- 1:15:13here and if they are an nn. linear
- 1:15:16module then we're going to make sure to
- 1:15:17initialize the weight using a normal
- 1:15:19with the standard deviation of
- 1:15:210.02 if there's a bias in this layer we
- 1:15:24will make sure to initialize that to
- 1:15:25zero note that zero initialization for
- 1:15:28the bias is not actually the pyto
- 1:15:29default um by default the bias here is
- 1:15:33initialized with a uniform so uh that's
- 1:15:36interesting so we make sure to use zero
- 1:15:38and for the embedding we're just going
- 1:15:40to use 0.02 and um keep it the same um
- 1:15:43so we're not going to change it to 0.01
- 1:15:45for positional because it's about the
- 1:15:47same and then if you look through our
- 1:15:49model the only other layer that requires
- 1:15:51initialization and that has parameters
- 1:15:53is the layer norm and the fighter defer
- 1:15:55initialization sets the scale in the
- 1:15:57layer Norm to be one and the offset in
- 1:16:00the layer Norm to be zero so that's
- 1:16:01exactly what we want and so we're just
- 1:16:03going to uh keep it that way and so this
- 1:16:06is the default initialization if we are
- 1:16:09following the um where is it the uh gpt2
- 1:16:14uh source code that they released I
- 1:16:17would like to point out by the way that
- 1:16:19um typically the standard deviation here
- 1:16:21on this initialization if you follow the
- 1:16:23Javier initialization would be one of
- 1:16:24over the square root of the number of
- 1:16:27features that are incoming into this
- 1:16:28layer but if you'll notice actually 0.02
- 1:16:32is basically consistent with that
- 1:16:34because the the model sizes inside these
- 1:16:36Transformers for gpt2 are roughly 768
- 1:16:391600 Etc so 1 over the square root of
- 1:16:41for example 768 gives us
- 1:16:440.03 if we plug in 600 1,600 we get
- 1:16:490.02 if we plug in three times that
- 1:16:520.014 Etc so basically 0.02 is roughly
- 1:16:56in the vicinity of reasonable values for
- 1:16:59the for um for these initializations
- 1:17:02anyway so so it's not uh completely
- 1:17:05crazy to be hard coding 0.02 here uh but
- 1:17:08you'd like typically uh some something
- 1:17:11that grows with the model size instead
- 1:17:13but we will keep this because that is
- 1:17:15the gpt2 initialization per their source
- 1:17:17code but we are not fully done yet on
- 1:17:19initialization because there's one more
- 1:17:20caveat here so
- 1:17:23here a mod initialization which accounts
- 1:17:26for the accumulation on the residual
- 1:17:27path with model depth is used we scale
- 1:17:30the weight of residual layers of
- 1:17:31initialization by factor of one over squ
- 1:17:33of n where n is the number of residual
- 1:17:35layers so this is what gbt2 paper says
- 1:17:38so we have not implemented that yet and
- 1:17:41uh we can do so now now I'd like to
- 1:17:43actually kind of like motivate a little
- 1:17:44bit what they mean here I think um so
- 1:17:47here's roughly what they
- 1:17:49mean if you start out with zeros in your
- 1:17:52residual stream remember that each
- 1:17:54residual stream is a is of this form
- 1:17:57where we continue adding to it X is X
- 1:18:00plus something some kind of contribution
- 1:18:02so every single block of the residual uh
- 1:18:05Network contributes some uh amount and
- 1:18:09it gets added and so what ends up
- 1:18:11happening is that the variance of the
- 1:18:15activations in the residual stream grows
- 1:18:18so here's a small example if we start at
- 1:18:19zero and then we for 100 times uh we
- 1:18:23have sort of this residual stream of of
- 1:18:25768 uh zeros and then 100 times we add
- 1:18:30um random which is a normal distribution
- 1:18:33zero mean one standard deviation if we
- 1:18:36add to it then by the end the residual
- 1:18:37stream has grown to have standard
- 1:18:39deviation of 10 and that's just because
- 1:18:42um we're always adding um these numbers
- 1:18:47and so this scaling factor that they use
- 1:18:50here exactly compensates for that growth
- 1:18:53so if we take n and we basically um
- 1:18:57scale down every one of these
- 1:18:59contributions into the residual stream
- 1:19:00by one over theare Ro of n so 1 over
- 1:19:03theun of n is n to the 0.5
- 1:19:07right because n the5 is the square root
- 1:19:11and then one over the square root is n.5
- 1:19:14if we scale it in this way then we see
- 1:19:16that we actually get um
- 1:19:20one
- 1:19:21so this is a way to control the growth
- 1:19:24of of activations inside the residual
- 1:19:26stream in the forward pass and so we'd
- 1:19:29like to initialize in the same way where
- 1:19:31these weights that are at the end of
- 1:19:33each block so this C uh layer uh the gbt
- 1:19:38paper proposes to scale down those
- 1:19:40weights by one over the square root of
- 1:19:42the number of residual
- 1:19:43layers so one crude way to implement
- 1:19:46this is the following I don't know if
- 1:19:48this is uh pyro sanctioned but it works
- 1:19:50for me is we'll do in the
- 1:19:53initialization see that s that do
- 1:19:56special nanog
- 1:19:58GPT uh scale in it is one so we're
- 1:20:04setting um kind of like a flag for this
- 1:20:06module there must be a better way in py
- 1:20:08torch right but I don't
- 1:20:11know okay so we're basically attaching
- 1:20:13this flag and trying to make sure that
- 1:20:16it doesn't conflict with anything
- 1:20:17previously and then when we come down
- 1:20:20here this STD should be 0.02 by default
- 1:20:25but then if
- 1:20:27haat um module of this thing
- 1:20:31then STD *
- 1:20:34equals
- 1:20:36um copal is not guessing correctly uh so
- 1:20:39we want one over the square root of the
- 1:20:41number of layers so
- 1:20:44um the number of residual layers here is
- 1:20:47twice
- 1:20:48times Salt out config layers and then
- 1:20:52this times .5 so we want to scale down
- 1:20:57that standard deviation and this should
- 1:20:59be um correct and Implement that I
- 1:21:03should clarify by the way that the two
- 1:21:04times number of layers comes from the
- 1:21:06fact that every single one of our layers
- 1:21:07in the Transformer actually has two
- 1:21:09blocks that add to the ridal pathway
- 1:21:11right we have the attention and then the
- 1:21:13MLP so that's where the two times comes
- 1:21:16from and the other thing to mention is
- 1:21:18that uh what's slightly awkward but
- 1:21:21we're not going to fix it is that um
- 1:21:23because we are weight sharing the wte
- 1:21:26and the LM head in this iteration of our
- 1:21:29old subm modules we're going to actually
- 1:21:31come around to that tensor twice so
- 1:21:33we're going to first initialize it as an
- 1:21:34embedding with 0.02 and then we're going
- 1:21:37to come back around it again in a linear
- 1:21:39and initialize it again using 0.02 and
- 1:21:42it's going to be 0.02 because the LM
- 1:21:44head is of course not not scaled so it's
- 1:21:46not going to come here it's just it's
- 1:21:48going to be basically initialized twice
- 1:21:50using the identical same initialization
- 1:21:52but that's okay and then scrolling over
- 1:21:56here I added uh some code here so that
- 1:21:59we have
- 1:22:00reproducibility um to set the seeds and
- 1:22:03now we should be able to python train
- 1:22:05gpt2 pi and let this running and as far
- 1:22:09as I know this is the gpt2
- 1:22:11initialization uh in the way we've
- 1:22:12implemented it right now so this
- 1:22:16looks uh reasonable to me okay so at
- 1:22:19this point we have the gpt2 model we
- 1:22:21have some confidence that it's correctly
- 1:22:23implemented we've initialized it
- 1:22:24properly and we have a data loader
- 1:22:26that's iterating through data batches
- 1:22:27and we can train so now comes the fun
- 1:22:30part I'd like us to speed up the
- 1:22:31training by a lot so we're getting our
- 1:22:33money's worth with respect to the
- 1:22:34hardware that we are uh using here and
- 1:22:38uh we're going to speed up the training
- 1:22:39by quite a bit uh now you always want to
- 1:22:42start with what Hardware do you have
- 1:22:44what does it offer and are you fully
- 1:22:45utilizing it so in my case if we go to
- 1:22:48Nvidia
- 1:22:49SMI we can see
- 1:22:53that I have eight gpus and each one of
- 1:22:57those gpus is an a100 sxm 80 gb so this
- 1:23:01is the GPU that I have available to me
- 1:23:03in this box now when I look when I use
- 1:23:07um to spin up these kinds of Boxes by
- 1:23:09the way my favorite place to go to is
- 1:23:11Lambda Labs um they do sponsor my
- 1:23:14development and that of my projects uh
- 1:23:17but I this is my favorite place to go
- 1:23:20and this is where you can spin up one of
- 1:23:21these machines and you pay per hour and
- 1:23:23it's very very simple
- 1:23:25so I like to spin them up and then
- 1:23:26connect vsod to it and that's how I
- 1:23:28develop now when we look at the A1 100s
- 1:23:30that are available here a100 80 GB sxm
- 1:23:35is the um GPU that I have here and we
- 1:23:39have a bunch of numbers here for um how
- 1:23:41many calculations you can expect out of
- 1:23:43this GPU so when I come over here
- 1:23:46and I break in right after here so
- 1:23:50python
- 1:23:51trity so I'm breaking in right after we
- 1:23:53calculate the loit and
- 1:23:55laws and the interesting thing I'd like
- 1:23:57you to note is when I do lit. dtype this
- 1:24:02prints a torch. FL 32 so by default iny
- 1:24:06torch when you create tensors um and
- 1:24:08this is the case for all the activations
- 1:24:10and for the parameters of the network
- 1:24:11and so on by default everything is in
- 1:24:13float 32 that means that every single
- 1:24:17number activation or weight and so on is
- 1:24:20using a float representation that has 32
- 1:24:23bits and uh that's actually quite a bit
- 1:24:26of memory and it turns out empirically
- 1:24:27that for deep learning as a
- 1:24:28computational workload this is way too
- 1:24:30much and deep learning and the training
- 1:24:32of these networks can tolerate
- 1:24:34significantly lower precisions um not
- 1:24:37all computational workflows can tolerate
- 1:24:39small Precision so for example um if we
- 1:24:43go back to to the data sheet you'll see
- 1:24:45that actually these gpus support up to
- 1:24:48fp64 and this is quite useful I
- 1:24:50understand for a lot of um scientific
- 1:24:52Computing applications and there really
- 1:24:54need this uh but we don't need that much
- 1:24:56Precision for deep learning training So
- 1:24:59currently we are here
- 1:25:01fp32 and with this code as it is right
- 1:25:04now we expect to get at at most 19.5
- 1:25:08Tera flops of performance that means
- 1:25:10we're doing 19.5 trillion operations
- 1:25:13floating Point operations so this is
- 1:25:15floating Point multiply add most um most
- 1:25:20likely and so these are the floating
- 1:25:23Point operations
- 1:25:25uh now notice that if we are willing to
- 1:25:27go down in Precision so tf32 is a lower
- 1:25:31Precision format we're going to see in a
- 1:25:32second you can actually get an 8X
- 1:25:34Improvement here and if you're willing
- 1:25:36to go down to float 16 or B float 16 you
- 1:25:39can actually get time 16x performance
- 1:25:42all the way to 312 Tera flops you see
- 1:25:45here that Nvidia likes to site numbers
- 1:25:47that have an asterisk here this asterisk
- 1:25:50uh says with sparsity uh but we are not
- 1:25:52going to be using sparsity in R code and
- 1:25:55I don't know that this is very widely
- 1:25:56used in the industry right now so most
- 1:25:58people look at this number here uh
- 1:26:01without sparcity and you'll notice that
- 1:26:03we could have got even more here but
- 1:26:05this is int 8 and int 8 is used for
- 1:26:08inference not for training uh because
- 1:26:11int 8 has a um it basically has um
- 1:26:17uniform
- 1:26:18spacing um and uh we actually require a
- 1:26:21float so that we get a better match to
- 1:26:24the uh normal distributions that occur
- 1:26:28during training of neural networks where
- 1:26:29both activations and weights are
- 1:26:31distributed as a normal distribution and
- 1:26:33so uh floating points are really
- 1:26:35important to to match that uh
- 1:26:38representation so we're not typically
- 1:26:40using int 8 uh for training but we are
- 1:26:42using it for inference and if we bring
- 1:26:45down the Precision we can get a lot more
- 1:26:47Terra flops out of the tensor course
- 1:26:49available in the gpus we'll talk about
- 1:26:51that in a second but in addition to that
- 1:26:53if all of these numbers have fewer bits
- 1:26:56of representation it's going to be much
- 1:26:58easier to move them around and that's
- 1:27:00where we start to get into the memory
- 1:27:02bandwidth and the memory of the model so
- 1:27:04not only do we have a finite capacity of
- 1:27:06the number of bits that our GPU can
- 1:27:08store but in addition to that there's a
- 1:27:11speed with which you can access this
- 1:27:13memory um and you have a certain memory
- 1:27:16bandwidth it's a very precious resource
- 1:27:19and in fact many of the deep learning uh
- 1:27:21work workloads for training are memory
- 1:27:23bound and what that means is actually
- 1:27:25that the tensor cores that do all these
- 1:27:27extremely fast multiplications most of
- 1:27:29the time they're waiting around they're
- 1:27:31idle um because we can't feed them with
- 1:27:34data fast enough we can't load the data
- 1:27:37fast enough from memory so typical
- 1:27:38utilizations of your Hardware if you're
- 1:27:40getting 60% uh utilization you're
- 1:27:43actually doing extremely well um so half
- 1:27:46of the time in a well-tuned application
- 1:27:48your tensor cores are not doing
- 1:27:50multiplies because the data is not
- 1:27:51available so the memory bandwidth here
- 1:27:53is extremely important as well and if we
- 1:27:55come down in the Precision for all the
- 1:27:58floats all the numbers weights and
- 1:28:00activations suddenly require less memory
- 1:28:02so we can store more and we can access
- 1:28:05it faster so everything speeds up and
- 1:28:07it's amazing and now let's reap the
- 1:28:09benefits of it um and let's first look
- 1:28:12at the tensor float 32
- 1:28:14format okay so first of all what are
- 1:28:16tensor cores well tensor course tensor
- 1:28:19core is just an instruction in the a100
- 1:28:22architecture right so so what it does is
- 1:28:25it does basically a little 4x4 Matrix
- 1:28:27multiply so uh this is just matrix
- 1:28:30multiplication here of 4x4 matrices and
- 1:28:35there are multiple configurations as to
- 1:28:38what Precision any of these matrices are
- 1:28:40it in what Precision the internal
- 1:28:42accumulate happens and then what is the
- 1:28:45output Precision input precisions Etc so
- 1:28:47there's a few switches but it's
- 1:28:48basically a 4x4 multiply and then
- 1:28:51anytime we have any operations that
- 1:28:53require Magic multiplication uh they get
- 1:28:55broken up into these into this
- 1:28:58instruction of little 4x4 multiply and
- 1:29:00so everything gets broken up into this
- 1:29:02instruction because it's the fastest way
- 1:29:04to multiply matrices and it turns out
- 1:29:06that most of the computational work that
- 1:29:08we're doing up above uh all of it really
- 1:29:10is matrix multiplication most of the
- 1:29:12work computationally happens in the
- 1:29:14linear layers um linear linear Etc
- 1:29:20there's a few things sandwiched in
- 1:29:21between so there's some additions in
- 1:29:23residuals there's some G nonlinearities
- 1:29:25there's some layer Norms Etc but if you
- 1:29:28just time them you'll see that these are
- 1:29:30nothing like basically the in
- 1:29:32Transformer is just a bunch of Matrix
- 1:29:34multiplications really um and especially
- 1:29:37at this small scale 124 million
- 1:29:39parameter model actually the biggest
- 1:29:42matrix multiplication by far is the
- 1:29:44classifier layer at the top that is a
- 1:29:46massive Matrix multiply of going from
- 1:29:49768 to
- 1:29:5050257 and that Matrix multiply dominates
- 1:29:53anything else that happens in that
- 1:29:55Network roughly speaking so it's Matrix
- 1:29:58multiplies that become a lot faster
- 1:30:00which are hidden inside our linear
- 1:30:02layers and they're accelerated through
- 1:30:05tensor course now the best reference I
- 1:30:07would say for tensor course is basically
- 1:30:09just go to the um a 100 architecture
- 1:30:13white paper and then it's pretty
- 1:30:15detailed and but I think people it's
- 1:30:18like relatively readable mostly if you
- 1:30:20half understand what's happening um so
- 1:30:23figure 9 tensor float
- 1:30:2632 so this is the explanation basically
- 1:30:28for tf32 and what happens here and you
- 1:30:31see that there's many configuration
- 1:30:32options here available so the input
- 1:30:35operands and what precisions are they in
- 1:30:37the accumulator and um what um basically
- 1:30:41the um the internal representation
- 1:30:44within the instruction when you do the
- 1:30:46accumulate of this matrix
- 1:30:48multiplication so the intermediate plus
- 1:30:51equals um of the intermediate little
- 1:30:53vector multiplies here that all happens
- 1:30:55in
- 1:30:57fp32 and then uh this is an aex
- 1:31:00improvement as I mentioned to the Ops
- 1:31:01that we get so tf32 specifically we're
- 1:31:04looking at this row here and the way
- 1:31:06this works
- 1:31:07is
- 1:31:10um normally fp32 has 32 bits
- 1:31:14tf32 is the exact same bits we have one
- 1:31:18sign bit we have eight exponent bits
- 1:31:21except the mantisa bits get cropped in
- 1:31:24the float and so basically um we end up
- 1:31:27with just 19 bits instead of 32 bits
- 1:31:30because the last 133 bits get truncated
- 1:31:33they get dropped um and all this is
- 1:31:36internal to the instruction so none of
- 1:31:38it is visible to anything in our pytorch
- 1:31:41uh none of our pytorch code will change
- 1:31:43all of the numbers will look identical
- 1:31:45it's just that when you call the tensor
- 1:31:47core um instruction internally in the
- 1:31:50hardware it will crop out these 13 bits
- 1:31:54and that allows it to uh calculate this
- 1:31:57little Matrix multiply significantly
- 1:31:59faster 8X faster now of course this
- 1:32:02speed up comes at a cost and the cost is
- 1:32:04that we are reducing the Precision our
- 1:32:07accumulate is still an fp32 our output
- 1:32:09is fp32 our inputs are fp32 but
- 1:32:12internally things get truncated in the
- 1:32:14operand to perform the operation faster
- 1:32:17and so our results are starting to be a
- 1:32:19bit more approximate but empirically
- 1:32:21when you actually train with this you
- 1:32:22basically can't tell the difference
- 1:32:24so the reason I like tf32 is because if
- 1:32:26you can tolerate a little bit of a
- 1:32:28Precision fudge um then this is free
- 1:32:32like none of your codes sees this it's
- 1:32:34fully internal to the operation and the
- 1:32:36operation to you just go 8X faster and
- 1:32:39it's a bit more approximate and so it's
- 1:32:42a pretty sweet spot I would say in
- 1:32:43optimization and uh let's see what that
- 1:32:46looks like first so I've set up our Cod
- 1:32:48to just time the uh iterations so import
- 1:32:51time I changed the hyper parameters so
- 1:32:54that we have something a bit more that
- 1:32:55reflects uh kind of workload that we
- 1:32:57want to run uh because we want to do a
- 1:32:59fairly large run at the end of this so
- 1:33:01let's use batch size 16 and let's now
- 1:33:04use the actual gpt2 um maximum sequence
- 1:33:07length of 10,24
- 1:33:08tokens uh so this is the
- 1:33:11configuration and then for 50 iterations
- 1:33:15I'm just doing something very lazy here
- 1:33:17I'm doing time. time to get the current
- 1:33:19time and then this is the optimization
- 1:33:22Loop and now I want to time how long
- 1:33:24this takes now one issue with working
- 1:33:28with gpus is that as your
- 1:33:32CPU um when your CPU runs it's just
- 1:33:35scheduling work on GPU it's ordering
- 1:33:38some work right and so it send a request
- 1:33:40and then it continues running and so we
- 1:33:43can actually it can happen sometimes
- 1:33:44that we sort of um speed through this
- 1:33:48and we queue up a lot of kernels to run
- 1:33:50on the GPU and then the CPU sort of like
- 1:33:52gets here and takes time at time but
- 1:33:54actually the GPU is still running
- 1:33:56because it takes it time to actually
- 1:33:57work through the work that was scheduled
- 1:34:00to run and so you're just building up a
- 1:34:03queue for the GPU and so actually if you
- 1:34:05need to you want to wait toat data
- 1:34:07synchronize and this will wait for the
- 1:34:10GPU to finish all the work that was
- 1:34:12scheduled to run up above here and then
- 1:34:15we can actually take the time so
- 1:34:17basically we're waiting for the GPU to
- 1:34:19stop this iteration take time and then
- 1:34:22we're going to just print it so
- 1:34:24so here I'm going to run the training
- 1:34:26Loop and here on the right I'm watching
- 1:34:29Nvidia SMI so we start off at zero um
- 1:34:33we're not using the GPU and then by
- 1:34:35default P will use gpu0 so we see that
- 1:34:37it gets filled up and we're using 35 GB
- 1:34:40out of 80 gabt
- 1:34:42available and then here on the left we
- 1:34:45see that because we've cranked up the
- 1:34:47batch
- 1:34:48size now it's only 20 batches to do a
- 1:34:51single Epoch on our tiny Shakespeare
- 1:34:54and we see that we're seeing roughly a
- 1:34:55th000 milliseconds per iteration here
- 1:34:58right
- 1:35:00so the first iteration sometimes is
- 1:35:02slower and that's because pytorch might
- 1:35:04be doing a lot of initializations here
- 1:35:06on the very first iteration and so it's
- 1:35:08probably initializing all these uh
- 1:35:09tensors and buffers to hold all the
- 1:35:11gradients and I'm not 100% sure all the
- 1:35:13work that happens here but uh this could
- 1:35:16be a slower iteration when you're timing
- 1:35:18your logic you always want to be careful
- 1:35:19with that but basically we're seeing a
- 1:35:21th000 milliseconds per iteration
- 1:35:24um and so this will run for roughly 50
- 1:35:26seconds as we have it right now so
- 1:35:29that's our Baseline in flo 32 one more
- 1:35:32thing I wanted to mention is that if
- 1:35:35this doesn't fit into your GPU and
- 1:35:36you're getting out of memory errors then
- 1:35:38start decreasing your batch size until
- 1:35:40things fit so instead of 16 try eight or
- 1:35:42four or whatever you need to fit um the
- 1:35:46batch into your GPU and if you have a
- 1:35:48bigger GPU you can actually potentially
- 1:35:49get away with 32 and so on uh by default
- 1:35:52you want to basically max out has Max
- 1:35:54Max out the batch size that fits on your
- 1:35:56GPU and you want to keep it nice numbers
- 1:35:59so use numbers that have lots of powers
- 1:36:01of two in them so 16 is a good number 8
- 1:36:0524 32 48 These are nice numbers but
- 1:36:09don't use something like 17 uh because
- 1:36:11that will run very inefficiently on a
- 1:36:12GPU uh and we're going to see that a bit
- 1:36:14later as well so for now let's just
- 1:36:17stick with
- 1:36:1816124 and uh the one thing that I added
- 1:36:22also here and I ran it again is I'm
- 1:36:25calculating a tokens per second
- 1:36:27throughput during training
- 1:36:29because we might end up changing the
- 1:36:31backat size around over time but tokens
- 1:36:34per second is the objective measure that
- 1:36:35we actually really care about how many
- 1:36:37tokens of data are we training on and
- 1:36:39what is the throughput of tokens that
- 1:36:41we're getting in our optimization so
- 1:36:43right now we're processing and training
- 1:36:44on 163,000 tokens per second roughly and
- 1:36:48that's a bit more objective
- 1:36:50metric okay so let's now enable tf32 now
- 1:36:53luckily pytorch makes this fairly easy
- 1:36:56for us and uh to enable tf32 you just
- 1:36:59need to do a single line and is this and
- 1:37:02when we go to the py documentation here
- 1:37:04for this function basically this tells
- 1:37:07pych what kind of kernels to run and by
- 1:37:10default I believe it is highest highest
- 1:37:13Precision for mat M and that means that
- 1:37:15everything happens in float 32 just like
- 1:37:18it did before but if we set it to high
- 1:37:20as we do right now Matrix
- 1:37:22multiplications will not use tensor flow
- 1:37:2432 when it's
- 1:37:26available my GPU is a100 so it's an
- 1:37:30ampere series and therefore tf32 is
- 1:37:33available if you have an older GPU this
- 1:37:35might not be available for you but for
- 1:37:38my GPU it's available and so what I
- 1:37:39expect P to do is that every single
- 1:37:41place where we see an nn. linear inside
- 1:37:44there there's a matrix multiplication
- 1:37:46and I expect that matrix multiplication
- 1:37:48now to be um running on tensor course
- 1:37:51utilizing the TF 32%
- 1:37:55so this is the single line of change
- 1:37:58that is I believe necessary and let's
- 1:37:59rerun this now we saw that um in terms
- 1:38:03of the throughput that is promised to us
- 1:38:05we're supposed to be getting 8X roughly
- 1:38:08so let's see what
- 1:38:10happens and that 8X came from here right
- 1:38:15um 8X and it also came from looking at
- 1:38:20it um here 156 T flops instead of of
- 1:38:2419.5 okay so what actually happened uh
- 1:38:27so we're seeing that our throughput
- 1:38:29roughly 3x not aex so we are going we're
- 1:38:35from 1,000 milliseconds we're going down
- 1:38:37to 300 milliseconds and our throughput
- 1:38:39is now about 50,000 tokens per second so
- 1:38:41we have a roughly 3x instead of 8X so
- 1:38:43what happened and basically What's
- 1:38:46Happening Here is again a lot of these
- 1:38:48workloads are memory bound and so even
- 1:38:51though the
- 1:38:52tf32 offers in principle a lot faster
- 1:38:57throughput all of these numbers
- 1:38:59everywhere are still float 32s and it's
- 1:39:01float 32 numbers that are being shipped
- 1:39:03all over the place through the memory
- 1:39:05system and is just costing us way too
- 1:39:07much time to shuttle around all this
- 1:39:08data and so even though we've made the
- 1:39:10multiply itself much faster uh we are
- 1:39:13memory bound and we're not actually
- 1:39:14seeing the full benefit uh that would
- 1:39:16come from uh this napkin math here uh
- 1:39:19that said we are getting one a 3X faster
- 1:39:22throughput and this is free um single
- 1:39:26line of code in P torch all your
- 1:39:28variables are still float 32 everywhere
- 1:39:30it just runs faster and it's slightly
- 1:39:32more approximate but we're not going to
- 1:39:34notice it basically uh so that's
- 1:39:37tf32 okay so let's now continue so we've
- 1:39:41exercised this row and um we saw that we
- 1:39:44can crop out some of the Precision
- 1:39:46inside the operation itself but we saw
- 1:39:49that we're still memory bound we're
- 1:39:50still moving around all these floats
- 1:39:52right otherwise and we're paying that
- 1:39:53cost because of this so let's now
- 1:39:56decrease the amount of stuff that we're
- 1:39:57going to be moving around and we're
- 1:39:59going to do that by dropping down to B
- 1:40:01float 16 so we're only going to be
- 1:40:04maintaining 16 bits per float and we're
- 1:40:07going to use the B flat 16 and I'll
- 1:40:08explain in a bit uh fp16 difference and
- 1:40:12uh we're going to be in this row so when
- 1:40:14we go back to the documentation here for
- 1:40:17the a
- 1:40:18100 um we see here the precisions that
- 1:40:23are are available and this is the
- 1:40:25original fp32 the tf32 crops out the
- 1:40:28Precision and then here in
- 1:40:30bf16 you see that it is very similar to
- 1:40:33tf32 but it's even more aggressive in
- 1:40:36cropping off of the Precision the
- 1:40:38mantisa of this float so the important
- 1:40:40thing with B float 16 is that the
- 1:40:42exponent bits and the sign bit of course
- 1:40:45remain unchanged so if you're familiar
- 1:40:47with your float numbers and I think this
- 1:40:49should should probably be an entire
- 1:40:52video by itself
- 1:40:53the exponent sets the range that you can
- 1:40:56represent of your numbers and the
- 1:40:58Precision is how much Precision you have
- 1:41:00for your numbers and so the range of
- 1:41:04numbers is identical but we can we have
- 1:41:07fewer possibilities within that range
- 1:41:10because we are truncating the Mena so we
- 1:41:12have less Precision in that
- 1:41:14range what that means is that things are
- 1:41:17actually fairly nice because we have the
- 1:41:19original range of numbers that are
- 1:41:21representable in float but we just have
- 1:41:24less Precision for it and the difference
- 1:41:27with fp16 is that they actually touch
- 1:41:29and change the range so fp16 cannot
- 1:41:32represent the full range of fp32 it has
- 1:41:35a reduced range and that's where you
- 1:41:37start to actually run into issues
- 1:41:39because now you need uh these gradient
- 1:41:41scalers and things like that and I'm not
- 1:41:43going to go into the detail of that in
- 1:41:45this video because that's a whole video
- 1:41:48by itself but fb16 actually historically
- 1:41:50came first that was available in the
- 1:41:52Volta series before Amper and so fp16
- 1:41:56came first and everyone started to train
- 1:41:58in fp16 but everyone had to use all
- 1:42:00these gradient scaling operations which
- 1:42:02are kind of annoying and it's an
- 1:42:03additional source of state and
- 1:42:05complexity and the reason for that was
- 1:42:07because the exponent range was reduced
- 1:42:09in fp16 so that's the i e fp16 spec and
- 1:42:13then they came out with bf16 and the
- 1:42:15Ampere and they made it much simpler
- 1:42:18because we're just truncating manessa we
- 1:42:20have the exact same range and we do not
- 1:42:21need gradient scalers so everything is
- 1:42:24much much simpler now when we do use
- 1:42:26bf16 though we are impacting the numbers
- 1:42:30that we might be seeing in our pytorch
- 1:42:32code these this change is not just local
- 1:42:35to the operation itself so let's see how
- 1:42:37that works
- 1:42:39um there's some documentation here that
- 1:42:43so I think this is probably the best
- 1:42:44best page to explain how to use mixed
- 1:42:46Precision in pytorch um because there
- 1:42:49are many other tutorials and so on even
- 1:42:51within pitor documentation that are a
- 1:42:53lot more confusing and so I recommend
- 1:42:55specifically this one because there's
- 1:42:57five other copies that I would not
- 1:42:59recommend and then when we come
- 1:43:02here ignore everything about everything
- 1:43:05ignore everything about gradient
- 1:43:07scalers and only look at torch.
- 1:43:10AutoCast and basically also this comes
- 1:43:13to a single line of code at the end so
- 1:43:15this is the context manager that we
- 1:43:18want and we want to use that in our
- 1:43:21Network when you click into the torch.
- 1:43:25AutoCast autocasting it has a few more
- 1:43:28uh a bit more guideline for you so it's
- 1:43:30telling you do not call B flat 16 on any
- 1:43:34of your tensors just use AutoCast and
- 1:43:36only surround the uh forward pass of the
- 1:43:39model and the loss calculation and
- 1:43:41that's the only two things that you
- 1:43:43should be surrounding leave the backward
- 1:43:45and the optimizer step alone so that's
- 1:43:47the guidance that comes from the P team
- 1:43:49so we're going to follow that guidance
- 1:43:51and for us because the L calculation is
- 1:43:53inside of the model forward pass for us
- 1:43:56we are going to be doing
- 1:43:58this and then we don't want to be using
- 1:44:00torch Flo 16 because if we do that we
- 1:44:02need to start using gradient scalers as
- 1:44:04well so we are going to be using B float
- 1:44:0616 this is only possible to do an ampere
- 1:44:09uh but this means that the changes are
- 1:44:11extremely minimal like basically just
- 1:44:13this one line of
- 1:44:14code um let me first break
- 1:44:19in to here before we actually run this
- 1:44:22so right after logits I'd like to show
- 1:44:25you that different from the tf32 that we
- 1:44:28saw this is actually going to impact our
- 1:44:31tensors
- 1:44:32so this Lis tensor if we now look at
- 1:44:36this and we look at the dtype we
- 1:44:38suddenly see that this is now B float
- 1:44:4016 uh it's not float 32 anymore so our
- 1:44:43activations have been changed the
- 1:44:45activations tensor is now B FL 16 but
- 1:44:48not everything has changed so model.
- 1:44:51Transformer
- 1:44:55wte uh this is the weight uh token
- 1:44:57embedding table it has a weight inside
- 1:45:00it and the dtype of this weight this
- 1:45:02parameter is still torch float 32 so our
- 1:45:06parameters seem to still be in float 32
- 1:45:09but our activations the loits are now in
- 1:45:11P 16 so clearly this is why we get the
- 1:45:14mixed Precision some things pytorch is
- 1:45:16keeping inlow 32 some things pytorch is
- 1:45:19converting to lower Precision um and
- 1:45:23what gets converted at what point is not
- 1:45:26super clear I remember scrolling
- 1:45:30down is it
- 1:45:34here okay I can't find
- 1:45:37it I I thought it was here okay there we
- 1:45:41go so there are a few docks on when
- 1:45:44you're using this AutoCast what gets
- 1:45:46converted to B FL 16 and and when so for
- 1:45:49example only these Matrix multiply like
- 1:45:51operations get converted to float 16 but
- 1:45:54a lot of operations remain in float 32
- 1:45:56so in particular a lot of normalizations
- 1:45:58like layer norms and things like that
- 1:46:00not all of those layers might be
- 1:46:01converted um so only some layers
- 1:46:05selectively would be running B flat 16
- 1:46:07but things like softmax uh layer Norms
- 1:46:10uh log um log soft Max so loss function
- 1:46:14calculations a lot of those things might
- 1:46:15remain in float 32 because they are more
- 1:46:17susceptible to Precision changes major
- 1:46:20multiplies are fairly um
- 1:46:23robust to Precision changes uh so some
- 1:46:26parts of the network are um impacted
- 1:46:29more or less by the Precision
- 1:46:31change um so basically only some parts
- 1:46:34of the of the model are running in
- 1:46:35reduced Precision let's take it for a
- 1:46:38spin and let's actually see what kind of
- 1:46:41improvement we achieve
- 1:46:48here okay so we used to be 333
- 1:46:51milliseconds we're now 300
- 1:46:53and we used to be somewhere around
- 1:46:5450,000 tokens per second we're now at 55
- 1:46:57so we're definitely running faster but
- 1:46:59maybe not a lot faster and that's
- 1:47:02because there are still many many
- 1:47:03bottlenecks in our gbt2 we're just
- 1:47:05getting started but we have dropped down
- 1:47:07the precision as far as we can with my
- 1:47:09current GPU which is a100 we're using
- 1:47:12pytorch AutoCast unfortunately I don't
- 1:47:15actually exactly know what pytorch
- 1:47:17AutoCast do uh does I don't actually
- 1:47:19know exactly what's in B flat 16 what's
- 1:47:22in float 32
- 1:47:23we could go in and we could start to
- 1:47:24scrutinize it um but these are the kinds
- 1:47:27of rules that pytorch has internally and
- 1:47:29unfortunately they don't documented very
- 1:47:31well uh so we're not going to go into
- 1:47:34that into in too much detail but for now
- 1:47:36we are training in B flow 16 we do not
- 1:47:39need a gradient scaler and the reason
- 1:47:40things are running faster is because um
- 1:47:44we are able to run tensor course in B FL
- 1:47:4716 now that means we are in this row but
- 1:47:52uh we are also paying in Precision for
- 1:47:53this uh so um we expect slightly less
- 1:47:57accurate results with respect to the
- 1:47:58original fp32 but empirically in many
- 1:48:01cases this is a worth it uh kind of
- 1:48:04tradeoff because it allows you to run
- 1:48:06faster and you could for example train
- 1:48:07longer and make up for the uh for that
- 1:48:10Precision decrease so um that's b46 for
- 1:48:15now okay so as we can see we are
- 1:48:17currently at about 300 milliseconds uh
- 1:48:19per iteration and we're now going to
- 1:48:21reach for some really heavy weapons in
- 1:48:23the pie torch Arsenal and in particular
- 1:48:25we're going to introduce torch. compile
- 1:48:27so torch. compile is really quite
- 1:48:29incredible infrastructure from the
- 1:48:31pytorch team and it's basically a
- 1:48:32compiler for neural networks like it's
- 1:48:35almost like GCC for CN C++ code this is
- 1:48:38just this GCC of neural nuts so came out
- 1:48:42a while ago and extremely simple to use
- 1:48:46um the way to use torch compile is to do
- 1:48:48this it's a single line of code to
- 1:48:50compile your model and return it now
- 1:48:54this line of code will cost you
- 1:48:55compilation time but as you might guess
- 1:48:57it's going to make the code a lot faster
- 1:48:59so let's actually run that because this
- 1:49:01will take some time to run but currently
- 1:49:03remember we're at 300 milliseconds and
- 1:49:05we'll see what happens now while this is
- 1:49:08running I'd like to explain a little bit
- 1:49:10of what torch. compile does under the
- 1:49:11hood uh so feel free to read this page
- 1:49:15of P torch but basically there's no real
- 1:49:17good reason for you to not use torch
- 1:49:19compile in your pie torch I kind of feel
- 1:49:21like you should be using almost by
- 1:49:23default if you're not uh unless you're
- 1:49:25debugging and you want your code to run
- 1:49:26really fast and there's one line here in
- 1:49:29torch compile that I found that actually
- 1:49:31kind of like gets to why this is faster
- 1:49:33speed up mainly comes from reducing
- 1:49:35python overhead and GPU read wrs so let
- 1:49:38me unpack that a little bit um okay here
- 1:49:41we are okay so we went from 300
- 1:49:43milliseconds we're now running at 129
- 1:49:46milliseconds so this is uh 300 129 about
- 1:49:512.3x Improvement from a single line of
- 1:49:53code in py torch uh so quite incredible
- 1:49:56so what is happening what's happening
- 1:49:57under the hood well when you pass the
- 1:49:59model to torch
- 1:50:01compile what we have here in this NN
- 1:50:04module this is really just the
- 1:50:05algorithmic description of what we'd
- 1:50:08like to happen in our Network and torch
- 1:50:11compile will analyze the entire thing
- 1:50:14and it will look at what operations You'
- 1:50:15like to use and with the benefit of
- 1:50:18knowing exactly what's going to happen
- 1:50:20it doesn't have to run in What's called
- 1:50:22the e mode it doesn't have to just kind
- 1:50:24of like go layer by layer like the
- 1:50:26python interpreter normally would start
- 1:50:29at the
- 1:50:31forward and the python interpreter will
- 1:50:33go okay let's do this operation and then
- 1:50:36let's do that operation and it kind of
- 1:50:38materializes all the operations as it
- 1:50:40goes through uh so these um calculations
- 1:50:43are dispatched and run in this order and
- 1:50:45the python interpreter and this code
- 1:50:47doesn't know what kind of operations are
- 1:50:49going to happen later but torch compile
- 1:50:51sees your entire code at the same time
- 1:50:53and it's able to know what operations
- 1:50:56you intend to run and it will kind of
- 1:50:58optimize that process the first thing it
- 1:51:00will do is will it will take out the
- 1:51:01python interpreter from the forward pass
- 1:51:03entirely and it will kind of compile
- 1:51:05this entire neural net as a single
- 1:51:07object with no python interpreter
- 1:51:09involved so it knows exactly what's
- 1:51:11going to run and we'll just run that and
- 1:51:12it's all going to be running in
- 1:51:14efficient
- 1:51:15code uh the second thing that happens is
- 1:51:18uh this read write that they mentioned
- 1:51:21very briefly so a good example of that I
- 1:51:23think is the G nonlinearity that we've
- 1:51:25been looking at so here we use the n and
- 1:51:28G now this here is me uh basically just
- 1:51:32breaking up the inang Galu uh which you
- 1:51:35remember has this formula so this here
- 1:51:37is the equivalent implementation to
- 1:51:39what's happening inside g algorithmic l
- 1:51:41it's
- 1:51:42identical Now by default if uh we just
- 1:51:46we using this instead of ending. G here
- 1:51:48what would happen without torch compile
- 1:51:51well the python interpreter would make
- 1:51:52its way here and then it would be okay
- 1:51:54well there's an input well let me first
- 1:51:58let me raise this input to the third
- 1:51:59power and it's going to dispatch a
- 1:52:01kernel that takes your input and raises
- 1:52:03it to the third power and that kernel
- 1:52:05will run and when this kernel runs what
- 1:52:08ends up happening is this input is
- 1:52:11stored in the memory of the GPU so
- 1:52:13here's a helpful example of the layout
- 1:52:16of what's happening right you have your
- 1:52:18CPU this is in every single computer
- 1:52:21there's a few cores in there and you
- 1:52:23have your uh Ram uh your memory and the
- 1:52:26CPU can talk to the memory and this is
- 1:52:28all well known but now we've added the
- 1:52:30GPU and the GPU is a slightly different
- 1:52:32architecture of course they can
- 1:52:33communicate and it's different in that
- 1:52:35it's got a lot more course than a CPU
- 1:52:38all of those cores are individually a
- 1:52:40lot simpler too but it also has memory
- 1:52:43right this high bandwidth memory I'm
- 1:52:47sorry if I'm botching it hbm I don't
- 1:52:49even know what that stands for I'm just
- 1:52:51realizing that
- 1:52:53but uh this is the memory and it's very
- 1:52:54equivalent to uh RAM basically in the
- 1:52:58computer and what's happening is that
- 1:53:00input is living in the memory and when
- 1:53:02you do input
- 1:53:05cubed this has to travel to the GPU to
- 1:53:09the course and to all the caches and
- 1:53:12registers on the actual chip of this
- 1:53:15GPU and it has to calculate the all the
- 1:53:17elements to the third and then it saves
- 1:53:19the result back to the memory and it's
- 1:53:22this uh travel time that actually causes
- 1:53:25a lot of issues so here remember this
- 1:53:28memory bandwidth we can communicate
- 1:53:30about 2 terabytes per second which is a
- 1:53:31lot but also we have to Traverse this
- 1:53:35link and it's very slow so here on the
- 1:53:37GPU we're on chip and everything is
- 1:53:39super fast within the chip but going to
- 1:53:41the memory is extremely expensive takes
- 1:53:43extremely long amount of time and so we
- 1:53:46load the input do the calculations and
- 1:53:48load back the output and this round trip
- 1:53:51takes a lot of time
- 1:53:53and now right after we do that we
- 1:53:54multiply by this constant so what
- 1:53:57happens then is we dispatch another
- 1:53:59kernel and then the result travels back
- 1:54:02all the elements get multiplied by a
- 1:54:03constant and then the results travel
- 1:54:06back to the memory and then we take the
- 1:54:09result and we add back input and so this
- 1:54:12entire thing again travels to the GPU
- 1:54:15adds the inputs and gets written back so
- 1:54:18we're making all these round trips from
- 1:54:20the memory to actually where the comput
- 1:54:22happens because all the tensor cores and
- 1:54:24alus and everything like that is all
- 1:54:26stored on the chip in the GPU so we're
- 1:54:28doing a ton of round trips and pytorch
- 1:54:31uh without using torch compile doesn't
- 1:54:33know to optimize this because it doesn't
- 1:54:36know what kind of operations you're
- 1:54:37running later you're just telling it
- 1:54:39raise the power to the third then do
- 1:54:41this then do that and it will just do
- 1:54:43that in that sequence but torch compile
- 1:54:45sees your entire code it will come here
- 1:54:47and it will realize wait all of these
- 1:54:49are elementwise operations and actually
- 1:54:52what I'm going to do is I'm going to do
- 1:54:53a single trip of input to the GPU then
- 1:54:56for every single element I'm going to do
- 1:54:58all of these operations while that
- 1:55:00memory is on the GPU or chunks of it
- 1:55:04rather and then I'm going to write back
- 1:55:06a single time so we're not going to have
- 1:55:07these round trips and that's one example
- 1:55:09of what's called kernel fusion and is a
- 1:55:11major way in which everything is sped up
- 1:55:14so basically if you have your benefit of
- 1:55:15onet and you know exactly what you're
- 1:55:17going to compute you can optimize your
- 1:55:19round trips to the memory and you're not
- 1:55:21going to pay the the memory bandwidth
- 1:55:23cost and that's fundamentally what makes
- 1:55:25some of these operations a lot faster
- 1:55:27and what they mean by read writes
- 1:55:30here so let me erase this because we are
- 1:55:32not using it and yeah we should be using
- 1:55:36torch compile and our code is now
- 1:55:39significantly faster and we're doing
- 1:55:40about
- 1:55:42125,000 tokens per second but we still
- 1:55:45have a long way to go before we move on
- 1:55:47I wanted to supplement the discussion a
- 1:55:49little bit with a few more figures uh
- 1:55:51because this is a complic topic but it's
- 1:55:53worth understanding on a high level uh
- 1:55:55what's happening here and I could
- 1:55:56probably spend an entire video of like
- 1:55:58two hours on this but just the preview
- 1:56:00of that basically so this chip here that
- 1:56:03is uh the GPU this chip is where all the
- 1:56:06calculations happen mostly but this chip
- 1:56:09also does have some memory in it but
- 1:56:12most of the memory by far is here in the
- 1:56:15high bandwidth memory hbm and is
- 1:56:18connected they're connected um but these
- 1:56:20are two separate chips basically
- 1:56:23now here this is a zoom in of kind of
- 1:56:26this cartoon diagram of a GPU and what
- 1:56:30we're seeing here is number one you see
- 1:56:31this hbm I I realize it's probably very
- 1:56:34small for you but on the sides here it
- 1:56:35says hbm and so that that's the links to
- 1:56:38the hbm now the hbm is again off chip on
- 1:56:42the chip there are a large number of
- 1:56:45these streaming
- 1:56:46multiprocessors uh every one of these is
- 1:56:48an SM there's 120 of them in total and
- 1:56:51this is where the a lot of the
- 1:56:52calculations happen and this is a zoom
- 1:56:54in of a single individual as it has
- 1:56:57these four quadrants and see for example
- 1:56:59tensor core this is where a lot of the
- 1:57:00Matrix multiply stuff happens but
- 1:57:02there's all these other units to do all
- 1:57:04different kinds of calculations for fp64
- 1:57:07fp32 and for integers and so on now so
- 1:57:11we have all this uh logic here to do the
- 1:57:13calculations but in addition to that on
- 1:57:15the chip there is memory sprinkled
- 1:57:17throughout the chip so L2 cache is some
- 1:57:21amount of memory that lives on the chip
- 1:57:23and then on the SMS themselves there's
- 1:57:25L1 cache I realized it's probably very
- 1:57:28small for you but this blue bar is L1
- 1:57:31and there's also registers um and so
- 1:57:34there is memory stored here but the way
- 1:57:36this memory is stored is very different
- 1:57:38from the way memory is stored in hbm uh
- 1:57:41this is a very different implementation
- 1:57:44uh using um just in terms of like what
- 1:57:47the Silicon looks like it's a very
- 1:57:48different
- 1:57:49implementation um so here you would
- 1:57:52using transistors and capacitors and
- 1:57:54here it's a very different
- 1:57:55implementation uh with SRAM and what
- 1:57:57that looks like but long story short is
- 1:58:01um there is um memory inside the chip
- 1:58:05but it's not a lot of memory that's the
- 1:58:07critical point so this is some C this is
- 1:58:09a example diagram of a slightly
- 1:58:11different GPU just like here where it
- 1:58:14shows that for example typical numbers
- 1:58:16for CPU Dam memory which is this thing
- 1:58:19here you might have one tab of this
- 1:58:22right but it would be extremely
- 1:58:23expensive to access especially for a GPU
- 1:58:25you have to go through the CPU here now
- 1:58:28next we have the hbm so we have tens of
- 1:58:30gigabytes of hbm memory on a typical GPU
- 1:58:33here but it's as I mentioned very
- 1:58:35expensive to access and then on the chip
- 1:58:38itself everything is extremely fast
- 1:58:40within the chip but we only have couple
- 1:58:4210 megabytes of memory collectively
- 1:58:45throughout the Chip And so there's just
- 1:58:48not enough space because the memory is
- 1:58:50very expensive on the chip and so
- 1:58:52there's not a lot of it but it is
- 1:58:53lightning fast to access in relative
- 1:58:55terms and so basically whenever we have
- 1:58:58these kernels um the more accurate
- 1:59:01picture of what's Happening Here is that
- 1:59:03we take these inputs which live by
- 1:59:05default on the global memory and now we
- 1:59:08need to perform some calculation so we
- 1:59:10start streaming the data from the um
- 1:59:12Global memory to the uh chip we perform
- 1:59:16the calculations on the chip and then
- 1:59:18stream it back and store it back to the
- 1:59:19global memory right and so if we are if
- 1:59:23we don't have torch compile we are
- 1:59:24streaming the data through the chip
- 1:59:26doing the calculations and saving to the
- 1:59:27memory and we're doing those round trips
- 1:59:29many many
- 1:59:30times but uh if it's torch compiled then
- 1:59:33we start streaming the memory as before
- 1:59:35but then while we're on the chip we're
- 1:59:37we're we have a chunk of the uh data
- 1:59:40that we're trying to process so that
- 1:59:42chunk now lives on the chip while it's
- 1:59:44on the chip it's extremely fast to
- 1:59:46operate on so if we have kernel Fusion
- 1:59:48we can do all the operations right there
- 1:59:49in an element-wise fashion and those are
- 1:59:52very cheap and then we do a single round
- 1:59:54trip back to the global memory so
- 1:59:58operator Fusion basically allows you to
- 2:00:00keep your chunk of data on the Chip And
- 2:00:02do lots of calculations on it before you
- 2:00:04write it back and that gives huge
- 2:00:06savings and that's why torch compile
- 2:00:09ends up being a lot faster or that's one
- 2:00:11of the major
- 2:00:12reasons uh so again just a very brief
- 2:00:14intro to the memory hierarchy and
- 2:00:16roughly what torch compile does for you
- 2:00:19now torch compile is amazing but there
- 2:00:21are operations torch compile will not
- 2:00:23find and an amazing example of that is
- 2:00:26Flash attention to which we turn next so
- 2:00:29flash attention comes from this paper
- 2:00:30from uh Stanford in
- 2:00:332022 and it's this incredible algorithm
- 2:00:36for performing attention so um and
- 2:00:39running it a lot faster so flash
- 2:00:41attention will come here and we will
- 2:00:44take out these four
- 2:00:46lines and Flash attention implements
- 2:00:48these four lines really really quickly
- 2:00:51and how does it do that well flash
- 2:00:53attention is a kernel Fusion operation
- 2:00:57so you see here we have um in this
- 2:00:59diagram they're showing P torch and you
- 2:01:02have these four operations uh they're
- 2:01:04including Dropout but we are not using
- 2:01:06Dropout here so we just have these four
- 2:01:08lines of code here and instead of those
- 2:01:11we are fusing them into a single fused
- 2:01:13kernel of flash attention so it's an
- 2:01:16it's a it's a kernel Fusion algorithm
- 2:01:19but it's a kernel Fusion that torch
- 2:01:20compile cannot find
- 2:01:22and the reason that it cannot find it is
- 2:01:24that it um requires an algorithmic
- 2:01:26rewrite of how attention is actually
- 2:01:28implemented here in this case and what's
- 2:01:31remarkable about it is that uh flash
- 2:01:33attention actually if you just count the
- 2:01:35number of flops flash attention does
- 2:01:37more flops than this attention here but
- 2:01:41flash attention is actually
- 2:01:42significantly faster in fact they site
- 2:01:457. six times faster potentially and
- 2:01:48that's because it is very mindful of the
- 2:01:51memory hierarchy as I described it just
- 2:01:53now and so it's very mindful about
- 2:01:55what's in high bandwidth memory what's
- 2:01:57in the shared memory and it is very
- 2:02:00careful with how it orchestrates the
- 2:02:02computation such that we have fewer
- 2:02:04reads and writes to the high bandwidth
- 2:02:06memory and so even though we're doing
- 2:02:08more flops the expensive part is they
- 2:02:10load and store into hbm and that's what
- 2:02:12they avoid and so in particular they do
- 2:02:15not ever materialize this end byend
- 2:02:17attention Matrix this ATT here a flash
- 2:02:21attention is designed such that this
- 2:02:23Matrix never gets materialized at any
- 2:02:25point and it never gets read or written
- 2:02:28to the hbm and this is a very large
- 2:02:30Matrix right so um because this is where
- 2:02:32all the queries and keys interact and
- 2:02:34we're sort of getting
- 2:02:36um for each head for each batch element
- 2:02:40we're getting a t BYT Matrix of
- 2:02:42attention which is a Million numbers
- 2:02:45even for a single head at a single batch
- 2:02:47index at like so so basically this is a
- 2:02:50ton of memory and and this is never
- 2:02:52materialized and the way that this is
- 2:02:54achieved is that basically the
- 2:02:57fundamental algorithmic rewrite here
- 2:02:58relies on this online softmax trick
- 2:03:02which was proposed previously and I'll
- 2:03:03show you the paper in a bit and the
- 2:03:05online softmax trick coming from a
- 2:03:07previous paper um shows how you can
- 2:03:10incrementally evaluate a soft Max
- 2:03:14without having to sort of realize all of
- 2:03:16the inputs to the softmax to do the
- 2:03:18normalization and you do that by having
- 2:03:19these intermediate variables M and L and
- 2:03:22there's an update to them that allows
- 2:03:24you to evaluate the softmax in an online
- 2:03:26manner um now flash attention actually
- 2:03:30so recently flash attention 2 came out
- 2:03:32as well so I have that paper up here as
- 2:03:34well uh that has additional gains to how
- 2:03:36it calculates flash attention and the
- 2:03:38original paper that this is based on
- 2:03:40basically is this online normalizer
- 2:03:42calculation for softmax and remarkably
- 2:03:45it came out of Nvidia and it came out of
- 2:03:46it like really early 2018 so this is 4
- 2:03:50years before flash attention
- 2:03:52and this paper says that we propose a
- 2:03:55way to compute the classical softmax
- 2:03:57with fewer memory accesses and
- 2:03:59hypothesize that this reduction in
- 2:04:00memory accesses should improve softmax
- 2:04:02performance on actual hardware and so
- 2:04:05they are extremely correct in this
- 2:04:08hypothesis but it's really fascinating
- 2:04:10to me that they're from Nvidia and that
- 2:04:12they had this realization but they
- 2:04:13didn't actually take it to the actual
- 2:04:15flash attention that had to come four
- 2:04:18years later from Stanford so I don't
- 2:04:20fully understand the historical how this
- 2:04:22happened historically um but they do
- 2:04:24basically propose this online update to
- 2:04:26the softmax uh right here and this is
- 2:04:29fundamentally what they reuse here to
- 2:04:31calculate the softmax in a streaming
- 2:04:33Manner and then they realize they can
- 2:04:35actually fuse all the other operations
- 2:04:37with the online sofx calculation into a
- 2:04:40single fused kernel flash attention and
- 2:04:42that's what we are about to use so great
- 2:04:45example I think of being aware of um
- 2:04:47memory hierarchy the fact that flops
- 2:04:49don't matter uh the entire memory access
- 2:04:52pattern matters and that torch compile
- 2:04:54is amazing but there are many
- 2:04:55optimizations that are still available
- 2:04:57to us that potentially torch compile
- 2:04:59cannot find maybe maybe one day it could
- 2:05:01but right now it seems like a lot to ask
- 2:05:04so here's what we're going to do we're
- 2:05:05going to use Flash attention and the way
- 2:05:09to do that basically in pytorch is we
- 2:05:11are going to comment out these four
- 2:05:14lines and we're going to replace them
- 2:05:15with a single line and here we are
- 2:05:18calling this compound operation in
- 2:05:20pytorch called scale that product
- 2:05:22attention and uh pytorch will call flash
- 2:05:27attention when you use it in this way
- 2:05:31I'm not actually 100% sure why torch
- 2:05:32compile doesn't realize that these four
- 2:05:34lines should just call flash attention
- 2:05:36in this exact way we have to do it again
- 2:05:38for it which in my opinion is a little
- 2:05:40bit odd but um here we are so you have
- 2:05:46to use this compound up and uh let's
- 2:05:49wait for a few moments before torch comp
- 2:05:51compile gets around to it and then let's
- 2:05:53remember that we achieved 6.05 661 I
- 2:05:58have it here that's the loss we were
- 2:06:00expecting to see and we took 130
- 2:06:03milliseconds uh before this change so
- 2:06:05we're expecting to see the exact same
- 2:06:07result by iteration 49 but we expect to
- 2:06:10see faster runtime because Flash
- 2:06:13attention is just a an algorithmic
- 2:06:14rewrite and it's a faster kernel but it
- 2:06:16doesn't actually change any of the
- 2:06:17computation and we should have the exact
- 2:06:19same optimization so okay so we're a lot
- 2:06:21faster we're at about 95 milliseconds
- 2:06:24and we achiev
- 2:06:286.58 okay so they're basically identical
- 2:06:31up to a floating Point fudge Factor so
- 2:06:34it's the identical computation but it's
- 2:06:36significantly faster going from 130 to
- 2:06:39roughly 90
- 2:06:4096 and so this is um 96 divide
- 2:06:44130ish so this is maybe 27 is%
- 2:06:48Improvement um so uh really interesting
- 2:06:52and that is Flash retention okay we are
- 2:06:54now getting to one of my favorite
- 2:06:57optimizations and it is simultaneously
- 2:06:59the dumbest and the most brilliant
- 2:07:02optimization and it's always a little
- 2:07:03bit surprising to me um anyway so
- 2:07:06basically I mentioned a few minutes ago
- 2:07:08that there are some numbers that are
- 2:07:10nice and some numbers that are ugly so
- 2:07:1364 is a beautiful nice number 128 is
- 2:07:17even nicer 256 is beautiful what makes
- 2:07:20these numbers beautiful is that there
- 2:07:21are many powers of two inside them you
- 2:07:23can divide by two many times and uh
- 2:07:26examples of ugly numbers are like 13 and
- 2:07:2817 and something like that prime numbers
- 2:07:30numbers that are not even and so on and
- 2:07:32so pretty much you always want to use
- 2:07:34nice numbers in all of your code that
- 2:07:36deals with neural networks or Cuda
- 2:07:38because everything in Cuda Works in sort
- 2:07:40of like powers of two and lots of
- 2:07:42kernels are written in terms of powers
- 2:07:45of Two And there are lots of blocks of
- 2:07:47sizes 16 and uh 64 and so on so
- 2:07:50everything is written in those terms and
- 2:07:52you always have special case handling
- 2:07:54for all kinds of uh logic that U when
- 2:07:57your inputs are not made of nice numbers
- 2:08:00so let's see what that looks like
- 2:08:01basically scan your code and look for
- 2:08:03ugly numbers is roughly theistic so
- 2:08:06three times is kind of ugly um I'm not
- 2:08:10100% sure maybe this can be improved but
- 2:08:12this is uh this is ugly and not
- 2:08:15ideal um four times is nice so that's uh
- 2:08:20that's nice
- 2:08:221024 is very nice that's a power of two
- 2:08:2512 is a little bit suspicious um not too
- 2:08:28many powers of two 768 is great 50, 257
- 2:08:32is a really really ugly number um it's
- 2:08:36first of all it's odd so uh and there's
- 2:08:38no not too many powers of two in there
- 2:08:40so this is a very ugly number and it's
- 2:08:43highly suspicious and then when we
- 2:08:45scroll down all these numbers are nice
- 2:08:48and then here we have mostly nice
- 2:08:50numbers except for 25 so in this
- 2:08:53configuration of gpt2 XL a number of
- 2:08:55heads is 25 uh that's a really ugly
- 2:08:57number that's an odd number and um
- 2:09:00actually this did cause a lot of
- 2:09:01headaches for us recently when we're
- 2:09:02trying to optimize some kernels uh to
- 2:09:04run this fast um and required a bunch of
- 2:09:07special case handling so basically these
- 2:09:10numbers are we have some ugly numbers
- 2:09:12and some of them are easier to fix than
- 2:09:13others and in particular the voap size
- 2:09:15being 50257 that's a very ugly number
- 2:09:18very suspicious and we want to fix it
- 2:09:20now when you when you fix these things
- 2:09:23uh one of the easy ways to do that is
- 2:09:24you basically um increase the number
- 2:09:27until it's the nearest power of two that
- 2:09:29you like so here's a much nicer number
- 2:09:32it's
- 2:09:3350304 and why is that because 50304 can
- 2:09:37be divided by 8 or by 16 or by 32
- 2:09:4364 it can even be divided by 128 I think
- 2:09:46yeah so it's a very nice number um so
- 2:09:49what we're going to do here is the GPT
- 2:09:51config and you see that we initialized B
- 2:09:53cap size to
- 2:09:5450257 Let's override just
- 2:09:58that um element to be
- 2:10:0150304 okay so everything else stays the
- 2:10:05same we're just increasing our
- 2:10:06vocabulary size so we're adding it's
- 2:10:09almost like we're adding fake tokens uh
- 2:10:12so that book up size has powers of two
- 2:10:14inside it now actually what I'm doing
- 2:10:16here by the way is I'm increasing the
- 2:10:18amount of computation that our network
- 2:10:19will be doing if you just count the the
- 2:10:21flops on like do the math of how many
- 2:10:23flops we're doing we're going to be
- 2:10:25doing more flops and we still have to
- 2:10:27think through whether this doesn't break
- 2:10:30anything but if I just run this uh let's
- 2:10:33see what we get uh currently this ran in
- 2:10:35maybe
- 2:10:3896.5 milliseconds per step I'm just kind
- 2:10:41of like eyeballing it and let's see what
- 2:10:43kind of a result we're going to
- 2:10:46get uh while this is compiling let's
- 2:10:49think through whether our code actually
- 2:10:51works okay when we increase the vocap
- 2:10:53size like this let's look at where vocap
- 2:10:55size is actually
- 2:10:57used so we swing up to the inet and we
- 2:11:00see that it's used inside the embedding
- 2:11:01table of course so all the way at the
- 2:11:03bottom of the Transformer and it's used
- 2:11:05at the classifier layer all the way at
- 2:11:06the top of the Transformer so in two
- 2:11:08places and let's take a look and we're
- 2:11:11running at 93 so 93 milliseconds instead
- 2:11:14of
- 2:11:1596.5 so we are seeing a roughly yeah 4%
- 2:11:19Improvement here uh by doing more
- 2:11:22calculations and the reason for this is
- 2:11:25we fixed we've made an ugly number into
- 2:11:28a nice number let's I'm going to come
- 2:11:30into the explanation for that a little
- 2:11:32bit again but for now let's just
- 2:11:34convince ourselves that we're not
- 2:11:35breaking anything when we do this so
- 2:11:36first of all we've made the the wte the
- 2:11:39embedding table for the tokens we've
- 2:11:41made it larger it's almost like we
- 2:11:43introduced more tokens at the bottom and
- 2:11:46these tokens are never used because the
- 2:11:48gbt tokenizer only has tokens up to
- 2:11:50$50,000
- 2:11:51256 and so we'll never index into the
- 2:11:55rows that we've added so we're wasting a
- 2:11:57little bit of space here by creating
- 2:11:59memory that's never going to be accessed
- 2:12:01never going to be used Etc now that's
- 2:12:03not fully correct because this wte
- 2:12:06weight ends up being shared and ends up
- 2:12:08being used in the classifier here at the
- 2:12:10end so what is that doing to the
- 2:12:13classifier right here well what what
- 2:12:15that's doing is we're predicting
- 2:12:16additional Dimensions at the classifier
- 2:12:18now and we're predicting probabilities
- 2:12:20for tokens that will of course never be
- 2:12:21present in the training set um and so
- 2:12:25therefore the network has to learn that
- 2:12:27these probabilities uh have to be driven
- 2:12:29to zero and so the logits that the
- 2:12:31network produces have to drive those
- 2:12:33dimensions of the output to negative
- 2:12:35Infinity but it that's no different from
- 2:12:38all the other tokens that are already in
- 2:12:39our data set um or rather that are not
- 2:12:42in our data set so Shakespeare only
- 2:12:45probably uses let's say a th000 tokens
- 2:12:46out of 50,000 to 57 tokens so most of
- 2:12:49the tokens are already being driven to
- 2:12:51zero probability by the optimization we'
- 2:12:53just introduced a few more tokens now
- 2:12:55that in a similar manner will never be
- 2:12:57used and have to be driven to zero in
- 2:12:59probability um so functionally though
- 2:13:02nothing breaks we're using a bit more
- 2:13:05extra um memory but otherwise this is a
- 2:13:08harmless operation as far as I can tell
- 2:13:11but and we're adding calculation but
- 2:13:12it's running faster and it's running
- 2:13:14faster because as I mentioned in Cuda so
- 2:13:17many kernels use uh block tiles and
- 2:13:21these block towels are usually nice
- 2:13:22numbers uh so powers of two so
- 2:13:25calculations are done in like chunks of
- 2:13:2664 or chunks of 32 and when your um when
- 2:13:31your desired calculation doesn't neatly
- 2:13:32fit into those block tiles um there are
- 2:13:36all kinds of boundary kernels that can
- 2:13:38kick in to like do the last part so
- 2:13:42basically in a lot of kernels they will
- 2:13:44chunk at up your input and they will do
- 2:13:46the nice part first and then they have a
- 2:13:47whole second second phase where they
- 2:13:50come back to any that like uh remains uh
- 2:13:54and then they process the remaining part
- 2:13:56and the kernels for that could be very
- 2:13:57inefficient and so you're basically um
- 2:14:00spinning up all this extra compute and
- 2:14:02is extremely inefficient so you might as
- 2:14:04well pad your inputs and um make it fit
- 2:14:07nicely and usually that empiric lens up
- 2:14:10actually running faster um so this is
- 2:14:13another example of a 4% Improvement that
- 2:14:16we've added and this is something that
- 2:14:18also torch compile did not find for us
- 2:14:21you would hope that torch compile at
- 2:14:22some point could figure an optimization
- 2:14:24like this out uh but for now uh this is
- 2:14:27it and I also have to point out that
- 2:14:28we're using pytorch nightly so that's
- 2:14:30why we're only seeing 4% if you're using
- 2:14:33pytorch 2.3.1 or earlier you would
- 2:14:36actually see something like 30%
- 2:14:37Improvement just from this change from
- 2:14:39changing it to from 50,000 to 57 to
- 2:14:4350304 so again one of my favorite
- 2:14:47examples also of having to understand
- 2:14:49the under the hood and how it all works
- 2:14:51and to know what kinds of things to
- 2:14:52Tinker with to push the performance of
- 2:14:54your code okay so at this point we have
- 2:14:56improved the performance by about 11x
- 2:14:58right because we started at about 1,000
- 2:15:00milliseconds per step and we're now down
- 2:15:02to like 93 milliseconds so that's uh
- 2:15:05quite good and we're uh doing a much
- 2:15:08better job of utilizing our GPU
- 2:15:09resources so I'm going to now turn to
- 2:15:12more algorithmic changes uh and
- 2:15:14improvements to the actual optimization
- 2:15:16itself and what we would like to do is
- 2:15:18we would like to follow the hyper
- 2:15:19parameters that are mentioned in the GP
- 2:15:20G2 or gpt2 gpt3 paper now sadly gpt2 is
- 2:15:26uh doesn't actually say too much it's
- 2:15:28very nice of them that they released the
- 2:15:30model weights and the code but the paper
- 2:15:32itself is extremely vague as to the
- 2:15:33optimization details uh the code itself
- 2:15:36that they released as well the code
- 2:15:38we've been looking at this is just the
- 2:15:40inference code so there's no training
- 2:15:41code here and very few hyp parameters so
- 2:15:44this doesn't also tell us too much so
- 2:15:46for that we have to turn to the gpt3
- 2:15:48paper and um in the depending of the
- 2:15:51gpt3 paper um they have a lot more hyper
- 2:15:55parameters here for us to use and the
- 2:15:57gpt3 paper in general is a lot more
- 2:15:59detailed as to uh all of the you know
- 2:16:02small details that go into the model
- 2:16:04training but gpt3 U models were never
- 2:16:07released so gbt2 we have the weights but
- 2:16:10no details and gpt3 we have lots of
- 2:16:11details but no weights so um but roughly
- 2:16:15speaking gpt2 and gpt3 architectures are
- 2:16:17very very similar and um basically there
- 2:16:21are very few changes the context length
- 2:16:23was expanded from 1024 to 2048 and
- 2:16:25that's kind of like the major change uh
- 2:16:28and some of the hyper parameters around
- 2:16:29the Transformer have changed but
- 2:16:31otherwise they're pretty much the same
- 2:16:32model it's just that gpt3 was trained
- 2:16:34for a lot longer on a bigger data set
- 2:16:36and uh has a lot more thorough
- 2:16:38evaluations uh and the gpt3 model is 175
- 2:16:42billion instead of 1.6 billion um in the
- 2:16:46gpt2 so long story short we're going to
- 2:16:49go to gp3 paper to follow along some the
- 2:16:51hyper parameters so to train all the
- 2:16:54versions of gpt3 we use atom with beta 1
- 2:16:56beta 2 of9 and .95 so let's swing over
- 2:17:00here and make sure that the betas
- 2:17:02parameter which you can see here
- 2:17:04defaults to 0.9 and
- 2:17:06999 is actually set to 0.9 and
- 2:17:11.95 and then the Epsilon parameter uh
- 2:17:14you can see is the default is 1 in8 and
- 2:17:17this is also one in8 let's just uh put
- 2:17:19it in so that works
- 2:17:22expit uh now next up they say we clip
- 2:17:25the gra Global Norm of the gradient at
- 2:17:271.0 so what this is referring to is that
- 2:17:30once we calculate the gradients right
- 2:17:32after l. backward um we basically have
- 2:17:35the gradients at all the parameter
- 2:17:37tensors and what people like to do is
- 2:17:40basically uh clip them to have some kind
- 2:17:42of a maximum Norm so in pytor this is
- 2:17:45fairly easy to do uh it's one line of
- 2:17:48code here that we have to insert right
- 2:17:50after we calcul Cal the gradients and
- 2:17:52what this utility function is doing is
- 2:17:55um it's calculating the global Norm of
- 2:17:58the parameters so every single par um
- 2:18:01gradient on all the parameters you
- 2:18:03square it and you add it all up and you
- 2:18:05take a big square root of that and
- 2:18:07that's the norm of the parameter V
- 2:18:10Vector basically it's the it's the
- 2:18:12length of it if you if you'd like to
- 2:18:14look at it that way and we are basically
- 2:18:16making sure that its length is no more
- 2:18:18than 1.0 and we're going to clip it
- 2:18:21and the reason that people like to use
- 2:18:23this is that uh sometimes you can get
- 2:18:25unlucky during your optimization maybe
- 2:18:27it's a bad data batch or something like
- 2:18:28that and if you get very unlucky in the
- 2:18:31batch you might get really high loss and
- 2:18:33really high loss could lead to a really
- 2:18:35high gradient and this could basically
- 2:18:38uh shock your model and shock the
- 2:18:40optimization so people like to use a
- 2:18:42gradient Norm clipping uh to prevent the
- 2:18:45model from um basically getting too big
- 2:18:49of shocks in terms of the gradient
- 2:18:50magnet ude and uh the upper bound it in
- 2:18:53this way it's a bit of a hacky solution
- 2:18:55it's about like a patch on top of like
- 2:18:57deeper issues uh but uh people still do
- 2:19:00it fairly frequently now the clip grad
- 2:19:03Norm Returns the norm of the gradient
- 2:19:05which I like to always visualize uh
- 2:19:08because um it is useful information and
- 2:19:11sometimes you can look at the norm of
- 2:19:13the gradient and if it's well behaved
- 2:19:15things are good if it's climbing things
- 2:19:17are bad and they're destabilizing during
- 2:19:19training sometimes you could get a spike
- 2:19:21in the norm and that means there's some
- 2:19:22kind of an issue or an instability so
- 2:19:25the norm here will be a
- 2:19:28norm uh and let's do a uh 4f or
- 2:19:33something like
- 2:19:34that and I believe this is just a float
- 2:19:37and so we should be able to uh print
- 2:19:40that uh so that's Global gradient
- 2:19:44clipping now they go into the details of
- 2:19:46the learning rate uh scheduler so they
- 2:19:49don't just use a fixed learning rate
- 2:19:51like we do here for 3 E4 but there's
- 2:19:54actually basically a cosine DK learning
- 2:19:57rate schedule um it's got a warm-up and
- 2:20:00it's got a cosine DEC to 10% over some
- 2:20:04Horizon
- 2:20:06um and so we're going to implement uh
- 2:20:09this in a second I just like to see Norm
- 2:20:11printed here okay there we go so what
- 2:20:14happened here is the norm is actually
- 2:20:16really high in the beginning 30 or so
- 2:20:19and you see that as we continue training
- 2:20:21it kind of like
- 2:20:22stabilizes um at values below one um and
- 2:20:27this is not that crazy uncommon for the
- 2:20:30norm to be high in the very first few
- 2:20:31stages basically What's Happening Here
- 2:20:33is the model is completely random and so
- 2:20:35there's a ton of learning happening very
- 2:20:37early in the network but that learning
- 2:20:39is kind of like um you know it's mostly
- 2:20:41learning the biases of the output tokens
- 2:20:44and so it's a bit of an unstable time uh
- 2:20:46but the network usually stabilizes in a
- 2:20:48very few iterations so this looks very
- 2:20:50relatively reasonable to me except
- 2:20:52usually I would expect this looks a
- 2:20:54little bit funky that we go from 28 to 6
- 2:20:56to 2 and then to 10 um it's not
- 2:20:59completely insane but it's just kind of
- 2:21:01a little bit
- 2:21:02funky um okay so let's now get to the
- 2:21:05learning rate schuer so the learning
- 2:21:07rate schedule that's used here in gpt3
- 2:21:09is what's called a cosine Decay learning
- 2:21:12schedule with warmup and the way this
- 2:21:14looks is that the learning rate is
- 2:21:17basically starts right at around zero
- 2:21:19linearly rank s up over some amount of
- 2:21:21time and then comes down with this
- 2:21:24cosine sort of form and comes down to
- 2:21:27some kind of a minimum learning rate
- 2:21:28that's up to you so here the minimum
- 2:21:30learning rate is zero but uh here in the
- 2:21:33paper they said that they use cosine
- 2:21:35Decay for learning rate down to 10% of
- 2:21:37its value over the first 260 billion
- 2:21:40tokens and then training continues 10%
- 2:21:43after and there's a linear warmup over
- 2:21:46the first 375 million tokens so that's
- 2:21:50about the learn R so let's now implement
- 2:21:52this uh so I already implemented it here
- 2:21:55and the way this works is let me scroll
- 2:21:58down first here I changed our training
- 2:22:00Loop a little bit so this was a 4i in
- 2:22:02Max steps I just change it to step now
- 2:22:04so that we have the notion of a step is
- 2:22:07a single optimization step in the in the
- 2:22:09for Loop and then here I get the LR for
- 2:22:13this step of the optimization using a
- 2:22:15new function I call get LR and then in
- 2:22:18pytorch to set the learning rate I think
- 2:22:20this is is the way to set the learning
- 2:22:21rate it's a little bit gnarly um because
- 2:22:24you have to basically there's a notion
- 2:22:25of different par parameter groups that
- 2:22:27could exist in the optimizer and so you
- 2:22:28actually have to iterate over them even
- 2:22:30though we currently have a single param
- 2:22:32group only um and you have to set the LR
- 2:22:34in this for Loop kind of style is is my
- 2:22:37impression right now so we have this
- 2:22:39look of LR we set the learning rate and
- 2:22:42then on the bottom I'm also printing it
- 2:22:45uh so that's all the changes I made to
- 2:22:47this Loop and then of course the get LR
- 2:22:49is my scheduler now it's worth pointing
- 2:22:51out that pytorch actually has learning
- 2:22:53rate schedulers and you can use them and
- 2:22:55I believe there's a cosine learning rate
- 2:22:57schedule in pytorch I just don't really
- 2:22:59love using that code because honestly
- 2:23:02it's like five lines of code and I fully
- 2:23:06understand what's happening inside these
- 2:23:07lines so I don't love to use
- 2:23:09abstractions where they're kind of in
- 2:23:11screwable and then I don't know what
- 2:23:13they're doing so personal style so the
- 2:23:16max learning rate here is let's say 3 E4
- 2:23:19but we're going to see that in gpt3
- 2:23:22here they have a table of what the
- 2:23:25maximum learning rate is for every model
- 2:23:28size so um for for this one basically 12
- 2:23:3412 layer 768 gpt3 so the gpt3 small is
- 2:23:37roughly like a GPT
- 2:23:402124m we see that here they use a
- 2:23:42learning rate of 6 E4 so we could
- 2:23:44actually go higher um in fact we may
- 2:23:46want to try to follow that and just set
- 2:23:48the max LR here at six
- 2:23:51uh then the that's the maximum learning
- 2:23:53rate the minum learning rate is uh 10%
- 2:23:55of that per description in the paper
- 2:23:58some number of steps that we're going to
- 2:24:00warm up over and then the maximum steps
- 2:24:02of the optimization which I now use also
- 2:24:05in the for Loop down here and then you
- 2:24:07can go over this code if you like it's
- 2:24:09not U it's not terribly inside Flor
- 2:24:11interesting I'm just uh modulating based
- 2:24:13on the iteration number which learning
- 2:24:16rate uh there should be so this is the
- 2:24:18warm-up region um
- 2:24:21this is the region after the
- 2:24:22optimization and then this is the region
- 2:24:24sort of in between and this is where I
- 2:24:26calculate the cosine learning rate
- 2:24:28schedule and you can step through this
- 2:24:29in detail if you'd like uh but this is
- 2:24:32basically implementing this
- 2:24:33curve and I ran this already and this is
- 2:24:38what that looks
- 2:24:40like um so when we now run we start at
- 2:24:45um some very low number now note that we
- 2:24:47don't start exactly at zero because that
- 2:24:49would be not useful to update with a
- 2:24:50learning rate of zero that's why there's
- 2:24:52an it+ one so that on the zeroth
- 2:24:54iteration we are not using exactly zero
- 2:24:57we're using something very very low then
- 2:24:59we linearly warm up to maximum learning
- 2:25:02rate which in this case was 34 when I
- 2:25:04ran it but now would be 6 E4 and then it
- 2:25:07starts to decay all the way down to um 3
- 2:25:11E5 which was at the time 10% of the
- 2:25:14original learning rate now one thing we
- 2:25:16are not following exactly is that they
- 2:25:18mentioned that um
- 2:25:21let me see if I can find it
- 2:25:23again we're not exactly following what
- 2:25:26they did
- 2:25:28because uh they mentioned that their
- 2:25:30training Horizon is 300 billion tokens
- 2:25:33and they come down to 10% of the initial
- 2:25:35learning rate of at 260 billion and then
- 2:25:37they train after 260 with 10% so
- 2:25:41basically their Decay time is less than
- 2:25:43the max steps time whereas for us
- 2:25:45they're exactly equal so it's not
- 2:25:47exactly faithful but it's um it's an
- 2:25:51okay um this is okay for us and for our
- 2:25:53purposes right now and um we're just
- 2:25:57going to use this ourselves I don't
- 2:25:58think it makes too too big of a
- 2:26:00difference honestly I should point out
- 2:26:02that what learning rate schedule you use
- 2:26:04is totally up to you there's many
- 2:26:05different types um coign learning rate
- 2:26:08has been popularized a lot by gpt2 and
- 2:26:10gpt3 but people have come up with all
- 2:26:12kinds of uh other learning rate
- 2:26:14schedules um and this is kind of like an
- 2:26:16active area of uh research as to which
- 2:26:18one is the most effective at train these
- 2:26:20networks okay next up the paper talks
- 2:26:23about the gradual batch size increase so
- 2:26:26there's a ramp on the batch size that is
- 2:26:29linear and you start with very small
- 2:26:31batch size and you ramp up to a big
- 2:26:32batch size over time uh we're going to
- 2:26:35actually skip this and we're not going
- 2:26:36to work with it and the reason I don't
- 2:26:38love to use it is that it complicates a
- 2:26:41lot of the arithmetic because you are
- 2:26:42changing the number of tokens that
- 2:26:43you're processing at every single step
- 2:26:45of the optimization and I like to keep
- 2:26:47that math very very simple also my
- 2:26:49understanding is that that this is not
- 2:26:50like a major um Improvement and also my
- 2:26:54understanding is that this is not like
- 2:26:55an algorithmic optimization Improvement
- 2:26:57it's more of a systems and speed
- 2:26:59Improvement and roughly speaking this is
- 2:27:02because uh in the early stages of the
- 2:27:05optimization uh again the model is in a
- 2:27:07very atypical setting and mostly what
- 2:27:10you're learning is that um you're mostly
- 2:27:13learning to ignore the tokens uh that
- 2:27:15don't come up in your training set very
- 2:27:16often you're learning very simple biases
- 2:27:19and and that kind of a thing and so
- 2:27:23every single example that you put
- 2:27:24through your network is basically just
- 2:27:26telling you use these tokens and don't
- 2:27:28use these tokens and so the gradients
- 2:27:30from every single example are actually
- 2:27:31extremely highly correlated they all
- 2:27:33look roughly the same in the in the OR
- 2:27:36original parts of the optimization
- 2:27:38because they're all just telling you
- 2:27:39that these tokens don't appear and these
- 2:27:40tokens do appear and so because the
- 2:27:43gradients are all very similar and
- 2:27:45they're highly correlated then why are
- 2:27:46you doing batch sizes of like Millions
- 2:27:49when if you do a batch size of 32k
- 2:27:51you're basically getting the exact same
- 2:27:53gradient early on in the training and
- 2:27:55then later in the optimization once
- 2:27:57you've learned all the simple stuff
- 2:28:00that's where the actual work starts and
- 2:28:01that's where the gradients become more
- 2:28:02decorrelated per examples and that's
- 2:28:04where they actually offer you sort of
- 2:28:07statistical power in some sense um so
- 2:28:10we're going to skip this just because it
- 2:28:12kind of complicates things and we're
- 2:28:14going to go
- 2:28:15to uh data are sampled without
- 2:28:18replacement during training um so until
- 2:28:21an Epoch boundary is reached so without
- 2:28:23replacement means that they're not
- 2:28:24sampling from some fixed pool and then
- 2:28:27uh take a sequence train on it but then
- 2:28:31also like return the sequence to the
- 2:28:32pool they are exhausting a pool so when
- 2:28:34they draw a sequence it's it's gone
- 2:28:37until the next Epoch of training uh so
- 2:28:39we're already doing that because our
- 2:28:41data loader um iterates over chunks of
- 2:28:44data so there's no replacement they
- 2:28:47don't become eligible to be drawn again
- 2:28:49until the next P so we're basically
- 2:28:51already doing
- 2:28:53that um all models use a weight decay of
- 2:28:560.1 to provide a small amount of
- 2:28:59regularization so let's Implement a
- 2:29:01weight Decay and you see here that I've
- 2:29:03already kind of made the changes and in
- 2:29:04particular instead of creating the
- 2:29:06optimizer right here um I I'm creating a
- 2:29:10new configure optimizers function inside
- 2:29:12the model and I'm passing in some of the
- 2:29:14hyper parameters instead so let's look
- 2:29:17at the configure optimizers which is
- 2:29:18supposed to return the optimizer
- 2:29:24object okay so it looks complicated but
- 2:29:27it's actually really simple and it's
- 2:29:29just um we're just being very careful
- 2:29:31and there's a few settings here to go
- 2:29:32through the most important thing with
- 2:29:34respect to this line is that you see
- 2:29:36there's a weight Decay parameter here
- 2:29:38and I'm passing that
- 2:29:41into um well I'm passing that into
- 2:29:44something called optim groups that
- 2:29:46eventually ends up going into the addom
- 2:29:47W Optimizer um and the weight Decay
- 2:29:50that's by default used in Addam W here
- 2:29:53is 0.01 so it's it's u 10 times lower
- 2:29:57than what's used in gpt3 paper here um
- 2:30:01so the weight dek basically ends up
- 2:30:02making its way into the ADD and W
- 2:30:04through the optimizer groups now what
- 2:30:05else is going on here in this uh
- 2:30:07function so the two things that are
- 2:30:09happening here that are important is
- 2:30:10that I'm splitting up the parameters
- 2:30:12into those that should be weight decayed
- 2:30:14and those that should not be weight
- 2:30:15decayed so in particular it is common to
- 2:30:18not weight decay uh biases and any other
- 2:30:22sort of one-dimensional tensors so the
- 2:30:25one-dimensional tensors are in the no
- 2:30:27Decay prams and these are also things
- 2:30:30like uh layer Norm scales and biases it
- 2:30:33doesn't really make sense to weight
- 2:30:34Decay those you mostly want to weight
- 2:30:36Decay uh the weights that participate in
- 2:30:39Matrix multiplications and you want to
- 2:30:41potentially weight Decay the
- 2:30:43embeddings and uh We've covered in
- 2:30:46previous video why it makes sense to
- 2:30:47Decay the weights because you can sort
- 2:30:49of the it as a regularization because
- 2:30:51when you're pulling down all the weights
- 2:30:53you're forcing the optimization to use
- 2:30:55more of the weights um and you're not
- 2:30:57allowing any one of the weights
- 2:30:59individually to be way too large um
- 2:31:02you're forcing you're forcing the
- 2:31:03network to kind of like distribute the
- 2:31:05work across more channels because
- 2:31:07there's sort of like a pull of gravity
- 2:31:09on the weights
- 2:31:11themselves um so that's why we are
- 2:31:13separating it in those ways here we're
- 2:31:16only decaying the embeddings and the
- 2:31:18mmal participating ways
- 2:31:21uh we're printing the number of uh
- 2:31:22parameters that we decaying and not most
- 2:31:24of the parameters will be decayed and
- 2:31:26then one more thing that we're doing
- 2:31:27here is I'm doing another optimization
- 2:31:31here and previous add and W did not have
- 2:31:34this option but later parts of pytorch
- 2:31:37introduced it and that's why I'm
- 2:31:38guarding it with an inspect do signature
- 2:31:41which is basically checking if this
- 2:31:43fused um quar is present inside atom W
- 2:31:48and then if it is present I'm going to
- 2:31:50end up using it and passing it in here
- 2:31:53because some earlier versions do not
- 2:31:55have fused equals so here's adamw fused
- 2:31:58equals it did not used to exist and it
- 2:32:00was added later and there's some docks
- 2:32:03here for what's happening and basically
- 2:32:05they say that by default they do not use
- 2:32:07fused because it is relatively new and
- 2:32:10we want to give it sufficient big time
- 2:32:12so by default they don't use fused but
- 2:32:13fused is a lot faster when it is
- 2:32:15available and when you're running on
- 2:32:17Cuda and what that does is in instead of
- 2:32:20iterating in a for Loop over all the
- 2:32:22parameter tensors and updating them that
- 2:32:25would launch a lot of kernels right and
- 2:32:27so a fused just means that it's a um all
- 2:32:30those kernels are fused into a single
- 2:32:31kernel you get rid of a lot of overhead
- 2:32:34and you a single time on all the
- 2:32:36parameters call a uh kernel that updates
- 2:32:39them and so it's just basically a kernel
- 2:32:42Fusion for the atom W update instead of
- 2:32:44iterating over all the
- 2:32:47tensors so that's the configure
- 2:32:48optimizers function that I like to use
- 2:32:51and we can rerun and we're not going to
- 2:32:53see any major differences from what we
- 2:32:55saw before but we are going to see some
- 2:32:57prints uh coming from here so let's just
- 2:33:00take a look at what they look
- 2:33:01like so we see that number of Decay
- 2:33:04tensors is 50 and it's most of the
- 2:33:06parameters and number of non- deay
- 2:33:08tensors is 98 and these are the biases
- 2:33:10and the layer Norm parameters mostly and
- 2:33:13that's there's only 100,000 of those so
- 2:33:15most of it is decayed and then we are
- 2:33:18using the fused implementation of ATM W
- 2:33:20which will be a lot faster so if you
- 2:33:22have it available I would advise you to
- 2:33:24use it I'm not actually 100% sure why
- 2:33:26they don't default to it it seems fairly
- 2:33:28benign and
- 2:33:29harmless and also because we are using
- 2:33:31the fused implementation I think this is
- 2:33:34why we have dropped um notice that the
- 2:33:37running time used to be 93 milliseconds
- 2:33:39per step and we're now down to 90
- 2:33:41milliseconds per step because of using
- 2:33:43the fused atom W Optimizer so in a
- 2:33:46single commit here we are introducing
- 2:33:48fused atom getting improvements on the
- 2:33:51time and we're adding or changing the
- 2:33:54weight Decay but we're only weight
- 2:33:56decaying the two dimensional parameters
- 2:33:58the embeddings and the matrices that
- 2:34:00participate in linear so that is this
- 2:34:03and we can take this out and uh yeah
- 2:34:06that is it for this line one more quick
- 2:34:10note before we continue here I just want
- 2:34:11to point out that the relationship
- 2:34:13between weight Decay learning rate batch
- 2:34:15size the atom parameters beta 1 beta 2
- 2:34:18the Epsilon and so on these are very
- 2:34:20complicated uh mathematical
- 2:34:22relationships in the optimization
- 2:34:24literature and um for the most part I'm
- 2:34:27in this video I'm just trying to copy
- 2:34:29paste the settings that open AI used but
- 2:34:31this is a complicated topic uh quite
- 2:34:33deep and um yeah in this video I just
- 2:34:36want to copy the parameters because it's
- 2:34:38a whole different video to really talk
- 2:34:39about that in detail and give it a
- 2:34:41proper Justice instead of just high
- 2:34:42level
- 2:34:43intuitions uh now the next thing that I
- 2:34:45want to move on to is that uh this
- 2:34:48paragraph here by the way we're going to
- 2:34:49turn back around to when we improve our
- 2:34:51data loader for now I want to swing back
- 2:34:54around
- 2:34:56to this
- 2:35:01table where you will notice that um for
- 2:35:04different models we of course have
- 2:35:06different U hyper parameters for the
- 2:35:08Transformer that dictate the size of the
- 2:35:10Transformer Network we also have a
- 2:35:12different learning rate so we're seeing
- 2:35:13the pattern that the bigger networks are
- 2:35:14trained with slightly lower learning
- 2:35:16rates and we also see this batch size
- 2:35:20where in in the small networks they use
- 2:35:22a smaller batch size and in the bigger
- 2:35:23networks they use a bigger batch size
- 2:35:26now the problem with for us is we can't
- 2:35:28just use 0.5 million batch size because
- 2:35:31uh if I just try to come in here and I
- 2:35:33try to set uh this uh B where is my
- 2:35:38b
- 2:35:40um b
- 2:35:44equals where where do I call the DAT
- 2:35:46okay b equal 16 if I try to set um
- 2:35:51well well we have to be careful it's not
- 2:35:520.5 million because this is the badge
- 2:35:54size in the number of tokens every
- 2:35:56single one of our rows is24 tokens so
- 2:36:000.5 E6 1 million divide 1024 this would
- 2:36:04need about a
- 2:36:06488 match size so the problem is I can't
- 2:36:09come in here and set this to 488 uh
- 2:36:12because my GPU would explode um this
- 2:36:15would not fit for sure and so but we
- 2:36:18still want to use this batch size
- 2:36:20because again as I mentioned the batch
- 2:36:22size is correlated with all the other
- 2:36:24optimization hyper parameters and the
- 2:36:26learning rates and so on so we want to
- 2:36:28have a faithful representation of all
- 2:36:29the hyper parameters and therefore we
- 2:36:31need to uh use a bat size of .5 million
- 2:36:34roughly but the question is how do we
- 2:36:37use .5 million if we only have a small
- 2:36:39GPU well for that we need to use what's
- 2:36:41called gradient accumulation uh so we're
- 2:36:44going to turn to that next and it allows
- 2:36:46us to simulate in a Serial way any
- 2:36:48arbitrary batch size that we set and so
- 2:36:51we can do a batch size of .5 million we
- 2:36:54just have to run longer and we have to
- 2:36:56process multiple sequences and basically
- 2:36:59add up all the gradients from them to
- 2:37:02simulate a batch size of .5 million so
- 2:37:04let's turn to that next okay so I
- 2:37:05started the implementation right here
- 2:37:07just by adding these lines of code and
- 2:37:09basically what I did is first I set the
- 2:37:12total batch size that we desire so this
- 2:37:14is exactly .5 million and I used a nice
- 2:37:17number a power of two uh because 2 to
- 2:37:19the 19 is 524 288 so it's roughly .5
- 2:37:23million it's a nice number now our micro
- 2:37:26batch size as we call it now is 16 so
- 2:37:29this is going to be we still have B BYT
- 2:37:32in the SE that go into the Transformer
- 2:37:34and do forward backward but we're not
- 2:37:36going to do an update right we're going
- 2:37:38to do many forward backwards we're going
- 2:37:40to and those gradients are all going to
- 2:37:42plus equals on the parameter gradients
- 2:37:44they're all going to add up so we're
- 2:37:46going to do forward backward grad akum
- 2:37:48steps number of times and then we're
- 2:37:50going to do a single update once all
- 2:37:52that is
- 2:37:53accumulated so in particular our micro
- 2:37:55batch size is just now controlling how
- 2:37:58many tokens how many rows we're
- 2:37:59processing in a single go over a forward
- 2:38:02backward so um here we are doing 16 *
- 2:38:06124 we're doing 16
- 2:38:09384 um tokens per forward backward and
- 2:38:14we are supposed to be doing 2 to the 19
- 2:38:17whoops what am I doing 2 to the
- 2:38:2019 in total so the grat Aon will be
- 2:38:2632 uh so therefore gr AUM here will work
- 2:38:28out to 32 and we have to do 32 forward
- 2:38:32backward um and then a single update now
- 2:38:35we see that we have about 100
- 2:38:37milliseconds for a singer forward
- 2:38:38backward so doing 32 of them will be
- 2:38:41will make every step roughly 3 seconds
- 2:38:44just napkin
- 2:38:46math so that's grum steps but now we
- 2:38:48actually have to Implement that so we're
- 2:38:50going to swing over to our training Loop
- 2:38:54because now this part
- 2:38:56here and this part here the forward and
- 2:38:59the backward we have to now repeat this
- 2:39:0132 times before we do everything else
- 2:39:04that follows so let's uh see how we can
- 2:39:06Implement that so let's come over here
- 2:39:09and actually we do have to load a new
- 2:39:10batch every single time so let me move
- 2:39:12that over here and now this is where we
- 2:39:14have the inner loop so for micro step in
- 2:39:18range graum
- 2:39:20steps we do this and remember that l.
- 2:39:24backward always deposits gradients so
- 2:39:26we're doing inside losta backward
- 2:39:27there's always a plus equals on the
- 2:39:29gradients so in every single L of
- 2:39:31backward gradients will add up on the
- 2:39:33gradient
- 2:39:35tensors um so we lost that backward and
- 2:39:38then we get all the gradients over there
- 2:39:41and then we normalize and everything
- 2:39:43else should just follow um so we're very
- 2:39:47close but actually there's like subtle
- 2:39:50and deep issue here and this is actually
- 2:39:52incorrect so invite I invite you to
- 2:39:54think about why this is not yet
- 2:39:56sufficient um and uh let me fix it then
- 2:39:59okay so I brought back the jupyter
- 2:40:01notebook so we can think about this
- 2:40:02carefully in a simple toy setting and
- 2:40:05see what's happening so let's create a
- 2:40:07very simple neural nut that takes a 16
- 2:40:10Vector of 16 numbers and returns a
- 2:40:11single
- 2:40:12number and then here I'm creating some
- 2:40:15random uh examples X and some targets uh
- 2:40:19y Y and then we are using the mean
- 2:40:21squared loss uh here to calculate the
- 2:40:25loss so basically what this is is four
- 2:40:28individual examples and we're just doing
- 2:40:30Simple regression with the mean squared
- 2:40:31loss over those four
- 2:40:34examples now when we calculate the loss
- 2:40:36and we lost that backward and look at
- 2:40:38the gradient this is the gradient that
- 2:40:40we
- 2:40:41achieve now the loss objective here
- 2:40:44notice that in MSE loss the default for
- 2:40:46the loss function is reduction is mean
- 2:40:49so we're we're calculating the average
- 2:40:52mean loss um the the mean loss here over
- 2:40:56the four examples so this is the exact
- 2:40:59loss objective and this is the average
- 2:41:02the one over four because there are four
- 2:41:03independent examples here and then we
- 2:41:06have the four examples and their mean
- 2:41:08squared error the squared error and then
- 2:41:11this makes it the mean squared error so
- 2:41:14therefore uh we are we calculate the
- 2:41:16squared error and then we normalize it
- 2:41:18to make it the mean over the examples
- 2:41:20and there's four examples here so now
- 2:41:22when we come to the gradient
- 2:41:24accumulation version of it this uh this
- 2:41:28here is the gradient accumulation
- 2:41:30version of it where we have grad acum
- 2:41:32steps of four and I reset the gradient
- 2:41:35we've grum steps of four and now I'm
- 2:41:38evaluating all the examples individually
- 2:41:39instead and calling L that backward on
- 2:41:41them many times and then we're looking
- 2:41:43at the gradient that we achieve from
- 2:41:44that so basically now we forward our
- 2:41:47function calculate the exact same loss
- 2:41:49do a backward and we do that four times
- 2:41:52and when we look at the gradient uh
- 2:41:54you'll notice that the gradients don't
- 2:41:57match so here we uh did a single batch
- 2:42:00of four and here we did uh four gradient
- 2:42:03accumulation steps of batch size one and
- 2:42:06the gradients are not the same and
- 2:42:08basically the the reason that they're
- 2:42:09not the same is exactly because this
- 2:42:11mean squared error gets lost this one
- 2:42:14quarter in this loss gets lost because
- 2:42:16what happens here is the loss of
- 2:42:19objective for every one of the loops is
- 2:42:22just a mean squ error um which in this
- 2:42:25case because there's only a single
- 2:42:26example is just this term here so that
- 2:42:28was the loss in the zeroth eration same
- 2:42:30in the first third and so on and then
- 2:42:33when you do the loss. backward we're
- 2:42:35accumulating gradients and what happens
- 2:42:38is that accumulation in the gradient is
- 2:42:40basically equivalent to doing a sum in
- 2:42:43the
- 2:42:45loss so our loss actually here is this
- 2:42:49without the factor of one quarter
- 2:42:51outside of it so we're missing the
- 2:42:54normalizer and therefore our gradients
- 2:42:56are off and so the way to fix this or
- 2:42:58one of them is basically we can actually
- 2:43:00come here and we can say loss equals
- 2:43:02loss divide
- 2:43:044 and what happens now is that we're
- 2:43:07introducing we're we're scaling our loss
- 2:43:09we're introducing a one quarter in front
- 2:43:11of all of these
- 2:43:14places so all the individual losses are
- 2:43:17now scaled by one quarter and and then
- 2:43:19when we backward all of these accumulate
- 2:43:22with a sum but now there's a one quarter
- 2:43:24inside every one of these components and
- 2:43:26now our losses will be
- 2:43:28equivalent so when I run this you see
- 2:43:32that the U gradients are now identical
- 2:43:35so long story short with this simple
- 2:43:37example uh when you step through it you
- 2:43:39can see that basically the reason that
- 2:43:41this is not correct is because in the
- 2:43:44same way as here in the MSE loss the
- 2:43:46loss that we're calculating here in the
- 2:43:50model is using a reduction of mean as
- 2:43:54well uh so where's the loss after that
- 2:43:57cross
- 2:43:58entropy and by default the reduction uh
- 2:44:01here in Cross entropy is also I don't
- 2:44:03know why they don't show it but it's the
- 2:44:05mean uh the mean uh loss at all the B
- 2:44:08BYT elements
- 2:44:10right so there's a reduction by mean in
- 2:44:13there and if we're just doing this
- 2:44:15gradient accumulation here we're missing
- 2:44:16that and so the way to fix this is to
- 2:44:19simply compensate for the number of
- 2:44:21gradient accumulation steps and we can
- 2:44:23in the same way divide this loss so in
- 2:44:25particular here the number of steps that
- 2:44:26we're doing is loss equals loss divide
- 2:44:31gradient accumulation steps so even uh
- 2:44:33co-pilot s gets the modification but in
- 2:44:36the same way exactly we are scaling down
- 2:44:38the loss so that when we do loss that
- 2:44:40backward which basically corresponds to
- 2:44:42a sum in the objective we are summing up
- 2:44:45the already
- 2:44:46normalized um loss and and therefore
- 2:44:49when we sum up the losses divided by
- 2:44:51grum steps we are recovering the
- 2:44:53additional normalizer uh and so now
- 2:44:56these two will be now this will be
- 2:44:59equivalent to the original uh sort of
- 2:45:01optimization because the gradient will
- 2:45:03come out the same okay so I had to do a
- 2:45:05few more touch-ups and I launched
- 2:45:07launched the optimization here so in
- 2:45:09particular one thing we want to do
- 2:45:10because we want to print things nicely
- 2:45:13is well first of all we need to create
- 2:45:15like an accumulator over the loss we
- 2:45:16can't just print the loss because we'd
- 2:45:18be printing only the final loss at the
- 2:45:20final micro step so instead we have loss
- 2:45:22ofon which I initialize at zero and then
- 2:45:25I accumulate a uh the loss into it and
- 2:45:28I'm using detach so that um uh I'm
- 2:45:31detaching the tensor uh from the graph
- 2:45:35and I'm just trying to keep track of the
- 2:45:36values so I'm making these Leaf nodes
- 2:45:38when I add them so that's lakum and then
- 2:45:42we're printing that here instead of loss
- 2:45:43and then in addition to that I had to
- 2:45:46account for the grum steps inside the
- 2:45:48tokens processed because now the tokens
- 2:45:50processed per step is B * T * gradient
- 2:45:54accumulation so long story short here we
- 2:45:57have the optimization it looks uh
- 2:45:59reasonable right we're starting at a
- 2:46:00good spot we calculated the grum steps
- 2:46:03to be
- 2:46:0432 and uh we're getting about 3 seconds
- 2:46:07here
- 2:46:08right
- 2:46:10um
- 2:46:12and so this looks pretty good now if
- 2:46:14you'd like to verify that uh your
- 2:46:16optimization and the implementation here
- 2:46:18is correct and your working on a side
- 2:46:20well now because we have the total patch
- 2:46:21size and the gradient accumulation steps
- 2:46:24our setting of B is purely a performance
- 2:46:26optimization kind of setting so if you
- 2:46:29have a big GPU you can actually increase
- 2:46:31this to 32 and you'll probably go a bit
- 2:46:33faster if you have a very small GPU you
- 2:46:35can try eight or four but in any case
- 2:46:37you should be getting the exact same
- 2:46:38optimization and the same answers up to
- 2:46:41like a floating Point error because the
- 2:46:43gradient accumulation kicks in and um
- 2:46:46and can um handle everything serially as
- 2:46:48an
- 2:46:49Neary so uh that's it for gradient
- 2:46:51accumulation I think okay so now is the
- 2:46:53time to bring out the heavy weapons uh
- 2:46:56you've noticed that so far we've only
- 2:46:57been using a single GPU for training but
- 2:47:00actually I am paying for eight gpus here
- 2:47:02and so uh we should be putting all of
- 2:47:04them to work and in particular they are
- 2:47:06going to collaborate and uh you know
- 2:47:09optimize over tokens at the same time
- 2:47:12and communicate so that um uh they're
- 2:47:15all kind of collaborating on the
- 2:47:16optimization for this we are going to be
- 2:47:18using the distributed data parallel from
- 2:47:20pytorch there's also a legacy data
- 2:47:22parallel which I recommend you not use
- 2:47:24and that's kind of like you know Legacy
- 2:47:27distributed data parallel Works in a
- 2:47:28very simple way we have eight gpus so
- 2:47:31we're going to uh launch eight processes
- 2:47:35and each process is going to be assigned
- 2:47:36to GPU and for each process the training
- 2:47:40Loop and everything we've worked on so
- 2:47:41far is going to look pretty much the
- 2:47:42same H GPU as far as it's concerned is
- 2:47:45just working on exactly what we've built
- 2:47:47so far but now Secret L there's eight of
- 2:47:49them and they're all going to be
- 2:47:51processing slightly different parts of
- 2:47:52the data and we're going to add one more
- 2:47:56part where once they all calculate their
- 2:47:58gradients there's one more part where we
- 2:48:00do a average of those
- 2:48:03gradients and so that's how they're
- 2:48:05going to be collaborating on uh the
- 2:48:07computational workload here so to use
- 2:48:10all eight of them we're not going to be
- 2:48:12launching our script anymore with just
- 2:48:14um pytorch train
- 2:48:16gbt2 piy we're going to be running it
- 2:48:19with a special command called torrun in
- 2:48:21pytorch we'll see that in a bit and
- 2:48:23torrun uh when it runs our python script
- 2:48:26we'll actually make sure to run eight
- 2:48:28eight of them in parallel and it creates
- 2:48:32these environmental variables where each
- 2:48:34of these processes can look up which uh
- 2:48:37basically which one of the processes it
- 2:48:40is so for example torron will set rank
- 2:48:43local Rank and World size environmental
- 2:48:46variables and so this is a bad way to
- 2:48:48detect whether uh DDP is running so if
- 2:48:51we're using torch run if DDP is
- 2:48:54running then uh we have to make sure
- 2:48:57that K is available because I don't know
- 2:48:58that you can run this on CPU anymore or
- 2:49:01that that makes sense to do um this is
- 2:49:05some um setup code here the important
- 2:49:07part is that there's a world size which
- 2:49:10for us will be eight that's the total
- 2:49:11number of processes running there's a
- 2:49:14rank which is um each process will
- 2:49:17basically run the ex exact same code at
- 2:49:19the exact same time roughly but all the
- 2:49:22process the only difference between
- 2:49:24these processes is that they all have a
- 2:49:26different dtp rank so the um gpu0 will
- 2:49:30have DDP rank of zero GPU 1 will have uh
- 2:49:33rank of one Etc so otherwise they're all
- 2:49:36running the exact same script it's just
- 2:49:38that DDP rank will be a slightly
- 2:49:40different integer and that is the way
- 2:49:42for us to coordinate that they don't for
- 2:49:44example run on the same data we want to
- 2:49:46we want them to run on different parts
- 2:49:47of the data and so on
- 2:49:49now local rank is something that is only
- 2:49:52used in a multi- node setting we only
- 2:49:54have a single node with ag gpus and so
- 2:49:57local rank is the rank of the GPU on a
- 2:50:00single node so from 0 to seven as an
- 2:50:04example but for us we're mostly going to
- 2:50:06be running on a single box so the things
- 2:50:08we care about are Rank and World size
- 2:50:10this is eight and this will be whatever
- 2:50:12it is depending on the GPU uh that uh
- 2:50:15that this particular instantiation of
- 2:50:17the script runs on
- 2:50:19now here we make sure that according to
- 2:50:23the local rank we are setting the device
- 2:50:27to be Cuda colon and colon indicates
- 2:50:30which GPU to use if there are more than
- 2:50:32one gpus so depending on the local rank
- 2:50:36of this process it's going to use just
- 2:50:39the appropriate GPU so there's no
- 2:50:40collisions on which GPU is being used by
- 2:50:42which
- 2:50:43process and finally there's a Boolean
- 2:50:45variable that I like to create which is
- 2:50:47the DDP rank equ equal Z so the master
- 2:50:50process is arbitrarily process number
- 2:50:53zero and it does a lot of the printing
- 2:50:55logging checkpointing Etc and the other
- 2:50:57processes are thought of mostly as a
- 2:50:59compute processes that are assisting and
- 2:51:01so Master process zero will have some
- 2:51:03additional work to do all the other
- 2:51:05processes will uh will mostly just be
- 2:51:06doing forward
- 2:51:08backwards and if we're not using DDP and
- 2:51:10none of these variables are set we
- 2:51:12revert back to single GPU training so
- 2:51:14that means that we only have rank zero
- 2:51:16the world size is just one uh and and we
- 2:51:19are the master process and we try to
- 2:51:21autodetect the device and this is world
- 2:51:24as
- 2:51:25normal so so far all we've done is we've
- 2:51:27initialized
- 2:51:28DDP and uh in the case where we're
- 2:51:31running with torrun which we'll see in a
- 2:51:33bit there's going to be eight copies
- 2:51:35running in parallel each one of them
- 2:51:37will have a different Rank and now we
- 2:51:39have to make sure that everything
- 2:51:41happens uh correctly afterwards so the
- 2:51:44tricky thing with running multiple
- 2:51:45processes is you always have to imagine
- 2:51:48that there's going to be eight processes
- 2:51:50running in parallel so as you read the
- 2:51:52code now you have to imagine there's
- 2:51:54eight you know eight python interpreters
- 2:51:57running down these lines of code and the
- 2:51:59only difference between them is that
- 2:52:01they have a different DDP rank so they
- 2:52:03all come here they all pick the exact
- 2:52:05same seed they all make all of these
- 2:52:08calculations completely unaware of the
- 2:52:10other copies running roughly speaking
- 2:52:12right so they all make the exact same
- 2:52:14calculations and now we have to adjust
- 2:52:16these calculations to take into account
- 2:52:19that there's actually like a certain
- 2:52:21world size and certain ranks so in
- 2:52:24particular these micro batches and
- 2:52:26sequence lengths these are all just per
- 2:52:28GPU right so now there's going to be num
- 2:52:31processes of them running in parallel so
- 2:52:34we have to adjust this right because the
- 2:52:36grum steps now is going to be total B
- 2:52:39size divide B * T time U DDP R
- 2:52:43size because each um process will will
- 2:52:48do B * T and there's this many of
- 2:52:51them and so in addition to that we we
- 2:52:54want to make sure that this fits nicely
- 2:52:56into total batch size which for us it
- 2:52:58will because 16 * 124 * 8 8 gpus is
- 2:53:04131 uh K and so
- 2:53:08524288 this means that our gratum will
- 2:53:10be four with the current settings right
- 2:53:13so there's going to be 16 * 124 process
- 2:53:16on each GPU and then there's a GP pus so
- 2:53:18we're going to be doing
- 2:53:20131,000 tokens in a single forward
- 2:53:23backward on the 8
- 2:53:26gpus so we want to make sure that this
- 2:53:28fits nicely so that we can derive a nice
- 2:53:30gradient accumulation
- 2:53:32steps and uh yeah let's just adjust the
- 2:53:36comments here times uh DDP World size
- 2:53:41okay so each GPU calculates this now
- 2:53:45this is where we start to get run into
- 2:53:46issues right so we are each process is
- 2:53:49going to come by a print and they're all
- 2:53:51going to print so we're going to have
- 2:53:53eight copies of these prints so one way
- 2:53:56to deal with this is exactly this master
- 2:53:58process variable that we have so if
- 2:54:00Master process then guard this and
- 2:54:03that's just so that we just print this a
- 2:54:05single time because otherwise all the
- 2:54:07processes would have computed the exact
- 2:54:08same variables and there's no need to
- 2:54:10print this eight
- 2:54:11times um before getting into the data
- 2:54:14loader and we're going to have to
- 2:54:15refactor it obviously maybe at this
- 2:54:18point is uh we should do some prints and
- 2:54:21uh just take it out for a spin and exit
- 2:54:23at this point so import
- 2:54:26sis and S start exit and print IM
- 2:54:33GPU um DDP
- 2:54:38rank IM GPU DDP Rank and that um
- 2:54:43print
- 2:54:46by so uh so now let's try to run this
- 2:54:49and just see how this works so let's
- 2:54:51take it for a spin just so we see what
- 2:54:52it looks like so normally we use to
- 2:54:54launch python train gpd2 P like this now
- 2:54:57we're going to run with torch run and
- 2:54:59this is what it looks like so torch run
- 2:55:02Standalone number of processes for
- 2:55:04example is eight for us because we have
- 2:55:05eight gpus uh and then change of2 Pi so
- 2:55:09this is what the command would look like
- 2:55:11and torch run again we'll run eight of
- 2:55:13these so let's just see what happens so
- 2:55:16first
- 2:55:18it gets a little busy so there's a lot
- 2:55:20going on here so first of all there's
- 2:55:22some warnings from distributed and I
- 2:55:24don't actually know that these mean
- 2:55:26anything I think this is just like the
- 2:55:28code is setting up and the processes are
- 2:55:29coming online and we're seeing some
- 2:55:31preliminary failure to collect while the
- 2:55:33processes come up I'm not 100% sure
- 2:55:36about that but we start to then get into
- 2:55:39actual prints
- 2:55:41so all the processes went down and then
- 2:55:44the first print actually comes from
- 2:55:46process 5 uh just by chance and then it
- 2:55:50printed so process 5 basically got here
- 2:55:52first it said I'm process on GPU 5 buy
- 2:55:56and then this these prints come from the
- 2:56:00master
- 2:56:01process so process 5 just finished first
- 2:56:04for whatever reason it just depends on
- 2:56:05how the operating system scheduled the
- 2:56:07processes to run uh then gpu0 ended then
- 2:56:10GPU 3 and two and then uh probably
- 2:56:14process 5 or something like that has uh
- 2:56:17exited and and DDP really doesn't like
- 2:56:19that because we didn't properly dispose
- 2:56:21of uh the multi-gpus um setting and so
- 2:56:27process group has not been destroyed
- 2:56:28before we destruct uh so it really
- 2:56:31doesn't like that and in an actual
- 2:56:33application we would want to call
- 2:56:34destroy process group uh so that we
- 2:56:37clean up DDP properly and so it doesn't
- 2:56:40like that too much and then the rest of
- 2:56:41the gpus finish and that's it so
- 2:56:45basically we can't guarantee when these
- 2:56:46processes are running it's totally
- 2:56:48but they are running in parallel we
- 2:56:50don't want them to be printing um and
- 2:56:54next up let's erase
- 2:56:57this next up we want to make sure that
- 2:56:59when we create data loader light we need
- 2:57:01to now make it aware of this
- 2:57:03multi-process um setting because we
- 2:57:06don't want all the processes to be
- 2:57:07loading the exact same data we want
- 2:57:10every process to get its own chunk of
- 2:57:11data so that they're all working on
- 2:57:13different parts of the data set of
- 2:57:14course so let's adjust that so one
- 2:57:17particular particularly simple and a
- 2:57:19naive way to do this is we have to make
- 2:57:21sure that we pass in the rank and the
- 2:57:23size to the data
- 2:57:25loader and then when we come up here we
- 2:57:28see that we now take Rank and processes
- 2:57:29and we save them now the current
- 2:57:32position will not be zero uh because
- 2:57:35what we want is we want to stride out
- 2:57:37all the processes so one way to do this
- 2:57:40is we basically take S.B times salt. T
- 2:57:43and then multiply it by the process
- 2:57:46rank so proc process rank 0 will start
- 2:57:49at zero but process rank one now starts
- 2:57:52at B * T process rank two is starts at 2
- 2:57:55* B * D Etc so that is the
- 2:57:59initialization now we still they still
- 2:58:01do this identically but now when we
- 2:58:04advance we don't Advance by B * T we
- 2:58:06advance by B * T times number of
- 2:58:10processes right so basically um the
- 2:58:14total number of tokens that we're um
- 2:58:16consuming is B * T * number processes
- 2:58:19and they all go off to a different Rank
- 2:58:23and the position has to advance by the
- 2:58:24entire
- 2:58:26chunk and then here B * T time uh s. num
- 2:58:30processes + one would be to exceed
- 2:58:33number of tokens then we're going to
- 2:58:35Loop and when we Loop we want to of
- 2:58:37course Loop in the exact same way so we
- 2:58:39sort of like reset back uh so this is
- 2:58:42the simplest change that I can uh find
- 2:58:45for kind of a very simple distributed
- 2:58:47data Lo light and um you can notice that
- 2:58:50if process rank is zero and non
- 2:58:52processes is one then uh the whole thing
- 2:58:54will be identical to what we had before
- 2:58:56but now we can have actually multiple
- 2:58:58processes uh running and this should
- 2:59:00work
- 2:59:01fine um so that's the data loader okay
- 2:59:05so next up once they've all initialized
- 2:59:07the data loader they come here and they
- 2:59:09all create a GPT model uh so we create
- 2:59:13eight GPT models on eight processes but
- 2:59:15because the seeds are fixed here they
- 2:59:17all create the same identical model they
- 2:59:20all move it to the device of their Rank
- 2:59:22and they all compile the model and
- 2:59:25because the models are identical there
- 2:59:26are eight identical compilations
- 2:59:28happening in parallel but that's okay
- 2:59:31now none of this uh changes because that
- 2:59:33is on a per step basis and we're
- 2:59:34currently working kind of within step
- 2:59:36because we need to um just uh all the
- 2:59:39all the changes we're making are kind of
- 2:59:41like a within step
- 2:59:42changes now the important thing here is
- 2:59:44when we construct the M model we
- 2:59:47actually have a bit of work to to do
- 2:59:48here get loits is deprecated so uh
- 2:59:50create
- 2:59:52model we need to actually wrap the model
- 2:59:55into the distributed data parallel
- 2:59:58container so um this is how we wrap the
- 3:00:01model into the DDP container and these
- 3:00:04are the docs for DDP and they're quite
- 3:00:07extensive and there's a lot of caveats
- 3:00:09and a lot of things to be careful with
- 3:00:10because everything complexifies times 10
- 3:00:12when multiple processes are involved but
- 3:00:15roughly speaking this device IDs I
- 3:00:17believe has to be passed in now
- 3:00:18unfortunately the docs for what device
- 3:00:20IDs is is is extremely unclear uh so
- 3:00:24when you actually like come here this
- 3:00:26comment for what device IDs is is
- 3:00:29roughly
- 3:00:30nonsensical um but I'm pretty sure it's
- 3:00:33supposed to be the DDP local rank so not
- 3:00:35the DDP rank the local rank uh so this
- 3:00:39is what you pass in here this wraps the
- 3:00:41model and in particular what DDP does
- 3:00:43for you is in a forward pass it actually
- 3:00:45behaves identically so um my
- 3:00:48understanding of it is nothing should be
- 3:00:49changed in the forward pass but in the
- 3:00:51backward pass as you are doing the
- 3:00:53backward pass um in the simpl setting
- 3:00:56once the backp passes over on each
- 3:00:59independent GPU each independent GPU has
- 3:01:02the gradient for all the parameters and
- 3:01:05what DDP does for you is once the
- 3:01:06backward pass is over it will call
- 3:01:09what's called all reduce and it
- 3:01:11basically does an average across all the
- 3:01:14uh ranks of their gradients and and then
- 3:01:18it will deposit that average on every
- 3:01:20single rank so every sing Single rank
- 3:01:22will end up with the average on it and
- 3:01:25so basically that's the communication it
- 3:01:27just synchronizes and averages the
- 3:01:28gradients and that's what DDP offers you
- 3:01:31now DDP actually is a little bit more um
- 3:01:34it is a little bit more involved than
- 3:01:35that because as you are doing the
- 3:01:37backward pass through the layers of the
- 3:01:38Transformer it actually can dispatch
- 3:01:41Communications for the gradient while
- 3:01:43the backward pass is still happening so
- 3:01:45there's overlap of the uh communication
- 3:01:47of the gradient and the synchronization
- 3:01:48of them and uh the backward pass and uh
- 3:01:52this is just more efficient and um uh to
- 3:01:55do it that way so that's what DDP does
- 3:01:57for you um forward is unchanged and
- 3:02:00backward is mostly unchanged and we're
- 3:02:02tacking on this average as we'll see in
- 3:02:04a bit okay so now let's go to the uh
- 3:02:08optimization nothing here changes let's
- 3:02:11go to the optimization here the inner
- 3:02:12loop and think through the
- 3:02:13synchronization of uh these gradients in
- 3:02:15the DP so basically by default what
- 3:02:18happens as I mentioned is when you do l.
- 3:02:20backward here it will do the backward
- 3:02:22pass and then it will synchronize the
- 3:02:24gradients um the problem here is because
- 3:02:28of the gradient accumulation steps Loop
- 3:02:30here we don't actually want to do the
- 3:02:33synchronization after every single La
- 3:02:35step backward because we are just
- 3:02:37depositing gradients and we're doing
- 3:02:39that serially and we just want them
- 3:02:40adding up and we don't want to
- 3:02:42synchronize every single time that would
- 3:02:44be extremely wasteful so basically we
- 3:02:46want to add them up and then on the the
- 3:02:48very last uh it's only on the very last
- 3:02:50step when micro when micro step becomes
- 3:02:53gratak steps minus one only at that last
- 3:02:55step do we want to actually do the
- 3:02:58alberu uh to average up the gradients so
- 3:03:02to do that we come here and um the
- 3:03:05official sanctioned way by the way is to
- 3:03:07do this no sync context manager so
- 3:03:10pytorch says this is a context manager
- 3:03:13to disable gradient synchronization
- 3:03:14across DDP processes So within this
- 3:03:17context gradient will be
- 3:03:19accumulated and basically when you do no
- 3:03:21sync there will be no communication so
- 3:03:24they are telling us to do with DDP no
- 3:03:26sync uh do the gradient accumulation
- 3:03:29accumulate grats and then they are
- 3:03:30asking us to do DDP again with another
- 3:03:32input and that backward and I just
- 3:03:35really don't love this I I just really
- 3:03:37don't like it uh the fact that you have
- 3:03:39to copy paste your code here and use a
- 3:03:40context manager and this is just super
- 3:03:42ugly so when I went to this source code
- 3:03:45here you can see that when you enter
- 3:03:48you simply toggle this variable this
- 3:03:51require backward grat sync and this is
- 3:03:54uh being toggled around and changed and
- 3:03:58this is the variable that basically uh
- 3:04:01if you step through it is being toggled
- 3:04:03to determine if the gradient is going to
- 3:04:05be synchronized so I actually just kind
- 3:04:07of like to use that directly uh so
- 3:04:10instead what I like to do is the
- 3:04:13following right here before the L back
- 3:04:15backward if we are using the DDP then um
- 3:04:20then basically we only want to
- 3:04:23synchronize we only want this variable
- 3:04:25to be true when it is the final
- 3:04:28iteration in all the other iterations
- 3:04:31inside the micr steps we want to be
- 3:04:33false so I just toggle it like this so
- 3:04:36required backward graph sync should only
- 3:04:38turn on when the micro step is the last
- 3:04:41step and so I'm toggling this variable
- 3:04:44directly and I hope that that impacts
- 3:04:47last St backwards
- 3:04:48and this is a naughty thing to do
- 3:04:49because you know they could probably
- 3:04:51change the DDP and this variable will go
- 3:04:53away but for now I believe this this
- 3:04:55works and it allows me to avoid the use
- 3:04:57of context managers and code duplication
- 3:05:00I'm just toggling the variable and then
- 3:05:01Lop backward will not synchronize most
- 3:05:03of the steps and it will synchronize the
- 3:05:04very last step and so once this is over
- 3:05:08uh and we come out every single um rank
- 3:05:13will suddenly magically have the average
- 3:05:17of all the gradients that were stored on
- 3:05:20all the ranks so now we have to think
- 3:05:22through whether that is what we want and
- 3:05:24also um if this suffices and whether how
- 3:05:29it works with the loss and what is loss
- 3:05:31AUM so let's think through through that
- 3:05:33now and the problem I'm getting at is
- 3:05:35that we've averaged the gradients which
- 3:05:37is great but the loss AUM has not been
- 3:05:40impacted yet and the and this is outside
- 3:05:43of the DDP container so that is not
- 3:05:45being averaged um and so here when when
- 3:05:47we are printing Los AUM well presumably
- 3:05:49we're only going to be printing on the
- 3:05:51master process uh rank zero and it's
- 3:05:53just going to be printing the losses
- 3:05:55that it saw on its process but instead
- 3:05:57we want it to print the loss over all
- 3:06:00the processes and the average of that
- 3:06:02loss because we did average of gradients
- 3:06:04so we want the average of loss as well
- 3:06:06so simply here after this uh this is the
- 3:06:09code that I've used in the past um and
- 3:06:13instead of LF we want
- 3:06:15Lum so if
- 3:06:18DDP again then this is a p torch
- 3:06:22distributed I import it where do I
- 3:06:24import
- 3:06:26it uh oh gosh so this file is starting
- 3:06:30to get out of control huh so if uh so
- 3:06:33import torch. distributed as dist
- 3:06:36so dist.
- 3:06:38ALU and we're doing the average on Lum
- 3:06:42and so this lakum tensor exists on all
- 3:06:44the ranks when we call all use of
- 3:06:46average it creates the average of those
- 3:06:48numbers and it deposits that average on
- 3:06:51all the ranks so all the ranks after
- 3:06:53this um call will now contain L AUM uh
- 3:06:57averaged up and so when we print here on
- 3:07:00the master process the L AUM is
- 3:07:02identical in all the other ranks as well
- 3:07:04so here if Master process
- 3:07:07oops we want to print like this okay and
- 3:07:10finally we have to be careful because
- 3:07:12we're not processing even more tokens so
- 3:07:15times DDP World size
- 3:07:18that's number of tokens that we've
- 3:07:19processed up
- 3:07:21above
- 3:07:24and everything else should be fine uh
- 3:07:27the only other thing to be careful with
- 3:07:29is as I mentioned you want to destroy
- 3:07:31the process group so that we are nice to
- 3:07:33nickel and it's not going to uh to uh to
- 3:07:35DDP and it's not going to complain to us
- 3:07:38uh when we exit
- 3:07:40here so that should be it let's try to
- 3:07:43take it for a spin okay so I launched
- 3:07:44the script and it should be uh printing
- 3:07:46here imminently we're now training with
- 3:07:488 gpus at the same time so the gradient
- 3:07:51accumulation steps is not 32 it is now
- 3:07:53divide 8 and it's just four uh so um
- 3:07:58otherwise this is what the optimization
- 3:07:59now looks like and wow we're going
- 3:08:01really fast so we're processing 1.5
- 3:08:04million tokens uh per second now so
- 3:08:09these are some serious numbers and the
- 3:08:11tiny shakespare data set is so tiny that
- 3:08:12we're just doing like so many Epoch over
- 3:08:15it most likely but this is roughly what
- 3:08:17looks like um one thing that I had to
- 3:08:20fix by the way is that this was model.
- 3:08:23configure optimizers which Now doesn't
- 3:08:25work because model now is a DDP model so
- 3:08:27instead this has to become raw
- 3:08:29model. configure optimizers where raw
- 3:08:32model is something I create here so
- 3:08:35right after I wrap the model into DDP uh
- 3:08:38I have to create the raw model which in
- 3:08:40the case of DDP is a model. module is
- 3:08:43where it stores the raw and then module
- 3:08:46of gpt2 as we have it which contains the
- 3:08:49uh configure optimizers function that we
- 3:08:51want to call so that's one thing that I
- 3:08:53have to fix otherwise this seems to run
- 3:08:56now one thing you'll notice is that when
- 3:08:57you actually compare this run and the
- 3:08:59numbers in it to the just running a
- 3:09:01single GPU you'll notice that this is
- 3:09:04single GPU run with 32 gratum the
- 3:09:06numbers won't exactly match
- 3:09:09up and uh that's kind of a boring reason
- 3:09:11for why that happens uh the reason for
- 3:09:13that is that in the data loader we're
- 3:09:15basically just iterating through batches
- 3:09:17and slightly different way because now
- 3:09:18we're looking for an entire page of data
- 3:09:21and if that page uh for all the gpus if
- 3:09:24that chunk exceeds the number of tokens
- 3:09:26we just Loop and so actually the single
- 3:09:29GPU and the H GPU process will end up um
- 3:09:33resetting in a slightly different Manner
- 3:09:35and so our batches are slightly
- 3:09:36different and so we get slightly
- 3:09:38different numbers but one way to
- 3:09:39convince yourself that this is okay it
- 3:09:42just make the total batch size much
- 3:09:43smaller and the b and a t and then um
- 3:09:48so I think I used uh 4 * 124 * 8 so I
- 3:09:52used 32768 as a total patch size and
- 3:09:55then um so I made sure that the single
- 3:09:57GPU will do eight creting accumulation
- 3:10:00steps and then the multi-gpu and then
- 3:10:02you're reducing the boundary effects of
- 3:10:04the data loader and you'll see that the
- 3:10:06numbers match up so long story short
- 3:10:08we're now going really really fast the
- 3:10:10optimization is mostly consistent with
- 3:10:12gpt2 and three hyper parameters and uh
- 3:10:16we have outgrown our tiny Shakespeare
- 3:10:18file and we want to upgrade it so let's
- 3:10:20move to next to that next so let's now
- 3:10:22take a look at what data sets were used
- 3:10:23by gpt2 and gpt3 so gbt2 used this web
- 3:10:27Text data set that was never released um
- 3:10:30there's an attempt at reproducing it
- 3:10:32called open web text uh so basically
- 3:10:34roughly speaking what they say here in
- 3:10:35the paper is that they scraped all
- 3:10:37outbound links from Reddit and then uh
- 3:10:41with at least three Karma and that was
- 3:10:43kind of like their starting point and
- 3:10:44they collected all the web P all the web
- 3:10:45pages and all the text in them and so
- 3:10:48this was 45 million links and this ended
- 3:10:50up being 40 GB of text so uh so that's
- 3:10:54roughly what gpt2 says about its data
- 3:10:57set so it's basically outbound links
- 3:10:58from Reddit now when we go over to gpt3
- 3:11:01there's a training data set section and
- 3:11:03that's where they start to talk about um
- 3:11:05common coll which is a lot more uh used
- 3:11:09actually I think even gpt2 talked about
- 3:11:11common coll um but basically it's not a
- 3:11:14very high quality data set all by itself
- 3:11:16because it is extremely noisy this is a
- 3:11:18completely random subset of the internet
- 3:11:20and it's much worse than you think so
- 3:11:22people go into Great Lengths to filter
- 3:11:24common craw because there's good stuff
- 3:11:26in it but most of it is just like ad
- 3:11:27spam random tables and numbers and stock
- 3:11:30tickers and uh it's just total mess
- 3:11:35so that's why people like to train on
- 3:11:38these data mixtures that they curate and
- 3:11:41uh are careful with so a large chunk of
- 3:11:44these data mixtures typically will be
- 3:11:45common C like for example 50% of the
- 3:11:47tokens will be comic but then here in
- 3:11:50gpt3 they're also using web text to from
- 3:11:52before so that's Reddit outbound but
- 3:11:54they're also adding for example books
- 3:11:56and they're adding Wikipedia there's
- 3:11:58many other things you can decide to add
- 3:12:00now this data set for gpt3 was also
- 3:12:02never released so today some of the data
- 3:12:05sets that I'm familiar with that are
- 3:12:06quite good and would be representative
- 3:12:08of something along these lines are
- 3:12:10number one the red pajama data set or
- 3:12:12more specifically for example the slim
- 3:12:14pajama subset of the red pajama data set
- 3:12:17which is a cleaned and D duplicated
- 3:12:19version of it and just to give you a
- 3:12:21sense again it's a bunch of common crawl
- 3:12:24um C4 which is also as far as I know
- 3:12:27more common craw but processed
- 3:12:28differently and then we have GitHub
- 3:12:30books archive Wikipedia stack exchange
- 3:12:33these are the kinds of data sets that
- 3:12:35would go into these data mixtures now
- 3:12:37specifically the one that I like that
- 3:12:38came out recently is called Fine web
- 3:12:41data set uh so this is an attempt to
- 3:12:43basically collect really high quality
- 3:12:45common coll data and filter it in this
- 3:12:48case to 15 trillion tokens and then in
- 3:12:51addition to that more recently
- 3:12:52huggingface released this fine web edu
- 3:12:55subset which is 1.3 trillion of
- 3:12:58educational and 5.4 trillion of high
- 3:13:01educational content so basically they're
- 3:13:03trying to filter common C to very high
- 3:13:06quality educational subsets and uh this
- 3:13:09is the one that we will use there's a
- 3:13:11long uh web page here on fine web and
- 3:13:14they go into a ton of detail about how
- 3:13:16they process the data which is really
- 3:13:17fascinating reading by the way and I
- 3:13:19would definitely recommend if you're
- 3:13:20interested into Data mixtures and so on
- 3:13:22and how data gets processed at these
- 3:13:24scales a look at this uh page and more
- 3:13:27specifically we'll be working with the
- 3:13:28fine web edu I think and it's basically
- 3:13:32educational content from the
- 3:13:34internet uh they show that training on
- 3:13:36educational content in in their metrics
- 3:13:39um uh works really really well and we're
- 3:13:43going to use this sample 10 billion
- 3:13:46tokens subsample of it because we're not
- 3:13:49going to be training on trillions of
- 3:13:50tokens uh we're just going to train on
- 3:13:52uh 10 billion sample of the fine web edu
- 3:13:56because empirically in my previous few
- 3:13:58experiments this actually suffices to
- 3:14:00really get close to gpt2 Performance and
- 3:14:02it's um simple enough to work with and
- 3:14:04so let's work with the sample 10 uh BT
- 3:14:07so our goal will be to download it
- 3:14:10process it and make sure that our data
- 3:14:12loader can work with it so let's get to
- 3:14:15that okay so I introduced another um
- 3:14:18file here that will basically download
- 3:14:21Fine web edu from huging face data sets
- 3:14:24it will pre-process and pre- tokenize
- 3:14:26all of the data and it will save data
- 3:14:28shards to a uh folder on um local disk
- 3:14:34and so while this is running uh just
- 3:14:38wanted to briefly mention that you can
- 3:14:40kind of look through the data set viewer
- 3:14:41here just to get a sense of what's in
- 3:14:43here and it's kind of interesting I mean
- 3:14:45it's a it basically looks like it's
- 3:14:47working fairly well like it's talking
- 3:14:48about nuclear energy in France it's
- 3:14:51talking
- 3:14:52about Mexican
- 3:14:54America some mac PJs Etc so actually it
- 3:14:58seems like their filters are working
- 3:14:59pretty well uh the filters here by the
- 3:15:01way were applied automatically using um
- 3:15:04llama 370b I believe and so uh basically
- 3:15:08llms are judging which content is
- 3:15:10educational and that ends up making it
- 3:15:11through the filter uh so that's pretty
- 3:15:13cool now in terms of the script itself
- 3:15:16I'm not going to go through the full
- 3:15:17script because it's not as interesting
- 3:15:19and not as llm Centric but when you run
- 3:15:22this basically number one we're going to
- 3:15:24load the data set uh which this is all
- 3:15:26huging face code running this you're
- 3:15:28going to need to uh pip install data
- 3:15:31sets um so it's downloading the data set
- 3:15:35then it is tokenizing all of the
- 3:15:37documents inside this data set now when
- 3:15:39we tokenize the documents you'll notice
- 3:15:42that um to tokenize a single document uh
- 3:15:46we first
- 3:15:47start the tokens with the end of text
- 3:15:49token and this is a special token in the
- 3:15:51gpt2 tokenizer as you know so
- 3:15:5450256 is the ID of the end of text and
- 3:15:57this is what begins a document even
- 3:15:59though it's called end of text but this
- 3:16:01is uh the first token that begins a
- 3:16:03document then we extend with all of the
- 3:16:06tokens of that document then we create a
- 3:16:08numpy array out of that we make sure
- 3:16:11that all the tokens are between
- 3:16:14oh okay let me debug this
- 3:16:17okay so apologies for that uh it just
- 3:16:19had to do with me using a float division
- 3:16:21in Python it must be integer division so
- 3:16:23that this is an INT and everything is
- 3:16:25nice um okay but basically the
- 3:16:28tokenization here is relatively
- 3:16:29straightforward returns tokens in mp.
- 3:16:32un6 uh we're using .16 to save a little
- 3:16:35bit of space because 2 to the 16us 1 is
- 3:16:3965,000 so the gpt2 max token ID is well
- 3:16:43below that and then here there's a bunch
- 3:16:45of multiprocessing code and it's
- 3:16:47honestly not that exciting so I'm not
- 3:16:48going to step through it but we're
- 3:16:50loading the data set we're tokenizing it
- 3:16:52and we're saving everything to shards
- 3:16:55and the shards are numpy files uh so
- 3:16:58just storing a numpy array and uh which
- 3:17:01is very very similar to torch
- 3:17:03tensors and the first Shard 0000 is a
- 3:17:07Val a validation Shard and all the other
- 3:17:09shards are uh training shards and as I
- 3:17:12mentioned they all have 100 million
- 3:17:14tokens in them exactly um and and that
- 3:17:17just makes it easier to work with as to
- 3:17:20Shard the files because if we just have
- 3:17:22a single massive file sometimes they can
- 3:17:24be hard to work with on the disk and so
- 3:17:26sharting it is just kind of um nicer
- 3:17:28from that
- 3:17:30perspective and uh yeah so we'll just
- 3:17:32let this run this will be probably um
- 3:17:3630ish minutes or so and then we're going
- 3:17:38to come back to actually train on this
- 3:17:39data and we're going to be actually
- 3:17:41doing some legit pre-training in this
- 3:17:42case this is a good data set we're doing
- 3:17:45lots of tokens per second we have 8 gpus
- 3:17:48the code is ready and so we're actually
- 3:17:50going to be doing a serious training run
- 3:17:52so let's get P it back in a bit okay so
- 3:17:54we're back so uh if we LS edu fine web
- 3:17:58we see that there's now 100 charts in it
- 3:18:02um and that makes sense because each
- 3:18:03chart is 100 million tokens so 100
- 3:18:06charts of that is 10 billion tokens in
- 3:18:08total now swinging over to the main file
- 3:18:11I made some adjustments to our data
- 3:18:12loader again and that's because we're
- 3:18:14not running with uh Shakespeare anymore
- 3:18:17we want to use the fine web shards and
- 3:18:20so you'll see some code here that
- 3:18:21additionally basically can load these
- 3:18:23shards uh we load the um un6 numpy file
- 3:18:28we convert it to a torch. long tensor
- 3:18:30which is what a lot of the layers up top
- 3:18:32expect by default and then here we're
- 3:18:35just enumerating all the shards I also
- 3:18:38added a split to data load of light so
- 3:18:40we can uh load the split train but also
- 3:18:42the split Val uh the zero
- 3:18:44split and then we can load the shards
- 3:18:47and then here we also have not just the
- 3:18:49current position now but also the
- 3:18:51current Shard so we have a position
- 3:18:53inside A Shard and then when we uh run
- 3:18:55out of tokens in A Single Shard we first
- 3:18:58Advance The Shard and loop if we need to
- 3:19:01and then we get the tokens and readjust
- 3:19:03the position so this data loader will
- 3:19:06now iterate all the shards as well so I
- 3:19:09Chang that and then the other thing that
- 3:19:11I did while uh the data was processing
- 3:19:14is our train loader now has split train
- 3:19:17of course and down here I set up some I
- 3:19:20set up some numbers
- 3:19:21so we are doing 2 to the
- 3:19:249 uh tokens per uh per um per step and
- 3:19:31we want to do roughly 10 billion tokens
- 3:19:35um because that's how many unique tokens
- 3:19:36we have so if we did 10 billion tokens
- 3:19:39then divide that by 29 we see that this
- 3:19:41is 1973 steps so that's where that's
- 3:19:44from and then the GPT three paper says
- 3:19:47that they warm up the learning rate over
- 3:19:49375 million tokens so I came here and
- 3:19:53375 E6 tokens divide uh 2 to the
- 3:19:5719 is 715 steps so that's why warm-up
- 3:20:01steps is set to 715 so this will exactly
- 3:20:04match um the warm-up schedule that gpt3
- 3:20:07used and I think 715 by the way is very
- 3:20:10uh mild and this could be made
- 3:20:12significantly more aggressive probably
- 3:20:13even like 100 is good enough um
- 3:20:17but it's okay let's leave it for now so
- 3:20:18that we have the exact hyper parameters
- 3:20:20of gpt3 so I fix that and then um that's
- 3:20:25pretty much it we can we can run so we
- 3:20:28have our script
- 3:20:29here and we can
- 3:20:32launch and actually sorry let me do one
- 3:20:34more
- 3:20:38thing excuse
- 3:20:40me for my GPU I can actually fit more
- 3:20:43batch size and I believe I can fat I can
- 3:20:45fit 60 4 on my GPU as a micro bash size
- 3:20:50so let me try
- 3:20:54that I could be misremembering but that
- 3:20:57means 64 * 124 per GPU and then we have
- 3:21:00a gpus so that means we would not even
- 3:21:02be doing gradient accumulation if this
- 3:21:04fits because uh this just multi
- 3:21:06multiplies out to uh the full total bat
- 3:21:09size so no gradient
- 3:21:12accumulation and that would run pretty
- 3:21:14quickly if that fits
- 3:21:26let's go let's go I mean if this works
- 3:21:29then this is basically a serious
- 3:21:31pre-training run um we're not logging
- 3:21:34we're not evaluating the validation
- 3:21:35split we're not running any evaluations
- 3:21:37yet so it's not we haven't crossed our
- 3:21:39te's and dotted our eyes but uh if we
- 3:21:42let this run for a while we're going to
- 3:21:44actually get a pretty good model and the
- 3:21:46model that might even be on par with or
- 3:21:49better than gpt2 124 M okay so it looks
- 3:21:54like everything is going great we're
- 3:21:55processing 1.5 million tokens per
- 3:21:58second uh everything here looks good
- 3:22:03we're doing 330 milliseconds per
- 3:22:06iteration and we have to do a total
- 3:22:09of uh where are we printing that 1973 so
- 3:22:1319073 times 0.33
- 3:22:17is this many seconds this many minutes
- 3:22:20so this will run for 1.7
- 3:22:24hours uh so one and a half hour run uh
- 3:22:28like this and uh we don't even have to
- 3:22:30use gradient accumulation which is nice
- 3:22:31and you might not have that luxury in
- 3:22:33your GPU in that case just start
- 3:22:35decreasing the batch size until things
- 3:22:37fit but keep it to nice
- 3:22:39numbers um so that's pretty exciting
- 3:22:42we're currently warming up the learning
- 3:22:43rate so you see that it's still very low
- 3:22:45one4 so this will ramp up over the next
- 3:22:48few steps all the way to 6 e
- 3:22:50Nega uh 4
- 3:22:53here very cool so now what I'd like to
- 3:22:56do is uh let's cross the T and do our
- 3:22:58eyes let's evaluate on the validation
- 3:23:00split and let's try to figure out how we
- 3:23:02can run evals how we can do logging how
- 3:23:05we can visualize our losses and all the
- 3:23:07good stuff so let's get to that before
- 3:23:09we actually do the run okay so I've
- 3:23:11adjusted the code so that we're
- 3:23:13evaluating on the validation split so
- 3:23:15creating the Val loader just by passing
- 3:23:17in Split equals Val that will basically
- 3:23:19create a data loader just for the uh
- 3:23:21validation
- 3:23:22Shard um the other thing I did is in the
- 3:23:25data loader I introduced a new function
- 3:23:27reset which is called at init and it
- 3:23:29basically resets the data loader and
- 3:23:31that is very useful because when we come
- 3:23:34to the main training Loop now so this is
- 3:23:37the code that I've added and basically
- 3:23:39every 100th iteration including the
- 3:23:41zeroth iteration we put the model into
- 3:23:44evaluation mode we reset the Val loader
- 3:23:47and then um no gradients involved we're
- 3:23:50going to
- 3:23:52basically accumulate the gradients over
- 3:23:54say 20 steps and then average it all up
- 3:23:58and print out the validation loss and so
- 3:24:01that basically is the exact same logic
- 3:24:03as the training Loop roughly but there's
- 3:24:06no loss that backward it's only
- 3:24:07inference we're just measuring the loss
- 3:24:09we're adding it up everything else
- 3:24:11otherwise applies and is exactly as
- 3:24:13we've seen it before and so this will
- 3:24:15print the validation laws
- 3:24:16um every 100th iteration including on
- 3:24:19the very first
- 3:24:20iteration uh so that's nice that will
- 3:24:23tell us some amount some a little bit
- 3:24:25about how much we're overfitting that
- 3:24:27said like uh we have roughly Infinity
- 3:24:29data so we're mostly expecting our train
- 3:24:31and Val loss to be about the same but
- 3:24:33the other reason I'm kind of interested
- 3:24:35in this is because we can take the GPT
- 3:24:362124m as openi released it we can
- 3:24:39initialize from it and we can basically
- 3:24:41see what kind of loss it achieves on the
- 3:24:43validation loss as well and that gives
- 3:24:45us kind of an indication as to uh how
- 3:24:47much that model would generalize to 124
- 3:24:49M but it's not an sorry to fine web edu
- 3:24:52validation split that said it's not a
- 3:24:55super fair comparison to gpt2 because it
- 3:24:57was trained on a very different data
- 3:24:58distribution but it's still kind of like
- 3:25:00an interesting data point and in any
- 3:25:02case you would always want to have a
- 3:25:03validation split in a training run like
- 3:25:06this so that you can make sure that you
- 3:25:08are not um overfitting and this is
- 3:25:11especially a concern if we were to make
- 3:25:13more Epoch in our training data um so
- 3:25:16for example right now we're just doing a
- 3:25:18single Epoch but if we get to a point
- 3:25:20where we want to train on 10 epochs or
- 3:25:21something like that we would be really
- 3:25:23careful with maybe we are memorizing
- 3:25:26that data too much if we have a big
- 3:25:28enough model and our validation split
- 3:25:30would be one way to tell whether that is
- 3:25:32happening okay and in addition to that
- 3:25:34if you remember at bottom of our script
- 3:25:36we had all of this orphaned code for
- 3:25:37sampling from way back when so I deleted
- 3:25:40that code and I moved it up um to here
- 3:25:43so once in a while we simply value
- 3:25:45validation
- 3:25:46once in a while we sample we generate
- 3:25:49samples and then uh we do that only
- 3:25:52every 100 steps and we train on every
- 3:25:55single step so that's how I have a
- 3:25:56structure right now and I've been
- 3:25:58running this for 10,000 iterations so
- 3:26:00here are some samples on neration
- 3:26:021,000
- 3:26:05um hello I'm a language model and I'm
- 3:26:07not able to get more
- 3:26:09creative I'm a language model and
- 3:26:10languages file you're learning about
- 3:26:12here is or is the beginning of a
- 3:26:14computer
- 3:26:16okay so this is all like pretty uh this
- 3:26:19is still a garble uh but we're only at
- 3:26:21ration 1,000 and we've only just barely
- 3:26:24reached maximum learning rate uh so this
- 3:26:26is still learning uh we're about to get
- 3:26:28some more samples coming up in
- 3:26:321,00 okay
- 3:26:35um okay this is you know the model is
- 3:26:38still is still a young baby okay so uh
- 3:26:42basically all of this sampling code that
- 3:26:44I've put here everything should be
- 3:26:45familiar with to you and came from
- 3:26:47before the only thing that I did is I
- 3:26:49created a generator object in pytorch so
- 3:26:52that I have a direct control over the
- 3:26:54sampling of the random numbers don't
- 3:26:56because I don't want to impact the RNG
- 3:26:58state of the random number generator
- 3:27:00that is the global one used for training
- 3:27:02I want this to be completely outside of
- 3:27:04the training Loop and so I'm using a
- 3:27:07special sampling RNG and then I make
- 3:27:09sure to seed it that every single rank
- 3:27:12has a different seed and then I pass in
- 3:27:14here where we sort of consumer in the
- 3:27:17numbers in multinomial where the
- 3:27:18sampling happens I make sure to pass in
- 3:27:20the generator object there otherwise
- 3:27:22this is identical uh now the other thing
- 3:27:25is um you'll notice that we're running a
- 3:27:27bit slower that's because I actually had
- 3:27:29to disable torch. compile to get this to
- 3:27:32sample and um so we're running a bit
- 3:27:34slower so for some reason it works with
- 3:27:36no torch compile but when I torch
- 3:27:37compile my model I get a really scary
- 3:27:39error from pytorch and I have no idea
- 3:27:41how to resolve it right now so probably
- 3:27:43by the time you see this code released
- 3:27:45or something like that maybe it's fixed
- 3:27:47but for now I'm just going to do end
- 3:27:49false um and I'm going to bring back
- 3:27:51toor compile and you're not going to get
- 3:27:54samples and I I think I'll fix this
- 3:27:56later uh by the way um I will be
- 3:27:59releasing all this code and actually
- 3:28:01I've been very careful about making get
- 3:28:03commits every time we add something and
- 3:28:05so I'm going to release the entire repo
- 3:28:07that starts completely from scratch all
- 3:28:09the way to uh now and after this as well
- 3:28:12and so everything should be exactly
- 3:28:13documented in the git commit history um
- 3:28:16um and so I think that will be nice so
- 3:28:19hopefully by the time you go to GitHub
- 3:28:20uh this is removed and it's working and
- 3:28:22I will have fixed the bug okay so I have
- 3:28:24the optimization running here and it's
- 3:28:26stepping and we're on step 6,000 or so
- 3:28:28so we're about 30% through training now
- 3:28:31while this is training I would like to
- 3:28:32introduce one evaluation that we're
- 3:28:34going to use to supplement the
- 3:28:35validation set and that is the H swag
- 3:28:38eval so hos swag comes from this paper
- 3:28:42back in 2019 so it's a 5-year-old eval
- 3:28:44now and the way H swag works is there is
- 3:28:47basically a sentence completion data set
- 3:28:50so it's a multiple choice for every one
- 3:28:52of these questions we have uh basically
- 3:28:54a shared context like a woman is outside
- 3:28:57with a bucket and a dog the dog is
- 3:28:59running around trying to avoid bath she
- 3:29:02a Rises the bucket off with soap and
- 3:29:04blow dry the dog's head B uses a hose to
- 3:29:08keep it from getting soapy C gets the
- 3:29:11dog wet and it runs away again or D gets
- 3:29:14into a bathtub with the dog
- 3:29:16and so basically the idea is that these
- 3:29:19multiple choice are constructed so that
- 3:29:22one of them is a natural continuation of
- 3:29:25the um sentence and the others are
- 3:29:30not and uh the others might not make
- 3:29:32sense like uses the host to keep it from
- 3:29:34getting soaped that makes no sense and
- 3:29:36so what happens is that models that are
- 3:29:38not trained very well are not able to
- 3:29:40tell these apart but models that have a
- 3:29:43lot of World Knowledge and can tell uh
- 3:29:45which um and can tell a lot about the
- 3:29:48world will be able to create these
- 3:29:50completions and these sentences are
- 3:29:52sourced from activity net and from Wiki
- 3:29:55how and at the bottom of the uh
- 3:30:00paper there's kind of like a cool chart
- 3:30:03of the kinds of domains in Wiki house so
- 3:30:05there's a lot of sentences from
- 3:30:07computers and electronics and Homes and
- 3:30:09Garden and it has kind of a broad
- 3:30:11coverage of the kinds of things you need
- 3:30:13to know about the world in order to find
- 3:30:15the most likely completion and um the
- 3:30:19identity of that of that completion one
- 3:30:22more thing that's kind of interesting
- 3:30:23about H swag is the way it was
- 3:30:25constructed is that the incorrect um
- 3:30:28options are deliberately um
- 3:30:32adversarially sourced so they're not
- 3:30:34just random sentences they're actually
- 3:30:37sentences generated by language models
- 3:30:39and they're generated in such a way that
- 3:30:41language models basically find them
- 3:30:42difficult but humans find them easy and
- 3:30:45so they mentioned that humans have a 95%
- 3:30:47accuracy on this set but at the time the
- 3:30:49state-of-the-art language models had
- 3:30:51only 48% and so at the time this was a
- 3:30:54good Benchmark now you can read the
- 3:30:57details of this paper to to learn more
- 3:30:59um the thing to point out though is that
- 3:31:01this is 5 years ago and since then what
- 3:31:03happened to H swag is that it's been
- 3:31:05totally just uh
- 3:31:08um solved and so now the language models
- 3:31:11here are 96% so basically the 4% the
- 3:31:14last 4% is probably errors in the data
- 3:31:16set or the questions are really really
- 3:31:18hard and so basically this data set is
- 3:31:20kind of crushed with respect to language
- 3:31:22models but back then the best language
- 3:31:23model was only at about 50% uh but this
- 3:31:27is how far things got but still the the
- 3:31:30reason people like H swag and it's not
- 3:31:33used by the way in gpt2 but in gpt3
- 3:31:37there is H swag eval and lots of people
- 3:31:39use H
- 3:31:41swag and so for gpt3 we have results
- 3:31:45here
- 3:31:46that are cited so we know what percent
- 3:31:48accuracies gpt3 um attains at all these
- 3:31:51different model checkpoints for H swag
- 3:31:54eval and the reason people like it is
- 3:31:56because H swag is a smooth eval and it
- 3:31:59is an eval that offers quote unquote
- 3:32:01early signal uh so early signal means
- 3:32:04that even small language models are
- 3:32:06going to start at the random chance of
- 3:32:0825% but they're going to slowly improve
- 3:32:11and you're going to see 25 26 27 Etc and
- 3:32:15uh you can see slow Improvement even
- 3:32:17when the models are very small and it's
- 3:32:19very early so it's smooth it has early
- 3:32:23signal and um it's been around for a
- 3:32:26long time so that's why people kind of
- 3:32:28like this
- 3:32:29eval uh now the way that we're going to
- 3:32:32evaluate this is as
- 3:32:34follows as I mentioned we have a shared
- 3:32:37context and this is kind of like a
- 3:32:39multiple choice task but instead of
- 3:32:41giving the model a multiple choice
- 3:32:42question and asking it for A B C or D uh
- 3:32:46we can't do that because these models
- 3:32:47when they are so small as we are seeing
- 3:32:49here the models can't actually do
- 3:32:51multiple choice they don't understand
- 3:32:53the concept of associating a label to
- 3:32:55one of the options of multiple choice uh
- 3:32:58they don't understand that so we have to
- 3:32:59give it to them in a native form and the
- 3:33:01native form is a token completion so
- 3:33:05here's what we do we construct a batch
- 3:33:06of four rows and uh T tokens whatever
- 3:33:10that t happens to be then the shared
- 3:33:13context that is basically the context
- 3:33:15for the for choices the tokens of that
- 3:33:17are shared across all of the rows and
- 3:33:20then we have the four options so we kind
- 3:33:22of like lay them out and then only one
- 3:33:25of the options is correct in this case
- 3:33:26label three option three and so um this
- 3:33:30is the correct option and option one two
- 3:33:32and for are
- 3:33:33incorrect now these options might be of
- 3:33:36different lengths so what we do is we
- 3:33:38sort of like take the longest length and
- 3:33:40that's the size of the batch B BYT and
- 3:33:42then some of these uh here are going to
- 3:33:45be pded Dimensions so they're going to
- 3:33:47be unused and so we need the tokens we
- 3:33:51need the correct label and we need a
- 3:33:53mask that tells us which tokens are
- 3:33:55active and the mask is then zero for
- 3:33:58these uh padded areas so that's how we
- 3:34:01construct these batches and then in
- 3:34:04order to get the language model to
- 3:34:05predict A B C or D the way this works is
- 3:34:08basically we're just going to look at
- 3:34:10the tokens their probabilities and we're
- 3:34:12going to pick the option that gets the
- 3:34:15lowest or the highest average
- 3:34:18probability for the token so for the
- 3:34:22tokens because that is the most likely
- 3:34:25completion according to the language
- 3:34:27model so we're just going to look at the
- 3:34:29um probabilities here and average them
- 3:34:33up across the options and pick the one
- 3:34:35with the highest probability roughly
- 3:34:38speaking so this is how we're going to
- 3:34:40do H swag
- 3:34:42um and this is I believe also how uh
- 3:34:46gpt3 did it um this is how gpt3 did it
- 3:34:50as far as I know but you should note
- 3:34:52that some of the other evals where you
- 3:34:54might see H swag may not do it this way
- 3:34:57they may do it in a multiple choice
- 3:34:58format where you sort of uh give the the
- 3:35:00context a single time and then the four
- 3:35:02completions and so the model is able to
- 3:35:05see all the four options before it picks
- 3:35:07the best possible option and that's
- 3:35:08actually an easier task for a model
- 3:35:11because you get to see the other options
- 3:35:12when you're picking your choice um but
- 3:35:15unfortunately models at our size can't
- 3:35:17do that only models at a bigger size are
- 3:35:20able to do that and so our models are
- 3:35:22actually slightly handicapped in this
- 3:35:23way that they are not going to see the
- 3:35:25other options they're only going to see
- 3:35:27one option at a time and they just have
- 3:35:29to assign probabilities and the correct
- 3:35:31option has to win out in this metric all
- 3:35:34right so let's now implement this very
- 3:35:36briefly and incorporate it into our
- 3:35:38script okay so what I've done here is
- 3:35:40I've introduced a new file called hell
- 3:35:42swag. py that you can take a look into
- 3:35:45and I'm not going to to step through all
- 3:35:46of it because uh this is not exactly
- 3:35:48like deep code deep code it's kind of
- 3:35:51like a little bit tedious honestly
- 3:35:53because what's happening is I'm
- 3:35:54downloading hsac from GitHub and I'm
- 3:35:56rendering all of its examples and there
- 3:35:58are a total of 10,000 examples I am
- 3:36:00rendering them into this format um and
- 3:36:04so here at the end of this render
- 3:36:07example function you can see that I'm
- 3:36:09returning the
- 3:36:11tokens uh the tokens of this um 4xt
- 3:36:16uh array of Tokens The Mask which tells
- 3:36:19us which parts are the options and
- 3:36:21everything else is zero and the label
- 3:36:24that is the correct label and so that
- 3:36:26allows us to then iterate the examples
- 3:36:28and render them and I have an evaluate
- 3:36:30function here which can load a um gpt2
- 3:36:33from huging face and it runs the eval
- 3:36:36here um and it basically just calculates
- 3:36:40uh just as I described it predicts the
- 3:36:42option that has the lowest or the
- 3:36:45highest prob ility and the way to do
- 3:36:47that actually is we can basically
- 3:36:48evaluate the cross entropy loss so we're
- 3:36:51basically evaluating the loss of
- 3:36:53predicting the next token in a sequence
- 3:36:55and then we're looking at the row that
- 3:36:57has the lowest average loss and that's
- 3:37:01the uh option that we pick as the
- 3:37:04prediction and then we do some stats and
- 3:37:06prints and stuff like that so that is a
- 3:37:08way to evaluate L swag now if you go up
- 3:37:11here I'm showing that for GPT 2124m if
- 3:37:14you run this script you're going to see
- 3:37:16that H swag gets
- 3:37:1929.5% um so that's the performance we
- 3:37:22get here now remember that random Chan
- 3:37:23is 25% so we haven't gone too far and
- 3:37:27gpt2 XL which is the biggest the gpt2
- 3:37:31gets all the way up to 49% roughly so uh
- 3:37:34these are pretty low values considering
- 3:37:36that today's state-ofthe-art is more
- 3:37:37like 95% uh so these are definitely
- 3:37:40older models by now and then there's one
- 3:37:42more thing called Uther harness which is
- 3:37:44a very piece of infrastructure for
- 3:37:46running evals for language models and
- 3:37:48they get slightly different numbers and
- 3:37:50I'm not 100% sure what the discrepancy
- 3:37:52is for these um it could be that they
- 3:37:54actually do the multiple choice uh
- 3:37:57instead of just the completions and that
- 3:37:59could be the um uh the discrepancy but
- 3:38:02I'm not 100% sure about that i' have to
- 3:38:04take a look but for now our script
- 3:38:06reports 2955 and so that is the number
- 3:38:08that we'd like to beat if we are
- 3:38:10training a GPD 2124m from scratch and
- 3:38:13ourselves um
- 3:38:16so now I'm going to go into actually
- 3:38:19incorporating this eval into our main
- 3:38:22training script and um and basically
- 3:38:26because we want to evaluate it in a
- 3:38:28periodic manner so that we can track H
- 3:38:30swag and how it evolves over time and
- 3:38:32see when when and if we cross uh this
- 3:38:362955 um sort of region so let's now walk
- 3:38:41through some of the changes to train
- 3:38:42gpt2 thatp the first thing I did here is
- 3:38:45I actually made use compile optional
- 3:38:47kind of and I disabled it by default and
- 3:38:51the problem with that is the problem
- 3:38:53with compile is that unfortunately it
- 3:38:55does make our code faster but it
- 3:38:56actually breaks the evaluation code and
- 3:38:58the sampling code it gives me a very
- 3:39:00gnarly message and I don't know why so
- 3:39:02hopefully by the time you get to the
- 3:39:04codebase when I put it up on GitHub uh
- 3:39:06we're going to fix that by then but for
- 3:39:07now I'm running without torch compile
- 3:39:09which is why you see this be a bit
- 3:39:11slower so we're running without torch
- 3:39:13compile I also create cre a log
- 3:39:15directory log where we can place our
- 3:39:18log.txt which will record the train loss
- 3:39:22validation loss and the H swag
- 3:39:23accuracies so a very simple text file
- 3:39:25and we're going to uh open for writing
- 3:39:28so that it sort of starts empty and then
- 3:39:30we're going to append to
- 3:39:32it I created a simple variable that um
- 3:39:36helps tell us when we have a last step
- 3:39:39and then basically periodically inside
- 3:39:40this Loop every 250th iteration or at
- 3:39:44the last step we're going to evaluate
- 3:39:46the validation loss and then every 250th
- 3:39:50iteration um we are going to evaluate H
- 3:39:53swag but only if we are not using
- 3:39:56compile because compile breaks it so I'm
- 3:39:59going to come back to this code for
- 3:40:01evaluating H swag in a second and then
- 3:40:04every 250th iteration as well we're also
- 3:40:06going to sample from the model and so
- 3:40:08you should recognize this as our ancient
- 3:40:10code from way back when we started the
- 3:40:12video and we're just sampling from the
- 3:40:13model
- 3:40:15and then finally here um these are if
- 3:40:18we're not after we validate sample and
- 3:40:21evaluate hell swag we actually do a
- 3:40:23training step here and so this is one
- 3:40:26step of uh training and you should be
- 3:40:28pretty familiar with all of what this
- 3:40:30does and at the end here once we get our
- 3:40:32training laws we write it to the file so
- 3:40:35the only thing that changed that I
- 3:40:37really added is this entire section for
- 3:40:38H swag eval and the way this works is
- 3:40:41I'm trying to get all the gpus to
- 3:40:43collaborate on the H swag and so we're
- 3:40:45iterating all the examples and then each
- 3:40:48process only picks the examples that
- 3:40:52assigned to it so we sort of take I and
- 3:40:54moded by the world size and we have to
- 3:40:56make it equal to rank otherwise we
- 3:40:58continue and then we render an example
- 3:41:01put it on the GPU we get the low jits
- 3:41:04then I create a helper function that
- 3:41:05helps us basically predict the option
- 3:41:08with the lowest loss so this comes here
- 3:41:10the prediction and then if it's correct
- 3:41:12we sort of keep count and then if
- 3:41:15multiple processes were collaborating on
- 3:41:17all this then we need to synchronize
- 3:41:18their stats and so the way one way to do
- 3:41:21that is to package up our statistics
- 3:41:23here into tensors which we can then call
- 3:41:26this. alberon and
- 3:41:29sum and then here we sort of um unwrap
- 3:41:33them from tensors so that we just have
- 3:41:35ins and then here the master process
- 3:41:37will print and log the hellis swag
- 3:41:40accuracy
- 3:41:41so that's kind of the that's kind of it
- 3:41:45and that's what I'm running right here
- 3:41:47so you see this optimization here and uh
- 3:41:50we just had a generation and this is
- 3:41:52Step 10,000 out of about 20,000 right so
- 3:41:55we are halfway done and these are the
- 3:41:58kinds of samples that uh we are getting
- 3:41:59at this stage so let's take a look hello
- 3:42:02I'm a language model so I'd like to use
- 3:42:04it to generate some kinds of output
- 3:42:07hello I'm a language model and I'm a
- 3:42:08developer for a lot of
- 3:42:10companies Al language
- 3:42:12model uh let's see if I can find fun
- 3:42:17one
- 3:42:28um I don't know you can go through this
- 3:42:30yourself but certainly the predictions
- 3:42:32are getting less and less random uh it
- 3:42:34seems like the model is a little bit
- 3:42:35more self-aware and using language uh
- 3:42:38that is a bit
- 3:42:39more uh specific to it being language
- 3:42:43model hello I'm a language model and
- 3:42:45like how the language is used to
- 3:42:46communicate I'm a language model and I'm
- 3:42:48going to be speaking English and German
- 3:42:52okay I don't know so let's just wait
- 3:42:53until this optimization finishes and uh
- 3:42:56we'll see what kind of samples we get
- 3:42:57and we're also going to look at the
- 3:42:59train Val and the hway accuracy and see
- 3:43:03how we're doing with respect to
- 3:43:06gpt2 okay good morning so focusing For a
- 3:43:09Moment On The jupyter Notebook here on
- 3:43:11the right I created a new cell that
- 3:43:13basically allows us to visualize the the
- 3:43:15train Val and Hela and um the hel score
- 3:43:19and you can step through this it
- 3:43:21basically like parses the log file that
- 3:43:22we are writing and um a lot of this is
- 3:43:25just like boring ma plot lip code but
- 3:43:28basically this is what our optimization
- 3:43:30looks like
- 3:43:32so we ran for
- 3:43:3819,731 billion tokens which is whoops oh
- 3:43:41my gosh which is one Epoch of the sample
- 3:43:4410B of webd on the left we have the loss
- 3:43:48and the in blue we have the training
- 3:43:50loss in Orange we have the validation
- 3:43:52loss and red as a horizontal line we
- 3:43:55have the opening IG gpt2 124 M model
- 3:43:58checkpoint when it's just evaluated on
- 3:44:00the validation set of um of this fine
- 3:44:04web edu uh so you can see that we are
- 3:44:06surpassing this orange is below the red
- 3:44:09so we're surpassing the validation set
- 3:44:11of this data set and like I mentioned
- 3:44:13the data set distribution is very
- 3:44:15different from what gpt2 trained on so
- 3:44:16this is not an exactly fair comparison
- 3:44:19but it's a good cross check uh to uh to
- 3:44:22look at now we would ideally like
- 3:44:25something that is withheld and
- 3:44:27comparable and somewhat standard um and
- 3:44:30so for us that is helis swag and so on
- 3:44:33here we see the H swag progress we made
- 3:44:35from 25% all the way here in red we see
- 3:44:39the open gpt2 124 M model in red so it
- 3:44:44achieves this h bag here and the the
- 3:44:47gpt3 model 124 M which was trained on
- 3:44:50300 billion tokens achieves green so
- 3:44:54that's over here so you see that we
- 3:44:56basically surpassed the gbt2 24m uh
- 3:45:00model right here uh which is uh really
- 3:45:03nice now interestingly we were able to
- 3:45:07do so with only training on 10 billion
- 3:45:08tokens while gpt2 was trained on 100
- 3:45:11billion tokens so uh for some reason we
- 3:45:14were able to get away with significantly
- 3:45:16fewer tokens for training there are many
- 3:45:18possibilities to as to why we could
- 3:45:21match or surpass this accuracy um with
- 3:45:24only 10 million training so number one
- 3:45:27um it could be that opening gbt2 was
- 3:45:30trained on a much wider data
- 3:45:32distribution so in particular fine web
- 3:45:34edu is all English it's not multilingual
- 3:45:38and there's not that much math and code
- 3:45:40um and so math and code and multilingual
- 3:45:43could have been stealing capacity from
- 3:45:45the original gpt2 model and um basically
- 3:45:50that could be partially the reason why
- 3:45:52uh this is not working out there's many
- 3:45:54other reasons um so for example the H
- 3:45:57swag eval is fairly old uh maybe 5 years
- 3:45:59or so it is possible that aspects of H
- 3:46:02swag in some way or even identically
- 3:46:04have made it into the training Set uh of
- 3:46:07fine web we don't know for sure but if
- 3:46:10that was the case then we are basically
- 3:46:11looking at the training curve instead of
- 3:46:12the validation curve so long story short
- 3:46:15this is not a perfect eval and there's
- 3:46:16some caveats here uh but at least we
- 3:46:19have some confidence that that we're not
- 3:46:20doing something completely wrong and
- 3:46:23um and uh it's probably the case that
- 3:46:26when people try to create these data
- 3:46:27sets they try to make sure that test
- 3:46:29sets that are very common are not part
- 3:46:31of the training set for example uh when
- 3:46:33hugging face created the fine web BDU
- 3:46:35they use H swag as an eval so I would
- 3:46:37hope that they make sure that they D
- 3:46:39duplicate and that there's no hella swag
- 3:46:41in the training set but we can't be sure
- 3:46:45uh the other thing I wanted to address
- 3:46:46briefly is look at this loss curve this
- 3:46:48looks really this looks really wrong
- 3:46:50here I don't actually know 100% what
- 3:46:52this is and I suspect it's because the
- 3:46:55uh 10 billion sample of fine web edu was
- 3:46:58not properly shuffled um and there's
- 3:47:01some issue here uh with the data that I
- 3:47:04don't fully understand yet and there's
- 3:47:06some weird periodicity to it um and
- 3:47:08because we are in a very lazy way sort
- 3:47:10of serializing all the tokens and just
- 3:47:12iterating all them from scratch without
- 3:47:14doing any permutation or any random
- 3:47:16sampling ourselves I think we're
- 3:47:18inheriting some of the ordering that
- 3:47:21they have in the data set so uh this is
- 3:47:24not ideal but hopefully by the time you
- 3:47:26get to this repo uh some of these things
- 3:47:28by the way will hopefully be fixed and I
- 3:47:32will release this build n GPT repo and
- 3:47:35right now it looks a little ugly and
- 3:47:37preliminary uh so hopefully by the time
- 3:47:39you get here it's nicer but down here
- 3:47:41I'm going to show aada and I'm going to
- 3:47:44talk about about some of the things that
- 3:47:45happened after the video and I expect
- 3:47:48that we will have fixed uh the small
- 3:47:50issue uh but for now basically this
- 3:47:52shows that uh our training is not uh
- 3:47:55completely wrong and it shows that uh
- 3:47:58we're able to surpass the accuracy with
- 3:48:00only 10x the token budget um and
- 3:48:03possibly it could be also that the data
- 3:48:05set may have improved so uh the original
- 3:48:08uh gpt2 data set was web text it's
- 3:48:11possible that not a lot of care and
- 3:48:12attention went into the data set this
- 3:48:14was very early in llms whereas now
- 3:48:17there's a lot more scrutiny on good
- 3:48:18practices around uh D duplication
- 3:48:20filtering uh quality filtering and so on
- 3:48:23and it's possible that the data that
- 3:48:24we're training on is just of higher
- 3:48:25quality per token and that could be
- 3:48:27giving us a boost as well so a number of
- 3:48:30cave has to think about but for now uh
- 3:48:32we're pretty happy with this um and yeah
- 3:48:36now the next thing I was interested in
- 3:48:37is as you see it's a morning now so
- 3:48:39there was an overnight and I wanted to
- 3:48:41basically see how far I could push the
- 3:48:43result so uh to do an overnight run I
- 3:48:46basically did instead of one Epoch which
- 3:48:48took roughly two hours I just did a
- 3:48:50times four so that that would take eight
- 3:48:52hours while I was sleeping and so we did
- 3:48:54four Epoch or roughly 40 billion uh
- 3:48:56tokens of training and I was trying to
- 3:48:59see how far we could get um and so this
- 3:49:01was the only change and I reran the
- 3:49:03script and when I point uh and read the
- 3:49:05log file at uh at the 40b uh this is
- 3:49:08what the curve look
- 3:49:10like okay so to narrate this number one
- 3:49:13we are seeing this issue here here with
- 3:49:15the periodicity through the different
- 3:49:17Epoch and something really weird with
- 3:49:19the fine web edu data set and that is to
- 3:49:22be determined uh but otherwise we are
- 3:49:25seeing that the H swag actually went up
- 3:49:27by a lot and we almost we almost made it
- 3:49:31uh to the GPT 324m accuracy uh up here
- 3:49:35uh but not quite so uh it's too bad that
- 3:49:37I didn't sleep slightly longer um and uh
- 3:49:41I think if this was an uh five Epoch run
- 3:49:44we may have gotten here now one thing to
- 3:49:47point out is that if you're doing multi
- 3:49:49Epoch runs uh we're not actually being
- 3:49:51very careful in our data loader and
- 3:49:53we're not um I this data loader goes
- 3:49:56through the data in exactly the same
- 3:49:59format and exactly the same order and
- 3:50:01this is kind of suboptimal and you would
- 3:50:03want to look into extensions where you
- 3:50:05actually permute the data uh randomly
- 3:50:08you permute the documents around in
- 3:50:10Every Single Shard on every single new
- 3:50:12Epoch um and po even permute the
- 3:50:16shards and that would go a long way into
- 3:50:18decreasing the pricity and it's also
- 3:50:20better for the optimization so that
- 3:50:22you're not seeing things ident in the
- 3:50:23identical format and you're introducing
- 3:50:25some of the some uh Randomness in how
- 3:50:28the documents follow each other because
- 3:50:29you have to remember that in every
- 3:50:31single row these documents follow each
- 3:50:33other and then there's the end of text
- 3:50:34token and then the next document so the
- 3:50:36documents are currently glued together
- 3:50:39in the exact same identical manner but
- 3:50:41we actually want to break break up the
- 3:50:43documents and shuffle them around
- 3:50:45because the order of the documents
- 3:50:46shouldn't matter and they shouldn't um
- 3:50:49basically we want to break up that
- 3:50:50dependence because it's a kind of a
- 3:50:51spous correlation and so our data lad is
- 3:50:54not currently doing that and that's one
- 3:50:56Improvement uh you could think of
- 3:50:58making um the other thing to point out
- 3:51:01is we're almost matching gpt3 accuracy
- 3:51:03with only 40 billion tokens gpt3 trained
- 3:51:06on 300 billion tokens so again we're
- 3:51:08seeing about a 10x um Improvement here
- 3:51:11with respect to learning efficiency uh
- 3:51:14the other thing I wanted to and I don't
- 3:51:16actually know exactly what to attribute
- 3:51:18this to other than some of the things
- 3:51:19that I already mentioned previously for
- 3:51:21the previous run uh the other thing I
- 3:51:23wanted to briefly mention is uh the max
- 3:51:26LR here I saw some people already play
- 3:51:29with this a little bit in a previous
- 3:51:31related repository um and it turns out
- 3:51:33that you can actually almost like three
- 3:51:35xas so it's possible that the maximum
- 3:51:37learning rate can be a lot higher and
- 3:51:39for some reason the gpt3 hyper
- 3:51:40parameters that we are inheriting are
- 3:51:42actually extremely conservative and you
- 3:51:44can actually get away with a Higher
- 3:51:45Learning rate and it would train faster
- 3:51:47so a lot of these hyper parameters um
- 3:51:50are quite tunable and feel free to play
- 3:51:52with them and they're probably not set
- 3:51:54precisely correctly and um it's possible
- 3:51:59that you can get away with doing this
- 3:52:01basically and if you wanted to exactly
- 3:52:03be faithful to gpt3 you would also want
- 3:52:07to make the following difference you'd
- 3:52:10want to come here and the sequence
- 3:52:11length of gpt3 is 2x it's 20 48 instead
- 3:52:15of 1,24 so you would come here change
- 3:52:17this to 248 for T and then if you want
- 3:52:20the exact same number of tokens uh half
- 3:52:22a million per iteration or per step you
- 3:52:25want to then decrease this to 32 so they
- 3:52:28still multiply to half a mil so that
- 3:52:31would give your model sequence length
- 3:52:33equal to that of gpt3 and in that case
- 3:52:36basically the
- 3:52:37um the models would be roughly identical
- 3:52:40as far as I'm as far as I'm aware
- 3:52:42because again gpt2 and gpt3 are very
- 3:52:44very similar models now we can also look
- 3:52:47at some of the samples here from the
- 3:52:48model that was trained overnight so this
- 3:52:51is
- 3:52:52the optimization and you see that here
- 3:52:55we stepped all the way to
- 3:52:5776290 also or so and these are the hos
- 3:53:02mag we achieved was 33.2 4 and these are
- 3:53:06some of the samples from the model and
- 3:53:08you can see that if you read through
- 3:53:10this and pause the video briefly you can
- 3:53:11see that they are a lot more coherent uh
- 3:53:14so
- 3:53:15um and they're actually addressing the
- 3:53:17fact that it's a language model almost
- 3:53:21so uh hello I'm a language model and I
- 3:53:24try to be as accurate as
- 3:53:27possible um I'm a language model not a
- 3:53:29programming
- 3:53:31language I know how to communicate uh I
- 3:53:34use
- 3:53:35Python
- 3:53:37um I don't know if you pause this and
- 3:53:40look at it and then compare it to the
- 3:53:41one to the model that was only trained
- 3:53:43for 10 billion uh you will see that
- 3:53:45these are a lot more coherent and you
- 3:53:47can play with this uh
- 3:53:48yourself one more thing I added to The
- 3:53:50Code by the way is this chunk of code
- 3:53:52here so basically right after we
- 3:53:54evaluate the validation loss if we are
- 3:53:56the master process in addition to
- 3:53:58logging the validation loss every 5,000
- 3:54:01steps we're also going to save the
- 3:54:02checkpoint which is really just the
- 3:54:04state dictionary of the model and so
- 3:54:07checkpointing is nice just because uh
- 3:54:09you can save the model and later you can
- 3:54:11uh use it in some way if you wanted to
- 3:54:13resume the optimiz ation then in
- 3:54:15addition to saving the model we have to
- 3:54:17also save the optimizer State dict
- 3:54:20because remember that the optimizer has
- 3:54:21a few additional buffers because of adom
- 3:54:24so it's got the m and V and uh you need
- 3:54:28to also resume the optimizer properly
- 3:54:30you have to be careful with your RNG
- 3:54:31seeds uh random number generators and so
- 3:54:33on so if you wanted to exactly be able
- 3:54:35to resume optimization you have to think
- 3:54:37through the state of the of the training
- 3:54:40process but if you just want to save the
- 3:54:41model this is how you would do it and
- 3:54:43one one nice reason why you might want
- 3:54:45to do this is because you may want to
- 3:54:47evaluate the model a lot more carefully
- 3:54:50so here we are only kind of like winging
- 3:54:52the hell swag eval but you may want to
- 3:54:54use something um nicer like for example
- 3:54:57the Luther uh Luther evaluation hardness
- 3:55:01evaluation hardness hardness um so this
- 3:55:06is a way to also evaluate language
- 3:55:08models and um so it's possible that um
- 3:55:13you may want to use basically different
- 3:55:15infrastructure to more thoroughly
- 3:55:17evaluate the models on different um
- 3:55:20evaluations and compare it to the
- 3:55:21opening gbt2 model on many other um
- 3:55:25tasks like for example that involve math
- 3:55:26code or different languages and so on so
- 3:55:29this is a nice functionality to have as
- 3:55:30well
- 3:55:32um and then the other thing I wanted to
- 3:55:34mention is that everything we've built
- 3:55:36here this is only the pre-training step
- 3:55:39so um the GPT here is a it dreams
- 3:55:42documents it just predicts the next to
- 3:55:44you can't talk to it like you can talk
- 3:55:46to chat GPT uh chat GPT if you wanted to
- 3:55:49talk to the model we have to fine-tune
- 3:55:51it into the chat format and it's not
- 3:55:54actually like that complicated if you're
- 3:55:55looking at supervised fine-tuning or sft
- 3:55:58really what that means is we're just
- 3:55:59swapping out a data set into a data set
- 3:56:01that is a lot more conversational and
- 3:56:03there's a user assistant user assistant
- 3:56:04kind of structure and we just fine-tune
- 3:56:06on it and then we um we basically fill
- 3:56:09in the user tokens and we sample the
- 3:56:11assistant tokens it's not a lot more
- 3:56:13deeper than that uh but basically we
- 3:56:15swap out the data set and continue
- 3:56:17training uh but for now we're going to
- 3:56:19stop at uh pre-training one more thing
- 3:56:21that I wanted to briefly show you is
- 3:56:23that of course what we've built up today
- 3:56:25was building towards nanog GPT which is
- 3:56:27this repository from earlier uh but also
- 3:56:30there's actually another nanog GPT
- 3:56:32implementation and it's hiding in a more
- 3:56:34recent project that I've been working on
- 3:56:36called llm Doc and lm. C is a pure Cuda
- 3:56:41implementation of gpt2 or gpt3 training
- 3:56:44and it just directly uses uh Cuda and is
- 3:56:47written as Cuda now the nanog gbt here
- 3:56:51acts as reference code in pytorch to the
- 3:56:53C implementation so we're trying to
- 3:56:55exactly match up the two but we're
- 3:56:57hoping that the C Cuda is faster and of
- 3:56:59course currently that seems to be the
- 3:57:01case um because it is a direct optimized
- 3:57:04implementation so train gpt2 Pi in LL
- 3:57:06M.C is basically the nanog GPT and when
- 3:57:10you scroll through this file you'll find
- 3:57:12a lot of things that very much look like
- 3:57:16um things that we've built up in this
- 3:57:19lecture and then when you look at train
- 3:57:21gpt2 docu uh this is the C Cuda
- 3:57:25implementation so there's a lot of MPI
- 3:57:27nickel GPU Cuda
- 3:57:30cc++ and you have to be familiar with
- 3:57:32that but uh um when this is built up we
- 3:57:37can actually run the two side by side
- 3:57:39and they're going to produce the exact
- 3:57:40same results but lm. C actually runs
- 3:57:43faster so let's see that so on the left
- 3:57:45I have pytorch a nanog GPT looking thing
- 3:57:49on the right I have the llmc call and
- 3:57:52here I'm going to launch the
- 3:57:54two both of these are going to be
- 3:57:55running on a single GPU and here I'm
- 3:57:57putting the lm. C on GPU 1 and this one
- 3:58:00will grab uh gpu0 by default and
- 3:58:05then we can see here that lm. c
- 3:58:08compiled and then allocate space and
- 3:58:11it's
- 3:58:12stepping so
- 3:58:15basically uh meanwhile P torch is still
- 3:58:17compiling because torch compile is a bit
- 3:58:19slower here than the lm. C nbcc Cuda
- 3:58:24compile and so this program has already
- 3:58:26started running and uh we're still
- 3:58:28waiting here for torch compile now of
- 3:58:30course uh this is a very specific
- 3:58:33implementation to gpt2 and 3 a pytorch
- 3:58:35is a very general neural network
- 3:58:37framework so they're not exactly
- 3:58:38comparable but if you're only interested
- 3:58:39in training gpt2 and 3 lm. C is very
- 3:58:43fast it takes less space it's faster to
- 3:58:46start and it's faster per
- 3:58:49step and so P started to Stepping here
- 3:58:53and as you can see we're running at
- 3:58:54about 223,000 tokens per second here and
- 3:58:57about 185,000 tokens per second here um
- 3:59:03so quite a bit slower but I don't have
- 3:59:05full confidence that I exactly squeezed
- 3:59:08out all the juice from the pytorch
- 3:59:09implementation but the important thing
- 3:59:11here is notice that if I Aline up the
- 3:59:14steps you will see that the losses and
- 3:59:16Norms that are printed between these two
- 3:59:18are
- 3:59:19identical so on the left we have the pie
- 3:59:21torch and on the right this C
- 3:59:24implementation and they're the same
- 3:59:25except this one runs faster uh so that's
- 3:59:28kind of I wanted to show you also
- 3:59:30briefly lm. C and this is a parallel
- 3:59:33implementation and it's also something
- 3:59:35that you may want to uh play with or
- 3:59:37look at and um it's kind of interesting
- 3:59:39okay so at this point I should probably
- 3:59:40start wrapping up the video because I
- 3:59:42think it's getting way longer than I
- 3:59:44anticipated uh but we did Cover a lot of
- 3:59:46ground and we built everything from
- 3:59:48scratch so as a brief summary we were
- 3:59:50looking at the gpt2 and GPT 3
- 3:59:55papers we were looking at how you set up
- 3:59:57these training runs uh and all the
- 3:59:59considerations involved we wrote
- 4:00:01everything from scratch and then we saw
- 4:00:03that over the duration of either a
- 4:00:042-hour training run or an overnight run
- 4:00:07we can actually match the 124 million
- 4:00:09parameter checkpoints of gbt2 and gpt3
- 4:00:12uh to a very large extent
- 4:00:14um in principle the code that we wrote
- 4:00:16would be able to train even bigger
- 4:00:18models if you have the patients or the
- 4:00:19Computing resources uh and so you could
- 4:00:21potentially think about training some of
- 4:00:23the bigger checkpoints as well um there
- 4:00:26are a few remaining issues to address
- 4:00:28what's happening with the loss here
- 4:00:30which I suspect has to do with the fine
- 4:00:31web edu data sampling uh why can't we
- 4:00:34turn on Torch compile uh it currently
- 4:00:36breaks generation and H swag what's up
- 4:00:39with that in the data loader we should
- 4:00:41probably be permuting our data when we
- 4:00:43reach boundaries so there's a few more
- 4:00:45issues like that and I expect to be
- 4:00:47documenting some of those over time in
- 4:00:49the uh build n GPT repository here which
- 4:00:53I'm going to be releasing with this
- 4:00:55video if you have any questions or like
- 4:00:57to talk about anything that we covered
- 4:00:59please go to discussions tab uh so we
- 4:01:02can talk here uh or please go to issues
- 4:01:04or pull request pull requests um
- 4:01:07depending on what you'd like to
- 4:01:08contribute or also have a look at the uh
- 4:01:11Zero to Hero Discord and uh I'm going to
- 4:01:14be hanging out here on N GPT
- 4:01:17um otherwise for now I'm pretty happy
- 4:01:20about where we got um and I hope you
- 4:01:23enjoyed the video and I will see you
- 4:01:25later
About this transcript
This page contains the full transcript of Let's reproduce GPT-2 (124M) by Andrej Karpathy, generated from the public captions YouTube serves with the video. The transcript has 43,388 words across 6,049 segments, with the original timestamps preserved so you can click any line to jump to that moment in the embedded player.
What you can do with it
Use the transcript to take notes, quote the speaker, build a study guide, generate a summary with ChatGPT or Claude via the YouTube Summary tool, or export it as a timed subtitle file with YouTube to SRT. You can also re-open it in the transcriber to translate the transcript into 100+ languages.
Free YouTube transcript tool
YouTube2Text is a free YouTube transcript generator — no signup, no daily limit. Paste any YouTube link and get the full transcript instantly, with timestamps, click-to-jump, translation to 100+ languages, AI prompts for ChatGPT, Claude, and Gemini, and exports to TXT, SRT, VTT, or Markdown.