Building makemore Part 4: Becoming a Backprop Ninja — Transcript
Full transcript
- 0:00hi everyone so today we are once again
- 0:02continuing our implementation of make
- 0:04more now so far we've come up to here
- 0:07montalia perceptrons and our neural net
- 0:09looked like this and we were
- 0:11implementing this over the last few
- 0:12lectures
- 0:13now I'm sure everyone is very excited to
- 0:15go into recurring neural networks and
- 0:16all of their variants and how they work
- 0:18and the diagrams look cool and it's very
- 0:20exciting and interesting and we're going
- 0:21to get a better result but unfortunately
- 0:23I think we have to remain here for one
- 0:25more lecture and the reason for that is
- 0:28we've already trained this multilio
- 0:30perceptron right and we are getting
- 0:31pretty good loss and I think we have a
- 0:33pretty decent understanding of the
- 0:34architecture and how it works but the
- 0:37line of code here that I take an issue
- 0:39with is here lost up backward that is we
- 0:42are taking a pytorch auto grad and using
- 0:45it to calculate all of our gradients
- 0:46along the way and I would like to remove
- 0:48the use of lost at backward and I would
- 0:50like us to write our backward pass
- 0:52manually on the level of tensors and I
- 0:55think that this is a very useful
- 0:56exercise for the following reasons
- 0:58I actually have an entire blog post on
- 1:00this topic but I'd like to call back
- 1:02propagation a leaky abstraction
- 1:05and what I mean by that is back
- 1:07propagation does doesn't just make your
- 1:09neural networks just work magically it's
- 1:11not the case they can just Stack Up
- 1:12arbitrary Lego blocks of differentiable
- 1:14functions and just cross your fingers
- 1:16and back propagate and everything is
- 1:17great things don't just work
- 1:19automatically it is a leaky abstraction
- 1:22in the sense that you can shoot yourself
- 1:23in the foot if you do not understanding
- 1:25its internals it will magically not work
- 1:28or not work optimally and you will need
- 1:31to understand how it works under the
- 1:32hood if you're hoping to debug it and if
- 1:34you are hoping to address it in your
- 1:36neural nut
- 1:37um so this blog post here from a while
- 1:39ago goes into some of those examples so
- 1:42for example we've already covered them
- 1:43some of them already for example the
- 1:46flat tails of these functions and how
- 1:48you do not want to saturate them too
- 1:51much because your gradients will die the
- 1:53case of dead neurons which I've already
- 1:55covered as well
- 1:56the case of exploding or Vanishing
- 1:58gradients in the case of repair neural
- 2:00networks which we are about to cover
- 2:02and then also you will often come across
- 2:05some examples in the wild
- 2:07this is a snippet that I found uh in a
- 2:10random code base on the internet where
- 2:11they actually have like a very subtle
- 2:13but pretty major bug in their
- 2:15implementation and the bug points at the
- 2:18fact that the author of this code does
- 2:20not actually understand by propagation
- 2:21so they're trying to do here is they're
- 2:23trying to clip the loss at a certain
- 2:25maximum value but actually what they're
- 2:27trying to do is they're trying to
- 2:28collect the gradients to have a maximum
- 2:30value instead of trying to clip the loss
- 2:32at a maximum value and
- 2:34um indirectly they're basically causing
- 2:36some of the outliers to be actually
- 2:38ignored because when you clip a loss of
- 2:41an outlier you are setting its gradient
- 2:43to zero and so have a look through this
- 2:46and read through it but there's
- 2:48basically a bunch of subtle issues that
- 2:50you're going to avoid if you actually
- 2:51know what you're doing and that's why I
- 2:53don't think it's the case that because
- 2:55pytorch or other Frameworks offer
- 2:56autograd it is okay for us to ignore how
- 2:59it works
- 3:00now we've actually already covered
- 3:02covered autograd and we wrote micrograd
- 3:04but micrograd was an autograd engine
- 3:07only on the level of individual scalars
- 3:09so the atoms were single individual
- 3:11numbers and uh you know I don't think
- 3:13it's enough and I'd like us to basically
- 3:14think about back propagation on level of
- 3:16tensors as well and so in a summary I
- 3:19think it's a good exercise I think it is
- 3:21very very valuable you're going to
- 3:23become better at debugging neural
- 3:25networks and making sure that you
- 3:27understand what you're doing it is going
- 3:28to make everything fully explicit so
- 3:30you're not going to be nervous about
- 3:31what is hidden away from you and
- 3:33basically in general we're going to
- 3:34emerge stronger and so let's get into it
- 3:37a bit of a fun historical note here is
- 3:40that today writing your backward pass by
- 3:42hand and manually is not recommended and
- 3:43no one does it except for the purposes
- 3:45of exercise but about 10 years ago in
- 3:48deep learning this was fairly standard
- 3:49and in fact pervasive so at the time
- 3:52everyone used to write their own
- 3:53backward pass by hand manually including
- 3:55myself and it's just what you would do
- 3:57so we used to ride backward pass by hand
- 3:59and now everyone just calls lost that
- 4:01backward uh we've lost something I want
- 4:04to give you a few examples of this so
- 4:07here's a 2006 paper from Jeff Hinton and
- 4:11Russell selectinov in science that was
- 4:13influential at the time and this was
- 4:15training some architectures called
- 4:17restricted bolstery machines and
- 4:19basically it's an auto encoder trained
- 4:22here and this is from roughly 2010 I had
- 4:26a library for training researchable
- 4:27machines and this was at the time
- 4:30written in Matlab so python was not used
- 4:32for deep learning pervasively it was all
- 4:34Matlab and Matlab was this a scientific
- 4:36Computing package that everyone would
- 4:39use so we would write Matlab which is
- 4:41barely a programming language as well
- 4:44but I've had a very convenient tensor
- 4:46class and was this a Computing
- 4:48environment and you would run here it
- 4:49would all run on a CPU of course but you
- 4:51would have very nice plots to go with it
- 4:53and a built-in debugger and it was
- 4:54pretty nice now the code in this package
- 4:57in 2010 that I wrote for fitting
- 5:00research multiple machines to a large
- 5:03extent is recognizable but I wanted to
- 5:05show you how you would well I'm creating
- 5:07the data in the XY batches I'm
- 5:09initializing the neural nut so it's got
- 5:11weights and biases just like we're used
- 5:13to and then this is the training Loop
- 5:15where we actually do the forward pass
- 5:17and then here at this time they didn't
- 5:19even necessarily use back propagation to
- 5:21train neural networks so this in
- 5:23particular implements contrastive
- 5:25Divergence which estimates a gradient
- 5:28and then here we take that gradient and
- 5:30use it for a parameter update along the
- 5:32lines that we're used to
- 5:34um yeah here
- 5:36but you can see that basically people
- 5:38are meddling with these gradients uh
- 5:39directly and inline and themselves uh it
- 5:41wasn't that common to use an auto grad
- 5:43engine here's one more example from a
- 5:45paper of mine from 2014
- 5:47um called the fragmented embeddings
- 5:49and here what I was doing is I was
- 5:51aligning images and text
- 5:53um and so it's kind of like a clip if
- 5:55you're familiar with it but instead of
- 5:56working on the level of entire images
- 5:58and entire sentences it was working on
- 6:00the level of individual objects and
- 6:01little pieces of sentences and I was
- 6:03embedding them and then calculating very
- 6:05much like a clip-like loss and I dig up
- 6:08the code from 2014 of how I implemented
- 6:10this and it was already in numpy and
- 6:13python
- 6:14and here I'm planting the cost function
- 6:16and it was standard to implement not
- 6:19just the cost but also the backward pass
- 6:20manually so here I'm calculating the
- 6:23image embeddings sentence embeddings the
- 6:26loss function I calculate this course
- 6:28this is the loss function and then once
- 6:31I have the loss function I do the
- 6:32backward pass right here so I backward
- 6:34through the loss function and through
- 6:36the neural nut and I append
- 6:38regularization so everything was done by
- 6:41hand manually and you were just right
- 6:42out the backward pass and then you would
- 6:44use a gradient Checker to make sure that
- 6:46your numerical estimate of the gradient
- 6:47agrees with the one you calculated
- 6:49during back propagation so this was very
- 6:51standard for a long time but today of
- 6:53course it is standard to use an auto
- 6:55grad engine
- 6:56um but it was definitely useful and I
- 6:58think people sort of understood how
- 6:59these neural networks work on a very
- 7:01intuitive level and so I think it's a
- 7:03good exercise again and this is where we
- 7:04want to be okay so just as a reminder
- 7:06from our previous lecture this is The
- 7:08jupyter Notebook that we implemented at
- 7:09the time and
- 7:11we're going to keep everything the same
- 7:13so we're still going to have a two layer
- 7:15multiplayer perceptron with a batch
- 7:16normalization layer so the forward pass
- 7:18will be basically identical to this
- 7:20lecture but here we're going to get rid
- 7:22of lost and backward and instead we're
- 7:23going to write the backward pass
- 7:24manually
- 7:26now here's the starter code for this
- 7:27lecture we are becoming a back prop
- 7:29ninja in this notebook
- 7:31and the first few cells here are
- 7:34identical to what we are used to so we
- 7:36are doing some imports loading the data
- 7:37set and processing the data set none of
- 7:40this changed
- 7:41now here I'm introducing a utility
- 7:43function that we're going to use later
- 7:44to compare the gradients so in
- 7:46particular we are going to have the
- 7:47gradients that we estimate manually
- 7:49ourselves and we're going to have
- 7:50gradients that Pi torch calculates and
- 7:53we're going to be checking for
- 7:54correctness assuming of course that
- 7:55pytorch is correct
- 7:58um then here we have the initialization
- 8:00that we are quite used to so we have our
- 8:03embedding table for the characters the
- 8:05first layer second layer and the batch
- 8:06normalization in between
- 8:08and here's where we create all the
- 8:09parameters now you will note that I
- 8:11changed the initialization a little bit
- 8:13uh to be small numbers so normally you
- 8:16would set the biases to be all zero here
- 8:18I am setting them to be small random
- 8:20numbers and I'm doing this because
- 8:22if your variables are initialized to
- 8:24exactly zero sometimes what can happen
- 8:26is that can mask an incorrect
- 8:28implementation of a gradient
- 8:30um because uh when everything is zero it
- 8:32sort of like simplifies and gives you a
- 8:34much simpler expression of the gradient
- 8:35than you would otherwise get and so by
- 8:37making it small numbers I'm trying to
- 8:39unmask those potential errors in these
- 8:41calculations
- 8:43you also notice that I'm using uh B1 in
- 8:46the first layer I'm using a bias despite
- 8:48batch normalization right afterwards
- 8:50um so this would typically not be what
- 8:52you do because we talked about the fact
- 8:54that you don't need the bias but I'm
- 8:55doing this here just for fun
- 8:57um because we're going to have a
- 8:58gradient with respect to it and we can
- 9:00check that we are still calculating it
- 9:01correctly even though this bias is
- 9:03asparious
- 9:05so here I'm calculating a single batch
- 9:07and then here I'm doing a forward pass
- 9:10now you'll notice that the forward pass
- 9:11is significantly expanded from what we
- 9:13are used to here the forward pass was
- 9:15just
- 9:16um here
- 9:17now the reason that the forward pass is
- 9:19longer is for two reasons number one
- 9:22here we just had an F dot cross entropy
- 9:24but here I am bringing back a explicit
- 9:26implementation of the loss function
- 9:28and number two
- 9:29I've broken up the implementation into
- 9:32manageable chunks so we have a lot a lot
- 9:35more intermediate tensors along the way
- 9:37in the forward pass and that's because
- 9:38we are about to go backwards and
- 9:40calculate the gradients in this back
- 9:42propagation from the bottom to the top
- 9:45so we're going to go upwards and just
- 9:48like we have for example the lock props
- 9:49tensor in a forward pass in the backward
- 9:51pass we're going to have a d-lock probes
- 9:53which is going to store the derivative
- 9:55of the loss with respect to the lock
- 9:56props tensor and so we're going to be
- 9:58prepending D to every one of these
- 10:00tensors and calculating it along the way
- 10:02of this back propagation
- 10:04so as an example we have a b and raw
- 10:07here we're going to be calculating a DB
- 10:09in raw so here I'm telling pytorch that
- 10:12we want to retain the grad of all these
- 10:14intermediate values because here in
- 10:16exercise one we're going to calculate
- 10:18the backward pass so we're going to
- 10:20calculate all these D values D variables
- 10:22and use the CNP function I've introduced
- 10:25above to check our correctness with
- 10:26respect to what pi torch is telling us
- 10:29this is going to be exercise one uh
- 10:31where we sort of back propagate through
- 10:32this entire graph
- 10:34now just to give you a very quick
- 10:36preview of what's going to happen in
- 10:37exercise two and below here we have
- 10:40fully broken up the loss and back
- 10:43propagated through it manually in all
- 10:45the little Atomic pieces that make it up
- 10:47but here we're going to collapse the
- 10:49laws into a single cross-entropy call
- 10:50and instead we're going to analytically
- 10:53derive using math and paper and pencil
- 10:56the gradient of the loss with respect to
- 10:59the logits and instead of back
- 11:01propagating through all of its little
- 11:02chunks one at a time we're just going to
- 11:04analytically derive what that gradient
- 11:05is and we're going to implement that
- 11:07which is much more efficient as we'll
- 11:09see in the in a bit
- 11:10then we're going to do the exact same
- 11:12thing for patch normalization so instead
- 11:14of breaking up bass drum into all the
- 11:16old tiny components we're going to use
- 11:18uh pen and paper and Mathematics and
- 11:20calculus to derive the gradient through
- 11:22the bachelor Bachelor layer so we're
- 11:25going to calculate the backward
- 11:27passthrough bathroom layer in a much
- 11:28more efficient expression instead of
- 11:30backward propagating through all of its
- 11:31little pieces independently
- 11:33so there's going to be exercise three
- 11:36and then in exercise four we're going to
- 11:38put it all together and this is the full
- 11:40code of training this two layer MLP and
- 11:42we're going to basically insert our
- 11:44manual back prop and we're going to take
- 11:46out lost it backward and you will
- 11:48basically see that you can get all the
- 11:50same results using fully your own code
- 11:53and the only thing we're using from
- 11:55pytorch is the torch.tensor to make the
- 11:59calculations efficient but otherwise you
- 12:01will understand fully what it means to
- 12:03forward and backward and neural net and
- 12:04train it and I think that'll be awesome
- 12:06so let's get to it
- 12:08okay so I read all the cells of this
- 12:10notebook all the way up to here and I'm
- 12:13going to erase this and I'm going to
- 12:14start implementing backward pass
- 12:15starting with d lock problems so we want
- 12:18to understand what should go here to
- 12:20calculate the gradient of the loss with
- 12:22respect to all the elements of the log
- 12:23props tensor
- 12:25now I'm going to give away the answer
- 12:26here but I wanted to put a quick note
- 12:28here that I think would be most
- 12:30pedagogically useful for you is to
- 12:32actually go into the description of this
- 12:34video and find the link to this Jupiter
- 12:36notebook you can find it both on GitHub
- 12:38but you can also find Google collab with
- 12:40it so you don't have to install anything
- 12:41you'll just go to a website on Google
- 12:43collab and you can try to implement
- 12:45these derivatives or gradients yourself
- 12:47and then if you are not able to come to
- 12:50my video and see me do it and so work in
- 12:53Tandem and try it first yourself and
- 12:55then see me give away the answer and I
- 12:57think that'll be most valuable to you
- 12:59and that's how I recommend you go
- 13:00through this lecture
- 13:01so we are starting here with d-log props
- 13:03now d-lock props will hold the
- 13:06derivative of the loss with respect to
- 13:08all the elements of log props
- 13:11what is inside log blobs the shape of
- 13:13this is 32 by 27. so it's not going to
- 13:18surprise you that D log props should
- 13:19also be an array of size 32 by 27
- 13:21because we want the derivative loss with
- 13:23respect to all of its elements so the
- 13:26sizes of those are always going to be
- 13:27equal
- 13:29now how how does log props influence the
- 13:33loss okay loss is negative block probes
- 13:36indexed with range of N and YB and then
- 13:40the mean of that now just as a reminder
- 13:42YB is just a basically an array of all
- 13:47the correct indices
- 13:51um so what we're doing here is we're
- 13:52taking the lock props array of size 32
- 13:54by 27.
- 13:57right
- 13:58and then we are going in every single
- 14:00row and in each row we are plugging
- 14:03plucking out the index eight and then 14
- 14:06and 15 and so on so we're going down the
- 14:07rows that's the iterator range of N and
- 14:10then we are always plucking out the
- 14:12index of the column specified by this
- 14:15tensor YB so in the zeroth row we are
- 14:17taking the eighth column in the first
- 14:20row we're taking the 14th column Etc and
- 14:23so log props at this plugs out
- 14:26all those
- 14:28log probabilities of the correct next
- 14:30character in a sequence
- 14:32so that's what that does and the shape
- 14:34of this or the size of it is of course
- 14:3632 because our batch size is 32.
- 14:40so these elements get plugged out and
- 14:43then their mean and the negative of that
- 14:45becomes loss
- 14:47so I always like to work with simpler
- 14:49examples to understand the numerical
- 14:52form of derivative what's going on here
- 14:55is once we've plucked out these examples
- 14:58um we're taking the mean and then the
- 15:00negative so the loss basically
- 15:02I can write it this way is the negative
- 15:04of say a plus b plus c
- 15:07and the mean of those three numbers
- 15:09would be say negative would divide three
- 15:11that would be how we achieve the mean of
- 15:13three numbers ABC although we actually
- 15:15have 32 numbers here
- 15:16and so what is basically the loss by say
- 15:20like d a right
- 15:22well if we simplify this expression
- 15:24mathematically this is negative one over
- 15:26three of A and negative plus negative
- 15:28one over three of B
- 15:30plus negative 1 over 3 of c and so what
- 15:33is D loss by D A it's just negative one
- 15:35over three
- 15:36and so you can see that if we don't just
- 15:38have a b and c but we have 32 numbers
- 15:40then D loss by D
- 15:43um you know every one of those numbers
- 15:45is going to be one over N More generally
- 15:47because n is the um the size of the
- 15:50batch 32 in this case
- 15:53so D loss by
- 15:55um D Lock probs is negative 1 over n
- 15:59in all these places
- 16:02now what about the other elements inside
- 16:04lock problems because lock props is
- 16:05large array you see that lock problems
- 16:07at shape is 32 by 27. but only 32 of
- 16:11them participate in the loss calculation
- 16:13so what's the derivative of all the
- 16:15other most of the elements that do not
- 16:18get plucked out here
- 16:20while their loss intuitively is zero
- 16:22sorry they're gradient intuitively is
- 16:24zero and that's because they did not
- 16:25participate in the loss
- 16:27so most of these numbers inside this
- 16:29tensor does not feed into the loss and
- 16:32so if we were to change these numbers
- 16:33then the loss doesn't change which is
- 16:36the equivalent of way of saying that the
- 16:38derivative of the loss with respect to
- 16:39them is zero they don't impact it
- 16:43so here's a way to implement this
- 16:45derivative then we start out with
- 16:47torch.zeros of shape 32 by 27 or let's
- 16:50just say instead of doing this because
- 16:52we don't want to hard code numbers let's
- 16:54do torch.zeros like
- 16:57block probs so basically this is going
- 16:59to create an array of zeros exactly in
- 17:00the shape of log probs
- 17:02and then we need to set the derivative
- 17:05of negative 1 over n inside exactly
- 17:07these locations so here's what we can do
- 17:09the lock props indexed in The Identical
- 17:12way
- 17:14will be just set to negative one over
- 17:16zero divide n
- 17:19right just like we derived here
- 17:22so now let me erase all this reasoning
- 17:25and then this is the candidate
- 17:27derivative for D log props let's
- 17:29uncomment the first line and check that
- 17:31this is correct
- 17:34okay so CMP ran and let's go back to CMP
- 17:39and you see that what it's doing is it's
- 17:41calculating if
- 17:42the calculated value by us which is DT
- 17:46is exactly equal to T dot grad as
- 17:48calculated by pi torch and then this is
- 17:51making sure that all the elements are
- 17:52exactly equal and then converting this
- 17:54to a single Boolean value because we
- 17:57don't want the Boolean tensor we just
- 17:58want to Boolean value
- 18:00and then here we are making sure that
- 18:02okay if they're not exactly equal maybe
- 18:04they are approximately equal because of
- 18:06some floating Point issues but they're
- 18:07very very close
- 18:09so here we are using torch.allclose
- 18:10which has a little bit of a wiggle
- 18:13available because sometimes you can get
- 18:15very very close but if you use a
- 18:17slightly different calculation because a
- 18:19floating Point arithmetic you can get a
- 18:22slightly different result so this is
- 18:24checking if you get an approximately
- 18:25close result
- 18:27and then here we are checking the
- 18:28maximum uh basically the value that has
- 18:31the highest difference and what is the
- 18:34difference in the absolute value
- 18:35difference between those two and so we
- 18:37are printing whether we have an exact
- 18:39equality an approximate equality and
- 18:42what is the largest difference
- 18:45and so here
- 18:46we see that we actually have exact
- 18:48equality and so therefore of course we
- 18:50also have an approximate equality and
- 18:52the maximum difference is exactly zero
- 18:54so basically our d-log props is exactly
- 18:57equal to what pytors calculated to be
- 19:00lockprops.grad in its back propagation
- 19:03so so far we're working pretty well okay
- 19:06so let's now continue our back
- 19:07propagation
- 19:08we have that lock props depends on
- 19:10probes through a log
- 19:12so all the elements of probes are being
- 19:14element wise applied log to
- 19:17now if we want deep props then then
- 19:19remember your micrograph training
- 19:22we have like a log node it takes in
- 19:24probs and creates log probs and the
- 19:27props will be the local derivative of
- 19:30that individual Operation Log times the
- 19:33derivative loss with respect to its
- 19:34output which in this case is D log props
- 19:37so what is the local derivative of this
- 19:39operation well we are taking log element
- 19:41wise and we can come here and we can see
- 19:43well from alpha is your friend that d by
- 19:45DX of log of x is just simply one of our
- 19:47X
- 19:48so therefore in this case X is problems
- 19:51so we have d by DX is one over X which
- 19:54is one of our probes and then this is
- 19:56the local derivative and then times we
- 19:58want to chain it
- 20:00so this is chain rule
- 20:01times do log props
- 20:03let me uncomment this and let me run the
- 20:06cell in place and we see that the
- 20:08derivative of props as we calculated
- 20:10here is exactly correct
- 20:12and so notice here how this works probes
- 20:15that are props is going to be inverted
- 20:18and then element was multiplied here
- 20:20so if your probes is very very close to
- 20:23one that means you are your network is
- 20:25currently predicting the character
- 20:26correctly then this will become one over
- 20:28one and D log probes just gets passed
- 20:30through
- 20:31but if your probabilities are
- 20:33incorrectly assigned so if the correct
- 20:35character here is getting a very low
- 20:37probability then 1.0 dividing by it will
- 20:41boost this
- 20:43and then multiply by the log props so
- 20:45basically what this line is doing
- 20:46intuitively is it's taking the examples
- 20:49that have a very low probability
- 20:50currently assigned and it's boosting
- 20:52their gradient uh you can you can look
- 20:55at it that way next up is Count some imp
- 20:59so we want the river of this now let me
- 21:02just pause here and kind of introduce
- 21:05What's Happening Here in general because
- 21:06I know it's a little bit confusing we
- 21:08have the locusts that come out of the
- 21:09neural nut here what I'm doing is I'm
- 21:11finding the maximum in each row and I'm
- 21:15subtracting it for the purposes of
- 21:16numerical stability and we talked about
- 21:18how if you do not do this you run
- 21:20numerical issues if some of the logits
- 21:22take on two large values because we end
- 21:24up exponentiating them
- 21:26so this is done just for safety
- 21:28numerically then here's the
- 21:30exponentiation of all the sort of like
- 21:32logits to create our accounts and then
- 21:35we want to take the some of these counts
- 21:38and normalize so that all of the probes
- 21:40sum to one
- 21:41now here instead of using one over count
- 21:43sum I use uh raised to the power of
- 21:46negative one mathematically they are
- 21:47identical I just found that there's
- 21:49something wrong with the pytorch
- 21:50implementation of the backward pass of
- 21:52division
- 21:53um and it gives like a real result but
- 21:55that doesn't happen for star star native
- 21:58one that's why I'm using this formula
- 21:59instead but basically all that's
- 22:01happening here is we got the logits
- 22:04we're going to exponentiate all of them
- 22:05and want to normalize the counts to
- 22:07create our probabilities it's just that
- 22:09it's happening across multiple lines
- 22:12so now
- 22:14here
- 22:17we want to First Take the derivative we
- 22:20want to back propagate into account
- 22:21sumiv and then into counts as well
- 22:24so what should be the count sum M now we
- 22:28actually have to be careful here because
- 22:29we have to scrutinize and be careful
- 22:32with the shapes so counts that shape and
- 22:35then count some inverse shape
- 22:39are different
- 22:40so in particular counts as 32 by 27 but
- 22:43this count sum m is 32 by 1. and so in
- 22:47this multiplication here we also have an
- 22:49implicit broadcasting that pytorch will
- 22:52do because it needs to take this column
- 22:53tensor of 32 numbers and replicate it
- 22:55horizontally 27 times to align these two
- 22:58tensors so it can do an element twice
- 23:00multiply
- 23:01so really what this looks like is the
- 23:03following using a toy example again
- 23:06what we really have here is just props
- 23:08is counts times conservative so it's a C
- 23:10equals a times B
- 23:11but a is 3 by 3 and b is just three by
- 23:15one a column tensor and so pytorch
- 23:17internally replicated this elements of B
- 23:19and it did that across all the columns
- 23:22so for example B1 which is the first
- 23:24element of B would be replicated here
- 23:26across all the columns in this
- 23:27multiplication
- 23:29and now we're trying to back propagate
- 23:31through this operation to count some m
- 23:34so when we're calculating this
- 23:35derivative
- 23:37it's important to realize that these two
- 23:39this looks like a single operation but
- 23:41actually is two operations applied
- 23:44sequentially the first operation that
- 23:46pytorch did is it took this column
- 23:48tensor and replicated it across all the
- 23:52um across all the columns basically 27
- 23:54times so that's the first operation it's
- 23:55a replication and then the second
- 23:57operation is the multiplication so let's
- 23:59first background through the
- 24:01multiplication
- 24:02if these two arrays are of the same size
- 24:05and we just have a and b of both of them
- 24:08three by three then how do we mult how
- 24:11do we back propagate through a
- 24:12multiplication so if we just have
- 24:14scalars and not tensors then if you have
- 24:16C equals a times B then what is uh the
- 24:19order of the of C with respect to B well
- 24:21it's just a and so that's the local
- 24:23derivative
- 24:24so here in our case undoing the
- 24:27multiplication and back propagating
- 24:29through just the multiplication itself
- 24:30which is element wise is going to be the
- 24:33local derivative which in this case is
- 24:36simply counts because counts is the a
- 24:40so this is the local derivative and then
- 24:42times because the chain rule D props
- 24:46so this here is the derivative or the
- 24:48gradient but with respect to replicated
- 24:50B
- 24:52but we don't have a replicated B we just
- 24:54have a single B column so how do we now
- 24:56back propagate through the replication
- 24:59and intuitively this B1 is the same
- 25:02variable and it's just reused multiple
- 25:04times
- 25:04and so you can look at it
- 25:07as being equivalent to a case we've
- 25:09encountered in micrograd
- 25:10and so here I'm just pulling out a
- 25:12random graph we used in micrograd we had
- 25:14an example where a single node
- 25:17has its output feeding into two branches
- 25:19of basically the graph until the last
- 25:22function and we're talking about how the
- 25:25correct thing to do in the backward pass
- 25:26is we need to sum all the gradients that
- 25:29arrive at any one node so across these
- 25:31different branches the gradients would
- 25:33sum
- 25:34so if a node is used multiple times the
- 25:37gradients for all of its uses sum during
- 25:39back propagation
- 25:41so here B1 is used multiple times in all
- 25:44these columns and therefore the right
- 25:45thing to do here is to sum
- 25:48horizontally across all the rows so I'm
- 25:51going to sum in
- 25:52Dimension one but we want to retain this
- 25:55Dimension so that the uh so that counts
- 25:58some end and its gradient are going to
- 26:00be exactly the same shape so we want to
- 26:02make sure that we keep them as true so
- 26:04we don't lose this dimension and this
- 26:07will make the count sum M be exactly
- 26:08shape 32 by 1.
- 26:11so revealing this comparison as well and
- 26:14running this we see that we get an exact
- 26:17match
- 26:18so this derivative is exactly correct
- 26:22and let me erase
- 26:24this now let's also back propagate into
- 26:26counts which is the other variable here
- 26:29to create probes so from props to count
- 26:32some INF we just did that let's go into
- 26:33counts as well
- 26:35so decounts will be
- 26:39the chances are a so DC by d a is just B
- 26:43so therefore it's count summative
- 26:47um and then times chain rule the props
- 26:51now councilman is three two by One D
- 26:54probs is 32 by 27.
- 26:57so
- 26:59um those will broadcast fine and will
- 27:02give us decounts there's no additional
- 27:04summation required here
- 27:06um there will be a broadcasting that
- 27:08happens in this multiply here because
- 27:11count some M needs to be replicated
- 27:12again to correctly multiply D props but
- 27:16that's going to give the correct result
- 27:18so as far as the single operation is
- 27:20concerned so we back probably go from
- 27:23props to counts but we can't actually
- 27:25check the derivative counts uh I have it
- 27:29much later on and the reason for that is
- 27:31because count sum in depends on counts
- 27:34and so there's a second Branch here that
- 27:36we have to finish because can't summon
- 27:38back propagates into account sum and
- 27:40count sum will buy properly into counts
- 27:42and so counts is a node that is being
- 27:44used twice it's used right here in two
- 27:46props and it goes through this other
- 27:48Branch through count summative
- 27:50so even though we've calculated the
- 27:52first contribution of it we still have
- 27:54to calculate the second contribution of
- 27:55it later
- 27:57okay so we're continuing with this
- 27:58Branch we have the derivative for count
- 28:00sum if now we want the derivative of
- 28:02count sum so D count sum equals what is
- 28:05the local derivative of this operation
- 28:07so this is basically an element wise one
- 28:09over counts sum
- 28:11so count sum raised to the power of
- 28:13negative one is the same as one over
- 28:15count sum if we go to all from alpha we
- 28:17see that x to the negative one D by D by
- 28:20D by DX of it is basically Negative X to
- 28:23the negative 2. right one negative one
- 28:25over squared is the same as Negative X
- 28:27to the negative two
- 28:29so D count sum here will be local
- 28:32derivative is going to be negative
- 28:35um
- 28:36counts sum to the negative two that's
- 28:39the local derivative times chain rule
- 28:41which is D count sum in
- 28:46so that's D count sum
- 28:49let's uncomment this and check that I am
- 28:51correct okay so we have perfect equality
- 28:55and there's no sketchiness going on here
- 28:58with any shapes because these are of the
- 28:59same shape okay next up we want to back
- 29:02propagate through this line we have that
- 29:04count sum it's count.sum along the rows
- 29:07so I wrote out
- 29:09um some help here we have to keep in
- 29:11mind that counts of course is 32 by 27
- 29:13and count sum is 32 by 1. so in this
- 29:17back propagation we need to take this
- 29:19column of derivatives and transform it
- 29:22into a array of derivatives
- 29:24two-dimensional array
- 29:26so what is this operation doing we're
- 29:28taking in some kind of an input like say
- 29:31a three by three Matrix a and we are
- 29:32summing up the rows into a column tells
- 29:36her B1 b2b3 that is basically this
- 29:39so now we have the derivatives of the
- 29:41loss with respect to B all the elements
- 29:44of B
- 29:45and now we want to derivative loss with
- 29:47respect to all these little A's
- 29:49so how do the B's depend on the ace is
- 29:52basically what we're after what is the
- 29:54local derivative of this operation
- 29:56well we can see here that B1 only
- 29:58depends on these elements here the
- 30:01derivative of B1 with respect to all of
- 30:03these elements down here is zero but for
- 30:06these elements here like a11 a12 Etc the
- 30:09local derivative is one right so DB 1 by
- 30:13D A 1 1 for example is one so it's one
- 30:16one and one
- 30:18so when we have the derivative of loss
- 30:19with respect to B1
- 30:21did a local derivative of B1 with
- 30:23respect to these inputs is zeros here
- 30:25but it's one on these guys
- 30:27so in the chain rule
- 30:29we have the local derivative uh times
- 30:32sort of the derivative of B1 and so
- 30:35because the local derivative is one on
- 30:37these three elements the look of them
- 30:39are multiplying the derivative of B1
- 30:41will just be the derivative of B1 and so
- 30:45you can look at it as a router basically
- 30:47an addition is a router of gradient
- 30:50whatever gradient comes from above it
- 30:52just gets routed equally to all the
- 30:53elements that participate in that
- 30:55addition
- 30:56so in this case the derivative of B1
- 30:58will just flow equally to the derivative
- 31:00of a11 a12 and a13
- 31:03. so if we have a derivative of all the
- 31:05elements of B and in this column tensor
- 31:07which is D counts sum that we've
- 31:10calculated just now
- 31:11we basically see that what that amounts
- 31:14to is all of these are now flowing to
- 31:17all these elements of a and they're
- 31:19doing that horizontally
- 31:21so basically what we want is we want to
- 31:22take the decount sum of size 30 by 1 and
- 31:26we just want to replicate it 27 times
- 31:28horizontally to create 32 by 27 array
- 31:32so there's many ways to implement this
- 31:33operation you could of course just
- 31:35replicate the tensor but I think maybe
- 31:37one clean one is that the counts is
- 31:40simply torch dot once like
- 31:43so just an two-dimensional arrays of
- 31:45ones in the shape of counts so 32 by 27
- 31:49times D counts sum so this way we're
- 31:53letting the broadcasting here basically
- 31:56implement the replication you can look
- 31:58at it that way
- 31:59but then we have to also be careful
- 32:02because decounts was already calculated
- 32:05we calculated earlier here and that was
- 32:08just the first branch and we're now
- 32:09finishing the second Branch so we need
- 32:11to make sure that these gradients add so
- 32:13plus equals
- 32:14and then here
- 32:16um let's comment out the comparison and
- 32:20let's make sure crossing fingers that we
- 32:23have the correct result so pytorch
- 32:25agrees with us on this gradient as well
- 32:28okay hopefully we're getting a hang of
- 32:29this now counts as an element-wise X of
- 32:32Norm legits so now we want D Norm logits
- 32:36and because it's an element price
- 32:38operation everything is very simple what
- 32:40is the local derivative of e to the X
- 32:41it's famously just e to the x so this is
- 32:45the local derivative
- 32:48that is the local derivative now we
- 32:50already calculated it and it's inside
- 32:51counts so we may as well potentially
- 32:53just reuse counts that is the local
- 32:55derivative
- 32:56times uh D counts
- 33:01funny as that looks constant decount is
- 33:04derivative on the normal objects and now
- 33:07let's erase this and let's verify and it
- 33:10looks good
- 33:12so that's uh normal agents
- 33:14okay so we are here on this line now the
- 33:17normal objects
- 33:18we have that and we're trying to
- 33:20calculate the logits and deloget Maxes
- 33:22so back propagating through this line
- 33:25now we have to be careful here because
- 33:26the shapes again are not the same and so
- 33:29there's an implicit broadcasting
- 33:30Happening Here
- 33:32so normal jits has this shape 32 by 27
- 33:34logist does as well but logit Maxis is
- 33:37only 32 by one so there's a broadcasting
- 33:40here in the minus
- 33:42now here I try to sort of write out a
- 33:45two example again we basically have that
- 33:48this is our C equals a minus B
- 33:50and we see that because of the shape
- 33:52these are three by three but this one is
- 33:54just a column
- 33:55and so for example every element of C we
- 33:57have to look at how it uh came to be and
- 34:00every element of C is just the
- 34:01corresponding element of a minus uh
- 34:04basically that associated b
- 34:08so it's very clear now that the
- 34:10derivatives of every one of these c's
- 34:13with respect to their inputs are one for
- 34:16the corresponding a
- 34:18and it's a negative one for the
- 34:20corresponding B
- 34:22and so therefore
- 34:24um
- 34:25the derivatives on the C will flow
- 34:27equally to the corresponding Ace and
- 34:30then also to the corresponding base but
- 34:33then in addition to that the B's are
- 34:35broadcast so we'll have to do the
- 34:36additional sum just like we did before
- 34:39and of course the derivatives for B's
- 34:41will undergo a minus because the local
- 34:43derivative here is uh negative one
- 34:46so DC three two by D B3 is negative one
- 34:50so let's just Implement that basically
- 34:52delugits will be uh exactly copying the
- 34:56derivative on normal objects
- 34:58so
- 34:59delugits equals the norm logits and I'll
- 35:03do a DOT clone for safety so we're just
- 35:05making a copy
- 35:06and then we have that the loaded Maxis
- 35:09will be the negative of the non-legits
- 35:13because of the negative sign
- 35:15and then we have to be careful because
- 35:17logic Maxis is a column
- 35:20and so just like we saw before because
- 35:23we keep replicating the same elements
- 35:26across all the columns
- 35:28then in the backward pass because we
- 35:31keep reusing this these are all just
- 35:33like separate branches of use of that
- 35:35one variable and so therefore we have to
- 35:37do a Sum along one would keep them
- 35:39equals true so that we don't destroy
- 35:42this dimension
- 35:43and then the logic Maxes will be the
- 35:45same shape now we have to be careful
- 35:47because this deloaches is not the final
- 35:49deloaches and that's because not only do
- 35:52we get gradient signal into logits
- 35:54through here but the logic Maxes as a
- 35:56function of logits and that's a second
- 35:58Branch into logits so this is not yet
- 36:01our final derivative for logits we will
- 36:03come back later for the second branch
- 36:05for now the logic Maxis is the final
- 36:07derivative so let me uncomment this CMP
- 36:10here and let's just run this
- 36:12and logit Maxes hit by torch agrees with
- 36:15us
- 36:16so that was the derivative into through
- 36:19this line
- 36:21now before we move on I want to pause
- 36:22here briefly and I want to look at these
- 36:24logic Maxes and especially their
- 36:26gradients
- 36:27we've talked previously in the previous
- 36:28lecture that the only reason we're doing
- 36:31this is for the numerical stability of
- 36:33the softmax that we are implementing
- 36:34here and we talked about how if you take
- 36:37these logents for any one of these
- 36:39examples so one row of this logit's
- 36:41tensor if you add or subtract any value
- 36:44equally to all the elements then the
- 36:47value of the probes will be unchanged
- 36:49you're not changing soft Max the only
- 36:51thing that this is doing is it's making
- 36:53sure that X doesn't overflow and the
- 36:55reason we're using a Max is because then
- 36:57we are guaranteed that each row of
- 36:58logits the highest number is zero and so
- 37:01this will be safe
- 37:03and so
- 37:05um
- 37:06basically what that has repercussions
- 37:09if it is the case that changing logit
- 37:11Maxis does not change the props and
- 37:13therefore there's not change the loss
- 37:15then the gradient on logic masses should
- 37:17be zero right because saying those two
- 37:20things is the same
- 37:21so indeed we hope that this is very very
- 37:23small numbers so indeed we hope this is
- 37:25zero now because of floating Point uh
- 37:28sort of wonkiness
- 37:30um this doesn't come out exactly zero
- 37:31only in some of the rows it does but we
- 37:33get extremely small values like one e
- 37:35negative nine or ten and so this is
- 37:37telling us that the values of loaded
- 37:39Maxes are not impacting the loss as they
- 37:42shouldn't
- 37:43it feels kind of weird to back propagate
- 37:44through this branch honestly because
- 37:48if you have any implementation of like f
- 37:50dot cross entropy and pytorch and you
- 37:52you block together all these elements
- 37:54and you're not doing the back
- 37:54propagation piece by piece then you
- 37:57would probably assume that the
- 37:59derivative through here is exactly zero
- 38:01uh so you would be sort of
- 38:03um skipping this branch because it's
- 38:07only done for numerical stability but
- 38:09it's interesting to see that even if you
- 38:10break up everything into the full atoms
- 38:13and you still do the computation as
- 38:14you'd like with respect to numerical
- 38:16stability uh the correct thing happens
- 38:17and you still get a very very small
- 38:20gradients here
- 38:21um basically reflecting the fact that
- 38:23the values of these do not matter with
- 38:26respect to the final loss
- 38:27okay so let's now continue back
- 38:29propagation through this line here we've
- 38:31just calculated the logit Maxis and now
- 38:33we want to back prop into logits through
- 38:35this second branch
- 38:36now here of course we took legits and we
- 38:38took the max along all the rows and then
- 38:41we looked at its values here now the way
- 38:43this works is that in pytorch
- 38:47this thing here
- 38:49the max returns both the values and it
- 38:52Returns the indices at which those
- 38:53values to count the maximum value
- 38:55now in the forward pass we only used
- 38:57values because that's all we needed but
- 39:00in the backward pass it's extremely
- 39:01useful to know about where those maximum
- 39:04values occurred and we have the indices
- 39:06at which they occurred and this will of
- 39:08course helps us to help us do the back
- 39:10propagation because what should the
- 39:12backward pass be here in this case we
- 39:14have the largest tensor which is 32 by
- 39:1627 and in each row we find the maximum
- 39:18value and then that value gets plucked
- 39:20out into loaded Maxis and so intuitively
- 39:24um basically the derivative flowing
- 39:27through here then should be one
- 39:31times the look of derivatives is 1 for
- 39:34the appropriate entry that was plucked
- 39:35out
- 39:36and then times the global derivative of
- 39:39the logic axis
- 39:40so really what we're doing here if you
- 39:42think through it is we need to take the
- 39:44deloachet Maxis and we need to scatter
- 39:46it to the correct positions in these
- 39:50logits from where the maximum values
- 39:52came
- 39:53and so
- 39:54um
- 39:56I came up with one line of code sort of
- 39:58that does that let me just erase a bunch
- 39:59of stuff here so the line of uh you
- 40:02could do it kind of very similar to what
- 40:03we've done here where we create a zeros
- 40:05and then we populate uh the correct
- 40:07elements uh so we use the indices here
- 40:10and we would set them to be one but you
- 40:13can also use one hot
- 40:15so F dot one hot and then I'm taking the
- 40:18lowest of Max over the First Dimension
- 40:21dot indices and I'm telling uh pytorch
- 40:24that the dimension of every one of these
- 40:27tensors should be
- 40:29um
- 40:2927 and so what this is going to do
- 40:33is okay I apologize this is crazy filthy
- 40:37that I am sure of this
- 40:39it's really just a an array of where the
- 40:41Maxes came from in each row and that
- 40:44element is one and the all the other
- 40:45elements are zero so it's a one-half
- 40:47Vector in each row and these indices are
- 40:50now populating a single one in the
- 40:53proper place
- 40:54and then what I'm doing here is I'm
- 40:56multiplying by the logit Maxis and keep
- 40:58in mind that this is a column
- 41:01of 32 by 1. and so when I'm doing this
- 41:05times the logic Maxis the logic Maxes
- 41:08will broadcast and that column will you
- 41:10know get replicated and in an element
- 41:12wise multiply will ensure that each of
- 41:15these just gets routed to whichever one
- 41:17of these bits is turned on
- 41:19and so that's another way to implement
- 41:21uh this kind of a this kind of a
- 41:23operation and both of these can be used
- 41:26I just thought I would show an
- 41:28equivalent way to do it and I'm using
- 41:30plus equals because we already
- 41:31calculated the logits here and this is
- 41:33not the second branch
- 41:35so let's
- 41:37look at logits and make sure that this
- 41:39is correct
- 41:40and we see that we have exactly the
- 41:42correct answer
- 41:44next up we want to continue with logits
- 41:46here that is an outcome of a matrix
- 41:49multiplication and a bias offset in this
- 41:51linear layer
- 41:53so I've printed out the shapes of all
- 41:56these intermediate tensors we see that
- 41:58logits is of course 32 by 27 as we've
- 42:00just seen
- 42:01then the H here is 32 by 64. so these
- 42:05are 64 dimensional hidden States and
- 42:08then this W Matrix projects those 64
- 42:10dimensional vectors into 27 dimensions
- 42:12and then there's a 27 dimensional offset
- 42:15which is a one-dimensional vector
- 42:18now we should note that this plus here
- 42:20actually broadcasts because H multiplied
- 42:23by by W2 will give us a 32 by 27. and so
- 42:27then this plus B2 is a 27 dimensional
- 42:31lecture here
- 42:32now in the rules of broadcasting what's
- 42:33going to happen with this bias Vector is
- 42:35that this one-dimensional Vector of 27
- 42:37will get aligned with a padded dimension
- 42:41of one on the left and it will basically
- 42:43become a row vector and then it will get
- 42:45replicated vertically 32 times to make
- 42:48it 32 by 27 and then there's an
- 42:50element-wise multiply
- 42:52now
- 42:54the question is how do we back propagate
- 42:56from logits to the hidden States the
- 42:59weight Matrix W2 and the bias B2
- 43:02and you might think that we need to go
- 43:03to some Matrix calculus and then we have
- 43:07to look up the derivative for a matrix
- 43:09multiplication but actually you don't
- 43:11have to do any of that and you can go
- 43:12back to First principles and derive this
- 43:14yourself on a piece of paper and
- 43:17specifically what I like to do and I
- 43:18what I find works well for me is you
- 43:20find a specific small example that you
- 43:23then fully write out and then in the
- 43:25process of analyzing how that individual
- 43:27small example works you will understand
- 43:28the broader pattern and you'll be able
- 43:30to generalize and write out the full
- 43:32general formula for what how these
- 43:35derivatives flow in an expression like
- 43:37this so let's try that out
- 43:39so pardon the low budget production here
- 43:41but what I've done here is I'm writing
- 43:43it out on a piece of paper really what
- 43:45we are interested in is we have a
- 43:46multiply B plus C and that creates a d
- 43:50and we have the derivative of the loss
- 43:53with respect to D and we'd like to know
- 43:54what the derivative of the losses with
- 43:55respect to a b and c
- 43:57now these here are little
- 44:00two-dimensional examples of a matrix
- 44:01multiplication Two by Two Times a two by
- 44:03two
- 44:04plus a 2 a vector of just two elements
- 44:07C1 and C2 gives me a two by two
- 44:10now notice here that I have a bias
- 44:14Vector here called C and the bisex
- 44:17vector is C1 and C2 but as I described
- 44:19over here that bias Vector will become a
- 44:21row Vector in the broadcasting and will
- 44:23replicate vertically so that's what's
- 44:24happening here as well C1 C2 is
- 44:27replicated vertically and we see how we
- 44:29have two rows of C1 C2 as a result
- 44:33so now when I say write it out I just
- 44:35mean like this basically break up this
- 44:37matrix multiplication into the actual
- 44:40thing that that's going on under the
- 44:41hood so as a result of matrix
- 44:44multiplication and how it works d11 is
- 44:46the result of a DOT product between the
- 44:48first row of a and the First Column of B
- 44:51so a11 b11 plus a12 B21 plus C1
- 44:57and so on so forth for all the other
- 44:59elements of D and once you actually
- 45:02write it out it becomes obvious this is
- 45:03just a bunch of multipliers and
- 45:06um adds and we know from micrograd how
- 45:09to differentiate multiplies and adds and
- 45:11so this is not scary anymore it's not
- 45:13just matrix multiplication it's just uh
- 45:15tedious unfortunately but this is
- 45:17completely tractable we have DL by D for
- 45:20all of these and we want DL by uh all
- 45:23these little other variables so how do
- 45:25we achieve that and how do we actually
- 45:26get the gradients okay so the low budget
- 45:29production continues here
- 45:30so let's for example derive the
- 45:32derivative of the loss with respect to
- 45:34a11
- 45:36we see here that a11 occurs twice in our
- 45:38simple expression right here right here
- 45:40and influences d11 and D12
- 45:43. so this is so what is DL by d a one
- 45:46one well it's DL by d11 times the local
- 45:51derivative of d11 which in this case is
- 45:53just b11 because that's what's
- 45:55multiplying a11 here
- 45:57so uh and likewise here the local
- 46:00derivative of D12 with respect to a11 is
- 46:02just B12 and so B12 well in the chain
- 46:05rule therefore multiply the L by d 1 2.
- 46:08and then because a11 is used both to
- 46:11produce d11 and D12 we need to add up
- 46:15the contributions of both of those sort
- 46:18of chains that are running in parallel
- 46:20and that's why we get a plus just adding
- 46:22up those two
- 46:24um those two contributions and that
- 46:26gives us DL by d a one one we can do the
- 46:29exact same analysis for the other one
- 46:31for all the other elements of a and when
- 46:34you simply write it out it's just super
- 46:36simple
- 46:37um taking of gradients on you know
- 46:40expressions like this
- 46:42you find that
- 46:44this Matrix DL by D A that we're after
- 46:47right if we just arrange all the all of
- 46:49them in the same shape as a takes so a
- 46:52is just too much Matrix so d l by D A
- 46:55here will be also just the same shape
- 46:59tester with the derivatives now so deal
- 47:03by D a11 Etc
- 47:05and we see that actually we can express
- 47:06what we've written out here as a matrix
- 47:09multiplied
- 47:10and so it just so happens that D all by
- 47:13that all of these formulas that we've
- 47:15derived here by taking gradients can
- 47:17actually be expressed as a matrix
- 47:19multiplication and in particular we see
- 47:21that it is the matrix multiplication of
- 47:22these two array matrices
- 47:25so it is the um DL by D and then Matrix
- 47:30multiplying B but B transpose actually
- 47:32so you see that B21 and b12 have changed
- 47:37place
- 47:38whereas before we had of course b11 B12
- 47:41B2 on B22 so you see that this other
- 47:45Matrix B is transposed
- 47:47and so basically what we have long story
- 47:49short just by doing very simple
- 47:50reasoning here by breaking up the
- 47:52expression in the case of a very simple
- 47:54example is that DL by d a is which is
- 47:58this is simply equal to DL by DD Matrix
- 48:02multiplied with B transpose
- 48:05so that is what we have so far now we
- 48:08also want the derivative with respect to
- 48:10um B and C now
- 48:13for B I'm not actually doing the full
- 48:15derivation because honestly it's um it's
- 48:18not deep it's just uh annoying it's
- 48:20exhausting you can actually do this
- 48:22analysis yourself you'll also find that
- 48:24if you take this these expressions and
- 48:26you differentiate with respect to b
- 48:27instead of a you will find that DL by DB
- 48:30is also a matrix multiplication in this
- 48:33case you have to take the Matrix a and
- 48:35transpose it and Matrix multiply that
- 48:37with bl by DD
- 48:39and that's what gives you a deal by DB
- 48:42and then here for the offsets C1 and C2
- 48:46if you again just differentiate with
- 48:47respect to C1 you will find an
- 48:50expression like this
- 48:52and C2 an expression like this
- 48:55and basically you'll find the DL by DC
- 48:57is simply because they're just
- 48:59offsetting these Expressions you just
- 49:01have to take the deal by DD Matrix
- 49:04of the derivatives of D and you just
- 49:07have to sum across the columns and that
- 49:11gives you the derivatives for C
- 49:13so long story short
- 49:15the backward Paths of a matrix multiply
- 49:18is a matrix multiply
- 49:20and instead of just like we had D equals
- 49:22a times B plus C in the scalar case uh
- 49:25we sort of like arrive at something very
- 49:27very similar but now uh with a matrix
- 49:29multiplication instead of a scalar
- 49:31multiplication
- 49:32so the derivative of D with respect to a
- 49:36is
- 49:37DL by DD Matrix multiplied B trespose
- 49:41and here it's a transpose multiply deal
- 49:44by DD but in both cases it's a matrix
- 49:46multiplication with the derivative and
- 49:49the other term in the multiplication
- 49:53and for C it is a sum
- 49:55now I'll tell you a secret I can never
- 49:58remember the formulas that we just
- 50:00arrived for back proper gain information
- 50:01multiplication and I can back propagate
- 50:03through these Expressions just fine and
- 50:05the reason this works is because the
- 50:07dimensions have to work out
- 50:09uh so let me give you an example say I
- 50:11want to create DH
- 50:13then what should the H be number one I
- 50:16have to know that the shape of DH must
- 50:19be the same as the shape of H
- 50:21and the shape of H is 32 by 64. and then
- 50:24the other piece of information I know is
- 50:26that DH must be some kind of matrix
- 50:28multiplication of the logits with W2
- 50:32and delojits is 32 by 27 and W2 is a 64
- 50:37by 27. there is only a single way to
- 50:40make the shape work out in this case and
- 50:43it is indeed the correct result in
- 50:45particular here H needs to be 32 by 64.
- 50:48the only way to achieve that is to take
- 50:50a deluges
- 50:52and Matrix multiply it with you see how
- 50:55I have to take W2 but I have to
- 50:57transpose it to make the dimensions work
- 50:58out
- 50:59so w to transpose and it's the only way
- 51:02to make these to Matrix multiply those
- 51:04two pieces to make the shapes work out
- 51:06and that turns out to be the correct
- 51:08formula so if we come here we want DH
- 51:11which is d a and we see that d a is DL
- 51:15by DD Matrix multiply B transpose
- 51:18so that's Delo just multiply and B is W2
- 51:21so W2 transpose which is exactly what we
- 51:24have here so there's no need to remember
- 51:26these formulas similarly now if I want
- 51:30dw2 well I know that it must be a matrix
- 51:33multiplication of D logits and H
- 51:37and maybe there's a few transpose like
- 51:39there's one transpose in there as well
- 51:40and I don't know which way it is so I
- 51:42have to come to W2 and I see that its
- 51:44shape is 64 by 27
- 51:47and that has to come from some interest
- 51:49multiplication of these two
- 51:51and so to get a 64 by 27 I need to take
- 51:55um
- 51:56H I need to transpose it
- 51:59and then I need to Matrix multiply it
- 52:01um so that will become 64 by 32 and then
- 52:04I need to make sure to multiply with the
- 52:0532 by 27 and that's going to give me a
- 52:0764 by 27. so I need to make sure it's
- 52:09multiplied this with the logist that
- 52:11shape just like that that's the only way
- 52:13to make the dimensions work out and just
- 52:15use matrix multiplication and if we come
- 52:17here we see that that's exactly what's
- 52:19here so a transpose a for us is H
- 52:23multiplied with deloaches
- 52:25so that's W2 and then db2
- 52:30is just the um
- 52:33vertical sum and actually in the same
- 52:35way there's only one way to make the
- 52:37shapes work out I don't have to remember
- 52:38that it's a vertical Sum along the zero
- 52:40axis because that's the only way that
- 52:42this makes sense because B2 shape is 27
- 52:45so in order to get a um delugits
- 52:50here is 30 by 27 so knowing that it's
- 52:54just sum over deloaches in some
- 52:56Direction
- 52:59that direction must be zero because I
- 53:02need to eliminate this Dimension so it's
- 53:04this
- 53:06so this is so let's kind of like the
- 53:08hacky way let me copy paste and delete
- 53:10that and let me swing over here and this
- 53:13is our backward pass for the linear
- 53:14layer uh hopefully
- 53:17so now let's uncomment
- 53:19these three and we're checking that we
- 53:21got all the three derivatives correct
- 53:24and run
- 53:26and we see that h wh and B2 are all
- 53:30exactly correct so we back propagated
- 53:33through a linear layer
- 53:36now next up we have derivative for the h
- 53:39already and we need to back propagate
- 53:41through 10h into h preact
- 53:43so we want to derive DH preact
- 53:47and here we have to back propagate
- 53:48through a 10 H and we've already done
- 53:50this in micrograd and we remember that
- 53:5210h has a very simple backward formula
- 53:54now unfortunately if I just put in D by
- 53:56DX of 10 h of X into both from alpha it
- 53:59lets us down it tells us that it's a
- 54:00hyperbolic secant function squared of X
- 54:03it's not exactly helpful but luckily
- 54:06Google image search does not let us down
- 54:08and it gives us the simpler formula and
- 54:10in particular if you have that a is
- 54:12equal to 10 h of Z then d a by DZ by
- 54:16propagating through 10 H is just one
- 54:17minus a square and take note that 1
- 54:21minus a square a here is the output of
- 54:23the 10h not the input to the 10h Z so
- 54:27the D A by DZ is here formulated in
- 54:29terms of the output of that 10h
- 54:31and here also in Google image search we
- 54:34have the full derivation if you want to
- 54:35actually take the actual definition of
- 54:3810h and work through the math to figure
- 54:39out 1 minus standard square of Z
- 54:42so 1 minus a square is the local
- 54:45derivative in our case that is 1 minus
- 54:49uh the output of 10 H squared which here
- 54:52is H
- 54:53so it's h squared and that is the local
- 54:56derivative and then times the chain rule
- 54:58DH
- 55:00so that is going to be our candidate
- 55:02implementation so if we come here
- 55:05and then uncomment this let's hope for
- 55:08the best
- 55:09and we have the right answer
- 55:12okay next up we have DH preact and we
- 55:15want to back propagate into the gain the
- 55:17B and raw and the B and bias
- 55:19so here this is the bathroom parameters
- 55:21being gained in bias inside the bash
- 55:23term that take the B and raw that is
- 55:25exact unit caution and then scale it and
- 55:28shift it
- 55:29and these are the parameters of The
- 55:30Bachelor now here we have a
- 55:33multiplication but it's worth noting
- 55:35that this multiply is very very
- 55:36different from this Matrix multiply here
- 55:38Matrix multiply are DOT products between
- 55:41rows and Columns of these matrices
- 55:43involved this is an element twice
- 55:45multiply so things are quite a bit
- 55:46simpler
- 55:47now we do have to be careful with some
- 55:49of the broadcasting happening in this
- 55:51line of code though so you see how BN
- 55:53gain and B and bias are 1 by 64. but H
- 55:58preact and B and raw are 32 by 64.
- 56:02so we have to be careful with that and
- 56:04make sure that all the shapes work out
- 56:05fine and that the broadcasting is
- 56:06correctly back propagated
- 56:08so in particular let's start with the B
- 56:10and Gain so DB and gain should be
- 56:14and here this is again elementorized
- 56:17multiply and whenever we have a times b
- 56:19equals c we saw that the local
- 56:21derivative here is just if this is a the
- 56:23local derivative is just the B the other
- 56:25one so the local derivative is just B
- 56:27and raw and then times chain rule
- 56:31so DH preact
- 56:34so this is the candidate gradient now
- 56:38again we have to be careful because B
- 56:40and Gain Is of size 1 by 64. but this
- 56:44here would be 32 by 64.
- 56:48and so
- 56:49um the correct thing to do in this case
- 56:51of course is that b and gain here is a
- 56:53rule Vector of 64 numbers it gets
- 56:55replicated vertically in this operation
- 56:58and so therefore the correct thing to do
- 57:00is to sum because it's being replicated
- 57:03and therefore all the gradients in each
- 57:06of the rows that are now flowing
- 57:07backwards need to sum up to that same
- 57:10tensor DB and Gain so we have to sum
- 57:13across all the zero all the examples
- 57:16basically
- 57:17which is the direction in which this
- 57:19gets replicated
- 57:20and now we have to be also careful
- 57:21because we
- 57:23um being gain is of shape 1 by 64. so in
- 57:26fact I need to keep them as true
- 57:29otherwise I would just get 64.
- 57:31now I don't actually really remember why
- 57:34the being gain and the BN bias I made
- 57:36them be 1 by 64.
- 57:40um
- 57:41but the biases B1 and B2 I just made
- 57:44them be one-dimensional vectors they're
- 57:45not two-dimensional tensors so I can't
- 57:47recall exactly why I left the gain and
- 57:51the bias as two-dimensional but it
- 57:53doesn't really matter as long as you are
- 57:54consistent and you're keeping it the
- 57:55same
- 57:56so in this case we want to keep the
- 57:58dimension so that the tensor shapes work
- 58:01next up we have B and raw so DB and raw
- 58:05will be BN gain
- 58:09multiplying
- 58:11dhreact that's our chain rule now what
- 58:15about the
- 58:17um
- 58:18dimensions of this we have to be careful
- 58:20right so DH preact is 32 by 64. B and
- 58:24gain is 1 by 64. so it will just get
- 58:27replicated and to create this
- 58:29multiplication which is the correct
- 58:31thing because in a forward pass it also
- 58:33gets replicated in just the same way
- 58:35so in fact we don't need the brackets
- 58:37here we're done
- 58:38and the shapes are already correct
- 58:40and finally for the bias
- 58:43very similar this bias here is very very
- 58:46similar to the bias we saw when you
- 58:47layer in the linear layer and we see
- 58:49that the gradients from each preact will
- 58:51simply flow into the biases and add up
- 58:54because these are just these are just
- 58:55offsets
- 58:57and so basically we want this to be DH
- 58:59preact but it needs to Sum along the
- 59:01right Dimension and in this case similar
- 59:04to the gain we need to sum across the
- 59:06zeroth dimension the examples because of
- 59:09the way that the bias gets replicated
- 59:10vertically
- 59:11and we also want to have keep them as
- 59:14true
- 59:15and so this will basically take this and
- 59:17sum it up and give us a 1 by 64.
- 59:20so this is the candidate implementation
- 59:23it makes all the shapes work
- 59:25let me bring it up down here and then
- 59:28let me uncomment these three lines
- 59:32to check that we are getting the correct
- 59:33result for all the three tensors and
- 59:36indeed we see that all of that got back
- 59:38propagated correctly so now we get to
- 59:40the batch Norm layer we see how here
- 59:42being gay and being bias are the
- 59:44parameters so the back propagation ends
- 59:46but B and raw now is the output of the
- 59:50standardization
- 59:51so here what I'm doing of course is I'm
- 59:53breaking up the batch form into
- 59:54manageable pieces so we can back
- 59:55propagate through each line individually
- 59:57but basically what's happening is BN
- 1:00:00mean I is the sum
- 1:00:03so this is the B and mean I I apologize
- 1:00:06for the variable naming B and diff is x
- 1:00:10minus mu
- 1:00:11B and div 2 is x minus mu squared here
- 1:00:15inside the variance
- 1:00:16B and VAR is the variance so uh Sigma
- 1:00:20Square this is B and bar and it's
- 1:00:22basically the sum of squares
- 1:00:25so this is the x minus mu squared and
- 1:00:28then the sum now you'll notice one
- 1:00:30departure here
- 1:00:32here it is normalized as 1 over m
- 1:00:34uh which is number of examples here I'm
- 1:00:37normalizing as one over n minus 1
- 1:00:39instead of N and this is deliberate and
- 1:00:42I'll come back to that in a bit when we
- 1:00:43are at this line it is something called
- 1:00:45the bezels correction
- 1:00:47but this is how I want it in our case
- 1:00:51bienvar inv then becomes basically
- 1:00:53bienvar plus Epsilon Epsilon is one
- 1:00:56negative five and then it's one over
- 1:00:58square root
- 1:00:59is the same as raising to the power of
- 1:01:02negative 0.5 right because 0.5 is square
- 1:01:05root and then negative makes it one over
- 1:01:07square root
- 1:01:08so BM Bar M is a one over this uh
- 1:01:12denominator here and then we can see
- 1:01:14that b and raw which is the X hat here
- 1:01:16is equal to the BN diff the numerator
- 1:01:19multiplied by the
- 1:01:22um BN bar in
- 1:01:24and this line here that creates pre-h
- 1:01:27pre-act was the last piece we've already
- 1:01:29back propagated through it
- 1:01:31so now what we want to do is we are here
- 1:01:34and we have B and raw and we have to
- 1:01:35first back propagate into B and diff and
- 1:01:38B and Bar M
- 1:01:40so now we're here and we have DB and raw
- 1:01:43and we need to back propagate through
- 1:01:45this line
- 1:01:46now I've written out the shapes here and
- 1:01:49indeed bien VAR m is a shape 1 by 64. so
- 1:01:53there is a broadcasting happening here
- 1:01:55that we have to be careful with but it
- 1:01:57is just an element-wise simple
- 1:01:58multiplication by now we should be
- 1:02:00pretty comfortable with that to get DB
- 1:02:02and diff we know that this is just B and
- 1:02:05varm
- 1:02:06multiplied with
- 1:02:08DP and raw
- 1:02:11and conversely to get dbmring
- 1:02:15we need to take the end if
- 1:02:17and multiply that by DB and raw
- 1:02:22so this is the candidate but of course
- 1:02:24we need to make sure that broadcasting
- 1:02:26is obeyed so in particular B and VAR M
- 1:02:29multiplying with DB and raw
- 1:02:31will be okay and give us 32 by 64 as we
- 1:02:35expect
- 1:02:36but dbm VAR inv would be taking a 32 by
- 1:02:4064.
- 1:02:42multiplying it by 32 by 64. so this is a
- 1:02:4532 by 64. but of course DB this uh B and
- 1:02:49VAR in is only 1 by 64. so the second
- 1:02:52line here needs a sum across the
- 1:02:55examples and because there's this
- 1:02:57Dimension here we need to make sure that
- 1:03:00keep them is true
- 1:03:02so this is the candidate
- 1:03:04let's erase this and let's swing down
- 1:03:07here
- 1:03:09and implement it and then let's comment
- 1:03:11out dbm barif and DB and diff
- 1:03:16now we'll actually notice that DB and
- 1:03:18diff by the way is going to be incorrect
- 1:03:22so when I run this
- 1:03:24BMR m is correct B and diff is not
- 1:03:27correct and this is actually expected
- 1:03:30because we're not done with b and diff
- 1:03:34so in particular when we slide here we
- 1:03:36see here that b and raw as a function of
- 1:03:37B and diff but actually B and far of is
- 1:03:40a function of B of R which is a function
- 1:03:42of B and df2 which is a function of B
- 1:03:44and diff
- 1:03:45so it comes here so bdn diff
- 1:03:48um these variable names are crazy I'm
- 1:03:50sorry it branches out into two branches
- 1:03:53and we've only done one branch of it we
- 1:03:55have to continue our back propagation
- 1:03:57and eventually come back to B and diff
- 1:03:58and then we'll be able to do a plus
- 1:04:00equals and get the actual card gradient
- 1:04:02for now it is good to verify that CMP
- 1:04:05also works it doesn't just lie to us and
- 1:04:07tell us that everything is always
- 1:04:08correct it can in fact detect when your
- 1:04:11gradient is not correct so it's that's
- 1:04:13good to see as well okay so now we have
- 1:04:15the derivative here and we're trying to
- 1:04:17back propagate through this line
- 1:04:18and because we're raising to a power of
- 1:04:21negative 0.5 I brought up the power rule
- 1:04:23and we see that basically we have that
- 1:04:25the BM bar will now be we bring down the
- 1:04:28exponent so negative 0.5 times
- 1:04:31uh X which is this
- 1:04:34and now raised to the power of negative
- 1:04:360.5 minus 1 which is negative 1.5
- 1:04:39now we would have to also apply a small
- 1:04:42chain rule here in our head because we
- 1:04:45need to take further the derivative of B
- 1:04:48and VAR with respect to this expression
- 1:04:49here inside the bracket but because this
- 1:04:51is an elementalized operation and
- 1:04:53everything is fairly simple that's just
- 1:04:54one and so there's nothing to do there
- 1:04:57so this is the local derivative and then
- 1:05:00times the global derivative to create
- 1:05:01the chain rule this is just times the BM
- 1:05:04bar have
- 1:05:05so this is our candidate let me bring
- 1:05:08this down
- 1:05:10and uncommon to the check
- 1:05:14and we see that we have the correct
- 1:05:16result
- 1:05:17now before we propagate through the next
- 1:05:19line I want to briefly talk about the
- 1:05:20note here where I'm using the bezels
- 1:05:22correction dividing by n minus 1 instead
- 1:05:24of dividing by n when I normalize here
- 1:05:27the sum of squares
- 1:05:29now you'll notice that this is departure
- 1:05:31from the paper which uses one over n
- 1:05:33instead not one over n minus one their m
- 1:05:36is RN
- 1:05:38and
- 1:05:39um so it turns out that there are two
- 1:05:40ways of estimating variance of an array
- 1:05:43one is the biased estimate which is one
- 1:05:46over n and the other one is the unbiased
- 1:05:49estimate which is one over n minus one
- 1:05:51now confusingly in the paper this is uh
- 1:05:54not very clearly described and also it's
- 1:05:56a detail that kind of matters I think
- 1:05:58um they are using the biased version
- 1:06:00training time but later when they are
- 1:06:02talking about the inference they are
- 1:06:04mentioning that when they do the
- 1:06:06inference they are using the unbiased
- 1:06:08estimate which is the n minus one
- 1:06:10version in
- 1:06:12um
- 1:06:12basically for inference
- 1:06:15and to calibrate the running mean and
- 1:06:18the running variance basically and so
- 1:06:20they they actually introduce a trained
- 1:06:22test mismatch where in training they use
- 1:06:24the biased version and in the in test
- 1:06:26time they use the unbiased version I
- 1:06:28find this extremely confusing you can
- 1:06:30read more about the bezels correction
- 1:06:32and why uh dividing by n minus one gives
- 1:06:35you a better estimate of the variance in
- 1:06:37a case where you have population size or
- 1:06:39samples for the population
- 1:06:41that are very small and that is indeed
- 1:06:44the case for us because we are dealing
- 1:06:46with many patches and these mini matches
- 1:06:48are a small sample of a larger
- 1:06:50population which is the entire training
- 1:06:52set and so it just turns out that if you
- 1:06:55just estimate it using one over n that
- 1:06:57actually almost always underestimates
- 1:06:58the variance and it is a biased
- 1:07:00estimator and it is advised that you use
- 1:07:02the unbiased version and divide by n
- 1:07:04minus one and you can go through this
- 1:07:06article here that I liked that actually
- 1:07:08describes the full reasoning and I'll
- 1:07:09link it in the video description
- 1:07:12now when you calculate the torture
- 1:07:13variance
- 1:07:15you'll notice that they take the
- 1:07:16unbiased flag whether or not you want to
- 1:07:18divide by n or n minus one confusingly
- 1:07:21they do not mention what the default is
- 1:07:24for unbiased but I believe unbiased by
- 1:07:26default is true I'm not sure why the
- 1:07:29docs here don't cite that
- 1:07:31now in The Bachelor
- 1:07:331D the documentation again is kind of
- 1:07:35wrong and confusing it says that the
- 1:07:38standard deviation is calculated via the
- 1:07:39biased estimator
- 1:07:41but this is actually not exactly right
- 1:07:43and people have pointed out that it is
- 1:07:44not right in a number of issues since
- 1:07:46then because actually the rabbit hole is
- 1:07:49deeper and they follow the paper exactly
- 1:07:52and they use the biased version for
- 1:07:54training but when they're estimating the
- 1:07:56running standard deviation we are using
- 1:07:58the unbiased version so again there's
- 1:08:00the train test mismatch so long story
- 1:08:02short I'm not a fan of trained test
- 1:08:05discrepancies I basically kind of
- 1:08:07consider
- 1:08:08the fact that we use the bias version
- 1:08:10the training time and the unbiased test
- 1:08:13time I basically consider this to be a
- 1:08:14bug and I don't think that there's a
- 1:08:16good reason for that it's not really
- 1:08:18they don't really go into the detail of
- 1:08:19the reasoning behind it in this paper so
- 1:08:22that's why I basically prefer to use the
- 1:08:24bestless correction in my own work
- 1:08:26unfortunately Bastion does not take a
- 1:08:29keyword argument that tells you whether
- 1:08:30or not you want to use the unbiased
- 1:08:33version of the bias version in both
- 1:08:34train and test and so therefore anyone
- 1:08:36using batch normalization basically in
- 1:08:38my view has a bit of a bug in the code
- 1:08:41um
- 1:08:42and this turns out to be much less of a
- 1:08:44problem if your batch mini batch sizes
- 1:08:46are a bit larger but still I just might
- 1:08:48kind of uh unpardable so maybe someone
- 1:08:51can explain why this is okay but for now
- 1:08:53I prefer to use the unbiased version
- 1:08:55consistently both during training and at
- 1:08:57this time and that's why I'm using one
- 1:09:00over n minus one here
- 1:09:01okay so let's now actually back
- 1:09:03propagate through this line
- 1:09:05so
- 1:09:07the first thing that I always like to do
- 1:09:08is I like to scrutinize the shapes first
- 1:09:10so in particular here looking at the
- 1:09:12shapes of what's involved I see that b
- 1:09:14and VAR shape is 1 by 64. so it's a row
- 1:09:18vector and BND if two dot shape is 32 by
- 1:09:2164.
- 1:09:22so clearly here we're doing a sum over
- 1:09:25the zeroth axis to squash the first
- 1:09:28dimension of of the shapes here using a
- 1:09:32sum so that right away actually hints to
- 1:09:35me that there will be some kind of a
- 1:09:36replication or broadcasting in the
- 1:09:38backward pass and maybe you're noticing
- 1:09:40the pattern here but basically anytime
- 1:09:42you have a sum in the forward pass that
- 1:09:45turns into a replication or broadcasting
- 1:09:47in the backward pass along the same
- 1:09:49Dimension and conversely when we have a
- 1:09:52replication or a broadcasting in the
- 1:09:54forward pass that indicates a variable
- 1:09:57reuse and so in the backward pass that
- 1:09:59turns into a sum over the exact same
- 1:10:01dimension
- 1:10:02and so hopefully you're noticing that
- 1:10:04Duality that those two are kind of like
- 1:10:06the opposite of each other in the
- 1:10:07forward and backward pass
- 1:10:09now once we understand the shapes the
- 1:10:11next thing I like to do always is I like
- 1:10:12to look at a toy example in my head to
- 1:10:15sort of just like understand roughly how
- 1:10:16uh the variable the variable
- 1:10:18dependencies go in the mathematical
- 1:10:19formula
- 1:10:21so here we have a two-dimensional array
- 1:10:24of the end of two which we are scaling
- 1:10:26by a constant and then we are summing uh
- 1:10:29vertically over the columns so if we
- 1:10:32have a two by two Matrix a and then we
- 1:10:33sum over the columns and scale we would
- 1:10:36get a row Vector B1 B2 and B1 depends on
- 1:10:39a in this way whereas just sum they're
- 1:10:42scaled of a and B2 in this way where
- 1:10:45it's the second column sump and scale
- 1:10:48and so looking at this basically
- 1:10:52what we want to do now is we have the
- 1:10:53derivatives on B1 and B2 and we want to
- 1:10:55back propagate them into Ace and so it's
- 1:10:58clear that just differentiating in your
- 1:10:59head the local derivative here is one
- 1:11:01over n minus 1 times uh one
- 1:11:05uh for each one of these A's and um
- 1:11:09basically the derivative of B1 has to
- 1:11:11flow through The Columns of a
- 1:11:13scaled by one over n minus one
- 1:11:16and that's roughly What's Happening Here
- 1:11:18so intuitively the derivative flow tells
- 1:11:21us that DB and diff2
- 1:11:24will be the local derivative of this
- 1:11:27operation and there are many ways to do
- 1:11:29this by the way but I like to do
- 1:11:31something like this torch dot once like
- 1:11:33of bndf2 so I'll create a large array
- 1:11:37two-dimensional of ones
- 1:11:39and then I will scale it so 1.0 divided
- 1:11:42by n minus 1.
- 1:11:44so this is a array of
- 1:11:46um one over n minus one and that's sort
- 1:11:49of like the local derivative
- 1:11:50and now for the chain rule I will simply
- 1:11:53just multiply it by dbm bar
- 1:11:58and notice here what's going to happen
- 1:12:00this is 32 by 64 and this is just 1 by
- 1:12:0264. so I'm letting the broadcasting do
- 1:12:06the replication because internally in
- 1:12:08pytorch basically dbnbar which is 1 by
- 1:12:1164 row vector
- 1:12:13well in this multiplication get
- 1:12:15um copied vertically until the two are
- 1:12:18of the same shape and then there will be
- 1:12:19an element wise multiply and so that uh
- 1:12:22so that the broadcasting is basically
- 1:12:23doing the replication
- 1:12:25and I will end up with the derivatives
- 1:12:27of DB and diff2 here
- 1:12:30so this is the candidate solution let's
- 1:12:32bring it down here
- 1:12:33let's uncomment this line where we check
- 1:12:36it and let's hope for the best
- 1:12:39and indeed we see that this is the
- 1:12:41correct formula next up let's
- 1:12:43differentiate here and to be in this
- 1:12:45so here we have that b and diff is
- 1:12:48element y squared to create B and F2
- 1:12:50so this is a relatively simple
- 1:12:52derivative because it's a simple element
- 1:12:54wise operation so it's kind of like the
- 1:12:56scalar case and we have that DB and div
- 1:12:59should be if this is x squared then the
- 1:13:02derivative of this is 2x right so it's
- 1:13:04simply 2 times B and if that's the local
- 1:13:07derivative
- 1:13:08and then times chain Rule and the shape
- 1:13:11of these is the same they are of the
- 1:13:13same shape so times this
- 1:13:15so that's the backward pass for this
- 1:13:17variable let me bring that down here
- 1:13:20and now we have to be careful because we
- 1:13:22already calculated dbm depth right so
- 1:13:24this is just the end of the other uh you
- 1:13:27know other Branch coming back to B and
- 1:13:30diff
- 1:13:30because B and diff was already back
- 1:13:32propagated to way over here
- 1:13:34from being raw so we now completed the
- 1:13:37second branch and so that's why I have
- 1:13:39to do plus equals and if you recall we
- 1:13:42had an incorrect derivative for being
- 1:13:43diff before and I'm hoping that once we
- 1:13:46append this last missing piece we have
- 1:13:48the exact correctness so let's run
- 1:13:51ambient to be in div now actually shows
- 1:13:55the exact correct derivative
- 1:13:57um so that's comforting okay so let's
- 1:14:00now back propagate through this line
- 1:14:01here
- 1:14:03um the first thing we do of course is we
- 1:14:04check the shapes and I wrote them out
- 1:14:07here and basically the shape of this is
- 1:14:0832 by 64. hpbn is the same shape
- 1:14:12but B and mean I is a row Vector 1 by
- 1:14:1564. so this minus here will actually do
- 1:14:17broadcasting and so we have to be
- 1:14:19careful with that and as a hint to us
- 1:14:21again because of The Duality a
- 1:14:23broadcasting and the forward pass means
- 1:14:25a variable reuse and therefore there
- 1:14:27will be a sum in the backward pass
- 1:14:30so let's write out the backward pass
- 1:14:31here now
- 1:14:33um
- 1:14:34back propagate into the hpbn
- 1:14:37because this is these are the same shape
- 1:14:39then the local derivative for each one
- 1:14:41of the elements here is just one for the
- 1:14:43corresponding element in here
- 1:14:45so basically what this means is that the
- 1:14:47gradient just simply copies it's just a
- 1:14:50variable assignment it's quality so I'm
- 1:14:52just going to clone this tensor just for
- 1:14:54safety to create an exact copy of DB and
- 1:14:58div
- 1:15:00and then here to back propagate into
- 1:15:01this one what I'm inclined to do here is
- 1:15:07will basically be
- 1:15:09uh what is the local derivative well
- 1:15:12it's negative torch.1's like
- 1:15:16of the shape of uh B and diff
- 1:15:19right
- 1:15:22and then times
- 1:15:24the um
- 1:15:27the derivative here dbf
- 1:15:32and this here is the back propagation
- 1:15:34for the replicated B and mean I
- 1:15:37so I still have to back propagate
- 1:15:39through the uh replication in the
- 1:15:42broadcasting and I do that by doing a
- 1:15:43sum so I'm going to take this whole
- 1:15:45thing and I'm going to do a sum over the
- 1:15:47zeroth dimension which was the
- 1:15:49replication
- 1:15:53so if you scrutinize this by the way
- 1:15:55you'll notice that this is the same
- 1:15:57shape as that and so what I'm doing uh
- 1:16:00what I'm doing here doesn't actually
- 1:16:01make that much sense because it's just a
- 1:16:03array of ones multiplying DP and diff so
- 1:16:06in fact I can just do this
- 1:16:10um and that is equivalent
- 1:16:12so this is the candidate backward pass
- 1:16:15let me copy it here and then let me
- 1:16:18comment out this one and this one
- 1:16:22enter
- 1:16:24and it's wrong
- 1:16:27damn
- 1:16:29actually sorry this is supposed to be
- 1:16:31wrong and it's supposed to be wrong
- 1:16:33because
- 1:16:34we are back propagating from a b and
- 1:16:36diff into hpbn and but we're not done
- 1:16:39because B and mean I depends on hpbn and
- 1:16:43there will be a second portion of that
- 1:16:44derivative coming from this second
- 1:16:46Branch so we're not done yet and we
- 1:16:48expect it to be incorrect so there you
- 1:16:50go
- 1:16:50uh so let's now back propagate from uh B
- 1:16:53and mean I into hpbn
- 1:16:56um
- 1:16:57and so here again we have to be careful
- 1:16:58because there's a broadcasting along
- 1:17:01um or there's a Sum along the zeroth
- 1:17:03dimension so this will turn into
- 1:17:04broadcasting in the backward pass now
- 1:17:06and I'm going to go a little bit faster
- 1:17:08on this line because it is very similar
- 1:17:10to the line that we had before and
- 1:17:12multiplies in the past in fact
- 1:17:14so the hpbn
- 1:17:18will be
- 1:17:20the gradient will be scaled by 1 over n
- 1:17:22and then basically this gradient here on
- 1:17:25dbn mean I
- 1:17:27is going to be scaled by 1 over n and
- 1:17:30then it's going to flow across all the
- 1:17:32columns and deposit itself into the hpvn
- 1:17:35so what we want is this thing scaled by
- 1:17:381 over n
- 1:17:39only put the constant up front here
- 1:17:43um
- 1:17:45so scale down the gradient and now we
- 1:17:47need to replicate it across all the um
- 1:17:51across all the rows here so we I like to
- 1:17:55do that by torch.lunslike of basically
- 1:18:00um hpbn
- 1:18:03and I will let the broadcasting do the
- 1:18:05work of replication
- 1:18:09so
- 1:18:14like that
- 1:18:16so this is uh the hppn and hopefully
- 1:18:21we can plus equals that
- 1:18:27so this here is broadcasting
- 1:18:30um and then this is the scaling so this
- 1:18:32should be current
- 1:18:33okay
- 1:18:35so that completes the back propagation
- 1:18:37of the bathroom layer and we are now
- 1:18:38here let's back propagate through the
- 1:18:40linear layer one here now because
- 1:18:43everything is getting a little
- 1:18:44vertically crazy I copy pasted the line
- 1:18:46here and let's just back properly
- 1:18:48through this one line
- 1:18:50so first of course we inspect the shapes
- 1:18:52and we see that this is 32 by 64. MCAT
- 1:18:56is 32 by 30.
- 1:18:58W1 is 30 30 by 64 and B1 is just 64. so
- 1:19:04as I mentioned back propagating through
- 1:19:06linear layers is fairly easy just by
- 1:19:08matching the shapes so let's do that we
- 1:19:11have that dmcat
- 1:19:14should be
- 1:19:15um some matrix multiplication of dhbn
- 1:19:18with uh W1 and one transpose thrown in
- 1:19:21there so to make uh MCAT be 32 by 30
- 1:19:28I need to take dhpn
- 1:19:3232 by 64 and multiply it by w1.
- 1:19:36transpose
- 1:19:39to get the only one I need to end up
- 1:19:43with 30 by 64.
- 1:19:45so to get that I need to take uh MCAT
- 1:19:48transpose
- 1:19:51and multiply that by
- 1:19:53uh dhpion
- 1:19:58and finally to get DB1
- 1:20:01this is a addition and we saw that
- 1:20:04basically I need to just sum the
- 1:20:06elements in dhpbn along some Dimension
- 1:20:09and to make the dimensions work out I
- 1:20:12need to Sum along the zeroth axis here
- 1:20:14to eliminate this Dimension and we do
- 1:20:17not keep dims
- 1:20:19uh so that we want to just get a single
- 1:20:21one-dimensional lecture of 64.
- 1:20:23so these are the claimed derivatives
- 1:20:27let me put that here and let me
- 1:20:29uncomment three lines and cross our
- 1:20:32fingers
- 1:20:34everything is great okay so we now
- 1:20:36continue almost there we have the
- 1:20:37derivative of MCAT and we want to
- 1:20:39derivative we want to back propagate
- 1:20:41into m
- 1:20:43so I again copied this line over here
- 1:20:46so this is the forward pass and then
- 1:20:48this is the shapes so remember that the
- 1:20:51shape here was 32 by 30 and the original
- 1:20:53shape of M plus 32 by 3 by 10. so this
- 1:20:57layer in the forward pass as you recall
- 1:20:58did the concatenation of these three
- 1:21:0110-dimensional character vectors
- 1:21:04and so now we just want to undo that
- 1:21:06so this is actually relatively
- 1:21:08straightforward operation because uh the
- 1:21:11backward pass of the what is the view
- 1:21:12view is just a representation of the
- 1:21:15array it's just a logical form of how
- 1:21:17you interpret the array so let's just
- 1:21:18reinterpret it to be what it was before
- 1:21:21so in other words the end is not uh 32
- 1:21:25by 30. it is basically dmcat
- 1:21:29but if you view it as the original shape
- 1:21:34so just m dot shape
- 1:21:37uh you can you can pass in tuples into
- 1:21:39view
- 1:21:40and so this should just be okay
- 1:21:44we just re-represent that view and then
- 1:21:47we uncomment this line here and
- 1:21:49hopefully
- 1:21:51yeah so the derivative of M is correct
- 1:21:55so in this case we just have to
- 1:21:56re-represent the shape of those
- 1:21:57derivatives into the original View
- 1:21:59so now we are at the final line and the
- 1:22:01only thing that's left to back propagate
- 1:22:02through is this indexing operation here
- 1:22:05MSC at xB so as I did before I copy
- 1:22:09pasted this line here and let's look at
- 1:22:11the shapes of everything that's involved
- 1:22:12and remind ourselves how this worked
- 1:22:15so m.shape was 32 by 3 by 10.
- 1:22:19it says 32 examples and then we have
- 1:22:22three characters each one of them has a
- 1:22:2410 dimensional embedding
- 1:22:26and this was achieved by taking the
- 1:22:28lookup table C which have 27 possible
- 1:22:31characters
- 1:22:32each of them 10 dimensional and we
- 1:22:34looked up
- 1:22:35at the rows that were specified inside
- 1:22:39this tensor xB
- 1:22:41so XB is 32 by 3 and it's basically
- 1:22:43giving us for each example the Identity
- 1:22:45or the index of which character is part
- 1:22:49of that example
- 1:22:50and so here I'm showing the first five
- 1:22:52rows of three of this tensor xB
- 1:22:57and so we can see that for example here
- 1:22:58it was the first example in this batch
- 1:23:00is that the first character and the
- 1:23:02first character and the fourth character
- 1:23:04comes into the neural net
- 1:23:06and then we want to predict the next
- 1:23:08character in a sequence after the
- 1:23:10character is one one four
- 1:23:12so basically What's Happening Here is
- 1:23:14there are integers inside XB and each
- 1:23:18one of these integers is specifying
- 1:23:19which row of C we want to pluck out
- 1:23:22right and then we arrange those rows
- 1:23:25that we've plucked out into 32 by 3 by
- 1:23:2810 tensor and we just package them in we
- 1:23:30just package them into the sensor
- 1:23:33and now what's happening is that we have
- 1:23:35D amp
- 1:23:36so for every one of these uh basically
- 1:23:39plucked out rows we have their gradients
- 1:23:41now
- 1:23:42but they're arranged inside this 32 by 3
- 1:23:45by 10 tensor so all we have to do now is
- 1:23:48we just need to Route this gradient
- 1:23:49backwards through this assignment so we
- 1:23:52need to find which row of C that every
- 1:23:54one of these
- 1:23:56um 10 dimensional embeddings come from
- 1:23:59and then we need to deposit them into DC
- 1:24:03so we just need to undo the indexing and
- 1:24:06of course if any of these rows of C was
- 1:24:08used multiple times which almost
- 1:24:10certainly is the case like the row one
- 1:24:11and one was used multiple times then we
- 1:24:13have to remember that the gradients that
- 1:24:15arrive there have to add
- 1:24:18so for each occurrence we have to have
- 1:24:19an addition
- 1:24:21so let's now write this out and I don't
- 1:24:23actually know if like a much better way
- 1:24:24to do this than a for Loop unfortunately
- 1:24:26in Python
- 1:24:28um so maybe someone can come up with a
- 1:24:29vectorized efficient operation but for
- 1:24:32now let's just use for loops so let me
- 1:24:34create a torch.zeros like
- 1:24:37C to initialize uh just uh 27 by 10
- 1:24:40tensor of all zeros
- 1:24:43and then honestly 4K in range XB dot
- 1:24:46shape at zero
- 1:24:49maybe someone has a better way to do
- 1:24:51this but for J and range
- 1:24:53be that shape at one
- 1:24:55this is going to iterate over all the
- 1:24:58um all the elements of XB all these
- 1:25:01integers
- 1:25:03and then let's get the index at this
- 1:25:05position
- 1:25:06so the index is basically x b at KJ
- 1:25:11so that an example of that like is 11 or
- 1:25:1414 and so on
- 1:25:16and now in the forward pass we took
- 1:25:19and we basically took um
- 1:25:24the row of C at index and we deposited
- 1:25:27it into M at K of J
- 1:25:30that's what happened that's where they
- 1:25:32are packaged so now we need to go
- 1:25:34backwards and we just need to route
- 1:25:36DM at the position KJ
- 1:25:39we now have these derivatives
- 1:25:42for each position and it's 10
- 1:25:44dimensional
- 1:25:45and you just need to go into the correct
- 1:25:47row of C
- 1:25:49so DC rather at IX is this but plus
- 1:25:54equals
- 1:25:55because there could be multiple
- 1:25:56occurrences uh like the same row could
- 1:25:58have been used many many times and so
- 1:26:00all of those derivatives will just go
- 1:26:04backwards through the indexing and they
- 1:26:06will add
- 1:26:07so this is my candidate solution
- 1:26:12let's copy it here
- 1:26:16let's uncomment this and cross our
- 1:26:19fingers
- 1:26:20hey
- 1:26:21so that's it we've back propagated
- 1:26:24through
- 1:26:25this entire Beast
- 1:26:28so there we go totally makes sense
- 1:26:31so now we come to exercise two it
- 1:26:33basically turns out that in this first
- 1:26:34exercise we were doing way too much work
- 1:26:36uh we were back propagating way too much
- 1:26:39and it was all good practice and so on
- 1:26:40but it's not what you would do in
- 1:26:42practice and the reason for that is for
- 1:26:44example here I separated out this loss
- 1:26:47calculation over multiple lines and I
- 1:26:49broke it up all all to like its smallest
- 1:26:51atomic pieces and we back propagated
- 1:26:53through all of those individually
- 1:26:55but it turns out that if you just look
- 1:26:56at the mathematical expression for the
- 1:26:58loss
- 1:27:00um then actually you can do the
- 1:27:02differentiation on pen and paper and a
- 1:27:04lot of terms cancel and simplify and the
- 1:27:06mathematical expression you end up with
- 1:27:07can be significantly shorter and easier
- 1:27:10to implement than back propagating
- 1:27:11through all the little pieces of
- 1:27:12everything you've done
- 1:27:13so before we had this complicated
- 1:27:16forward paths going from logits to the
- 1:27:18loss
- 1:27:19but in pytorch everything can just be
- 1:27:21glued together into a single call at
- 1:27:22that cross entropy you just pass in
- 1:27:24logits and the labels and you get the
- 1:27:26exact same loss as I verify here so our
- 1:27:28previous loss and the fast loss coming
- 1:27:31from the chunk of operations as a single
- 1:27:33mathematical expression is the same but
- 1:27:36it's much much faster in a forward pass
- 1:27:38it's also much much faster in backward
- 1:27:40pass and the reason for that is if you
- 1:27:42just look at the mathematical form of
- 1:27:43this and differentiate again you will
- 1:27:45end up with a very small and short
- 1:27:46expression so that's what we want to do
- 1:27:48here we want to in a single operation or
- 1:27:51in a single go or like very quickly go
- 1:27:54directly to delojits
- 1:27:56and we need to implement the logits as a
- 1:27:59function of logits and yb's
- 1:28:02but it will be significantly shorter
- 1:28:04than whatever we did here where to get
- 1:28:06to deluggets we had to go all the way
- 1:28:08here
- 1:28:10so all of this work can be skipped in a
- 1:28:12much much simpler mathematical
- 1:28:13expression that you can Implement here
- 1:28:16so you can give it a shot yourself
- 1:28:18basically look at what exactly is the
- 1:28:21mathematical expression of loss and
- 1:28:23differentiate with respect to the logits
- 1:28:26so let me show you a hint you can of
- 1:28:29course try it fully yourself but if not
- 1:28:31I can give you some hint of how to get
- 1:28:33started mathematically
- 1:28:36so basically What's Happening Here is we
- 1:28:38have logits then there's a softmax that
- 1:28:41takes the logits and gives you
- 1:28:42probabilities then we are using the
- 1:28:44identity of the correct next character
- 1:28:46to pluck out a row of probabilities take
- 1:28:50the negative log of it to get our
- 1:28:51negative block probability and then we
- 1:28:54average up all the log probabilities or
- 1:28:56negative block probabilities to get our
- 1:28:58loss
- 1:28:59so basically what we have is for a
- 1:29:01single individual example rather we have
- 1:29:04that loss is equal to negative log
- 1:29:06probability uh where P here is kind of
- 1:29:09like thought of as a vector of all the
- 1:29:11probabilities so at the Y position where
- 1:29:14Y is the label
- 1:29:16and we have that P here of course is the
- 1:29:19softmax so the ith component of P of
- 1:29:23this probability Vector is just the
- 1:29:25softmax function so raising all the
- 1:29:28logits uh basically to the power of E
- 1:29:31and normalizing so everything comes to
- 1:29:341.
- 1:29:35now if you write out P of Y here you can
- 1:29:38just write out the soft Max and then
- 1:29:40basically what we're interested in is
- 1:29:41we're interested in the derivative of
- 1:29:43the loss with respect to the I logit
- 1:29:47and so basically it's a d by DLI of this
- 1:29:51expression here
- 1:29:52where we have L indexed with the
- 1:29:54specific label Y and on the bottom we
- 1:29:56have a sum over J of e to the L J and
- 1:29:58the negative block of all that so
- 1:30:00potentially give it a shot pen and paper
- 1:30:02and see if you can actually derive the
- 1:30:04expression for the loss by DLI and then
- 1:30:07we're going to implement it here okay so
- 1:30:09I'm going to give away the result here
- 1:30:11so this is some of the math I did to
- 1:30:13derive the gradients analytically and so
- 1:30:17we see here that I'm just applying the
- 1:30:19rules of calculus from your first or
- 1:30:20second year of bachelor's degree if you
- 1:30:22took it and we see that the expression
- 1:30:24is actually simplify quite a bit you
- 1:30:26have to separate out the analysis in the
- 1:30:27case where the ith index that you're
- 1:30:30interested in inside logits is either
- 1:30:32equal to the label or it's not equal to
- 1:30:34the label and then the expression
- 1:30:35simplify and cancel in a slightly
- 1:30:37different way and what we end up with is
- 1:30:39something very very simple
- 1:30:41and we either end up with basically
- 1:30:43pirai where p is again this Vector of
- 1:30:46probabilities after a soft Max or P at I
- 1:30:49minus 1 where we just simply subtract a
- 1:30:51one but in any case we just need to
- 1:30:53calculate the soft Max p e and then in
- 1:30:56the correct Dimension we need to
- 1:30:58subtract one and that's the gradient the
- 1:31:00form that it takes analytically so let's
- 1:31:03implement this basically and we have to
- 1:31:04keep in mind that this is only done for
- 1:31:06a single example but here we are working
- 1:31:08with batches of examples
- 1:31:09so we have to be careful of that and
- 1:31:12then the loss for a batch is the average
- 1:31:14loss over all the examples so in other
- 1:31:17words is the example for all the
- 1:31:18individual examples is the loss for each
- 1:31:20individual example summed up and then
- 1:31:22divided by n and we have to back
- 1:31:24propagate through that as well and be
- 1:31:26careful with it
- 1:31:28so deluggets is going to be of that soft
- 1:31:30Max
- 1:31:32uh pytorch has a softmax function that
- 1:31:35you can call and we want to apply the
- 1:31:36softmax on the logits and we want to go
- 1:31:39in the dimension that is one so
- 1:31:42basically we want to do the softmax
- 1:31:44along the rows of these logits
- 1:31:47then at the correct positions we need to
- 1:31:49subtract a 1. so delugits at iterating
- 1:31:52over all the rows
- 1:31:54and indexing into the columns
- 1:31:57provided by the correct labels inside YB
- 1:32:00we need to subtract one
- 1:32:03and then finally it's the average loss
- 1:32:05that is the loss and in the average
- 1:32:07there's a one over n of all the losses
- 1:32:09added up and so we need to also
- 1:32:12propagate through that division
- 1:32:14so the gradient has to be scaled down by
- 1:32:16by n as well because of the mean
- 1:32:19but this otherwise should be the result
- 1:32:22so now if we verify this
- 1:32:24we see that we don't get an exact match
- 1:32:26but at the same time the maximum
- 1:32:30difference from logits from pytorch and
- 1:32:33RD logits here is uh on the order of 5e
- 1:32:37negative 9. so it's a tiny tiny number
- 1:32:39so because of floating point wantiness
- 1:32:41we don't get the exact bitwise result
- 1:32:44but we basically get the correct answer
- 1:32:47approximately
- 1:32:49now I'd like to pause here briefly
- 1:32:51before we move on to the next exercise
- 1:32:52because I'd like us to get an intuitive
- 1:32:54sense of what the logits is because it
- 1:32:56has a beautiful and very simple
- 1:32:58explanation honestly
- 1:33:00um so here I'm taking the logits and I'm
- 1:33:03visualizing it and we can see that we
- 1:33:05have a batch of 32 examples of 27
- 1:33:07characters
- 1:33:08and what is the logits intuitively right
- 1:33:10the logits is the probabilities that the
- 1:33:13properties Matrix in the forward pass
- 1:33:15but then here these black squares are
- 1:33:17the positions of the correct indices
- 1:33:19where we subtracted a one
- 1:33:21and so uh what is this doing right these
- 1:33:24are the derivatives on the logits and so
- 1:33:27let's look at just the first row here
- 1:33:31so that's what I'm doing here I'm
- 1:33:33clocking the probabilities of these
- 1:33:34logits and then I'm taking just the
- 1:33:36first row and this is the probability
- 1:33:38row and then the logits of the first row
- 1:33:41and multiplying by n just for us so that
- 1:33:43we don't have the scaling by n in here
- 1:33:46and everything is more interpretable we
- 1:33:48see that it's exactly equal to the
- 1:33:50probability of course but then the
- 1:33:52position of the correct index has a
- 1:33:53minus equals one so minus one on that
- 1:33:56position
- 1:33:57and so notice that
- 1:33:59um if you take Delo Jets at zero and you
- 1:34:01sum it
- 1:34:03it actually sums to zero and so you
- 1:34:06should think of these uh gradients here
- 1:34:08at each cell as like a force
- 1:34:12um we are going to be basically pulling
- 1:34:15down on the probabilities of the
- 1:34:17incorrect characters and we're going to
- 1:34:19be pulling up on the probability at the
- 1:34:22correct index and that's what's
- 1:34:24basically happening in each row and thus
- 1:34:29the amount of push and pull is exactly
- 1:34:31equalized because the sum is zero so the
- 1:34:34amount to which we pull down in the
- 1:34:36probabilities and the demand that we
- 1:34:37push up on the probability of the
- 1:34:39correct character is equal
- 1:34:41so sort of the the repulsion and the
- 1:34:43attraction are equal and think of the
- 1:34:45neural app now as a like a massive uh
- 1:34:48pulley system or something like that
- 1:34:50we're up here on top of the logits and
- 1:34:52we're pulling up we're pulling down the
- 1:34:54properties of Incorrect and pulling up
- 1:34:55the property of the correct and in this
- 1:34:57complicated pulley system because
- 1:34:59everything is mathematically uh just
- 1:35:01determined just think of it as sort of
- 1:35:03like this tension translating to this
- 1:35:05complicating pulling mechanism and then
- 1:35:07eventually we get a tug on the weights
- 1:35:09and the biases and basically in each
- 1:35:11update we just kind of like tug in the
- 1:35:13direction that we like for each of these
- 1:35:15elements and the parameters are slowly
- 1:35:17given in to the tug and that's what
- 1:35:19training in neural net kind of like
- 1:35:20looks like on a high level
- 1:35:22and so I think the the forces of push
- 1:35:24and pull in these gradients are actually
- 1:35:26uh very intuitive here we're pushing and
- 1:35:29pulling on the correct answer and the
- 1:35:31incorrect answers and the amount of
- 1:35:33force that we're applying is actually
- 1:35:34proportional to uh the probabilities
- 1:35:37that came out in the forward pass
- 1:35:39and so for example if our probabilities
- 1:35:41came out exactly correct so they would
- 1:35:43have had zero everywhere except for one
- 1:35:45at the correct uh position then the the
- 1:35:48logits would be all a row of zeros for
- 1:35:51that example there would be no push and
- 1:35:52pull so the amount to which your
- 1:35:55prediction is incorrect is exactly the
- 1:35:58amount by which you're going to get a
- 1:35:59pull or a push in that dimension
- 1:36:01so if you have for example a very
- 1:36:04confidently mispredicted element here
- 1:36:05then
- 1:36:07um what's going to happen is that
- 1:36:08element is going to be pulled down very
- 1:36:10heavily and the correct answer is going
- 1:36:12to be pulled up to the same amount
- 1:36:14and the other characters are not going
- 1:36:16to be influenced too much
- 1:36:19so the amounts to which you mispredict
- 1:36:21is then proportional to the strength of
- 1:36:23the pole and that's happening
- 1:36:25independently in all the dimensions of
- 1:36:27this of this tensor and it's sort of
- 1:36:29very intuitive and varies to think
- 1:36:30through and that's basically the magic
- 1:36:32of the cross-entropy loss and what it's
- 1:36:34doing dynamically in the backward pass
- 1:36:36of the neural net so now we get to
- 1:36:38exercise number three which is a very
- 1:36:41fun exercise
- 1:36:42um depending on your definition of fun
- 1:36:43and we are going to do for batch
- 1:36:45normalization exactly what we did for
- 1:36:47cross entropy loss in exercise number
- 1:36:49two that is we are going to consider it
- 1:36:51as a glued single mathematical
- 1:36:52expression and back propagate through it
- 1:36:54in a very efficient manner because we
- 1:36:56are going to derive a much simpler
- 1:36:58formula for the backward path of batch
- 1:36:59normalization
- 1:37:01and we're going to do that using pen and
- 1:37:02paper
- 1:37:03so previously we've broken up
- 1:37:05bastionalization into all of the little
- 1:37:06intermediate pieces and all the atomic
- 1:37:08operations inside it and then we back
- 1:37:10propagate it through it one by one
- 1:37:13now we just have a single sort of
- 1:37:15forward pass of a batch form and it's
- 1:37:18all glued together
- 1:37:20and we see that we get the exact same
- 1:37:21result as before
- 1:37:23now for the backward pass we'd like to
- 1:37:25also Implement a single formula
- 1:37:27basically for back propagating through
- 1:37:29this entire operation that is the
- 1:37:30bachelorization
- 1:37:32so in the forward pass previously we
- 1:37:34took hpvn the hidden states of the
- 1:37:37pre-batch realization and created H
- 1:37:39preact which is the hidden States just
- 1:37:42before the activation
- 1:37:44in the bachelorization paper each pbn is
- 1:37:46X and each preact is y
- 1:37:49so in the backward pass what we'd like
- 1:37:51to do now is we have DH preact and we'd
- 1:37:54like to produce d h previous
- 1:37:56and we'd like to do that in a very
- 1:37:57efficient manner so that's the name of
- 1:38:00the game calculate the H previan given
- 1:38:02DH preact and for the purposes of this
- 1:38:05exercise we're going to ignore gamma and
- 1:38:07beta and their derivatives because they
- 1:38:09take on a very simple form in a very
- 1:38:11similar way to what we did up above
- 1:38:14so let's calculate this given that right
- 1:38:18here
- 1:38:18so to help you a little bit like I did
- 1:38:20before I started off the implementation
- 1:38:23here on pen and paper and I took two
- 1:38:26sheets of paper to derive the
- 1:38:28mathematical formulas for the backward
- 1:38:29pass
- 1:38:30and basically to set up the problem uh
- 1:38:33just write out the MU Sigma Square
- 1:38:35variance x i hat and Y I exactly as in
- 1:38:39the paper except for the bezel
- 1:38:40correction
- 1:38:41and then
- 1:38:42in a backward pass we have the
- 1:38:44derivative of the loss with respect to
- 1:38:46all the elements of Y and remember that
- 1:38:48Y is a vector there's there's multiple
- 1:38:50numbers here
- 1:38:52so we have all the derivatives with
- 1:38:54respect to all the Y's
- 1:38:56and then there's a demo and a beta and
- 1:38:59this is kind of like the compute graph
- 1:39:01the gamma and the beta there's the X hat
- 1:39:03and then the MU and the sigma squared
- 1:39:06and the X so we have DL by DYI and we
- 1:39:10won't DL by d x i for all the I's in
- 1:39:13these vectors
- 1:39:15so this is the compute graph and you
- 1:39:17have to be careful because I'm trying to
- 1:39:19note here that these are vectors so
- 1:39:22there's many nodes here inside x x hat
- 1:39:25and Y but mu and sigma sorry Sigma
- 1:39:29Square are just individual scalars
- 1:39:30single numbers so you have to be careful
- 1:39:33with that you have to imagine there's
- 1:39:34multiple nodes here or you're going to
- 1:39:35get your math wrong
- 1:39:38um so as an example I would suggest that
- 1:39:40you go in the following order one two
- 1:39:43three four in terms of the back
- 1:39:44propagation so back propagating to X hat
- 1:39:46then into Sigma Square then into mu and
- 1:39:49then into X
- 1:39:52um just like in a topological sort in
- 1:39:54micrograd we would go from right to left
- 1:39:55you're doing the exact same thing except
- 1:39:57you're doing it with symbols and on a
- 1:39:59piece of paper
- 1:40:01so for number one uh I'm not giving away
- 1:40:05too much if you want DL of d x i hat
- 1:40:09then we just take DL by DYI and multiply
- 1:40:12it by gamma because of this expression
- 1:40:15here where any individual Yi is just
- 1:40:17gamma times x i hat plus beta so it
- 1:40:21doesn't help you too much there but this
- 1:40:23gives you basically the derivatives for
- 1:40:25all the X hats and so now try to go
- 1:40:28through this computational graph and
- 1:40:31derive what is DL by D Sigma Square
- 1:40:35and then what is DL by B mu and then one
- 1:40:38is D L by DX
- 1:40:39eventually so give it a go and I'm going
- 1:40:42to be revealing the answer one piece at
- 1:40:44a time okay so to get DL by D Sigma
- 1:40:46Square we have to remember again like I
- 1:40:48mentioned that there are many excess X
- 1:40:51hats here
- 1:40:52and remember that Sigma square is just a
- 1:40:54single individual number here
- 1:40:55so when we look at the expression
- 1:40:59for the L by D Sigma Square
- 1:41:01we have that we have to actually
- 1:41:03consider all the possible paths that um
- 1:41:08we basically have that there's many X
- 1:41:10hats and they all feed off from they all
- 1:41:13depend on Sigma Square so Sigma square
- 1:41:15has a large fan out there's lots of
- 1:41:17arrows coming out from Sigma square into
- 1:41:19all the X hats
- 1:41:20and then there's a back propagating
- 1:41:22signal from each X hat into Sigma square
- 1:41:24and that's why we actually need to sum
- 1:41:26over all those I's from I equal to 1 to
- 1:41:29m
- 1:41:30of the DL by d x i hat which is the
- 1:41:35global gradient
- 1:41:36times the x i Hat by D Sigma Square
- 1:41:40which is the local gradient
- 1:41:42of this operation here
- 1:41:44and then mathematically I'm just working
- 1:41:46it out here and I'm simplifying and you
- 1:41:48get a certain expression for DL by D
- 1:41:51Sigma square and we're going to be using
- 1:41:52this expression when we back propagate
- 1:41:53into mu and then eventually into X so
- 1:41:56now let's continue our back propagation
- 1:41:58into mu so what is D L by D mu now again
- 1:42:01be careful that mu influences X hat and
- 1:42:04X hat is actually lots of values so for
- 1:42:07example if our mini batch size is 32 as
- 1:42:09it is in our example that we were
- 1:42:10working on then this is 32 numbers and
- 1:42:1332 arrows going back to mu and then mu
- 1:42:16going to Sigma square is just a single
- 1:42:18Arrow because Sigma square is a scalar
- 1:42:19so in total there are 33 arrows
- 1:42:22emanating from you and then all of them
- 1:42:25have gradients coming into mu and they
- 1:42:27all need to be summed up
- 1:42:29and so that's why when we look at the
- 1:42:31expression for DL by D mu I am summing
- 1:42:34up over all the gradients of DL by d x i
- 1:42:37hat times the x i Hat by being mu
- 1:42:40uh so that's the that's this arrow and
- 1:42:43that's 32 arrows here and then plus the
- 1:42:45one Arrow from here which is the L by
- 1:42:47the sigma Square Times the sigma squared
- 1:42:49by D mu
- 1:42:50so now we have to work out that
- 1:42:52expression and let me just reveal the
- 1:42:54rest of it
- 1:42:55uh simplifying here is not complicated
- 1:42:58the first term and you just get an
- 1:43:00expression here
- 1:43:01for the second term though there's
- 1:43:02something really interesting that
- 1:43:03happens
- 1:43:04when we look at the sigma squared by D
- 1:43:06mu and we simplify
- 1:43:08at one point if we assume that in a
- 1:43:11special case where mu is actually the
- 1:43:14average of X I's as it is in this case
- 1:43:17then if we plug that in then actually
- 1:43:20the gradient vanishes and becomes
- 1:43:22exactly zero and that makes the entire
- 1:43:24second term cancel
- 1:43:26and so these uh if you just have a
- 1:43:29mathematical expression like this and
- 1:43:30you look at D Sigma Square by D mu you
- 1:43:33would get some mathematical formula for
- 1:43:35how mu impacts Sigma Square
- 1:43:37but if it is the special case that Nu is
- 1:43:39actually equal to the average as it is
- 1:43:42in the case of pastoralization that
- 1:43:43gradient will actually vanish and become
- 1:43:45zero so the whole term cancels and we
- 1:43:48just get a fairly straightforward
- 1:43:49expression here for DL by D mu okay and
- 1:43:52now we get to the craziest part which is
- 1:43:54uh deriving DL by dxi which is
- 1:43:57ultimately what we're after
- 1:43:59now let's count
- 1:44:00first of all how many numbers are there
- 1:44:03inside X as I mentioned there are 32
- 1:44:05numbers there are 32 Little X I's and
- 1:44:08let's count the number of arrows
- 1:44:09emanating from each x i
- 1:44:11there's an arrow going to Mu an arrow
- 1:44:13going to Sigma Square
- 1:44:14and then there's an arrow going to X hat
- 1:44:16but this Arrow here let's scrutinize
- 1:44:19that a little bit
- 1:44:20each x i hat is just a function of x i
- 1:44:23and all the other scalars so x i hat
- 1:44:27only depends on x i and none of the
- 1:44:29other X's
- 1:44:30and so therefore there are actually in
- 1:44:32this single Arrow there are 32 arrows
- 1:44:34but those 32 arrows are going exactly
- 1:44:37parallel they don't interfere and
- 1:44:39they're just going parallel between x
- 1:44:40and x hat you can look at it that way
- 1:44:42and so how many arrows are emanating
- 1:44:44from each x i there are three arrows mu
- 1:44:47Sigma squared and the associated X hat
- 1:44:50and so in back propagation we now need
- 1:44:53to apply the chain rule and we need to
- 1:44:55add up those three contributions
- 1:44:57so here's what that looks like if I just
- 1:44:59write that out
- 1:45:02we have uh we're going through we're
- 1:45:04chaining through mu Sigma square and
- 1:45:06through X hat and those three terms are
- 1:45:09just here
- 1:45:10now we already have three of these we
- 1:45:13have d l by d x i hat
- 1:45:15we have DL by D mu which we derived here
- 1:45:17and we have DL by D Sigma Square which
- 1:45:19we derived here but we need three other
- 1:45:22terms here
- 1:45:23the this one this one and this one so I
- 1:45:26invite you to try to derive them it's
- 1:45:28not that complicated you're just looking
- 1:45:29at these Expressions here and
- 1:45:31differentiating with respect to x i
- 1:45:34so give it a shot but here's the result
- 1:45:39or at least what I got
- 1:45:41um
- 1:45:42yeah I'm just I'm just differentiating
- 1:45:44with respect to x i for all these
- 1:45:45expressions and honestly I don't think
- 1:45:47there's anything too tricky here it's
- 1:45:48basic calculus
- 1:45:50now it gets a little bit more tricky is
- 1:45:52we are now going to plug everything
- 1:45:53together so all of these terms
- 1:45:55multiplied with all of these terms and
- 1:45:57add it up according to this formula and
- 1:45:59that gets a little bit hairy so what
- 1:46:01ends up happening is
- 1:46:04uh
- 1:46:05you get a large expression and the thing
- 1:46:08to be very careful with here of course
- 1:46:09is we are working with a DL by dxi for
- 1:46:12specific I here but when we are plugging
- 1:46:15in some of these terms
- 1:46:17like say
- 1:46:18um
- 1:46:19this term here deal by D signal squared
- 1:46:22you see how the L by D Sigma squared I
- 1:46:24end up with an expression and I'm
- 1:46:26iterating over little I's here but I
- 1:46:29can't use I as the variable when I plug
- 1:46:31in here because this is a different I
- 1:46:33from this eye
- 1:46:35this I here is just a place or like a
- 1:46:37local variable for for a for Loop in
- 1:46:39here so here when I plug that in you
- 1:46:41notice that I rename the I to a j
- 1:46:43because I need to make sure that this J
- 1:46:45is not that this J is not this I this J
- 1:46:48is like like a little local iterator
- 1:46:50over 32 terms and so you have to be
- 1:46:53careful with that when you're plugging
- 1:46:54in the expressions from here to here you
- 1:46:56may have to rename eyes into J's and you
- 1:46:58have to be very careful what is actually
- 1:47:00an I with respect to the L by t x i
- 1:47:04so some of these are J's some of these
- 1:47:07are I's
- 1:47:08and then we simplify this expression
- 1:47:11and I guess like the big thing to notice
- 1:47:13here is a bunch of terms just kind of
- 1:47:15come out to the front and you can
- 1:47:16refactor them there's a sigma squared
- 1:47:18plus Epsilon raised to the power of
- 1:47:19negative three over two uh this Sigma
- 1:47:21squared plus Epsilon can be actually
- 1:47:23separated out into three terms each of
- 1:47:25them are Sigma squared plus Epsilon to
- 1:47:28the negative one over two so the three
- 1:47:30of them multiplied is equal to this and
- 1:47:33then those three terms can go different
- 1:47:35places because of the multiplication so
- 1:47:37one of them actually comes out to the
- 1:47:39front and will end up here outside one
- 1:47:42of them joins up with this term and one
- 1:47:45of them joins up with this other term
- 1:47:47and then when you simplify the
- 1:47:49expression you'll notice that some of
- 1:47:51these terms that are coming out are just
- 1:47:52the x i hats
- 1:47:54so you can simplify just by rewriting
- 1:47:56that
- 1:47:57and what we end up with at the end is a
- 1:47:58fairly simple mathematical expression
- 1:48:00over here that I cannot simplify further
- 1:48:02but basically you'll notice that it only
- 1:48:05uses the stuff we have and it derives
- 1:48:06the thing we need so we have the L by d
- 1:48:10y for all the I's and those are used
- 1:48:13plenty of times here and also in
- 1:48:15addition what we're using is these x i
- 1:48:17hats and XJ hats and they just come from
- 1:48:19the forward pass
- 1:48:20and otherwise this is a simple
- 1:48:22expression and it gives us DL by d x i
- 1:48:25for all the I's and that's ultimately
- 1:48:27what we're interested in
- 1:48:29so that's the end of Bachelor backward
- 1:48:32pass analytically let's now implement
- 1:48:34this final result
- 1:48:36okay so I implemented the expression
- 1:48:38into a single line of code here and you
- 1:48:41can see that the max diff is Tiny so
- 1:48:43this is the correct implementation of
- 1:48:44this formula now I'll just uh
- 1:48:48basically tell you that getting this
- 1:48:50formula here from this mathematical
- 1:48:52expression was not trivial and there's a
- 1:48:54lot going on packed into this one
- 1:48:56formula and this is a whole exercise by
- 1:48:58itself because you have to consider the
- 1:49:00fact that this formula here is just for
- 1:49:03a single neuron and a batch of 32
- 1:49:05examples but what I'm doing here is I'm
- 1:49:07actually we actually have 64 neurons and
- 1:49:10so this expression has to in parallel
- 1:49:11evaluate the bathroom backward pass for
- 1:49:14all of those 64 neurons in parallel
- 1:49:16independently so this has to happen
- 1:49:18basically in every single
- 1:49:20um
- 1:49:20column of the inputs here
- 1:49:24and in addition to that you see how
- 1:49:26there are a bunch of sums here and we
- 1:49:28need to make sure that when I do those
- 1:49:29sums that they broadcast correctly onto
- 1:49:31everything else that's here
- 1:49:33and so getting this expression is just
- 1:49:35like highly non-trivial and I invite you
- 1:49:36to basically look through it and step
- 1:49:37through it and it's a whole exercise to
- 1:49:39make sure that this this checks out but
- 1:49:43once all the shapes are green and once
- 1:49:45you convince yourself that it's correct
- 1:49:46you can also verify that Patrick's gets
- 1:49:48the exact same answer as well and so
- 1:49:50that gives you a lot of peace of mind
- 1:49:51that this mathematical formula is
- 1:49:53correctly implemented here and
- 1:49:55broadcasted correctly and replicated in
- 1:49:57parallel for all of the 64 neurons
- 1:50:00inside this bastrum layer okay and
- 1:50:03finally exercise number four asks you to
- 1:50:05put it all together and uh here we have
- 1:50:08a redefinition of the entire problem so
- 1:50:10you see that we reinitialize the neural
- 1:50:11nut from scratch and everything and then
- 1:50:13here instead of calling loss that
- 1:50:15backward we want to have the manual back
- 1:50:18propagation here as we derived It Up
- 1:50:20Above so go up copy paste all the chunks
- 1:50:23of code that we've already derived put
- 1:50:25them here and drive your own gradients
- 1:50:26and then optimize this neural nut
- 1:50:28basically using your own gradients all
- 1:50:31the way to the calibration of The
- 1:50:33Bachelor and the evaluation of the loss
- 1:50:34and I was able to achieve quite a good
- 1:50:36loss basically the same loss you would
- 1:50:38achieve before and that shouldn't be
- 1:50:40surprising because all we've done is
- 1:50:41we've really gotten to Lost That
- 1:50:44backward and we've pulled out all the
- 1:50:45code
- 1:50:46and inserted it here but those gradients
- 1:50:49are identical and everything is
- 1:50:50identical and the results are identical
- 1:50:52it's just that we have full visibility
- 1:50:54on exactly what goes on under the hood
- 1:50:56I'll plot that backward in this specific
- 1:50:58case and this is all of our code this is
- 1:51:02the full backward pass using basically
- 1:51:04the simplified backward pass for the
- 1:51:06cross entropy loss and the mass
- 1:51:08generalization so back propagating
- 1:51:10through cross entropy the second layer
- 1:51:13the 10 H nonlinearity the batch
- 1:51:15normalization
- 1:51:16uh through the first layer and through
- 1:51:19the embedding and so you see that this
- 1:51:21is only maybe what is this 20 lines of
- 1:51:23code or something like that and that's
- 1:51:25what gives us gradients and now we can
- 1:51:27potentially erase losses backward so the
- 1:51:30way I have the code set up is you should
- 1:51:31be able to run this entire cell once you
- 1:51:33fill this in and this will run for only
- 1:51:36100 iterations and then break
- 1:51:37and it breaks because it gives you an
- 1:51:39opportunity to check your gradients
- 1:51:41against pytorch
- 1:51:43so here our gradients we see are not
- 1:51:46exactly equal they are approximately
- 1:51:48equal and the differences are tiny
- 1:51:51wanting negative 9 or so and I don't
- 1:51:52exactly know where they're coming from
- 1:51:54to be honest
- 1:51:56um so once we have some confidence that
- 1:51:57the gradients are basically correct we
- 1:51:59can take out the gradient tracking
- 1:52:01we can disable this breaking statement
- 1:52:05and then we can
- 1:52:07basically disable lost of backward we
- 1:52:10don't need it anymore it feels amazing
- 1:52:13to say that
- 1:52:14and then here when we are doing the
- 1:52:16update we're not going to use P dot grad
- 1:52:18this is the old way of pytorch we don't
- 1:52:21have that anymore because we're not
- 1:52:22doing backward we are going to use this
- 1:52:25update where we you see that I'm
- 1:52:27iterating over
- 1:52:29I've arranged the grads to be in the
- 1:52:30same order as the parameters and I'm
- 1:52:32zipping them up the gradients and the
- 1:52:34parameters into p and grad and then here
- 1:52:37I'm going to step with just the grad
- 1:52:38that we derived manually
- 1:52:40so the last piece
- 1:52:43um is that none of this now requires
- 1:52:46gradients from pytorch and so one thing
- 1:52:49you can do here
- 1:52:51um
- 1:52:52is you can do with no grad and offset
- 1:52:56this whole code block
- 1:52:58and really what you're saying is you're
- 1:52:59telling Pat George that hey I'm not
- 1:53:00going to call backward on any of this
- 1:53:02and this allows pytorch to be a bit more
- 1:53:03efficient with all of it
- 1:53:05and then we should be able to just uh
- 1:53:07run this
- 1:53:09and
- 1:53:11it's running
- 1:53:13and you see that losses backward is
- 1:53:16commented out
- 1:53:18and we're optimizing
- 1:53:20so we're going to leave this run and uh
- 1:53:23hopefully we get a good result
- 1:53:25okay so I allowed the neural net to
- 1:53:27finish optimization
- 1:53:28then here I calibrate the bachelor
- 1:53:31parameters because I did not keep track
- 1:53:33of the running mean and very variants in
- 1:53:35their training Loop
- 1:53:37then here I ran the loss and you see
- 1:53:39that we actually obtained a pretty good
- 1:53:40loss very similar to what we've achieved
- 1:53:42before
- 1:53:43and then here I'm sampling from the
- 1:53:45model and we see some of the name like
- 1:53:47gibberish that we're sort of used to so
- 1:53:49basically the model worked and samples
- 1:53:52uh pretty decent results compared to
- 1:53:54what we were used to so everything is
- 1:53:56the same but of course the big deal is
- 1:53:58that we did not use lots of backward we
- 1:54:00did not use package Auto grad and we
- 1:54:02estimated our gradients ourselves by
- 1:54:04hand
- 1:54:05and so hopefully you're looking at this
- 1:54:06the backward pass of this neural net and
- 1:54:08you're thinking to yourself actually
- 1:54:10that's not too complicated
- 1:54:12um
- 1:54:13each one of these layers is like three
- 1:54:15lines of code or something like that and
- 1:54:17most of it is fairly straightforward
- 1:54:18potentially with the notable exception
- 1:54:20of the batch normalization backward pass
- 1:54:22otherwise it's pretty good okay and
- 1:54:25that's everything I wanted to cover for
- 1:54:26this lecture so hopefully you found this
- 1:54:29interesting and what I liked about it
- 1:54:31honestly is that it gave us a very nice
- 1:54:33diversity of layers to back propagate
- 1:54:34through and
- 1:54:36um I think it gives a pretty nice and
- 1:54:38comprehensive sense of how these
- 1:54:39backward passes are implemented and how
- 1:54:41they work and you'd be able to derive
- 1:54:43them yourself but of course in practice
- 1:54:45you probably don't want to and you want
- 1:54:46to use the pythonograd but hopefully you
- 1:54:49have some intuition about how gradients
- 1:54:51flow backwards through the neural net
- 1:54:52starting at the loss and how they flow
- 1:54:55through all the variables and all the
- 1:54:56intermediate results
- 1:54:58and if you understood a good chunk of it
- 1:55:00and if you have a sense of that then you
- 1:55:02can count yourself as one of these buff
- 1:55:03doji's on the left instead of the uh
- 1:55:06those on the right here now in the next
- 1:55:09lecture we're actually going to go to
- 1:55:10recurrent neural nuts lstms and all the
- 1:55:13other variants of RNs and we're going to
- 1:55:16start to complexify the architecture and
- 1:55:17start to achieve better uh log
- 1:55:19likelihoods and so I'm really looking
- 1:55:21forward to that and I'll see you then
About this transcript
This page contains the full transcript of Building makemore Part 4: Becoming a Backprop Ninja by Andrej Karpathy, generated from the public captions YouTube serves with the video. The transcript has 20,776 words across 3,097 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.