YouTube2Text

Building makemore Part 4: Becoming a Backprop Ninja — Transcript

by Andrej Karpathy · 20,776 words · 3,097 segments · language en · Watch on YouTube

Full transcript

  1. 0:00hi everyone so today we are once again
  2. 0:02continuing our implementation of make
  3. 0:04more now so far we've come up to here
  4. 0:07montalia perceptrons and our neural net
  5. 0:09looked like this and we were
  6. 0:11implementing this over the last few
  7. 0:12lectures
  8. 0:13now I'm sure everyone is very excited to
  9. 0:15go into recurring neural networks and
  10. 0:16all of their variants and how they work
  11. 0:18and the diagrams look cool and it's very
  12. 0:20exciting and interesting and we're going
  13. 0:21to get a better result but unfortunately
  14. 0:23I think we have to remain here for one
  15. 0:25more lecture and the reason for that is
  16. 0:28we've already trained this multilio
  17. 0:30perceptron right and we are getting
  18. 0:31pretty good loss and I think we have a
  19. 0:33pretty decent understanding of the
  20. 0:34architecture and how it works but the
  21. 0:37line of code here that I take an issue
  22. 0:39with is here lost up backward that is we
  23. 0:42are taking a pytorch auto grad and using
  24. 0:45it to calculate all of our gradients
  25. 0:46along the way and I would like to remove
  26. 0:48the use of lost at backward and I would
  27. 0:50like us to write our backward pass
  28. 0:52manually on the level of tensors and I
  29. 0:55think that this is a very useful
  30. 0:56exercise for the following reasons
  31. 0:58I actually have an entire blog post on
  32. 1:00this topic but I'd like to call back
  33. 1:02propagation a leaky abstraction
  34. 1:05and what I mean by that is back
  35. 1:07propagation does doesn't just make your
  36. 1:09neural networks just work magically it's
  37. 1:11not the case they can just Stack Up
  38. 1:12arbitrary Lego blocks of differentiable
  39. 1:14functions and just cross your fingers
  40. 1:16and back propagate and everything is
  41. 1:17great things don't just work
  42. 1:19automatically it is a leaky abstraction
  43. 1:22in the sense that you can shoot yourself
  44. 1:23in the foot if you do not understanding
  45. 1:25its internals it will magically not work
  46. 1:28or not work optimally and you will need
  47. 1:31to understand how it works under the
  48. 1:32hood if you're hoping to debug it and if
  49. 1:34you are hoping to address it in your
  50. 1:36neural nut
  51. 1:37um so this blog post here from a while
  52. 1:39ago goes into some of those examples so
  53. 1:42for example we've already covered them
  54. 1:43some of them already for example the
  55. 1:46flat tails of these functions and how
  56. 1:48you do not want to saturate them too
  57. 1:51much because your gradients will die the
  58. 1:53case of dead neurons which I've already
  59. 1:55covered as well
  60. 1:56the case of exploding or Vanishing
  61. 1:58gradients in the case of repair neural
  62. 2:00networks which we are about to cover
  63. 2:02and then also you will often come across
  64. 2:05some examples in the wild
  65. 2:07this is a snippet that I found uh in a
  66. 2:10random code base on the internet where
  67. 2:11they actually have like a very subtle
  68. 2:13but pretty major bug in their
  69. 2:15implementation and the bug points at the
  70. 2:18fact that the author of this code does
  71. 2:20not actually understand by propagation
  72. 2:21so they're trying to do here is they're
  73. 2:23trying to clip the loss at a certain
  74. 2:25maximum value but actually what they're
  75. 2:27trying to do is they're trying to
  76. 2:28collect the gradients to have a maximum
  77. 2:30value instead of trying to clip the loss
  78. 2:32at a maximum value and
  79. 2:34um indirectly they're basically causing
  80. 2:36some of the outliers to be actually
  81. 2:38ignored because when you clip a loss of
  82. 2:41an outlier you are setting its gradient
  83. 2:43to zero and so have a look through this
  84. 2:46and read through it but there's
  85. 2:48basically a bunch of subtle issues that
  86. 2:50you're going to avoid if you actually
  87. 2:51know what you're doing and that's why I
  88. 2:53don't think it's the case that because
  89. 2:55pytorch or other Frameworks offer
  90. 2:56autograd it is okay for us to ignore how
  91. 2:59it works
  92. 3:00now we've actually already covered
  93. 3:02covered autograd and we wrote micrograd
  94. 3:04but micrograd was an autograd engine
  95. 3:07only on the level of individual scalars
  96. 3:09so the atoms were single individual
  97. 3:11numbers and uh you know I don't think
  98. 3:13it's enough and I'd like us to basically
  99. 3:14think about back propagation on level of
  100. 3:16tensors as well and so in a summary I
  101. 3:19think it's a good exercise I think it is
  102. 3:21very very valuable you're going to
  103. 3:23become better at debugging neural
  104. 3:25networks and making sure that you
  105. 3:27understand what you're doing it is going
  106. 3:28to make everything fully explicit so
  107. 3:30you're not going to be nervous about
  108. 3:31what is hidden away from you and
  109. 3:33basically in general we're going to
  110. 3:34emerge stronger and so let's get into it
  111. 3:37a bit of a fun historical note here is
  112. 3:40that today writing your backward pass by
  113. 3:42hand and manually is not recommended and
  114. 3:43no one does it except for the purposes
  115. 3:45of exercise but about 10 years ago in
  116. 3:48deep learning this was fairly standard
  117. 3:49and in fact pervasive so at the time
  118. 3:52everyone used to write their own
  119. 3:53backward pass by hand manually including
  120. 3:55myself and it's just what you would do
  121. 3:57so we used to ride backward pass by hand
  122. 3:59and now everyone just calls lost that
  123. 4:01backward uh we've lost something I want
  124. 4:04to give you a few examples of this so
  125. 4:07here's a 2006 paper from Jeff Hinton and
  126. 4:11Russell selectinov in science that was
  127. 4:13influential at the time and this was
  128. 4:15training some architectures called
  129. 4:17restricted bolstery machines and
  130. 4:19basically it's an auto encoder trained
  131. 4:22here and this is from roughly 2010 I had
  132. 4:26a library for training researchable
  133. 4:27machines and this was at the time
  134. 4:30written in Matlab so python was not used
  135. 4:32for deep learning pervasively it was all
  136. 4:34Matlab and Matlab was this a scientific
  137. 4:36Computing package that everyone would
  138. 4:39use so we would write Matlab which is
  139. 4:41barely a programming language as well
  140. 4:44but I've had a very convenient tensor
  141. 4:46class and was this a Computing
  142. 4:48environment and you would run here it
  143. 4:49would all run on a CPU of course but you
  144. 4:51would have very nice plots to go with it
  145. 4:53and a built-in debugger and it was
  146. 4:54pretty nice now the code in this package
  147. 4:57in 2010 that I wrote for fitting
  148. 5:00research multiple machines to a large
  149. 5:03extent is recognizable but I wanted to
  150. 5:05show you how you would well I'm creating
  151. 5:07the data in the XY batches I'm
  152. 5:09initializing the neural nut so it's got
  153. 5:11weights and biases just like we're used
  154. 5:13to and then this is the training Loop
  155. 5:15where we actually do the forward pass
  156. 5:17and then here at this time they didn't
  157. 5:19even necessarily use back propagation to
  158. 5:21train neural networks so this in
  159. 5:23particular implements contrastive
  160. 5:25Divergence which estimates a gradient
  161. 5:28and then here we take that gradient and
  162. 5:30use it for a parameter update along the
  163. 5:32lines that we're used to
  164. 5:34um yeah here
  165. 5:36but you can see that basically people
  166. 5:38are meddling with these gradients uh
  167. 5:39directly and inline and themselves uh it
  168. 5:41wasn't that common to use an auto grad
  169. 5:43engine here's one more example from a
  170. 5:45paper of mine from 2014
  171. 5:47um called the fragmented embeddings
  172. 5:49and here what I was doing is I was
  173. 5:51aligning images and text
  174. 5:53um and so it's kind of like a clip if
  175. 5:55you're familiar with it but instead of
  176. 5:56working on the level of entire images
  177. 5:58and entire sentences it was working on
  178. 6:00the level of individual objects and
  179. 6:01little pieces of sentences and I was
  180. 6:03embedding them and then calculating very
  181. 6:05much like a clip-like loss and I dig up
  182. 6:08the code from 2014 of how I implemented
  183. 6:10this and it was already in numpy and
  184. 6:13python
  185. 6:14and here I'm planting the cost function
  186. 6:16and it was standard to implement not
  187. 6:19just the cost but also the backward pass
  188. 6:20manually so here I'm calculating the
  189. 6:23image embeddings sentence embeddings the
  190. 6:26loss function I calculate this course
  191. 6:28this is the loss function and then once
  192. 6:31I have the loss function I do the
  193. 6:32backward pass right here so I backward
  194. 6:34through the loss function and through
  195. 6:36the neural nut and I append
  196. 6:38regularization so everything was done by
  197. 6:41hand manually and you were just right
  198. 6:42out the backward pass and then you would
  199. 6:44use a gradient Checker to make sure that
  200. 6:46your numerical estimate of the gradient
  201. 6:47agrees with the one you calculated
  202. 6:49during back propagation so this was very
  203. 6:51standard for a long time but today of
  204. 6:53course it is standard to use an auto
  205. 6:55grad engine
  206. 6:56um but it was definitely useful and I
  207. 6:58think people sort of understood how
  208. 6:59these neural networks work on a very
  209. 7:01intuitive level and so I think it's a
  210. 7:03good exercise again and this is where we
  211. 7:04want to be okay so just as a reminder
  212. 7:06from our previous lecture this is The
  213. 7:08jupyter Notebook that we implemented at
  214. 7:09the time and
  215. 7:11we're going to keep everything the same
  216. 7:13so we're still going to have a two layer
  217. 7:15multiplayer perceptron with a batch
  218. 7:16normalization layer so the forward pass
  219. 7:18will be basically identical to this
  220. 7:20lecture but here we're going to get rid
  221. 7:22of lost and backward and instead we're
  222. 7:23going to write the backward pass
  223. 7:24manually
  224. 7:26now here's the starter code for this
  225. 7:27lecture we are becoming a back prop
  226. 7:29ninja in this notebook
  227. 7:31and the first few cells here are
  228. 7:34identical to what we are used to so we
  229. 7:36are doing some imports loading the data
  230. 7:37set and processing the data set none of
  231. 7:40this changed
  232. 7:41now here I'm introducing a utility
  233. 7:43function that we're going to use later
  234. 7:44to compare the gradients so in
  235. 7:46particular we are going to have the
  236. 7:47gradients that we estimate manually
  237. 7:49ourselves and we're going to have
  238. 7:50gradients that Pi torch calculates and
  239. 7:53we're going to be checking for
  240. 7:54correctness assuming of course that
  241. 7:55pytorch is correct
  242. 7:58um then here we have the initialization
  243. 8:00that we are quite used to so we have our
  244. 8:03embedding table for the characters the
  245. 8:05first layer second layer and the batch
  246. 8:06normalization in between
  247. 8:08and here's where we create all the
  248. 8:09parameters now you will note that I
  249. 8:11changed the initialization a little bit
  250. 8:13uh to be small numbers so normally you
  251. 8:16would set the biases to be all zero here
  252. 8:18I am setting them to be small random
  253. 8:20numbers and I'm doing this because
  254. 8:22if your variables are initialized to
  255. 8:24exactly zero sometimes what can happen
  256. 8:26is that can mask an incorrect
  257. 8:28implementation of a gradient
  258. 8:30um because uh when everything is zero it
  259. 8:32sort of like simplifies and gives you a
  260. 8:34much simpler expression of the gradient
  261. 8:35than you would otherwise get and so by
  262. 8:37making it small numbers I'm trying to
  263. 8:39unmask those potential errors in these
  264. 8:41calculations
  265. 8:43you also notice that I'm using uh B1 in
  266. 8:46the first layer I'm using a bias despite
  267. 8:48batch normalization right afterwards
  268. 8:50um so this would typically not be what
  269. 8:52you do because we talked about the fact
  270. 8:54that you don't need the bias but I'm
  271. 8:55doing this here just for fun
  272. 8:57um because we're going to have a
  273. 8:58gradient with respect to it and we can
  274. 9:00check that we are still calculating it
  275. 9:01correctly even though this bias is
  276. 9:03asparious
  277. 9:05so here I'm calculating a single batch
  278. 9:07and then here I'm doing a forward pass
  279. 9:10now you'll notice that the forward pass
  280. 9:11is significantly expanded from what we
  281. 9:13are used to here the forward pass was
  282. 9:15just
  283. 9:16um here
  284. 9:17now the reason that the forward pass is
  285. 9:19longer is for two reasons number one
  286. 9:22here we just had an F dot cross entropy
  287. 9:24but here I am bringing back a explicit
  288. 9:26implementation of the loss function
  289. 9:28and number two
  290. 9:29I've broken up the implementation into
  291. 9:32manageable chunks so we have a lot a lot
  292. 9:35more intermediate tensors along the way
  293. 9:37in the forward pass and that's because
  294. 9:38we are about to go backwards and
  295. 9:40calculate the gradients in this back
  296. 9:42propagation from the bottom to the top
  297. 9:45so we're going to go upwards and just
  298. 9:48like we have for example the lock props
  299. 9:49tensor in a forward pass in the backward
  300. 9:51pass we're going to have a d-lock probes
  301. 9:53which is going to store the derivative
  302. 9:55of the loss with respect to the lock
  303. 9:56props tensor and so we're going to be
  304. 9:58prepending D to every one of these
  305. 10:00tensors and calculating it along the way
  306. 10:02of this back propagation
  307. 10:04so as an example we have a b and raw
  308. 10:07here we're going to be calculating a DB
  309. 10:09in raw so here I'm telling pytorch that
  310. 10:12we want to retain the grad of all these
  311. 10:14intermediate values because here in
  312. 10:16exercise one we're going to calculate
  313. 10:18the backward pass so we're going to
  314. 10:20calculate all these D values D variables
  315. 10:22and use the CNP function I've introduced
  316. 10:25above to check our correctness with
  317. 10:26respect to what pi torch is telling us
  318. 10:29this is going to be exercise one uh
  319. 10:31where we sort of back propagate through
  320. 10:32this entire graph
  321. 10:34now just to give you a very quick
  322. 10:36preview of what's going to happen in
  323. 10:37exercise two and below here we have
  324. 10:40fully broken up the loss and back
  325. 10:43propagated through it manually in all
  326. 10:45the little Atomic pieces that make it up
  327. 10:47but here we're going to collapse the
  328. 10:49laws into a single cross-entropy call
  329. 10:50and instead we're going to analytically
  330. 10:53derive using math and paper and pencil
  331. 10:56the gradient of the loss with respect to
  332. 10:59the logits and instead of back
  333. 11:01propagating through all of its little
  334. 11:02chunks one at a time we're just going to
  335. 11:04analytically derive what that gradient
  336. 11:05is and we're going to implement that
  337. 11:07which is much more efficient as we'll
  338. 11:09see in the in a bit
  339. 11:10then we're going to do the exact same
  340. 11:12thing for patch normalization so instead
  341. 11:14of breaking up bass drum into all the
  342. 11:16old tiny components we're going to use
  343. 11:18uh pen and paper and Mathematics and
  344. 11:20calculus to derive the gradient through
  345. 11:22the bachelor Bachelor layer so we're
  346. 11:25going to calculate the backward
  347. 11:27passthrough bathroom layer in a much
  348. 11:28more efficient expression instead of
  349. 11:30backward propagating through all of its
  350. 11:31little pieces independently
  351. 11:33so there's going to be exercise three
  352. 11:36and then in exercise four we're going to
  353. 11:38put it all together and this is the full
  354. 11:40code of training this two layer MLP and
  355. 11:42we're going to basically insert our
  356. 11:44manual back prop and we're going to take
  357. 11:46out lost it backward and you will
  358. 11:48basically see that you can get all the
  359. 11:50same results using fully your own code
  360. 11:53and the only thing we're using from
  361. 11:55pytorch is the torch.tensor to make the
  362. 11:59calculations efficient but otherwise you
  363. 12:01will understand fully what it means to
  364. 12:03forward and backward and neural net and
  365. 12:04train it and I think that'll be awesome
  366. 12:06so let's get to it
  367. 12:08okay so I read all the cells of this
  368. 12:10notebook all the way up to here and I'm
  369. 12:13going to erase this and I'm going to
  370. 12:14start implementing backward pass
  371. 12:15starting with d lock problems so we want
  372. 12:18to understand what should go here to
  373. 12:20calculate the gradient of the loss with
  374. 12:22respect to all the elements of the log
  375. 12:23props tensor
  376. 12:25now I'm going to give away the answer
  377. 12:26here but I wanted to put a quick note
  378. 12:28here that I think would be most
  379. 12:30pedagogically useful for you is to
  380. 12:32actually go into the description of this
  381. 12:34video and find the link to this Jupiter
  382. 12:36notebook you can find it both on GitHub
  383. 12:38but you can also find Google collab with
  384. 12:40it so you don't have to install anything
  385. 12:41you'll just go to a website on Google
  386. 12:43collab and you can try to implement
  387. 12:45these derivatives or gradients yourself
  388. 12:47and then if you are not able to come to
  389. 12:50my video and see me do it and so work in
  390. 12:53Tandem and try it first yourself and
  391. 12:55then see me give away the answer and I
  392. 12:57think that'll be most valuable to you
  393. 12:59and that's how I recommend you go
  394. 13:00through this lecture
  395. 13:01so we are starting here with d-log props
  396. 13:03now d-lock props will hold the
  397. 13:06derivative of the loss with respect to
  398. 13:08all the elements of log props
  399. 13:11what is inside log blobs the shape of
  400. 13:13this is 32 by 27. so it's not going to
  401. 13:18surprise you that D log props should
  402. 13:19also be an array of size 32 by 27
  403. 13:21because we want the derivative loss with
  404. 13:23respect to all of its elements so the
  405. 13:26sizes of those are always going to be
  406. 13:27equal
  407. 13:29now how how does log props influence the
  408. 13:33loss okay loss is negative block probes
  409. 13:36indexed with range of N and YB and then
  410. 13:40the mean of that now just as a reminder
  411. 13:42YB is just a basically an array of all
  412. 13:47the correct indices
  413. 13:51um so what we're doing here is we're
  414. 13:52taking the lock props array of size 32
  415. 13:54by 27.
  416. 13:57right
  417. 13:58and then we are going in every single
  418. 14:00row and in each row we are plugging
  419. 14:03plucking out the index eight and then 14
  420. 14:06and 15 and so on so we're going down the
  421. 14:07rows that's the iterator range of N and
  422. 14:10then we are always plucking out the
  423. 14:12index of the column specified by this
  424. 14:15tensor YB so in the zeroth row we are
  425. 14:17taking the eighth column in the first
  426. 14:20row we're taking the 14th column Etc and
  427. 14:23so log props at this plugs out
  428. 14:26all those
  429. 14:28log probabilities of the correct next
  430. 14:30character in a sequence
  431. 14:32so that's what that does and the shape
  432. 14:34of this or the size of it is of course
  433. 14:3632 because our batch size is 32.
  434. 14:40so these elements get plugged out and
  435. 14:43then their mean and the negative of that
  436. 14:45becomes loss
  437. 14:47so I always like to work with simpler
  438. 14:49examples to understand the numerical
  439. 14:52form of derivative what's going on here
  440. 14:55is once we've plucked out these examples
  441. 14:58um we're taking the mean and then the
  442. 15:00negative so the loss basically
  443. 15:02I can write it this way is the negative
  444. 15:04of say a plus b plus c
  445. 15:07and the mean of those three numbers
  446. 15:09would be say negative would divide three
  447. 15:11that would be how we achieve the mean of
  448. 15:13three numbers ABC although we actually
  449. 15:15have 32 numbers here
  450. 15:16and so what is basically the loss by say
  451. 15:20like d a right
  452. 15:22well if we simplify this expression
  453. 15:24mathematically this is negative one over
  454. 15:26three of A and negative plus negative
  455. 15:28one over three of B
  456. 15:30plus negative 1 over 3 of c and so what
  457. 15:33is D loss by D A it's just negative one
  458. 15:35over three
  459. 15:36and so you can see that if we don't just
  460. 15:38have a b and c but we have 32 numbers
  461. 15:40then D loss by D
  462. 15:43um you know every one of those numbers
  463. 15:45is going to be one over N More generally
  464. 15:47because n is the um the size of the
  465. 15:50batch 32 in this case
  466. 15:53so D loss by
  467. 15:55um D Lock probs is negative 1 over n
  468. 15:59in all these places
  469. 16:02now what about the other elements inside
  470. 16:04lock problems because lock props is
  471. 16:05large array you see that lock problems
  472. 16:07at shape is 32 by 27. but only 32 of
  473. 16:11them participate in the loss calculation
  474. 16:13so what's the derivative of all the
  475. 16:15other most of the elements that do not
  476. 16:18get plucked out here
  477. 16:20while their loss intuitively is zero
  478. 16:22sorry they're gradient intuitively is
  479. 16:24zero and that's because they did not
  480. 16:25participate in the loss
  481. 16:27so most of these numbers inside this
  482. 16:29tensor does not feed into the loss and
  483. 16:32so if we were to change these numbers
  484. 16:33then the loss doesn't change which is
  485. 16:36the equivalent of way of saying that the
  486. 16:38derivative of the loss with respect to
  487. 16:39them is zero they don't impact it
  488. 16:43so here's a way to implement this
  489. 16:45derivative then we start out with
  490. 16:47torch.zeros of shape 32 by 27 or let's
  491. 16:50just say instead of doing this because
  492. 16:52we don't want to hard code numbers let's
  493. 16:54do torch.zeros like
  494. 16:57block probs so basically this is going
  495. 16:59to create an array of zeros exactly in
  496. 17:00the shape of log probs
  497. 17:02and then we need to set the derivative
  498. 17:05of negative 1 over n inside exactly
  499. 17:07these locations so here's what we can do
  500. 17:09the lock props indexed in The Identical
  501. 17:12way
  502. 17:14will be just set to negative one over
  503. 17:16zero divide n
  504. 17:19right just like we derived here
  505. 17:22so now let me erase all this reasoning
  506. 17:25and then this is the candidate
  507. 17:27derivative for D log props let's
  508. 17:29uncomment the first line and check that
  509. 17:31this is correct
  510. 17:34okay so CMP ran and let's go back to CMP
  511. 17:39and you see that what it's doing is it's
  512. 17:41calculating if
  513. 17:42the calculated value by us which is DT
  514. 17:46is exactly equal to T dot grad as
  515. 17:48calculated by pi torch and then this is
  516. 17:51making sure that all the elements are
  517. 17:52exactly equal and then converting this
  518. 17:54to a single Boolean value because we
  519. 17:57don't want the Boolean tensor we just
  520. 17:58want to Boolean value
  521. 18:00and then here we are making sure that
  522. 18:02okay if they're not exactly equal maybe
  523. 18:04they are approximately equal because of
  524. 18:06some floating Point issues but they're
  525. 18:07very very close
  526. 18:09so here we are using torch.allclose
  527. 18:10which has a little bit of a wiggle
  528. 18:13available because sometimes you can get
  529. 18:15very very close but if you use a
  530. 18:17slightly different calculation because a
  531. 18:19floating Point arithmetic you can get a
  532. 18:22slightly different result so this is
  533. 18:24checking if you get an approximately
  534. 18:25close result
  535. 18:27and then here we are checking the
  536. 18:28maximum uh basically the value that has
  537. 18:31the highest difference and what is the
  538. 18:34difference in the absolute value
  539. 18:35difference between those two and so we
  540. 18:37are printing whether we have an exact
  541. 18:39equality an approximate equality and
  542. 18:42what is the largest difference
  543. 18:45and so here
  544. 18:46we see that we actually have exact
  545. 18:48equality and so therefore of course we
  546. 18:50also have an approximate equality and
  547. 18:52the maximum difference is exactly zero
  548. 18:54so basically our d-log props is exactly
  549. 18:57equal to what pytors calculated to be
  550. 19:00lockprops.grad in its back propagation
  551. 19:03so so far we're working pretty well okay
  552. 19:06so let's now continue our back
  553. 19:07propagation
  554. 19:08we have that lock props depends on
  555. 19:10probes through a log
  556. 19:12so all the elements of probes are being
  557. 19:14element wise applied log to
  558. 19:17now if we want deep props then then
  559. 19:19remember your micrograph training
  560. 19:22we have like a log node it takes in
  561. 19:24probs and creates log probs and the
  562. 19:27props will be the local derivative of
  563. 19:30that individual Operation Log times the
  564. 19:33derivative loss with respect to its
  565. 19:34output which in this case is D log props
  566. 19:37so what is the local derivative of this
  567. 19:39operation well we are taking log element
  568. 19:41wise and we can come here and we can see
  569. 19:43well from alpha is your friend that d by
  570. 19:45DX of log of x is just simply one of our
  571. 19:47X
  572. 19:48so therefore in this case X is problems
  573. 19:51so we have d by DX is one over X which
  574. 19:54is one of our probes and then this is
  575. 19:56the local derivative and then times we
  576. 19:58want to chain it
  577. 20:00so this is chain rule
  578. 20:01times do log props
  579. 20:03let me uncomment this and let me run the
  580. 20:06cell in place and we see that the
  581. 20:08derivative of props as we calculated
  582. 20:10here is exactly correct
  583. 20:12and so notice here how this works probes
  584. 20:15that are props is going to be inverted
  585. 20:18and then element was multiplied here
  586. 20:20so if your probes is very very close to
  587. 20:23one that means you are your network is
  588. 20:25currently predicting the character
  589. 20:26correctly then this will become one over
  590. 20:28one and D log probes just gets passed
  591. 20:30through
  592. 20:31but if your probabilities are
  593. 20:33incorrectly assigned so if the correct
  594. 20:35character here is getting a very low
  595. 20:37probability then 1.0 dividing by it will
  596. 20:41boost this
  597. 20:43and then multiply by the log props so
  598. 20:45basically what this line is doing
  599. 20:46intuitively is it's taking the examples
  600. 20:49that have a very low probability
  601. 20:50currently assigned and it's boosting
  602. 20:52their gradient uh you can you can look
  603. 20:55at it that way next up is Count some imp
  604. 20:59so we want the river of this now let me
  605. 21:02just pause here and kind of introduce
  606. 21:05What's Happening Here in general because
  607. 21:06I know it's a little bit confusing we
  608. 21:08have the locusts that come out of the
  609. 21:09neural nut here what I'm doing is I'm
  610. 21:11finding the maximum in each row and I'm
  611. 21:15subtracting it for the purposes of
  612. 21:16numerical stability and we talked about
  613. 21:18how if you do not do this you run
  614. 21:20numerical issues if some of the logits
  615. 21:22take on two large values because we end
  616. 21:24up exponentiating them
  617. 21:26so this is done just for safety
  618. 21:28numerically then here's the
  619. 21:30exponentiation of all the sort of like
  620. 21:32logits to create our accounts and then
  621. 21:35we want to take the some of these counts
  622. 21:38and normalize so that all of the probes
  623. 21:40sum to one
  624. 21:41now here instead of using one over count
  625. 21:43sum I use uh raised to the power of
  626. 21:46negative one mathematically they are
  627. 21:47identical I just found that there's
  628. 21:49something wrong with the pytorch
  629. 21:50implementation of the backward pass of
  630. 21:52division
  631. 21:53um and it gives like a real result but
  632. 21:55that doesn't happen for star star native
  633. 21:58one that's why I'm using this formula
  634. 21:59instead but basically all that's
  635. 22:01happening here is we got the logits
  636. 22:04we're going to exponentiate all of them
  637. 22:05and want to normalize the counts to
  638. 22:07create our probabilities it's just that
  639. 22:09it's happening across multiple lines
  640. 22:12so now
  641. 22:14here
  642. 22:17we want to First Take the derivative we
  643. 22:20want to back propagate into account
  644. 22:21sumiv and then into counts as well
  645. 22:24so what should be the count sum M now we
  646. 22:28actually have to be careful here because
  647. 22:29we have to scrutinize and be careful
  648. 22:32with the shapes so counts that shape and
  649. 22:35then count some inverse shape
  650. 22:39are different
  651. 22:40so in particular counts as 32 by 27 but
  652. 22:43this count sum m is 32 by 1. and so in
  653. 22:47this multiplication here we also have an
  654. 22:49implicit broadcasting that pytorch will
  655. 22:52do because it needs to take this column
  656. 22:53tensor of 32 numbers and replicate it
  657. 22:55horizontally 27 times to align these two
  658. 22:58tensors so it can do an element twice
  659. 23:00multiply
  660. 23:01so really what this looks like is the
  661. 23:03following using a toy example again
  662. 23:06what we really have here is just props
  663. 23:08is counts times conservative so it's a C
  664. 23:10equals a times B
  665. 23:11but a is 3 by 3 and b is just three by
  666. 23:15one a column tensor and so pytorch
  667. 23:17internally replicated this elements of B
  668. 23:19and it did that across all the columns
  669. 23:22so for example B1 which is the first
  670. 23:24element of B would be replicated here
  671. 23:26across all the columns in this
  672. 23:27multiplication
  673. 23:29and now we're trying to back propagate
  674. 23:31through this operation to count some m
  675. 23:34so when we're calculating this
  676. 23:35derivative
  677. 23:37it's important to realize that these two
  678. 23:39this looks like a single operation but
  679. 23:41actually is two operations applied
  680. 23:44sequentially the first operation that
  681. 23:46pytorch did is it took this column
  682. 23:48tensor and replicated it across all the
  683. 23:52um across all the columns basically 27
  684. 23:54times so that's the first operation it's
  685. 23:55a replication and then the second
  686. 23:57operation is the multiplication so let's
  687. 23:59first background through the
  688. 24:01multiplication
  689. 24:02if these two arrays are of the same size
  690. 24:05and we just have a and b of both of them
  691. 24:08three by three then how do we mult how
  692. 24:11do we back propagate through a
  693. 24:12multiplication so if we just have
  694. 24:14scalars and not tensors then if you have
  695. 24:16C equals a times B then what is uh the
  696. 24:19order of the of C with respect to B well
  697. 24:21it's just a and so that's the local
  698. 24:23derivative
  699. 24:24so here in our case undoing the
  700. 24:27multiplication and back propagating
  701. 24:29through just the multiplication itself
  702. 24:30which is element wise is going to be the
  703. 24:33local derivative which in this case is
  704. 24:36simply counts because counts is the a
  705. 24:40so this is the local derivative and then
  706. 24:42times because the chain rule D props
  707. 24:46so this here is the derivative or the
  708. 24:48gradient but with respect to replicated
  709. 24:50B
  710. 24:52but we don't have a replicated B we just
  711. 24:54have a single B column so how do we now
  712. 24:56back propagate through the replication
  713. 24:59and intuitively this B1 is the same
  714. 25:02variable and it's just reused multiple
  715. 25:04times
  716. 25:04and so you can look at it
  717. 25:07as being equivalent to a case we've
  718. 25:09encountered in micrograd
  719. 25:10and so here I'm just pulling out a
  720. 25:12random graph we used in micrograd we had
  721. 25:14an example where a single node
  722. 25:17has its output feeding into two branches
  723. 25:19of basically the graph until the last
  724. 25:22function and we're talking about how the
  725. 25:25correct thing to do in the backward pass
  726. 25:26is we need to sum all the gradients that
  727. 25:29arrive at any one node so across these
  728. 25:31different branches the gradients would
  729. 25:33sum
  730. 25:34so if a node is used multiple times the
  731. 25:37gradients for all of its uses sum during
  732. 25:39back propagation
  733. 25:41so here B1 is used multiple times in all
  734. 25:44these columns and therefore the right
  735. 25:45thing to do here is to sum
  736. 25:48horizontally across all the rows so I'm
  737. 25:51going to sum in
  738. 25:52Dimension one but we want to retain this
  739. 25:55Dimension so that the uh so that counts
  740. 25:58some end and its gradient are going to
  741. 26:00be exactly the same shape so we want to
  742. 26:02make sure that we keep them as true so
  743. 26:04we don't lose this dimension and this
  744. 26:07will make the count sum M be exactly
  745. 26:08shape 32 by 1.
  746. 26:11so revealing this comparison as well and
  747. 26:14running this we see that we get an exact
  748. 26:17match
  749. 26:18so this derivative is exactly correct
  750. 26:22and let me erase
  751. 26:24this now let's also back propagate into
  752. 26:26counts which is the other variable here
  753. 26:29to create probes so from props to count
  754. 26:32some INF we just did that let's go into
  755. 26:33counts as well
  756. 26:35so decounts will be
  757. 26:39the chances are a so DC by d a is just B
  758. 26:43so therefore it's count summative
  759. 26:47um and then times chain rule the props
  760. 26:51now councilman is three two by One D
  761. 26:54probs is 32 by 27.
  762. 26:57so
  763. 26:59um those will broadcast fine and will
  764. 27:02give us decounts there's no additional
  765. 27:04summation required here
  766. 27:06um there will be a broadcasting that
  767. 27:08happens in this multiply here because
  768. 27:11count some M needs to be replicated
  769. 27:12again to correctly multiply D props but
  770. 27:16that's going to give the correct result
  771. 27:18so as far as the single operation is
  772. 27:20concerned so we back probably go from
  773. 27:23props to counts but we can't actually
  774. 27:25check the derivative counts uh I have it
  775. 27:29much later on and the reason for that is
  776. 27:31because count sum in depends on counts
  777. 27:34and so there's a second Branch here that
  778. 27:36we have to finish because can't summon
  779. 27:38back propagates into account sum and
  780. 27:40count sum will buy properly into counts
  781. 27:42and so counts is a node that is being
  782. 27:44used twice it's used right here in two
  783. 27:46props and it goes through this other
  784. 27:48Branch through count summative
  785. 27:50so even though we've calculated the
  786. 27:52first contribution of it we still have
  787. 27:54to calculate the second contribution of
  788. 27:55it later
  789. 27:57okay so we're continuing with this
  790. 27:58Branch we have the derivative for count
  791. 28:00sum if now we want the derivative of
  792. 28:02count sum so D count sum equals what is
  793. 28:05the local derivative of this operation
  794. 28:07so this is basically an element wise one
  795. 28:09over counts sum
  796. 28:11so count sum raised to the power of
  797. 28:13negative one is the same as one over
  798. 28:15count sum if we go to all from alpha we
  799. 28:17see that x to the negative one D by D by
  800. 28:20D by DX of it is basically Negative X to
  801. 28:23the negative 2. right one negative one
  802. 28:25over squared is the same as Negative X
  803. 28:27to the negative two
  804. 28:29so D count sum here will be local
  805. 28:32derivative is going to be negative
  806. 28:35um
  807. 28:36counts sum to the negative two that's
  808. 28:39the local derivative times chain rule
  809. 28:41which is D count sum in
  810. 28:46so that's D count sum
  811. 28:49let's uncomment this and check that I am
  812. 28:51correct okay so we have perfect equality
  813. 28:55and there's no sketchiness going on here
  814. 28:58with any shapes because these are of the
  815. 28:59same shape okay next up we want to back
  816. 29:02propagate through this line we have that
  817. 29:04count sum it's count.sum along the rows
  818. 29:07so I wrote out
  819. 29:09um some help here we have to keep in
  820. 29:11mind that counts of course is 32 by 27
  821. 29:13and count sum is 32 by 1. so in this
  822. 29:17back propagation we need to take this
  823. 29:19column of derivatives and transform it
  824. 29:22into a array of derivatives
  825. 29:24two-dimensional array
  826. 29:26so what is this operation doing we're
  827. 29:28taking in some kind of an input like say
  828. 29:31a three by three Matrix a and we are
  829. 29:32summing up the rows into a column tells
  830. 29:36her B1 b2b3 that is basically this
  831. 29:39so now we have the derivatives of the
  832. 29:41loss with respect to B all the elements
  833. 29:44of B
  834. 29:45and now we want to derivative loss with
  835. 29:47respect to all these little A's
  836. 29:49so how do the B's depend on the ace is
  837. 29:52basically what we're after what is the
  838. 29:54local derivative of this operation
  839. 29:56well we can see here that B1 only
  840. 29:58depends on these elements here the
  841. 30:01derivative of B1 with respect to all of
  842. 30:03these elements down here is zero but for
  843. 30:06these elements here like a11 a12 Etc the
  844. 30:09local derivative is one right so DB 1 by
  845. 30:13D A 1 1 for example is one so it's one
  846. 30:16one and one
  847. 30:18so when we have the derivative of loss
  848. 30:19with respect to B1
  849. 30:21did a local derivative of B1 with
  850. 30:23respect to these inputs is zeros here
  851. 30:25but it's one on these guys
  852. 30:27so in the chain rule
  853. 30:29we have the local derivative uh times
  854. 30:32sort of the derivative of B1 and so
  855. 30:35because the local derivative is one on
  856. 30:37these three elements the look of them
  857. 30:39are multiplying the derivative of B1
  858. 30:41will just be the derivative of B1 and so
  859. 30:45you can look at it as a router basically
  860. 30:47an addition is a router of gradient
  861. 30:50whatever gradient comes from above it
  862. 30:52just gets routed equally to all the
  863. 30:53elements that participate in that
  864. 30:55addition
  865. 30:56so in this case the derivative of B1
  866. 30:58will just flow equally to the derivative
  867. 31:00of a11 a12 and a13
  868. 31:03. so if we have a derivative of all the
  869. 31:05elements of B and in this column tensor
  870. 31:07which is D counts sum that we've
  871. 31:10calculated just now
  872. 31:11we basically see that what that amounts
  873. 31:14to is all of these are now flowing to
  874. 31:17all these elements of a and they're
  875. 31:19doing that horizontally
  876. 31:21so basically what we want is we want to
  877. 31:22take the decount sum of size 30 by 1 and
  878. 31:26we just want to replicate it 27 times
  879. 31:28horizontally to create 32 by 27 array
  880. 31:32so there's many ways to implement this
  881. 31:33operation you could of course just
  882. 31:35replicate the tensor but I think maybe
  883. 31:37one clean one is that the counts is
  884. 31:40simply torch dot once like
  885. 31:43so just an two-dimensional arrays of
  886. 31:45ones in the shape of counts so 32 by 27
  887. 31:49times D counts sum so this way we're
  888. 31:53letting the broadcasting here basically
  889. 31:56implement the replication you can look
  890. 31:58at it that way
  891. 31:59but then we have to also be careful
  892. 32:02because decounts was already calculated
  893. 32:05we calculated earlier here and that was
  894. 32:08just the first branch and we're now
  895. 32:09finishing the second Branch so we need
  896. 32:11to make sure that these gradients add so
  897. 32:13plus equals
  898. 32:14and then here
  899. 32:16um let's comment out the comparison and
  900. 32:20let's make sure crossing fingers that we
  901. 32:23have the correct result so pytorch
  902. 32:25agrees with us on this gradient as well
  903. 32:28okay hopefully we're getting a hang of
  904. 32:29this now counts as an element-wise X of
  905. 32:32Norm legits so now we want D Norm logits
  906. 32:36and because it's an element price
  907. 32:38operation everything is very simple what
  908. 32:40is the local derivative of e to the X
  909. 32:41it's famously just e to the x so this is
  910. 32:45the local derivative
  911. 32:48that is the local derivative now we
  912. 32:50already calculated it and it's inside
  913. 32:51counts so we may as well potentially
  914. 32:53just reuse counts that is the local
  915. 32:55derivative
  916. 32:56times uh D counts
  917. 33:01funny as that looks constant decount is
  918. 33:04derivative on the normal objects and now
  919. 33:07let's erase this and let's verify and it
  920. 33:10looks good
  921. 33:12so that's uh normal agents
  922. 33:14okay so we are here on this line now the
  923. 33:17normal objects
  924. 33:18we have that and we're trying to
  925. 33:20calculate the logits and deloget Maxes
  926. 33:22so back propagating through this line
  927. 33:25now we have to be careful here because
  928. 33:26the shapes again are not the same and so
  929. 33:29there's an implicit broadcasting
  930. 33:30Happening Here
  931. 33:32so normal jits has this shape 32 by 27
  932. 33:34logist does as well but logit Maxis is
  933. 33:37only 32 by one so there's a broadcasting
  934. 33:40here in the minus
  935. 33:42now here I try to sort of write out a
  936. 33:45two example again we basically have that
  937. 33:48this is our C equals a minus B
  938. 33:50and we see that because of the shape
  939. 33:52these are three by three but this one is
  940. 33:54just a column
  941. 33:55and so for example every element of C we
  942. 33:57have to look at how it uh came to be and
  943. 34:00every element of C is just the
  944. 34:01corresponding element of a minus uh
  945. 34:04basically that associated b
  946. 34:08so it's very clear now that the
  947. 34:10derivatives of every one of these c's
  948. 34:13with respect to their inputs are one for
  949. 34:16the corresponding a
  950. 34:18and it's a negative one for the
  951. 34:20corresponding B
  952. 34:22and so therefore
  953. 34:24um
  954. 34:25the derivatives on the C will flow
  955. 34:27equally to the corresponding Ace and
  956. 34:30then also to the corresponding base but
  957. 34:33then in addition to that the B's are
  958. 34:35broadcast so we'll have to do the
  959. 34:36additional sum just like we did before
  960. 34:39and of course the derivatives for B's
  961. 34:41will undergo a minus because the local
  962. 34:43derivative here is uh negative one
  963. 34:46so DC three two by D B3 is negative one
  964. 34:50so let's just Implement that basically
  965. 34:52delugits will be uh exactly copying the
  966. 34:56derivative on normal objects
  967. 34:58so
  968. 34:59delugits equals the norm logits and I'll
  969. 35:03do a DOT clone for safety so we're just
  970. 35:05making a copy
  971. 35:06and then we have that the loaded Maxis
  972. 35:09will be the negative of the non-legits
  973. 35:13because of the negative sign
  974. 35:15and then we have to be careful because
  975. 35:17logic Maxis is a column
  976. 35:20and so just like we saw before because
  977. 35:23we keep replicating the same elements
  978. 35:26across all the columns
  979. 35:28then in the backward pass because we
  980. 35:31keep reusing this these are all just
  981. 35:33like separate branches of use of that
  982. 35:35one variable and so therefore we have to
  983. 35:37do a Sum along one would keep them
  984. 35:39equals true so that we don't destroy
  985. 35:42this dimension
  986. 35:43and then the logic Maxes will be the
  987. 35:45same shape now we have to be careful
  988. 35:47because this deloaches is not the final
  989. 35:49deloaches and that's because not only do
  990. 35:52we get gradient signal into logits
  991. 35:54through here but the logic Maxes as a
  992. 35:56function of logits and that's a second
  993. 35:58Branch into logits so this is not yet
  994. 36:01our final derivative for logits we will
  995. 36:03come back later for the second branch
  996. 36:05for now the logic Maxis is the final
  997. 36:07derivative so let me uncomment this CMP
  998. 36:10here and let's just run this
  999. 36:12and logit Maxes hit by torch agrees with
  1000. 36:15us
  1001. 36:16so that was the derivative into through
  1002. 36:19this line
  1003. 36:21now before we move on I want to pause
  1004. 36:22here briefly and I want to look at these
  1005. 36:24logic Maxes and especially their
  1006. 36:26gradients
  1007. 36:27we've talked previously in the previous
  1008. 36:28lecture that the only reason we're doing
  1009. 36:31this is for the numerical stability of
  1010. 36:33the softmax that we are implementing
  1011. 36:34here and we talked about how if you take
  1012. 36:37these logents for any one of these
  1013. 36:39examples so one row of this logit's
  1014. 36:41tensor if you add or subtract any value
  1015. 36:44equally to all the elements then the
  1016. 36:47value of the probes will be unchanged
  1017. 36:49you're not changing soft Max the only
  1018. 36:51thing that this is doing is it's making
  1019. 36:53sure that X doesn't overflow and the
  1020. 36:55reason we're using a Max is because then
  1021. 36:57we are guaranteed that each row of
  1022. 36:58logits the highest number is zero and so
  1023. 37:01this will be safe
  1024. 37:03and so
  1025. 37:05um
  1026. 37:06basically what that has repercussions
  1027. 37:09if it is the case that changing logit
  1028. 37:11Maxis does not change the props and
  1029. 37:13therefore there's not change the loss
  1030. 37:15then the gradient on logic masses should
  1031. 37:17be zero right because saying those two
  1032. 37:20things is the same
  1033. 37:21so indeed we hope that this is very very
  1034. 37:23small numbers so indeed we hope this is
  1035. 37:25zero now because of floating Point uh
  1036. 37:28sort of wonkiness
  1037. 37:30um this doesn't come out exactly zero
  1038. 37:31only in some of the rows it does but we
  1039. 37:33get extremely small values like one e
  1040. 37:35negative nine or ten and so this is
  1041. 37:37telling us that the values of loaded
  1042. 37:39Maxes are not impacting the loss as they
  1043. 37:42shouldn't
  1044. 37:43it feels kind of weird to back propagate
  1045. 37:44through this branch honestly because
  1046. 37:48if you have any implementation of like f
  1047. 37:50dot cross entropy and pytorch and you
  1048. 37:52you block together all these elements
  1049. 37:54and you're not doing the back
  1050. 37:54propagation piece by piece then you
  1051. 37:57would probably assume that the
  1052. 37:59derivative through here is exactly zero
  1053. 38:01uh so you would be sort of
  1054. 38:03um skipping this branch because it's
  1055. 38:07only done for numerical stability but
  1056. 38:09it's interesting to see that even if you
  1057. 38:10break up everything into the full atoms
  1058. 38:13and you still do the computation as
  1059. 38:14you'd like with respect to numerical
  1060. 38:16stability uh the correct thing happens
  1061. 38:17and you still get a very very small
  1062. 38:20gradients here
  1063. 38:21um basically reflecting the fact that
  1064. 38:23the values of these do not matter with
  1065. 38:26respect to the final loss
  1066. 38:27okay so let's now continue back
  1067. 38:29propagation through this line here we've
  1068. 38:31just calculated the logit Maxis and now
  1069. 38:33we want to back prop into logits through
  1070. 38:35this second branch
  1071. 38:36now here of course we took legits and we
  1072. 38:38took the max along all the rows and then
  1073. 38:41we looked at its values here now the way
  1074. 38:43this works is that in pytorch
  1075. 38:47this thing here
  1076. 38:49the max returns both the values and it
  1077. 38:52Returns the indices at which those
  1078. 38:53values to count the maximum value
  1079. 38:55now in the forward pass we only used
  1080. 38:57values because that's all we needed but
  1081. 39:00in the backward pass it's extremely
  1082. 39:01useful to know about where those maximum
  1083. 39:04values occurred and we have the indices
  1084. 39:06at which they occurred and this will of
  1085. 39:08course helps us to help us do the back
  1086. 39:10propagation because what should the
  1087. 39:12backward pass be here in this case we
  1088. 39:14have the largest tensor which is 32 by
  1089. 39:1627 and in each row we find the maximum
  1090. 39:18value and then that value gets plucked
  1091. 39:20out into loaded Maxis and so intuitively
  1092. 39:24um basically the derivative flowing
  1093. 39:27through here then should be one
  1094. 39:31times the look of derivatives is 1 for
  1095. 39:34the appropriate entry that was plucked
  1096. 39:35out
  1097. 39:36and then times the global derivative of
  1098. 39:39the logic axis
  1099. 39:40so really what we're doing here if you
  1100. 39:42think through it is we need to take the
  1101. 39:44deloachet Maxis and we need to scatter
  1102. 39:46it to the correct positions in these
  1103. 39:50logits from where the maximum values
  1104. 39:52came
  1105. 39:53and so
  1106. 39:54um
  1107. 39:56I came up with one line of code sort of
  1108. 39:58that does that let me just erase a bunch
  1109. 39:59of stuff here so the line of uh you
  1110. 40:02could do it kind of very similar to what
  1111. 40:03we've done here where we create a zeros
  1112. 40:05and then we populate uh the correct
  1113. 40:07elements uh so we use the indices here
  1114. 40:10and we would set them to be one but you
  1115. 40:13can also use one hot
  1116. 40:15so F dot one hot and then I'm taking the
  1117. 40:18lowest of Max over the First Dimension
  1118. 40:21dot indices and I'm telling uh pytorch
  1119. 40:24that the dimension of every one of these
  1120. 40:27tensors should be
  1121. 40:29um
  1122. 40:2927 and so what this is going to do
  1123. 40:33is okay I apologize this is crazy filthy
  1124. 40:37that I am sure of this
  1125. 40:39it's really just a an array of where the
  1126. 40:41Maxes came from in each row and that
  1127. 40:44element is one and the all the other
  1128. 40:45elements are zero so it's a one-half
  1129. 40:47Vector in each row and these indices are
  1130. 40:50now populating a single one in the
  1131. 40:53proper place
  1132. 40:54and then what I'm doing here is I'm
  1133. 40:56multiplying by the logit Maxis and keep
  1134. 40:58in mind that this is a column
  1135. 41:01of 32 by 1. and so when I'm doing this
  1136. 41:05times the logic Maxis the logic Maxes
  1137. 41:08will broadcast and that column will you
  1138. 41:10know get replicated and in an element
  1139. 41:12wise multiply will ensure that each of
  1140. 41:15these just gets routed to whichever one
  1141. 41:17of these bits is turned on
  1142. 41:19and so that's another way to implement
  1143. 41:21uh this kind of a this kind of a
  1144. 41:23operation and both of these can be used
  1145. 41:26I just thought I would show an
  1146. 41:28equivalent way to do it and I'm using
  1147. 41:30plus equals because we already
  1148. 41:31calculated the logits here and this is
  1149. 41:33not the second branch
  1150. 41:35so let's
  1151. 41:37look at logits and make sure that this
  1152. 41:39is correct
  1153. 41:40and we see that we have exactly the
  1154. 41:42correct answer
  1155. 41:44next up we want to continue with logits
  1156. 41:46here that is an outcome of a matrix
  1157. 41:49multiplication and a bias offset in this
  1158. 41:51linear layer
  1159. 41:53so I've printed out the shapes of all
  1160. 41:56these intermediate tensors we see that
  1161. 41:58logits is of course 32 by 27 as we've
  1162. 42:00just seen
  1163. 42:01then the H here is 32 by 64. so these
  1164. 42:05are 64 dimensional hidden States and
  1165. 42:08then this W Matrix projects those 64
  1166. 42:10dimensional vectors into 27 dimensions
  1167. 42:12and then there's a 27 dimensional offset
  1168. 42:15which is a one-dimensional vector
  1169. 42:18now we should note that this plus here
  1170. 42:20actually broadcasts because H multiplied
  1171. 42:23by by W2 will give us a 32 by 27. and so
  1172. 42:27then this plus B2 is a 27 dimensional
  1173. 42:31lecture here
  1174. 42:32now in the rules of broadcasting what's
  1175. 42:33going to happen with this bias Vector is
  1176. 42:35that this one-dimensional Vector of 27
  1177. 42:37will get aligned with a padded dimension
  1178. 42:41of one on the left and it will basically
  1179. 42:43become a row vector and then it will get
  1180. 42:45replicated vertically 32 times to make
  1181. 42:48it 32 by 27 and then there's an
  1182. 42:50element-wise multiply
  1183. 42:52now
  1184. 42:54the question is how do we back propagate
  1185. 42:56from logits to the hidden States the
  1186. 42:59weight Matrix W2 and the bias B2
  1187. 43:02and you might think that we need to go
  1188. 43:03to some Matrix calculus and then we have
  1189. 43:07to look up the derivative for a matrix
  1190. 43:09multiplication but actually you don't
  1191. 43:11have to do any of that and you can go
  1192. 43:12back to First principles and derive this
  1193. 43:14yourself on a piece of paper and
  1194. 43:17specifically what I like to do and I
  1195. 43:18what I find works well for me is you
  1196. 43:20find a specific small example that you
  1197. 43:23then fully write out and then in the
  1198. 43:25process of analyzing how that individual
  1199. 43:27small example works you will understand
  1200. 43:28the broader pattern and you'll be able
  1201. 43:30to generalize and write out the full
  1202. 43:32general formula for what how these
  1203. 43:35derivatives flow in an expression like
  1204. 43:37this so let's try that out
  1205. 43:39so pardon the low budget production here
  1206. 43:41but what I've done here is I'm writing
  1207. 43:43it out on a piece of paper really what
  1208. 43:45we are interested in is we have a
  1209. 43:46multiply B plus C and that creates a d
  1210. 43:50and we have the derivative of the loss
  1211. 43:53with respect to D and we'd like to know
  1212. 43:54what the derivative of the losses with
  1213. 43:55respect to a b and c
  1214. 43:57now these here are little
  1215. 44:00two-dimensional examples of a matrix
  1216. 44:01multiplication Two by Two Times a two by
  1217. 44:03two
  1218. 44:04plus a 2 a vector of just two elements
  1219. 44:07C1 and C2 gives me a two by two
  1220. 44:10now notice here that I have a bias
  1221. 44:14Vector here called C and the bisex
  1222. 44:17vector is C1 and C2 but as I described
  1223. 44:19over here that bias Vector will become a
  1224. 44:21row Vector in the broadcasting and will
  1225. 44:23replicate vertically so that's what's
  1226. 44:24happening here as well C1 C2 is
  1227. 44:27replicated vertically and we see how we
  1228. 44:29have two rows of C1 C2 as a result
  1229. 44:33so now when I say write it out I just
  1230. 44:35mean like this basically break up this
  1231. 44:37matrix multiplication into the actual
  1232. 44:40thing that that's going on under the
  1233. 44:41hood so as a result of matrix
  1234. 44:44multiplication and how it works d11 is
  1235. 44:46the result of a DOT product between the
  1236. 44:48first row of a and the First Column of B
  1237. 44:51so a11 b11 plus a12 B21 plus C1
  1238. 44:57and so on so forth for all the other
  1239. 44:59elements of D and once you actually
  1240. 45:02write it out it becomes obvious this is
  1241. 45:03just a bunch of multipliers and
  1242. 45:06um adds and we know from micrograd how
  1243. 45:09to differentiate multiplies and adds and
  1244. 45:11so this is not scary anymore it's not
  1245. 45:13just matrix multiplication it's just uh
  1246. 45:15tedious unfortunately but this is
  1247. 45:17completely tractable we have DL by D for
  1248. 45:20all of these and we want DL by uh all
  1249. 45:23these little other variables so how do
  1250. 45:25we achieve that and how do we actually
  1251. 45:26get the gradients okay so the low budget
  1252. 45:29production continues here
  1253. 45:30so let's for example derive the
  1254. 45:32derivative of the loss with respect to
  1255. 45:34a11
  1256. 45:36we see here that a11 occurs twice in our
  1257. 45:38simple expression right here right here
  1258. 45:40and influences d11 and D12
  1259. 45:43. so this is so what is DL by d a one
  1260. 45:46one well it's DL by d11 times the local
  1261. 45:51derivative of d11 which in this case is
  1262. 45:53just b11 because that's what's
  1263. 45:55multiplying a11 here
  1264. 45:57so uh and likewise here the local
  1265. 46:00derivative of D12 with respect to a11 is
  1266. 46:02just B12 and so B12 well in the chain
  1267. 46:05rule therefore multiply the L by d 1 2.
  1268. 46:08and then because a11 is used both to
  1269. 46:11produce d11 and D12 we need to add up
  1270. 46:15the contributions of both of those sort
  1271. 46:18of chains that are running in parallel
  1272. 46:20and that's why we get a plus just adding
  1273. 46:22up those two
  1274. 46:24um those two contributions and that
  1275. 46:26gives us DL by d a one one we can do the
  1276. 46:29exact same analysis for the other one
  1277. 46:31for all the other elements of a and when
  1278. 46:34you simply write it out it's just super
  1279. 46:36simple
  1280. 46:37um taking of gradients on you know
  1281. 46:40expressions like this
  1282. 46:42you find that
  1283. 46:44this Matrix DL by D A that we're after
  1284. 46:47right if we just arrange all the all of
  1285. 46:49them in the same shape as a takes so a
  1286. 46:52is just too much Matrix so d l by D A
  1287. 46:55here will be also just the same shape
  1288. 46:59tester with the derivatives now so deal
  1289. 47:03by D a11 Etc
  1290. 47:05and we see that actually we can express
  1291. 47:06what we've written out here as a matrix
  1292. 47:09multiplied
  1293. 47:10and so it just so happens that D all by
  1294. 47:13that all of these formulas that we've
  1295. 47:15derived here by taking gradients can
  1296. 47:17actually be expressed as a matrix
  1297. 47:19multiplication and in particular we see
  1298. 47:21that it is the matrix multiplication of
  1299. 47:22these two array matrices
  1300. 47:25so it is the um DL by D and then Matrix
  1301. 47:30multiplying B but B transpose actually
  1302. 47:32so you see that B21 and b12 have changed
  1303. 47:37place
  1304. 47:38whereas before we had of course b11 B12
  1305. 47:41B2 on B22 so you see that this other
  1306. 47:45Matrix B is transposed
  1307. 47:47and so basically what we have long story
  1308. 47:49short just by doing very simple
  1309. 47:50reasoning here by breaking up the
  1310. 47:52expression in the case of a very simple
  1311. 47:54example is that DL by d a is which is
  1312. 47:58this is simply equal to DL by DD Matrix
  1313. 48:02multiplied with B transpose
  1314. 48:05so that is what we have so far now we
  1315. 48:08also want the derivative with respect to
  1316. 48:10um B and C now
  1317. 48:13for B I'm not actually doing the full
  1318. 48:15derivation because honestly it's um it's
  1319. 48:18not deep it's just uh annoying it's
  1320. 48:20exhausting you can actually do this
  1321. 48:22analysis yourself you'll also find that
  1322. 48:24if you take this these expressions and
  1323. 48:26you differentiate with respect to b
  1324. 48:27instead of a you will find that DL by DB
  1325. 48:30is also a matrix multiplication in this
  1326. 48:33case you have to take the Matrix a and
  1327. 48:35transpose it and Matrix multiply that
  1328. 48:37with bl by DD
  1329. 48:39and that's what gives you a deal by DB
  1330. 48:42and then here for the offsets C1 and C2
  1331. 48:46if you again just differentiate with
  1332. 48:47respect to C1 you will find an
  1333. 48:50expression like this
  1334. 48:52and C2 an expression like this
  1335. 48:55and basically you'll find the DL by DC
  1336. 48:57is simply because they're just
  1337. 48:59offsetting these Expressions you just
  1338. 49:01have to take the deal by DD Matrix
  1339. 49:04of the derivatives of D and you just
  1340. 49:07have to sum across the columns and that
  1341. 49:11gives you the derivatives for C
  1342. 49:13so long story short
  1343. 49:15the backward Paths of a matrix multiply
  1344. 49:18is a matrix multiply
  1345. 49:20and instead of just like we had D equals
  1346. 49:22a times B plus C in the scalar case uh
  1347. 49:25we sort of like arrive at something very
  1348. 49:27very similar but now uh with a matrix
  1349. 49:29multiplication instead of a scalar
  1350. 49:31multiplication
  1351. 49:32so the derivative of D with respect to a
  1352. 49:36is
  1353. 49:37DL by DD Matrix multiplied B trespose
  1354. 49:41and here it's a transpose multiply deal
  1355. 49:44by DD but in both cases it's a matrix
  1356. 49:46multiplication with the derivative and
  1357. 49:49the other term in the multiplication
  1358. 49:53and for C it is a sum
  1359. 49:55now I'll tell you a secret I can never
  1360. 49:58remember the formulas that we just
  1361. 50:00arrived for back proper gain information
  1362. 50:01multiplication and I can back propagate
  1363. 50:03through these Expressions just fine and
  1364. 50:05the reason this works is because the
  1365. 50:07dimensions have to work out
  1366. 50:09uh so let me give you an example say I
  1367. 50:11want to create DH
  1368. 50:13then what should the H be number one I
  1369. 50:16have to know that the shape of DH must
  1370. 50:19be the same as the shape of H
  1371. 50:21and the shape of H is 32 by 64. and then
  1372. 50:24the other piece of information I know is
  1373. 50:26that DH must be some kind of matrix
  1374. 50:28multiplication of the logits with W2
  1375. 50:32and delojits is 32 by 27 and W2 is a 64
  1376. 50:37by 27. there is only a single way to
  1377. 50:40make the shape work out in this case and
  1378. 50:43it is indeed the correct result in
  1379. 50:45particular here H needs to be 32 by 64.
  1380. 50:48the only way to achieve that is to take
  1381. 50:50a deluges
  1382. 50:52and Matrix multiply it with you see how
  1383. 50:55I have to take W2 but I have to
  1384. 50:57transpose it to make the dimensions work
  1385. 50:58out
  1386. 50:59so w to transpose and it's the only way
  1387. 51:02to make these to Matrix multiply those
  1388. 51:04two pieces to make the shapes work out
  1389. 51:06and that turns out to be the correct
  1390. 51:08formula so if we come here we want DH
  1391. 51:11which is d a and we see that d a is DL
  1392. 51:15by DD Matrix multiply B transpose
  1393. 51:18so that's Delo just multiply and B is W2
  1394. 51:21so W2 transpose which is exactly what we
  1395. 51:24have here so there's no need to remember
  1396. 51:26these formulas similarly now if I want
  1397. 51:30dw2 well I know that it must be a matrix
  1398. 51:33multiplication of D logits and H
  1399. 51:37and maybe there's a few transpose like
  1400. 51:39there's one transpose in there as well
  1401. 51:40and I don't know which way it is so I
  1402. 51:42have to come to W2 and I see that its
  1403. 51:44shape is 64 by 27
  1404. 51:47and that has to come from some interest
  1405. 51:49multiplication of these two
  1406. 51:51and so to get a 64 by 27 I need to take
  1407. 51:55um
  1408. 51:56H I need to transpose it
  1409. 51:59and then I need to Matrix multiply it
  1410. 52:01um so that will become 64 by 32 and then
  1411. 52:04I need to make sure to multiply with the
  1412. 52:0532 by 27 and that's going to give me a
  1413. 52:0764 by 27. so I need to make sure it's
  1414. 52:09multiplied this with the logist that
  1415. 52:11shape just like that that's the only way
  1416. 52:13to make the dimensions work out and just
  1417. 52:15use matrix multiplication and if we come
  1418. 52:17here we see that that's exactly what's
  1419. 52:19here so a transpose a for us is H
  1420. 52:23multiplied with deloaches
  1421. 52:25so that's W2 and then db2
  1422. 52:30is just the um
  1423. 52:33vertical sum and actually in the same
  1424. 52:35way there's only one way to make the
  1425. 52:37shapes work out I don't have to remember
  1426. 52:38that it's a vertical Sum along the zero
  1427. 52:40axis because that's the only way that
  1428. 52:42this makes sense because B2 shape is 27
  1429. 52:45so in order to get a um delugits
  1430. 52:50here is 30 by 27 so knowing that it's
  1431. 52:54just sum over deloaches in some
  1432. 52:56Direction
  1433. 52:59that direction must be zero because I
  1434. 53:02need to eliminate this Dimension so it's
  1435. 53:04this
  1436. 53:06so this is so let's kind of like the
  1437. 53:08hacky way let me copy paste and delete
  1438. 53:10that and let me swing over here and this
  1439. 53:13is our backward pass for the linear
  1440. 53:14layer uh hopefully
  1441. 53:17so now let's uncomment
  1442. 53:19these three and we're checking that we
  1443. 53:21got all the three derivatives correct
  1444. 53:24and run
  1445. 53:26and we see that h wh and B2 are all
  1446. 53:30exactly correct so we back propagated
  1447. 53:33through a linear layer
  1448. 53:36now next up we have derivative for the h
  1449. 53:39already and we need to back propagate
  1450. 53:41through 10h into h preact
  1451. 53:43so we want to derive DH preact
  1452. 53:47and here we have to back propagate
  1453. 53:48through a 10 H and we've already done
  1454. 53:50this in micrograd and we remember that
  1455. 53:5210h has a very simple backward formula
  1456. 53:54now unfortunately if I just put in D by
  1457. 53:56DX of 10 h of X into both from alpha it
  1458. 53:59lets us down it tells us that it's a
  1459. 54:00hyperbolic secant function squared of X
  1460. 54:03it's not exactly helpful but luckily
  1461. 54:06Google image search does not let us down
  1462. 54:08and it gives us the simpler formula and
  1463. 54:10in particular if you have that a is
  1464. 54:12equal to 10 h of Z then d a by DZ by
  1465. 54:16propagating through 10 H is just one
  1466. 54:17minus a square and take note that 1
  1467. 54:21minus a square a here is the output of
  1468. 54:23the 10h not the input to the 10h Z so
  1469. 54:27the D A by DZ is here formulated in
  1470. 54:29terms of the output of that 10h
  1471. 54:31and here also in Google image search we
  1472. 54:34have the full derivation if you want to
  1473. 54:35actually take the actual definition of
  1474. 54:3810h and work through the math to figure
  1475. 54:39out 1 minus standard square of Z
  1476. 54:42so 1 minus a square is the local
  1477. 54:45derivative in our case that is 1 minus
  1478. 54:49uh the output of 10 H squared which here
  1479. 54:52is H
  1480. 54:53so it's h squared and that is the local
  1481. 54:56derivative and then times the chain rule
  1482. 54:58DH
  1483. 55:00so that is going to be our candidate
  1484. 55:02implementation so if we come here
  1485. 55:05and then uncomment this let's hope for
  1486. 55:08the best
  1487. 55:09and we have the right answer
  1488. 55:12okay next up we have DH preact and we
  1489. 55:15want to back propagate into the gain the
  1490. 55:17B and raw and the B and bias
  1491. 55:19so here this is the bathroom parameters
  1492. 55:21being gained in bias inside the bash
  1493. 55:23term that take the B and raw that is
  1494. 55:25exact unit caution and then scale it and
  1495. 55:28shift it
  1496. 55:29and these are the parameters of The
  1497. 55:30Bachelor now here we have a
  1498. 55:33multiplication but it's worth noting
  1499. 55:35that this multiply is very very
  1500. 55:36different from this Matrix multiply here
  1501. 55:38Matrix multiply are DOT products between
  1502. 55:41rows and Columns of these matrices
  1503. 55:43involved this is an element twice
  1504. 55:45multiply so things are quite a bit
  1505. 55:46simpler
  1506. 55:47now we do have to be careful with some
  1507. 55:49of the broadcasting happening in this
  1508. 55:51line of code though so you see how BN
  1509. 55:53gain and B and bias are 1 by 64. but H
  1510. 55:58preact and B and raw are 32 by 64.
  1511. 56:02so we have to be careful with that and
  1512. 56:04make sure that all the shapes work out
  1513. 56:05fine and that the broadcasting is
  1514. 56:06correctly back propagated
  1515. 56:08so in particular let's start with the B
  1516. 56:10and Gain so DB and gain should be
  1517. 56:14and here this is again elementorized
  1518. 56:17multiply and whenever we have a times b
  1519. 56:19equals c we saw that the local
  1520. 56:21derivative here is just if this is a the
  1521. 56:23local derivative is just the B the other
  1522. 56:25one so the local derivative is just B
  1523. 56:27and raw and then times chain rule
  1524. 56:31so DH preact
  1525. 56:34so this is the candidate gradient now
  1526. 56:38again we have to be careful because B
  1527. 56:40and Gain Is of size 1 by 64. but this
  1528. 56:44here would be 32 by 64.
  1529. 56:48and so
  1530. 56:49um the correct thing to do in this case
  1531. 56:51of course is that b and gain here is a
  1532. 56:53rule Vector of 64 numbers it gets
  1533. 56:55replicated vertically in this operation
  1534. 56:58and so therefore the correct thing to do
  1535. 57:00is to sum because it's being replicated
  1536. 57:03and therefore all the gradients in each
  1537. 57:06of the rows that are now flowing
  1538. 57:07backwards need to sum up to that same
  1539. 57:10tensor DB and Gain so we have to sum
  1540. 57:13across all the zero all the examples
  1541. 57:16basically
  1542. 57:17which is the direction in which this
  1543. 57:19gets replicated
  1544. 57:20and now we have to be also careful
  1545. 57:21because we
  1546. 57:23um being gain is of shape 1 by 64. so in
  1547. 57:26fact I need to keep them as true
  1548. 57:29otherwise I would just get 64.
  1549. 57:31now I don't actually really remember why
  1550. 57:34the being gain and the BN bias I made
  1551. 57:36them be 1 by 64.
  1552. 57:40um
  1553. 57:41but the biases B1 and B2 I just made
  1554. 57:44them be one-dimensional vectors they're
  1555. 57:45not two-dimensional tensors so I can't
  1556. 57:47recall exactly why I left the gain and
  1557. 57:51the bias as two-dimensional but it
  1558. 57:53doesn't really matter as long as you are
  1559. 57:54consistent and you're keeping it the
  1560. 57:55same
  1561. 57:56so in this case we want to keep the
  1562. 57:58dimension so that the tensor shapes work
  1563. 58:01next up we have B and raw so DB and raw
  1564. 58:05will be BN gain
  1565. 58:09multiplying
  1566. 58:11dhreact that's our chain rule now what
  1567. 58:15about the
  1568. 58:17um
  1569. 58:18dimensions of this we have to be careful
  1570. 58:20right so DH preact is 32 by 64. B and
  1571. 58:24gain is 1 by 64. so it will just get
  1572. 58:27replicated and to create this
  1573. 58:29multiplication which is the correct
  1574. 58:31thing because in a forward pass it also
  1575. 58:33gets replicated in just the same way
  1576. 58:35so in fact we don't need the brackets
  1577. 58:37here we're done
  1578. 58:38and the shapes are already correct
  1579. 58:40and finally for the bias
  1580. 58:43very similar this bias here is very very
  1581. 58:46similar to the bias we saw when you
  1582. 58:47layer in the linear layer and we see
  1583. 58:49that the gradients from each preact will
  1584. 58:51simply flow into the biases and add up
  1585. 58:54because these are just these are just
  1586. 58:55offsets
  1587. 58:57and so basically we want this to be DH
  1588. 58:59preact but it needs to Sum along the
  1589. 59:01right Dimension and in this case similar
  1590. 59:04to the gain we need to sum across the
  1591. 59:06zeroth dimension the examples because of
  1592. 59:09the way that the bias gets replicated
  1593. 59:10vertically
  1594. 59:11and we also want to have keep them as
  1595. 59:14true
  1596. 59:15and so this will basically take this and
  1597. 59:17sum it up and give us a 1 by 64.
  1598. 59:20so this is the candidate implementation
  1599. 59:23it makes all the shapes work
  1600. 59:25let me bring it up down here and then
  1601. 59:28let me uncomment these three lines
  1602. 59:32to check that we are getting the correct
  1603. 59:33result for all the three tensors and
  1604. 59:36indeed we see that all of that got back
  1605. 59:38propagated correctly so now we get to
  1606. 59:40the batch Norm layer we see how here
  1607. 59:42being gay and being bias are the
  1608. 59:44parameters so the back propagation ends
  1609. 59:46but B and raw now is the output of the
  1610. 59:50standardization
  1611. 59:51so here what I'm doing of course is I'm
  1612. 59:53breaking up the batch form into
  1613. 59:54manageable pieces so we can back
  1614. 59:55propagate through each line individually
  1615. 59:57but basically what's happening is BN
  1616. 1:00:00mean I is the sum
  1617. 1:00:03so this is the B and mean I I apologize
  1618. 1:00:06for the variable naming B and diff is x
  1619. 1:00:10minus mu
  1620. 1:00:11B and div 2 is x minus mu squared here
  1621. 1:00:15inside the variance
  1622. 1:00:16B and VAR is the variance so uh Sigma
  1623. 1:00:20Square this is B and bar and it's
  1624. 1:00:22basically the sum of squares
  1625. 1:00:25so this is the x minus mu squared and
  1626. 1:00:28then the sum now you'll notice one
  1627. 1:00:30departure here
  1628. 1:00:32here it is normalized as 1 over m
  1629. 1:00:34uh which is number of examples here I'm
  1630. 1:00:37normalizing as one over n minus 1
  1631. 1:00:39instead of N and this is deliberate and
  1632. 1:00:42I'll come back to that in a bit when we
  1633. 1:00:43are at this line it is something called
  1634. 1:00:45the bezels correction
  1635. 1:00:47but this is how I want it in our case
  1636. 1:00:51bienvar inv then becomes basically
  1637. 1:00:53bienvar plus Epsilon Epsilon is one
  1638. 1:00:56negative five and then it's one over
  1639. 1:00:58square root
  1640. 1:00:59is the same as raising to the power of
  1641. 1:01:02negative 0.5 right because 0.5 is square
  1642. 1:01:05root and then negative makes it one over
  1643. 1:01:07square root
  1644. 1:01:08so BM Bar M is a one over this uh
  1645. 1:01:12denominator here and then we can see
  1646. 1:01:14that b and raw which is the X hat here
  1647. 1:01:16is equal to the BN diff the numerator
  1648. 1:01:19multiplied by the
  1649. 1:01:22um BN bar in
  1650. 1:01:24and this line here that creates pre-h
  1651. 1:01:27pre-act was the last piece we've already
  1652. 1:01:29back propagated through it
  1653. 1:01:31so now what we want to do is we are here
  1654. 1:01:34and we have B and raw and we have to
  1655. 1:01:35first back propagate into B and diff and
  1656. 1:01:38B and Bar M
  1657. 1:01:40so now we're here and we have DB and raw
  1658. 1:01:43and we need to back propagate through
  1659. 1:01:45this line
  1660. 1:01:46now I've written out the shapes here and
  1661. 1:01:49indeed bien VAR m is a shape 1 by 64. so
  1662. 1:01:53there is a broadcasting happening here
  1663. 1:01:55that we have to be careful with but it
  1664. 1:01:57is just an element-wise simple
  1665. 1:01:58multiplication by now we should be
  1666. 1:02:00pretty comfortable with that to get DB
  1667. 1:02:02and diff we know that this is just B and
  1668. 1:02:05varm
  1669. 1:02:06multiplied with
  1670. 1:02:08DP and raw
  1671. 1:02:11and conversely to get dbmring
  1672. 1:02:15we need to take the end if
  1673. 1:02:17and multiply that by DB and raw
  1674. 1:02:22so this is the candidate but of course
  1675. 1:02:24we need to make sure that broadcasting
  1676. 1:02:26is obeyed so in particular B and VAR M
  1677. 1:02:29multiplying with DB and raw
  1678. 1:02:31will be okay and give us 32 by 64 as we
  1679. 1:02:35expect
  1680. 1:02:36but dbm VAR inv would be taking a 32 by
  1681. 1:02:4064.
  1682. 1:02:42multiplying it by 32 by 64. so this is a
  1683. 1:02:4532 by 64. but of course DB this uh B and
  1684. 1:02:49VAR in is only 1 by 64. so the second
  1685. 1:02:52line here needs a sum across the
  1686. 1:02:55examples and because there's this
  1687. 1:02:57Dimension here we need to make sure that
  1688. 1:03:00keep them is true
  1689. 1:03:02so this is the candidate
  1690. 1:03:04let's erase this and let's swing down
  1691. 1:03:07here
  1692. 1:03:09and implement it and then let's comment
  1693. 1:03:11out dbm barif and DB and diff
  1694. 1:03:16now we'll actually notice that DB and
  1695. 1:03:18diff by the way is going to be incorrect
  1696. 1:03:22so when I run this
  1697. 1:03:24BMR m is correct B and diff is not
  1698. 1:03:27correct and this is actually expected
  1699. 1:03:30because we're not done with b and diff
  1700. 1:03:34so in particular when we slide here we
  1701. 1:03:36see here that b and raw as a function of
  1702. 1:03:37B and diff but actually B and far of is
  1703. 1:03:40a function of B of R which is a function
  1704. 1:03:42of B and df2 which is a function of B
  1705. 1:03:44and diff
  1706. 1:03:45so it comes here so bdn diff
  1707. 1:03:48um these variable names are crazy I'm
  1708. 1:03:50sorry it branches out into two branches
  1709. 1:03:53and we've only done one branch of it we
  1710. 1:03:55have to continue our back propagation
  1711. 1:03:57and eventually come back to B and diff
  1712. 1:03:58and then we'll be able to do a plus
  1713. 1:04:00equals and get the actual card gradient
  1714. 1:04:02for now it is good to verify that CMP
  1715. 1:04:05also works it doesn't just lie to us and
  1716. 1:04:07tell us that everything is always
  1717. 1:04:08correct it can in fact detect when your
  1718. 1:04:11gradient is not correct so it's that's
  1719. 1:04:13good to see as well okay so now we have
  1720. 1:04:15the derivative here and we're trying to
  1721. 1:04:17back propagate through this line
  1722. 1:04:18and because we're raising to a power of
  1723. 1:04:21negative 0.5 I brought up the power rule
  1724. 1:04:23and we see that basically we have that
  1725. 1:04:25the BM bar will now be we bring down the
  1726. 1:04:28exponent so negative 0.5 times
  1727. 1:04:31uh X which is this
  1728. 1:04:34and now raised to the power of negative
  1729. 1:04:360.5 minus 1 which is negative 1.5
  1730. 1:04:39now we would have to also apply a small
  1731. 1:04:42chain rule here in our head because we
  1732. 1:04:45need to take further the derivative of B
  1733. 1:04:48and VAR with respect to this expression
  1734. 1:04:49here inside the bracket but because this
  1735. 1:04:51is an elementalized operation and
  1736. 1:04:53everything is fairly simple that's just
  1737. 1:04:54one and so there's nothing to do there
  1738. 1:04:57so this is the local derivative and then
  1739. 1:05:00times the global derivative to create
  1740. 1:05:01the chain rule this is just times the BM
  1741. 1:05:04bar have
  1742. 1:05:05so this is our candidate let me bring
  1743. 1:05:08this down
  1744. 1:05:10and uncommon to the check
  1745. 1:05:14and we see that we have the correct
  1746. 1:05:16result
  1747. 1:05:17now before we propagate through the next
  1748. 1:05:19line I want to briefly talk about the
  1749. 1:05:20note here where I'm using the bezels
  1750. 1:05:22correction dividing by n minus 1 instead
  1751. 1:05:24of dividing by n when I normalize here
  1752. 1:05:27the sum of squares
  1753. 1:05:29now you'll notice that this is departure
  1754. 1:05:31from the paper which uses one over n
  1755. 1:05:33instead not one over n minus one their m
  1756. 1:05:36is RN
  1757. 1:05:38and
  1758. 1:05:39um so it turns out that there are two
  1759. 1:05:40ways of estimating variance of an array
  1760. 1:05:43one is the biased estimate which is one
  1761. 1:05:46over n and the other one is the unbiased
  1762. 1:05:49estimate which is one over n minus one
  1763. 1:05:51now confusingly in the paper this is uh
  1764. 1:05:54not very clearly described and also it's
  1765. 1:05:56a detail that kind of matters I think
  1766. 1:05:58um they are using the biased version
  1767. 1:06:00training time but later when they are
  1768. 1:06:02talking about the inference they are
  1769. 1:06:04mentioning that when they do the
  1770. 1:06:06inference they are using the unbiased
  1771. 1:06:08estimate which is the n minus one
  1772. 1:06:10version in
  1773. 1:06:12um
  1774. 1:06:12basically for inference
  1775. 1:06:15and to calibrate the running mean and
  1776. 1:06:18the running variance basically and so
  1777. 1:06:20they they actually introduce a trained
  1778. 1:06:22test mismatch where in training they use
  1779. 1:06:24the biased version and in the in test
  1780. 1:06:26time they use the unbiased version I
  1781. 1:06:28find this extremely confusing you can
  1782. 1:06:30read more about the bezels correction
  1783. 1:06:32and why uh dividing by n minus one gives
  1784. 1:06:35you a better estimate of the variance in
  1785. 1:06:37a case where you have population size or
  1786. 1:06:39samples for the population
  1787. 1:06:41that are very small and that is indeed
  1788. 1:06:44the case for us because we are dealing
  1789. 1:06:46with many patches and these mini matches
  1790. 1:06:48are a small sample of a larger
  1791. 1:06:50population which is the entire training
  1792. 1:06:52set and so it just turns out that if you
  1793. 1:06:55just estimate it using one over n that
  1794. 1:06:57actually almost always underestimates
  1795. 1:06:58the variance and it is a biased
  1796. 1:07:00estimator and it is advised that you use
  1797. 1:07:02the unbiased version and divide by n
  1798. 1:07:04minus one and you can go through this
  1799. 1:07:06article here that I liked that actually
  1800. 1:07:08describes the full reasoning and I'll
  1801. 1:07:09link it in the video description
  1802. 1:07:12now when you calculate the torture
  1803. 1:07:13variance
  1804. 1:07:15you'll notice that they take the
  1805. 1:07:16unbiased flag whether or not you want to
  1806. 1:07:18divide by n or n minus one confusingly
  1807. 1:07:21they do not mention what the default is
  1808. 1:07:24for unbiased but I believe unbiased by
  1809. 1:07:26default is true I'm not sure why the
  1810. 1:07:29docs here don't cite that
  1811. 1:07:31now in The Bachelor
  1812. 1:07:331D the documentation again is kind of
  1813. 1:07:35wrong and confusing it says that the
  1814. 1:07:38standard deviation is calculated via the
  1815. 1:07:39biased estimator
  1816. 1:07:41but this is actually not exactly right
  1817. 1:07:43and people have pointed out that it is
  1818. 1:07:44not right in a number of issues since
  1819. 1:07:46then because actually the rabbit hole is
  1820. 1:07:49deeper and they follow the paper exactly
  1821. 1:07:52and they use the biased version for
  1822. 1:07:54training but when they're estimating the
  1823. 1:07:56running standard deviation we are using
  1824. 1:07:58the unbiased version so again there's
  1825. 1:08:00the train test mismatch so long story
  1826. 1:08:02short I'm not a fan of trained test
  1827. 1:08:05discrepancies I basically kind of
  1828. 1:08:07consider
  1829. 1:08:08the fact that we use the bias version
  1830. 1:08:10the training time and the unbiased test
  1831. 1:08:13time I basically consider this to be a
  1832. 1:08:14bug and I don't think that there's a
  1833. 1:08:16good reason for that it's not really
  1834. 1:08:18they don't really go into the detail of
  1835. 1:08:19the reasoning behind it in this paper so
  1836. 1:08:22that's why I basically prefer to use the
  1837. 1:08:24bestless correction in my own work
  1838. 1:08:26unfortunately Bastion does not take a
  1839. 1:08:29keyword argument that tells you whether
  1840. 1:08:30or not you want to use the unbiased
  1841. 1:08:33version of the bias version in both
  1842. 1:08:34train and test and so therefore anyone
  1843. 1:08:36using batch normalization basically in
  1844. 1:08:38my view has a bit of a bug in the code
  1845. 1:08:41um
  1846. 1:08:42and this turns out to be much less of a
  1847. 1:08:44problem if your batch mini batch sizes
  1848. 1:08:46are a bit larger but still I just might
  1849. 1:08:48kind of uh unpardable so maybe someone
  1850. 1:08:51can explain why this is okay but for now
  1851. 1:08:53I prefer to use the unbiased version
  1852. 1:08:55consistently both during training and at
  1853. 1:08:57this time and that's why I'm using one
  1854. 1:09:00over n minus one here
  1855. 1:09:01okay so let's now actually back
  1856. 1:09:03propagate through this line
  1857. 1:09:05so
  1858. 1:09:07the first thing that I always like to do
  1859. 1:09:08is I like to scrutinize the shapes first
  1860. 1:09:10so in particular here looking at the
  1861. 1:09:12shapes of what's involved I see that b
  1862. 1:09:14and VAR shape is 1 by 64. so it's a row
  1863. 1:09:18vector and BND if two dot shape is 32 by
  1864. 1:09:2164.
  1865. 1:09:22so clearly here we're doing a sum over
  1866. 1:09:25the zeroth axis to squash the first
  1867. 1:09:28dimension of of the shapes here using a
  1868. 1:09:32sum so that right away actually hints to
  1869. 1:09:35me that there will be some kind of a
  1870. 1:09:36replication or broadcasting in the
  1871. 1:09:38backward pass and maybe you're noticing
  1872. 1:09:40the pattern here but basically anytime
  1873. 1:09:42you have a sum in the forward pass that
  1874. 1:09:45turns into a replication or broadcasting
  1875. 1:09:47in the backward pass along the same
  1876. 1:09:49Dimension and conversely when we have a
  1877. 1:09:52replication or a broadcasting in the
  1878. 1:09:54forward pass that indicates a variable
  1879. 1:09:57reuse and so in the backward pass that
  1880. 1:09:59turns into a sum over the exact same
  1881. 1:10:01dimension
  1882. 1:10:02and so hopefully you're noticing that
  1883. 1:10:04Duality that those two are kind of like
  1884. 1:10:06the opposite of each other in the
  1885. 1:10:07forward and backward pass
  1886. 1:10:09now once we understand the shapes the
  1887. 1:10:11next thing I like to do always is I like
  1888. 1:10:12to look at a toy example in my head to
  1889. 1:10:15sort of just like understand roughly how
  1890. 1:10:16uh the variable the variable
  1891. 1:10:18dependencies go in the mathematical
  1892. 1:10:19formula
  1893. 1:10:21so here we have a two-dimensional array
  1894. 1:10:24of the end of two which we are scaling
  1895. 1:10:26by a constant and then we are summing uh
  1896. 1:10:29vertically over the columns so if we
  1897. 1:10:32have a two by two Matrix a and then we
  1898. 1:10:33sum over the columns and scale we would
  1899. 1:10:36get a row Vector B1 B2 and B1 depends on
  1900. 1:10:39a in this way whereas just sum they're
  1901. 1:10:42scaled of a and B2 in this way where
  1902. 1:10:45it's the second column sump and scale
  1903. 1:10:48and so looking at this basically
  1904. 1:10:52what we want to do now is we have the
  1905. 1:10:53derivatives on B1 and B2 and we want to
  1906. 1:10:55back propagate them into Ace and so it's
  1907. 1:10:58clear that just differentiating in your
  1908. 1:10:59head the local derivative here is one
  1909. 1:11:01over n minus 1 times uh one
  1910. 1:11:05uh for each one of these A's and um
  1911. 1:11:09basically the derivative of B1 has to
  1912. 1:11:11flow through The Columns of a
  1913. 1:11:13scaled by one over n minus one
  1914. 1:11:16and that's roughly What's Happening Here
  1915. 1:11:18so intuitively the derivative flow tells
  1916. 1:11:21us that DB and diff2
  1917. 1:11:24will be the local derivative of this
  1918. 1:11:27operation and there are many ways to do
  1919. 1:11:29this by the way but I like to do
  1920. 1:11:31something like this torch dot once like
  1921. 1:11:33of bndf2 so I'll create a large array
  1922. 1:11:37two-dimensional of ones
  1923. 1:11:39and then I will scale it so 1.0 divided
  1924. 1:11:42by n minus 1.
  1925. 1:11:44so this is a array of
  1926. 1:11:46um one over n minus one and that's sort
  1927. 1:11:49of like the local derivative
  1928. 1:11:50and now for the chain rule I will simply
  1929. 1:11:53just multiply it by dbm bar
  1930. 1:11:58and notice here what's going to happen
  1931. 1:12:00this is 32 by 64 and this is just 1 by
  1932. 1:12:0264. so I'm letting the broadcasting do
  1933. 1:12:06the replication because internally in
  1934. 1:12:08pytorch basically dbnbar which is 1 by
  1935. 1:12:1164 row vector
  1936. 1:12:13well in this multiplication get
  1937. 1:12:15um copied vertically until the two are
  1938. 1:12:18of the same shape and then there will be
  1939. 1:12:19an element wise multiply and so that uh
  1940. 1:12:22so that the broadcasting is basically
  1941. 1:12:23doing the replication
  1942. 1:12:25and I will end up with the derivatives
  1943. 1:12:27of DB and diff2 here
  1944. 1:12:30so this is the candidate solution let's
  1945. 1:12:32bring it down here
  1946. 1:12:33let's uncomment this line where we check
  1947. 1:12:36it and let's hope for the best
  1948. 1:12:39and indeed we see that this is the
  1949. 1:12:41correct formula next up let's
  1950. 1:12:43differentiate here and to be in this
  1951. 1:12:45so here we have that b and diff is
  1952. 1:12:48element y squared to create B and F2
  1953. 1:12:50so this is a relatively simple
  1954. 1:12:52derivative because it's a simple element
  1955. 1:12:54wise operation so it's kind of like the
  1956. 1:12:56scalar case and we have that DB and div
  1957. 1:12:59should be if this is x squared then the
  1958. 1:13:02derivative of this is 2x right so it's
  1959. 1:13:04simply 2 times B and if that's the local
  1960. 1:13:07derivative
  1961. 1:13:08and then times chain Rule and the shape
  1962. 1:13:11of these is the same they are of the
  1963. 1:13:13same shape so times this
  1964. 1:13:15so that's the backward pass for this
  1965. 1:13:17variable let me bring that down here
  1966. 1:13:20and now we have to be careful because we
  1967. 1:13:22already calculated dbm depth right so
  1968. 1:13:24this is just the end of the other uh you
  1969. 1:13:27know other Branch coming back to B and
  1970. 1:13:30diff
  1971. 1:13:30because B and diff was already back
  1972. 1:13:32propagated to way over here
  1973. 1:13:34from being raw so we now completed the
  1974. 1:13:37second branch and so that's why I have
  1975. 1:13:39to do plus equals and if you recall we
  1976. 1:13:42had an incorrect derivative for being
  1977. 1:13:43diff before and I'm hoping that once we
  1978. 1:13:46append this last missing piece we have
  1979. 1:13:48the exact correctness so let's run
  1980. 1:13:51ambient to be in div now actually shows
  1981. 1:13:55the exact correct derivative
  1982. 1:13:57um so that's comforting okay so let's
  1983. 1:14:00now back propagate through this line
  1984. 1:14:01here
  1985. 1:14:03um the first thing we do of course is we
  1986. 1:14:04check the shapes and I wrote them out
  1987. 1:14:07here and basically the shape of this is
  1988. 1:14:0832 by 64. hpbn is the same shape
  1989. 1:14:12but B and mean I is a row Vector 1 by
  1990. 1:14:1564. so this minus here will actually do
  1991. 1:14:17broadcasting and so we have to be
  1992. 1:14:19careful with that and as a hint to us
  1993. 1:14:21again because of The Duality a
  1994. 1:14:23broadcasting and the forward pass means
  1995. 1:14:25a variable reuse and therefore there
  1996. 1:14:27will be a sum in the backward pass
  1997. 1:14:30so let's write out the backward pass
  1998. 1:14:31here now
  1999. 1:14:33um
  2000. 1:14:34back propagate into the hpbn
  2001. 1:14:37because this is these are the same shape
  2002. 1:14:39then the local derivative for each one
  2003. 1:14:41of the elements here is just one for the
  2004. 1:14:43corresponding element in here
  2005. 1:14:45so basically what this means is that the
  2006. 1:14:47gradient just simply copies it's just a
  2007. 1:14:50variable assignment it's quality so I'm
  2008. 1:14:52just going to clone this tensor just for
  2009. 1:14:54safety to create an exact copy of DB and
  2010. 1:14:58div
  2011. 1:15:00and then here to back propagate into
  2012. 1:15:01this one what I'm inclined to do here is
  2013. 1:15:07will basically be
  2014. 1:15:09uh what is the local derivative well
  2015. 1:15:12it's negative torch.1's like
  2016. 1:15:16of the shape of uh B and diff
  2017. 1:15:19right
  2018. 1:15:22and then times
  2019. 1:15:24the um
  2020. 1:15:27the derivative here dbf
  2021. 1:15:32and this here is the back propagation
  2022. 1:15:34for the replicated B and mean I
  2023. 1:15:37so I still have to back propagate
  2024. 1:15:39through the uh replication in the
  2025. 1:15:42broadcasting and I do that by doing a
  2026. 1:15:43sum so I'm going to take this whole
  2027. 1:15:45thing and I'm going to do a sum over the
  2028. 1:15:47zeroth dimension which was the
  2029. 1:15:49replication
  2030. 1:15:53so if you scrutinize this by the way
  2031. 1:15:55you'll notice that this is the same
  2032. 1:15:57shape as that and so what I'm doing uh
  2033. 1:16:00what I'm doing here doesn't actually
  2034. 1:16:01make that much sense because it's just a
  2035. 1:16:03array of ones multiplying DP and diff so
  2036. 1:16:06in fact I can just do this
  2037. 1:16:10um and that is equivalent
  2038. 1:16:12so this is the candidate backward pass
  2039. 1:16:15let me copy it here and then let me
  2040. 1:16:18comment out this one and this one
  2041. 1:16:22enter
  2042. 1:16:24and it's wrong
  2043. 1:16:27damn
  2044. 1:16:29actually sorry this is supposed to be
  2045. 1:16:31wrong and it's supposed to be wrong
  2046. 1:16:33because
  2047. 1:16:34we are back propagating from a b and
  2048. 1:16:36diff into hpbn and but we're not done
  2049. 1:16:39because B and mean I depends on hpbn and
  2050. 1:16:43there will be a second portion of that
  2051. 1:16:44derivative coming from this second
  2052. 1:16:46Branch so we're not done yet and we
  2053. 1:16:48expect it to be incorrect so there you
  2054. 1:16:50go
  2055. 1:16:50uh so let's now back propagate from uh B
  2056. 1:16:53and mean I into hpbn
  2057. 1:16:56um
  2058. 1:16:57and so here again we have to be careful
  2059. 1:16:58because there's a broadcasting along
  2060. 1:17:01um or there's a Sum along the zeroth
  2061. 1:17:03dimension so this will turn into
  2062. 1:17:04broadcasting in the backward pass now
  2063. 1:17:06and I'm going to go a little bit faster
  2064. 1:17:08on this line because it is very similar
  2065. 1:17:10to the line that we had before and
  2066. 1:17:12multiplies in the past in fact
  2067. 1:17:14so the hpbn
  2068. 1:17:18will be
  2069. 1:17:20the gradient will be scaled by 1 over n
  2070. 1:17:22and then basically this gradient here on
  2071. 1:17:25dbn mean I
  2072. 1:17:27is going to be scaled by 1 over n and
  2073. 1:17:30then it's going to flow across all the
  2074. 1:17:32columns and deposit itself into the hpvn
  2075. 1:17:35so what we want is this thing scaled by
  2076. 1:17:381 over n
  2077. 1:17:39only put the constant up front here
  2078. 1:17:43um
  2079. 1:17:45so scale down the gradient and now we
  2080. 1:17:47need to replicate it across all the um
  2081. 1:17:51across all the rows here so we I like to
  2082. 1:17:55do that by torch.lunslike of basically
  2083. 1:18:00um hpbn
  2084. 1:18:03and I will let the broadcasting do the
  2085. 1:18:05work of replication
  2086. 1:18:09so
  2087. 1:18:14like that
  2088. 1:18:16so this is uh the hppn and hopefully
  2089. 1:18:21we can plus equals that
  2090. 1:18:27so this here is broadcasting
  2091. 1:18:30um and then this is the scaling so this
  2092. 1:18:32should be current
  2093. 1:18:33okay
  2094. 1:18:35so that completes the back propagation
  2095. 1:18:37of the bathroom layer and we are now
  2096. 1:18:38here let's back propagate through the
  2097. 1:18:40linear layer one here now because
  2098. 1:18:43everything is getting a little
  2099. 1:18:44vertically crazy I copy pasted the line
  2100. 1:18:46here and let's just back properly
  2101. 1:18:48through this one line
  2102. 1:18:50so first of course we inspect the shapes
  2103. 1:18:52and we see that this is 32 by 64. MCAT
  2104. 1:18:56is 32 by 30.
  2105. 1:18:58W1 is 30 30 by 64 and B1 is just 64. so
  2106. 1:19:04as I mentioned back propagating through
  2107. 1:19:06linear layers is fairly easy just by
  2108. 1:19:08matching the shapes so let's do that we
  2109. 1:19:11have that dmcat
  2110. 1:19:14should be
  2111. 1:19:15um some matrix multiplication of dhbn
  2112. 1:19:18with uh W1 and one transpose thrown in
  2113. 1:19:21there so to make uh MCAT be 32 by 30
  2114. 1:19:28I need to take dhpn
  2115. 1:19:3232 by 64 and multiply it by w1.
  2116. 1:19:36transpose
  2117. 1:19:39to get the only one I need to end up
  2118. 1:19:43with 30 by 64.
  2119. 1:19:45so to get that I need to take uh MCAT
  2120. 1:19:48transpose
  2121. 1:19:51and multiply that by
  2122. 1:19:53uh dhpion
  2123. 1:19:58and finally to get DB1
  2124. 1:20:01this is a addition and we saw that
  2125. 1:20:04basically I need to just sum the
  2126. 1:20:06elements in dhpbn along some Dimension
  2127. 1:20:09and to make the dimensions work out I
  2128. 1:20:12need to Sum along the zeroth axis here
  2129. 1:20:14to eliminate this Dimension and we do
  2130. 1:20:17not keep dims
  2131. 1:20:19uh so that we want to just get a single
  2132. 1:20:21one-dimensional lecture of 64.
  2133. 1:20:23so these are the claimed derivatives
  2134. 1:20:27let me put that here and let me
  2135. 1:20:29uncomment three lines and cross our
  2136. 1:20:32fingers
  2137. 1:20:34everything is great okay so we now
  2138. 1:20:36continue almost there we have the
  2139. 1:20:37derivative of MCAT and we want to
  2140. 1:20:39derivative we want to back propagate
  2141. 1:20:41into m
  2142. 1:20:43so I again copied this line over here
  2143. 1:20:46so this is the forward pass and then
  2144. 1:20:48this is the shapes so remember that the
  2145. 1:20:51shape here was 32 by 30 and the original
  2146. 1:20:53shape of M plus 32 by 3 by 10. so this
  2147. 1:20:57layer in the forward pass as you recall
  2148. 1:20:58did the concatenation of these three
  2149. 1:21:0110-dimensional character vectors
  2150. 1:21:04and so now we just want to undo that
  2151. 1:21:06so this is actually relatively
  2152. 1:21:08straightforward operation because uh the
  2153. 1:21:11backward pass of the what is the view
  2154. 1:21:12view is just a representation of the
  2155. 1:21:15array it's just a logical form of how
  2156. 1:21:17you interpret the array so let's just
  2157. 1:21:18reinterpret it to be what it was before
  2158. 1:21:21so in other words the end is not uh 32
  2159. 1:21:25by 30. it is basically dmcat
  2160. 1:21:29but if you view it as the original shape
  2161. 1:21:34so just m dot shape
  2162. 1:21:37uh you can you can pass in tuples into
  2163. 1:21:39view
  2164. 1:21:40and so this should just be okay
  2165. 1:21:44we just re-represent that view and then
  2166. 1:21:47we uncomment this line here and
  2167. 1:21:49hopefully
  2168. 1:21:51yeah so the derivative of M is correct
  2169. 1:21:55so in this case we just have to
  2170. 1:21:56re-represent the shape of those
  2171. 1:21:57derivatives into the original View
  2172. 1:21:59so now we are at the final line and the
  2173. 1:22:01only thing that's left to back propagate
  2174. 1:22:02through is this indexing operation here
  2175. 1:22:05MSC at xB so as I did before I copy
  2176. 1:22:09pasted this line here and let's look at
  2177. 1:22:11the shapes of everything that's involved
  2178. 1:22:12and remind ourselves how this worked
  2179. 1:22:15so m.shape was 32 by 3 by 10.
  2180. 1:22:19it says 32 examples and then we have
  2181. 1:22:22three characters each one of them has a
  2182. 1:22:2410 dimensional embedding
  2183. 1:22:26and this was achieved by taking the
  2184. 1:22:28lookup table C which have 27 possible
  2185. 1:22:31characters
  2186. 1:22:32each of them 10 dimensional and we
  2187. 1:22:34looked up
  2188. 1:22:35at the rows that were specified inside
  2189. 1:22:39this tensor xB
  2190. 1:22:41so XB is 32 by 3 and it's basically
  2191. 1:22:43giving us for each example the Identity
  2192. 1:22:45or the index of which character is part
  2193. 1:22:49of that example
  2194. 1:22:50and so here I'm showing the first five
  2195. 1:22:52rows of three of this tensor xB
  2196. 1:22:57and so we can see that for example here
  2197. 1:22:58it was the first example in this batch
  2198. 1:23:00is that the first character and the
  2199. 1:23:02first character and the fourth character
  2200. 1:23:04comes into the neural net
  2201. 1:23:06and then we want to predict the next
  2202. 1:23:08character in a sequence after the
  2203. 1:23:10character is one one four
  2204. 1:23:12so basically What's Happening Here is
  2205. 1:23:14there are integers inside XB and each
  2206. 1:23:18one of these integers is specifying
  2207. 1:23:19which row of C we want to pluck out
  2208. 1:23:22right and then we arrange those rows
  2209. 1:23:25that we've plucked out into 32 by 3 by
  2210. 1:23:2810 tensor and we just package them in we
  2211. 1:23:30just package them into the sensor
  2212. 1:23:33and now what's happening is that we have
  2213. 1:23:35D amp
  2214. 1:23:36so for every one of these uh basically
  2215. 1:23:39plucked out rows we have their gradients
  2216. 1:23:41now
  2217. 1:23:42but they're arranged inside this 32 by 3
  2218. 1:23:45by 10 tensor so all we have to do now is
  2219. 1:23:48we just need to Route this gradient
  2220. 1:23:49backwards through this assignment so we
  2221. 1:23:52need to find which row of C that every
  2222. 1:23:54one of these
  2223. 1:23:56um 10 dimensional embeddings come from
  2224. 1:23:59and then we need to deposit them into DC
  2225. 1:24:03so we just need to undo the indexing and
  2226. 1:24:06of course if any of these rows of C was
  2227. 1:24:08used multiple times which almost
  2228. 1:24:10certainly is the case like the row one
  2229. 1:24:11and one was used multiple times then we
  2230. 1:24:13have to remember that the gradients that
  2231. 1:24:15arrive there have to add
  2232. 1:24:18so for each occurrence we have to have
  2233. 1:24:19an addition
  2234. 1:24:21so let's now write this out and I don't
  2235. 1:24:23actually know if like a much better way
  2236. 1:24:24to do this than a for Loop unfortunately
  2237. 1:24:26in Python
  2238. 1:24:28um so maybe someone can come up with a
  2239. 1:24:29vectorized efficient operation but for
  2240. 1:24:32now let's just use for loops so let me
  2241. 1:24:34create a torch.zeros like
  2242. 1:24:37C to initialize uh just uh 27 by 10
  2243. 1:24:40tensor of all zeros
  2244. 1:24:43and then honestly 4K in range XB dot
  2245. 1:24:46shape at zero
  2246. 1:24:49maybe someone has a better way to do
  2247. 1:24:51this but for J and range
  2248. 1:24:53be that shape at one
  2249. 1:24:55this is going to iterate over all the
  2250. 1:24:58um all the elements of XB all these
  2251. 1:25:01integers
  2252. 1:25:03and then let's get the index at this
  2253. 1:25:05position
  2254. 1:25:06so the index is basically x b at KJ
  2255. 1:25:11so that an example of that like is 11 or
  2256. 1:25:1414 and so on
  2257. 1:25:16and now in the forward pass we took
  2258. 1:25:19and we basically took um
  2259. 1:25:24the row of C at index and we deposited
  2260. 1:25:27it into M at K of J
  2261. 1:25:30that's what happened that's where they
  2262. 1:25:32are packaged so now we need to go
  2263. 1:25:34backwards and we just need to route
  2264. 1:25:36DM at the position KJ
  2265. 1:25:39we now have these derivatives
  2266. 1:25:42for each position and it's 10
  2267. 1:25:44dimensional
  2268. 1:25:45and you just need to go into the correct
  2269. 1:25:47row of C
  2270. 1:25:49so DC rather at IX is this but plus
  2271. 1:25:54equals
  2272. 1:25:55because there could be multiple
  2273. 1:25:56occurrences uh like the same row could
  2274. 1:25:58have been used many many times and so
  2275. 1:26:00all of those derivatives will just go
  2276. 1:26:04backwards through the indexing and they
  2277. 1:26:06will add
  2278. 1:26:07so this is my candidate solution
  2279. 1:26:12let's copy it here
  2280. 1:26:16let's uncomment this and cross our
  2281. 1:26:19fingers
  2282. 1:26:20hey
  2283. 1:26:21so that's it we've back propagated
  2284. 1:26:24through
  2285. 1:26:25this entire Beast
  2286. 1:26:28so there we go totally makes sense
  2287. 1:26:31so now we come to exercise two it
  2288. 1:26:33basically turns out that in this first
  2289. 1:26:34exercise we were doing way too much work
  2290. 1:26:36uh we were back propagating way too much
  2291. 1:26:39and it was all good practice and so on
  2292. 1:26:40but it's not what you would do in
  2293. 1:26:42practice and the reason for that is for
  2294. 1:26:44example here I separated out this loss
  2295. 1:26:47calculation over multiple lines and I
  2296. 1:26:49broke it up all all to like its smallest
  2297. 1:26:51atomic pieces and we back propagated
  2298. 1:26:53through all of those individually
  2299. 1:26:55but it turns out that if you just look
  2300. 1:26:56at the mathematical expression for the
  2301. 1:26:58loss
  2302. 1:27:00um then actually you can do the
  2303. 1:27:02differentiation on pen and paper and a
  2304. 1:27:04lot of terms cancel and simplify and the
  2305. 1:27:06mathematical expression you end up with
  2306. 1:27:07can be significantly shorter and easier
  2307. 1:27:10to implement than back propagating
  2308. 1:27:11through all the little pieces of
  2309. 1:27:12everything you've done
  2310. 1:27:13so before we had this complicated
  2311. 1:27:16forward paths going from logits to the
  2312. 1:27:18loss
  2313. 1:27:19but in pytorch everything can just be
  2314. 1:27:21glued together into a single call at
  2315. 1:27:22that cross entropy you just pass in
  2316. 1:27:24logits and the labels and you get the
  2317. 1:27:26exact same loss as I verify here so our
  2318. 1:27:28previous loss and the fast loss coming
  2319. 1:27:31from the chunk of operations as a single
  2320. 1:27:33mathematical expression is the same but
  2321. 1:27:36it's much much faster in a forward pass
  2322. 1:27:38it's also much much faster in backward
  2323. 1:27:40pass and the reason for that is if you
  2324. 1:27:42just look at the mathematical form of
  2325. 1:27:43this and differentiate again you will
  2326. 1:27:45end up with a very small and short
  2327. 1:27:46expression so that's what we want to do
  2328. 1:27:48here we want to in a single operation or
  2329. 1:27:51in a single go or like very quickly go
  2330. 1:27:54directly to delojits
  2331. 1:27:56and we need to implement the logits as a
  2332. 1:27:59function of logits and yb's
  2333. 1:28:02but it will be significantly shorter
  2334. 1:28:04than whatever we did here where to get
  2335. 1:28:06to deluggets we had to go all the way
  2336. 1:28:08here
  2337. 1:28:10so all of this work can be skipped in a
  2338. 1:28:12much much simpler mathematical
  2339. 1:28:13expression that you can Implement here
  2340. 1:28:16so you can give it a shot yourself
  2341. 1:28:18basically look at what exactly is the
  2342. 1:28:21mathematical expression of loss and
  2343. 1:28:23differentiate with respect to the logits
  2344. 1:28:26so let me show you a hint you can of
  2345. 1:28:29course try it fully yourself but if not
  2346. 1:28:31I can give you some hint of how to get
  2347. 1:28:33started mathematically
  2348. 1:28:36so basically What's Happening Here is we
  2349. 1:28:38have logits then there's a softmax that
  2350. 1:28:41takes the logits and gives you
  2351. 1:28:42probabilities then we are using the
  2352. 1:28:44identity of the correct next character
  2353. 1:28:46to pluck out a row of probabilities take
  2354. 1:28:50the negative log of it to get our
  2355. 1:28:51negative block probability and then we
  2356. 1:28:54average up all the log probabilities or
  2357. 1:28:56negative block probabilities to get our
  2358. 1:28:58loss
  2359. 1:28:59so basically what we have is for a
  2360. 1:29:01single individual example rather we have
  2361. 1:29:04that loss is equal to negative log
  2362. 1:29:06probability uh where P here is kind of
  2363. 1:29:09like thought of as a vector of all the
  2364. 1:29:11probabilities so at the Y position where
  2365. 1:29:14Y is the label
  2366. 1:29:16and we have that P here of course is the
  2367. 1:29:19softmax so the ith component of P of
  2368. 1:29:23this probability Vector is just the
  2369. 1:29:25softmax function so raising all the
  2370. 1:29:28logits uh basically to the power of E
  2371. 1:29:31and normalizing so everything comes to
  2372. 1:29:341.
  2373. 1:29:35now if you write out P of Y here you can
  2374. 1:29:38just write out the soft Max and then
  2375. 1:29:40basically what we're interested in is
  2376. 1:29:41we're interested in the derivative of
  2377. 1:29:43the loss with respect to the I logit
  2378. 1:29:47and so basically it's a d by DLI of this
  2379. 1:29:51expression here
  2380. 1:29:52where we have L indexed with the
  2381. 1:29:54specific label Y and on the bottom we
  2382. 1:29:56have a sum over J of e to the L J and
  2383. 1:29:58the negative block of all that so
  2384. 1:30:00potentially give it a shot pen and paper
  2385. 1:30:02and see if you can actually derive the
  2386. 1:30:04expression for the loss by DLI and then
  2387. 1:30:07we're going to implement it here okay so
  2388. 1:30:09I'm going to give away the result here
  2389. 1:30:11so this is some of the math I did to
  2390. 1:30:13derive the gradients analytically and so
  2391. 1:30:17we see here that I'm just applying the
  2392. 1:30:19rules of calculus from your first or
  2393. 1:30:20second year of bachelor's degree if you
  2394. 1:30:22took it and we see that the expression
  2395. 1:30:24is actually simplify quite a bit you
  2396. 1:30:26have to separate out the analysis in the
  2397. 1:30:27case where the ith index that you're
  2398. 1:30:30interested in inside logits is either
  2399. 1:30:32equal to the label or it's not equal to
  2400. 1:30:34the label and then the expression
  2401. 1:30:35simplify and cancel in a slightly
  2402. 1:30:37different way and what we end up with is
  2403. 1:30:39something very very simple
  2404. 1:30:41and we either end up with basically
  2405. 1:30:43pirai where p is again this Vector of
  2406. 1:30:46probabilities after a soft Max or P at I
  2407. 1:30:49minus 1 where we just simply subtract a
  2408. 1:30:51one but in any case we just need to
  2409. 1:30:53calculate the soft Max p e and then in
  2410. 1:30:56the correct Dimension we need to
  2411. 1:30:58subtract one and that's the gradient the
  2412. 1:31:00form that it takes analytically so let's
  2413. 1:31:03implement this basically and we have to
  2414. 1:31:04keep in mind that this is only done for
  2415. 1:31:06a single example but here we are working
  2416. 1:31:08with batches of examples
  2417. 1:31:09so we have to be careful of that and
  2418. 1:31:12then the loss for a batch is the average
  2419. 1:31:14loss over all the examples so in other
  2420. 1:31:17words is the example for all the
  2421. 1:31:18individual examples is the loss for each
  2422. 1:31:20individual example summed up and then
  2423. 1:31:22divided by n and we have to back
  2424. 1:31:24propagate through that as well and be
  2425. 1:31:26careful with it
  2426. 1:31:28so deluggets is going to be of that soft
  2427. 1:31:30Max
  2428. 1:31:32uh pytorch has a softmax function that
  2429. 1:31:35you can call and we want to apply the
  2430. 1:31:36softmax on the logits and we want to go
  2431. 1:31:39in the dimension that is one so
  2432. 1:31:42basically we want to do the softmax
  2433. 1:31:44along the rows of these logits
  2434. 1:31:47then at the correct positions we need to
  2435. 1:31:49subtract a 1. so delugits at iterating
  2436. 1:31:52over all the rows
  2437. 1:31:54and indexing into the columns
  2438. 1:31:57provided by the correct labels inside YB
  2439. 1:32:00we need to subtract one
  2440. 1:32:03and then finally it's the average loss
  2441. 1:32:05that is the loss and in the average
  2442. 1:32:07there's a one over n of all the losses
  2443. 1:32:09added up and so we need to also
  2444. 1:32:12propagate through that division
  2445. 1:32:14so the gradient has to be scaled down by
  2446. 1:32:16by n as well because of the mean
  2447. 1:32:19but this otherwise should be the result
  2448. 1:32:22so now if we verify this
  2449. 1:32:24we see that we don't get an exact match
  2450. 1:32:26but at the same time the maximum
  2451. 1:32:30difference from logits from pytorch and
  2452. 1:32:33RD logits here is uh on the order of 5e
  2453. 1:32:37negative 9. so it's a tiny tiny number
  2454. 1:32:39so because of floating point wantiness
  2455. 1:32:41we don't get the exact bitwise result
  2456. 1:32:44but we basically get the correct answer
  2457. 1:32:47approximately
  2458. 1:32:49now I'd like to pause here briefly
  2459. 1:32:51before we move on to the next exercise
  2460. 1:32:52because I'd like us to get an intuitive
  2461. 1:32:54sense of what the logits is because it
  2462. 1:32:56has a beautiful and very simple
  2463. 1:32:58explanation honestly
  2464. 1:33:00um so here I'm taking the logits and I'm
  2465. 1:33:03visualizing it and we can see that we
  2466. 1:33:05have a batch of 32 examples of 27
  2467. 1:33:07characters
  2468. 1:33:08and what is the logits intuitively right
  2469. 1:33:10the logits is the probabilities that the
  2470. 1:33:13properties Matrix in the forward pass
  2471. 1:33:15but then here these black squares are
  2472. 1:33:17the positions of the correct indices
  2473. 1:33:19where we subtracted a one
  2474. 1:33:21and so uh what is this doing right these
  2475. 1:33:24are the derivatives on the logits and so
  2476. 1:33:27let's look at just the first row here
  2477. 1:33:31so that's what I'm doing here I'm
  2478. 1:33:33clocking the probabilities of these
  2479. 1:33:34logits and then I'm taking just the
  2480. 1:33:36first row and this is the probability
  2481. 1:33:38row and then the logits of the first row
  2482. 1:33:41and multiplying by n just for us so that
  2483. 1:33:43we don't have the scaling by n in here
  2484. 1:33:46and everything is more interpretable we
  2485. 1:33:48see that it's exactly equal to the
  2486. 1:33:50probability of course but then the
  2487. 1:33:52position of the correct index has a
  2488. 1:33:53minus equals one so minus one on that
  2489. 1:33:56position
  2490. 1:33:57and so notice that
  2491. 1:33:59um if you take Delo Jets at zero and you
  2492. 1:34:01sum it
  2493. 1:34:03it actually sums to zero and so you
  2494. 1:34:06should think of these uh gradients here
  2495. 1:34:08at each cell as like a force
  2496. 1:34:12um we are going to be basically pulling
  2497. 1:34:15down on the probabilities of the
  2498. 1:34:17incorrect characters and we're going to
  2499. 1:34:19be pulling up on the probability at the
  2500. 1:34:22correct index and that's what's
  2501. 1:34:24basically happening in each row and thus
  2502. 1:34:29the amount of push and pull is exactly
  2503. 1:34:31equalized because the sum is zero so the
  2504. 1:34:34amount to which we pull down in the
  2505. 1:34:36probabilities and the demand that we
  2506. 1:34:37push up on the probability of the
  2507. 1:34:39correct character is equal
  2508. 1:34:41so sort of the the repulsion and the
  2509. 1:34:43attraction are equal and think of the
  2510. 1:34:45neural app now as a like a massive uh
  2511. 1:34:48pulley system or something like that
  2512. 1:34:50we're up here on top of the logits and
  2513. 1:34:52we're pulling up we're pulling down the
  2514. 1:34:54properties of Incorrect and pulling up
  2515. 1:34:55the property of the correct and in this
  2516. 1:34:57complicated pulley system because
  2517. 1:34:59everything is mathematically uh just
  2518. 1:35:01determined just think of it as sort of
  2519. 1:35:03like this tension translating to this
  2520. 1:35:05complicating pulling mechanism and then
  2521. 1:35:07eventually we get a tug on the weights
  2522. 1:35:09and the biases and basically in each
  2523. 1:35:11update we just kind of like tug in the
  2524. 1:35:13direction that we like for each of these
  2525. 1:35:15elements and the parameters are slowly
  2526. 1:35:17given in to the tug and that's what
  2527. 1:35:19training in neural net kind of like
  2528. 1:35:20looks like on a high level
  2529. 1:35:22and so I think the the forces of push
  2530. 1:35:24and pull in these gradients are actually
  2531. 1:35:26uh very intuitive here we're pushing and
  2532. 1:35:29pulling on the correct answer and the
  2533. 1:35:31incorrect answers and the amount of
  2534. 1:35:33force that we're applying is actually
  2535. 1:35:34proportional to uh the probabilities
  2536. 1:35:37that came out in the forward pass
  2537. 1:35:39and so for example if our probabilities
  2538. 1:35:41came out exactly correct so they would
  2539. 1:35:43have had zero everywhere except for one
  2540. 1:35:45at the correct uh position then the the
  2541. 1:35:48logits would be all a row of zeros for
  2542. 1:35:51that example there would be no push and
  2543. 1:35:52pull so the amount to which your
  2544. 1:35:55prediction is incorrect is exactly the
  2545. 1:35:58amount by which you're going to get a
  2546. 1:35:59pull or a push in that dimension
  2547. 1:36:01so if you have for example a very
  2548. 1:36:04confidently mispredicted element here
  2549. 1:36:05then
  2550. 1:36:07um what's going to happen is that
  2551. 1:36:08element is going to be pulled down very
  2552. 1:36:10heavily and the correct answer is going
  2553. 1:36:12to be pulled up to the same amount
  2554. 1:36:14and the other characters are not going
  2555. 1:36:16to be influenced too much
  2556. 1:36:19so the amounts to which you mispredict
  2557. 1:36:21is then proportional to the strength of
  2558. 1:36:23the pole and that's happening
  2559. 1:36:25independently in all the dimensions of
  2560. 1:36:27this of this tensor and it's sort of
  2561. 1:36:29very intuitive and varies to think
  2562. 1:36:30through and that's basically the magic
  2563. 1:36:32of the cross-entropy loss and what it's
  2564. 1:36:34doing dynamically in the backward pass
  2565. 1:36:36of the neural net so now we get to
  2566. 1:36:38exercise number three which is a very
  2567. 1:36:41fun exercise
  2568. 1:36:42um depending on your definition of fun
  2569. 1:36:43and we are going to do for batch
  2570. 1:36:45normalization exactly what we did for
  2571. 1:36:47cross entropy loss in exercise number
  2572. 1:36:49two that is we are going to consider it
  2573. 1:36:51as a glued single mathematical
  2574. 1:36:52expression and back propagate through it
  2575. 1:36:54in a very efficient manner because we
  2576. 1:36:56are going to derive a much simpler
  2577. 1:36:58formula for the backward path of batch
  2578. 1:36:59normalization
  2579. 1:37:01and we're going to do that using pen and
  2580. 1:37:02paper
  2581. 1:37:03so previously we've broken up
  2582. 1:37:05bastionalization into all of the little
  2583. 1:37:06intermediate pieces and all the atomic
  2584. 1:37:08operations inside it and then we back
  2585. 1:37:10propagate it through it one by one
  2586. 1:37:13now we just have a single sort of
  2587. 1:37:15forward pass of a batch form and it's
  2588. 1:37:18all glued together
  2589. 1:37:20and we see that we get the exact same
  2590. 1:37:21result as before
  2591. 1:37:23now for the backward pass we'd like to
  2592. 1:37:25also Implement a single formula
  2593. 1:37:27basically for back propagating through
  2594. 1:37:29this entire operation that is the
  2595. 1:37:30bachelorization
  2596. 1:37:32so in the forward pass previously we
  2597. 1:37:34took hpvn the hidden states of the
  2598. 1:37:37pre-batch realization and created H
  2599. 1:37:39preact which is the hidden States just
  2600. 1:37:42before the activation
  2601. 1:37:44in the bachelorization paper each pbn is
  2602. 1:37:46X and each preact is y
  2603. 1:37:49so in the backward pass what we'd like
  2604. 1:37:51to do now is we have DH preact and we'd
  2605. 1:37:54like to produce d h previous
  2606. 1:37:56and we'd like to do that in a very
  2607. 1:37:57efficient manner so that's the name of
  2608. 1:38:00the game calculate the H previan given
  2609. 1:38:02DH preact and for the purposes of this
  2610. 1:38:05exercise we're going to ignore gamma and
  2611. 1:38:07beta and their derivatives because they
  2612. 1:38:09take on a very simple form in a very
  2613. 1:38:11similar way to what we did up above
  2614. 1:38:14so let's calculate this given that right
  2615. 1:38:18here
  2616. 1:38:18so to help you a little bit like I did
  2617. 1:38:20before I started off the implementation
  2618. 1:38:23here on pen and paper and I took two
  2619. 1:38:26sheets of paper to derive the
  2620. 1:38:28mathematical formulas for the backward
  2621. 1:38:29pass
  2622. 1:38:30and basically to set up the problem uh
  2623. 1:38:33just write out the MU Sigma Square
  2624. 1:38:35variance x i hat and Y I exactly as in
  2625. 1:38:39the paper except for the bezel
  2626. 1:38:40correction
  2627. 1:38:41and then
  2628. 1:38:42in a backward pass we have the
  2629. 1:38:44derivative of the loss with respect to
  2630. 1:38:46all the elements of Y and remember that
  2631. 1:38:48Y is a vector there's there's multiple
  2632. 1:38:50numbers here
  2633. 1:38:52so we have all the derivatives with
  2634. 1:38:54respect to all the Y's
  2635. 1:38:56and then there's a demo and a beta and
  2636. 1:38:59this is kind of like the compute graph
  2637. 1:39:01the gamma and the beta there's the X hat
  2638. 1:39:03and then the MU and the sigma squared
  2639. 1:39:06and the X so we have DL by DYI and we
  2640. 1:39:10won't DL by d x i for all the I's in
  2641. 1:39:13these vectors
  2642. 1:39:15so this is the compute graph and you
  2643. 1:39:17have to be careful because I'm trying to
  2644. 1:39:19note here that these are vectors so
  2645. 1:39:22there's many nodes here inside x x hat
  2646. 1:39:25and Y but mu and sigma sorry Sigma
  2647. 1:39:29Square are just individual scalars
  2648. 1:39:30single numbers so you have to be careful
  2649. 1:39:33with that you have to imagine there's
  2650. 1:39:34multiple nodes here or you're going to
  2651. 1:39:35get your math wrong
  2652. 1:39:38um so as an example I would suggest that
  2653. 1:39:40you go in the following order one two
  2654. 1:39:43three four in terms of the back
  2655. 1:39:44propagation so back propagating to X hat
  2656. 1:39:46then into Sigma Square then into mu and
  2657. 1:39:49then into X
  2658. 1:39:52um just like in a topological sort in
  2659. 1:39:54micrograd we would go from right to left
  2660. 1:39:55you're doing the exact same thing except
  2661. 1:39:57you're doing it with symbols and on a
  2662. 1:39:59piece of paper
  2663. 1:40:01so for number one uh I'm not giving away
  2664. 1:40:05too much if you want DL of d x i hat
  2665. 1:40:09then we just take DL by DYI and multiply
  2666. 1:40:12it by gamma because of this expression
  2667. 1:40:15here where any individual Yi is just
  2668. 1:40:17gamma times x i hat plus beta so it
  2669. 1:40:21doesn't help you too much there but this
  2670. 1:40:23gives you basically the derivatives for
  2671. 1:40:25all the X hats and so now try to go
  2672. 1:40:28through this computational graph and
  2673. 1:40:31derive what is DL by D Sigma Square
  2674. 1:40:35and then what is DL by B mu and then one
  2675. 1:40:38is D L by DX
  2676. 1:40:39eventually so give it a go and I'm going
  2677. 1:40:42to be revealing the answer one piece at
  2678. 1:40:44a time okay so to get DL by D Sigma
  2679. 1:40:46Square we have to remember again like I
  2680. 1:40:48mentioned that there are many excess X
  2681. 1:40:51hats here
  2682. 1:40:52and remember that Sigma square is just a
  2683. 1:40:54single individual number here
  2684. 1:40:55so when we look at the expression
  2685. 1:40:59for the L by D Sigma Square
  2686. 1:41:01we have that we have to actually
  2687. 1:41:03consider all the possible paths that um
  2688. 1:41:08we basically have that there's many X
  2689. 1:41:10hats and they all feed off from they all
  2690. 1:41:13depend on Sigma Square so Sigma square
  2691. 1:41:15has a large fan out there's lots of
  2692. 1:41:17arrows coming out from Sigma square into
  2693. 1:41:19all the X hats
  2694. 1:41:20and then there's a back propagating
  2695. 1:41:22signal from each X hat into Sigma square
  2696. 1:41:24and that's why we actually need to sum
  2697. 1:41:26over all those I's from I equal to 1 to
  2698. 1:41:29m
  2699. 1:41:30of the DL by d x i hat which is the
  2700. 1:41:35global gradient
  2701. 1:41:36times the x i Hat by D Sigma Square
  2702. 1:41:40which is the local gradient
  2703. 1:41:42of this operation here
  2704. 1:41:44and then mathematically I'm just working
  2705. 1:41:46it out here and I'm simplifying and you
  2706. 1:41:48get a certain expression for DL by D
  2707. 1:41:51Sigma square and we're going to be using
  2708. 1:41:52this expression when we back propagate
  2709. 1:41:53into mu and then eventually into X so
  2710. 1:41:56now let's continue our back propagation
  2711. 1:41:58into mu so what is D L by D mu now again
  2712. 1:42:01be careful that mu influences X hat and
  2713. 1:42:04X hat is actually lots of values so for
  2714. 1:42:07example if our mini batch size is 32 as
  2715. 1:42:09it is in our example that we were
  2716. 1:42:10working on then this is 32 numbers and
  2717. 1:42:1332 arrows going back to mu and then mu
  2718. 1:42:16going to Sigma square is just a single
  2719. 1:42:18Arrow because Sigma square is a scalar
  2720. 1:42:19so in total there are 33 arrows
  2721. 1:42:22emanating from you and then all of them
  2722. 1:42:25have gradients coming into mu and they
  2723. 1:42:27all need to be summed up
  2724. 1:42:29and so that's why when we look at the
  2725. 1:42:31expression for DL by D mu I am summing
  2726. 1:42:34up over all the gradients of DL by d x i
  2727. 1:42:37hat times the x i Hat by being mu
  2728. 1:42:40uh so that's the that's this arrow and
  2729. 1:42:43that's 32 arrows here and then plus the
  2730. 1:42:45one Arrow from here which is the L by
  2731. 1:42:47the sigma Square Times the sigma squared
  2732. 1:42:49by D mu
  2733. 1:42:50so now we have to work out that
  2734. 1:42:52expression and let me just reveal the
  2735. 1:42:54rest of it
  2736. 1:42:55uh simplifying here is not complicated
  2737. 1:42:58the first term and you just get an
  2738. 1:43:00expression here
  2739. 1:43:01for the second term though there's
  2740. 1:43:02something really interesting that
  2741. 1:43:03happens
  2742. 1:43:04when we look at the sigma squared by D
  2743. 1:43:06mu and we simplify
  2744. 1:43:08at one point if we assume that in a
  2745. 1:43:11special case where mu is actually the
  2746. 1:43:14average of X I's as it is in this case
  2747. 1:43:17then if we plug that in then actually
  2748. 1:43:20the gradient vanishes and becomes
  2749. 1:43:22exactly zero and that makes the entire
  2750. 1:43:24second term cancel
  2751. 1:43:26and so these uh if you just have a
  2752. 1:43:29mathematical expression like this and
  2753. 1:43:30you look at D Sigma Square by D mu you
  2754. 1:43:33would get some mathematical formula for
  2755. 1:43:35how mu impacts Sigma Square
  2756. 1:43:37but if it is the special case that Nu is
  2757. 1:43:39actually equal to the average as it is
  2758. 1:43:42in the case of pastoralization that
  2759. 1:43:43gradient will actually vanish and become
  2760. 1:43:45zero so the whole term cancels and we
  2761. 1:43:48just get a fairly straightforward
  2762. 1:43:49expression here for DL by D mu okay and
  2763. 1:43:52now we get to the craziest part which is
  2764. 1:43:54uh deriving DL by dxi which is
  2765. 1:43:57ultimately what we're after
  2766. 1:43:59now let's count
  2767. 1:44:00first of all how many numbers are there
  2768. 1:44:03inside X as I mentioned there are 32
  2769. 1:44:05numbers there are 32 Little X I's and
  2770. 1:44:08let's count the number of arrows
  2771. 1:44:09emanating from each x i
  2772. 1:44:11there's an arrow going to Mu an arrow
  2773. 1:44:13going to Sigma Square
  2774. 1:44:14and then there's an arrow going to X hat
  2775. 1:44:16but this Arrow here let's scrutinize
  2776. 1:44:19that a little bit
  2777. 1:44:20each x i hat is just a function of x i
  2778. 1:44:23and all the other scalars so x i hat
  2779. 1:44:27only depends on x i and none of the
  2780. 1:44:29other X's
  2781. 1:44:30and so therefore there are actually in
  2782. 1:44:32this single Arrow there are 32 arrows
  2783. 1:44:34but those 32 arrows are going exactly
  2784. 1:44:37parallel they don't interfere and
  2785. 1:44:39they're just going parallel between x
  2786. 1:44:40and x hat you can look at it that way
  2787. 1:44:42and so how many arrows are emanating
  2788. 1:44:44from each x i there are three arrows mu
  2789. 1:44:47Sigma squared and the associated X hat
  2790. 1:44:50and so in back propagation we now need
  2791. 1:44:53to apply the chain rule and we need to
  2792. 1:44:55add up those three contributions
  2793. 1:44:57so here's what that looks like if I just
  2794. 1:44:59write that out
  2795. 1:45:02we have uh we're going through we're
  2796. 1:45:04chaining through mu Sigma square and
  2797. 1:45:06through X hat and those three terms are
  2798. 1:45:09just here
  2799. 1:45:10now we already have three of these we
  2800. 1:45:13have d l by d x i hat
  2801. 1:45:15we have DL by D mu which we derived here
  2802. 1:45:17and we have DL by D Sigma Square which
  2803. 1:45:19we derived here but we need three other
  2804. 1:45:22terms here
  2805. 1:45:23the this one this one and this one so I
  2806. 1:45:26invite you to try to derive them it's
  2807. 1:45:28not that complicated you're just looking
  2808. 1:45:29at these Expressions here and
  2809. 1:45:31differentiating with respect to x i
  2810. 1:45:34so give it a shot but here's the result
  2811. 1:45:39or at least what I got
  2812. 1:45:41um
  2813. 1:45:42yeah I'm just I'm just differentiating
  2814. 1:45:44with respect to x i for all these
  2815. 1:45:45expressions and honestly I don't think
  2816. 1:45:47there's anything too tricky here it's
  2817. 1:45:48basic calculus
  2818. 1:45:50now it gets a little bit more tricky is
  2819. 1:45:52we are now going to plug everything
  2820. 1:45:53together so all of these terms
  2821. 1:45:55multiplied with all of these terms and
  2822. 1:45:57add it up according to this formula and
  2823. 1:45:59that gets a little bit hairy so what
  2824. 1:46:01ends up happening is
  2825. 1:46:04uh
  2826. 1:46:05you get a large expression and the thing
  2827. 1:46:08to be very careful with here of course
  2828. 1:46:09is we are working with a DL by dxi for
  2829. 1:46:12specific I here but when we are plugging
  2830. 1:46:15in some of these terms
  2831. 1:46:17like say
  2832. 1:46:18um
  2833. 1:46:19this term here deal by D signal squared
  2834. 1:46:22you see how the L by D Sigma squared I
  2835. 1:46:24end up with an expression and I'm
  2836. 1:46:26iterating over little I's here but I
  2837. 1:46:29can't use I as the variable when I plug
  2838. 1:46:31in here because this is a different I
  2839. 1:46:33from this eye
  2840. 1:46:35this I here is just a place or like a
  2841. 1:46:37local variable for for a for Loop in
  2842. 1:46:39here so here when I plug that in you
  2843. 1:46:41notice that I rename the I to a j
  2844. 1:46:43because I need to make sure that this J
  2845. 1:46:45is not that this J is not this I this J
  2846. 1:46:48is like like a little local iterator
  2847. 1:46:50over 32 terms and so you have to be
  2848. 1:46:53careful with that when you're plugging
  2849. 1:46:54in the expressions from here to here you
  2850. 1:46:56may have to rename eyes into J's and you
  2851. 1:46:58have to be very careful what is actually
  2852. 1:47:00an I with respect to the L by t x i
  2853. 1:47:04so some of these are J's some of these
  2854. 1:47:07are I's
  2855. 1:47:08and then we simplify this expression
  2856. 1:47:11and I guess like the big thing to notice
  2857. 1:47:13here is a bunch of terms just kind of
  2858. 1:47:15come out to the front and you can
  2859. 1:47:16refactor them there's a sigma squared
  2860. 1:47:18plus Epsilon raised to the power of
  2861. 1:47:19negative three over two uh this Sigma
  2862. 1:47:21squared plus Epsilon can be actually
  2863. 1:47:23separated out into three terms each of
  2864. 1:47:25them are Sigma squared plus Epsilon to
  2865. 1:47:28the negative one over two so the three
  2866. 1:47:30of them multiplied is equal to this and
  2867. 1:47:33then those three terms can go different
  2868. 1:47:35places because of the multiplication so
  2869. 1:47:37one of them actually comes out to the
  2870. 1:47:39front and will end up here outside one
  2871. 1:47:42of them joins up with this term and one
  2872. 1:47:45of them joins up with this other term
  2873. 1:47:47and then when you simplify the
  2874. 1:47:49expression you'll notice that some of
  2875. 1:47:51these terms that are coming out are just
  2876. 1:47:52the x i hats
  2877. 1:47:54so you can simplify just by rewriting
  2878. 1:47:56that
  2879. 1:47:57and what we end up with at the end is a
  2880. 1:47:58fairly simple mathematical expression
  2881. 1:48:00over here that I cannot simplify further
  2882. 1:48:02but basically you'll notice that it only
  2883. 1:48:05uses the stuff we have and it derives
  2884. 1:48:06the thing we need so we have the L by d
  2885. 1:48:10y for all the I's and those are used
  2886. 1:48:13plenty of times here and also in
  2887. 1:48:15addition what we're using is these x i
  2888. 1:48:17hats and XJ hats and they just come from
  2889. 1:48:19the forward pass
  2890. 1:48:20and otherwise this is a simple
  2891. 1:48:22expression and it gives us DL by d x i
  2892. 1:48:25for all the I's and that's ultimately
  2893. 1:48:27what we're interested in
  2894. 1:48:29so that's the end of Bachelor backward
  2895. 1:48:32pass analytically let's now implement
  2896. 1:48:34this final result
  2897. 1:48:36okay so I implemented the expression
  2898. 1:48:38into a single line of code here and you
  2899. 1:48:41can see that the max diff is Tiny so
  2900. 1:48:43this is the correct implementation of
  2901. 1:48:44this formula now I'll just uh
  2902. 1:48:48basically tell you that getting this
  2903. 1:48:50formula here from this mathematical
  2904. 1:48:52expression was not trivial and there's a
  2905. 1:48:54lot going on packed into this one
  2906. 1:48:56formula and this is a whole exercise by
  2907. 1:48:58itself because you have to consider the
  2908. 1:49:00fact that this formula here is just for
  2909. 1:49:03a single neuron and a batch of 32
  2910. 1:49:05examples but what I'm doing here is I'm
  2911. 1:49:07actually we actually have 64 neurons and
  2912. 1:49:10so this expression has to in parallel
  2913. 1:49:11evaluate the bathroom backward pass for
  2914. 1:49:14all of those 64 neurons in parallel
  2915. 1:49:16independently so this has to happen
  2916. 1:49:18basically in every single
  2917. 1:49:20um
  2918. 1:49:20column of the inputs here
  2919. 1:49:24and in addition to that you see how
  2920. 1:49:26there are a bunch of sums here and we
  2921. 1:49:28need to make sure that when I do those
  2922. 1:49:29sums that they broadcast correctly onto
  2923. 1:49:31everything else that's here
  2924. 1:49:33and so getting this expression is just
  2925. 1:49:35like highly non-trivial and I invite you
  2926. 1:49:36to basically look through it and step
  2927. 1:49:37through it and it's a whole exercise to
  2928. 1:49:39make sure that this this checks out but
  2929. 1:49:43once all the shapes are green and once
  2930. 1:49:45you convince yourself that it's correct
  2931. 1:49:46you can also verify that Patrick's gets
  2932. 1:49:48the exact same answer as well and so
  2933. 1:49:50that gives you a lot of peace of mind
  2934. 1:49:51that this mathematical formula is
  2935. 1:49:53correctly implemented here and
  2936. 1:49:55broadcasted correctly and replicated in
  2937. 1:49:57parallel for all of the 64 neurons
  2938. 1:50:00inside this bastrum layer okay and
  2939. 1:50:03finally exercise number four asks you to
  2940. 1:50:05put it all together and uh here we have
  2941. 1:50:08a redefinition of the entire problem so
  2942. 1:50:10you see that we reinitialize the neural
  2943. 1:50:11nut from scratch and everything and then
  2944. 1:50:13here instead of calling loss that
  2945. 1:50:15backward we want to have the manual back
  2946. 1:50:18propagation here as we derived It Up
  2947. 1:50:20Above so go up copy paste all the chunks
  2948. 1:50:23of code that we've already derived put
  2949. 1:50:25them here and drive your own gradients
  2950. 1:50:26and then optimize this neural nut
  2951. 1:50:28basically using your own gradients all
  2952. 1:50:31the way to the calibration of The
  2953. 1:50:33Bachelor and the evaluation of the loss
  2954. 1:50:34and I was able to achieve quite a good
  2955. 1:50:36loss basically the same loss you would
  2956. 1:50:38achieve before and that shouldn't be
  2957. 1:50:40surprising because all we've done is
  2958. 1:50:41we've really gotten to Lost That
  2959. 1:50:44backward and we've pulled out all the
  2960. 1:50:45code
  2961. 1:50:46and inserted it here but those gradients
  2962. 1:50:49are identical and everything is
  2963. 1:50:50identical and the results are identical
  2964. 1:50:52it's just that we have full visibility
  2965. 1:50:54on exactly what goes on under the hood
  2966. 1:50:56I'll plot that backward in this specific
  2967. 1:50:58case and this is all of our code this is
  2968. 1:51:02the full backward pass using basically
  2969. 1:51:04the simplified backward pass for the
  2970. 1:51:06cross entropy loss and the mass
  2971. 1:51:08generalization so back propagating
  2972. 1:51:10through cross entropy the second layer
  2973. 1:51:13the 10 H nonlinearity the batch
  2974. 1:51:15normalization
  2975. 1:51:16uh through the first layer and through
  2976. 1:51:19the embedding and so you see that this
  2977. 1:51:21is only maybe what is this 20 lines of
  2978. 1:51:23code or something like that and that's
  2979. 1:51:25what gives us gradients and now we can
  2980. 1:51:27potentially erase losses backward so the
  2981. 1:51:30way I have the code set up is you should
  2982. 1:51:31be able to run this entire cell once you
  2983. 1:51:33fill this in and this will run for only
  2984. 1:51:36100 iterations and then break
  2985. 1:51:37and it breaks because it gives you an
  2986. 1:51:39opportunity to check your gradients
  2987. 1:51:41against pytorch
  2988. 1:51:43so here our gradients we see are not
  2989. 1:51:46exactly equal they are approximately
  2990. 1:51:48equal and the differences are tiny
  2991. 1:51:51wanting negative 9 or so and I don't
  2992. 1:51:52exactly know where they're coming from
  2993. 1:51:54to be honest
  2994. 1:51:56um so once we have some confidence that
  2995. 1:51:57the gradients are basically correct we
  2996. 1:51:59can take out the gradient tracking
  2997. 1:52:01we can disable this breaking statement
  2998. 1:52:05and then we can
  2999. 1:52:07basically disable lost of backward we
  3000. 1:52:10don't need it anymore it feels amazing
  3001. 1:52:13to say that
  3002. 1:52:14and then here when we are doing the
  3003. 1:52:16update we're not going to use P dot grad
  3004. 1:52:18this is the old way of pytorch we don't
  3005. 1:52:21have that anymore because we're not
  3006. 1:52:22doing backward we are going to use this
  3007. 1:52:25update where we you see that I'm
  3008. 1:52:27iterating over
  3009. 1:52:29I've arranged the grads to be in the
  3010. 1:52:30same order as the parameters and I'm
  3011. 1:52:32zipping them up the gradients and the
  3012. 1:52:34parameters into p and grad and then here
  3013. 1:52:37I'm going to step with just the grad
  3014. 1:52:38that we derived manually
  3015. 1:52:40so the last piece
  3016. 1:52:43um is that none of this now requires
  3017. 1:52:46gradients from pytorch and so one thing
  3018. 1:52:49you can do here
  3019. 1:52:51um
  3020. 1:52:52is you can do with no grad and offset
  3021. 1:52:56this whole code block
  3022. 1:52:58and really what you're saying is you're
  3023. 1:52:59telling Pat George that hey I'm not
  3024. 1:53:00going to call backward on any of this
  3025. 1:53:02and this allows pytorch to be a bit more
  3026. 1:53:03efficient with all of it
  3027. 1:53:05and then we should be able to just uh
  3028. 1:53:07run this
  3029. 1:53:09and
  3030. 1:53:11it's running
  3031. 1:53:13and you see that losses backward is
  3032. 1:53:16commented out
  3033. 1:53:18and we're optimizing
  3034. 1:53:20so we're going to leave this run and uh
  3035. 1:53:23hopefully we get a good result
  3036. 1:53:25okay so I allowed the neural net to
  3037. 1:53:27finish optimization
  3038. 1:53:28then here I calibrate the bachelor
  3039. 1:53:31parameters because I did not keep track
  3040. 1:53:33of the running mean and very variants in
  3041. 1:53:35their training Loop
  3042. 1:53:37then here I ran the loss and you see
  3043. 1:53:39that we actually obtained a pretty good
  3044. 1:53:40loss very similar to what we've achieved
  3045. 1:53:42before
  3046. 1:53:43and then here I'm sampling from the
  3047. 1:53:45model and we see some of the name like
  3048. 1:53:47gibberish that we're sort of used to so
  3049. 1:53:49basically the model worked and samples
  3050. 1:53:52uh pretty decent results compared to
  3051. 1:53:54what we were used to so everything is
  3052. 1:53:56the same but of course the big deal is
  3053. 1:53:58that we did not use lots of backward we
  3054. 1:54:00did not use package Auto grad and we
  3055. 1:54:02estimated our gradients ourselves by
  3056. 1:54:04hand
  3057. 1:54:05and so hopefully you're looking at this
  3058. 1:54:06the backward pass of this neural net and
  3059. 1:54:08you're thinking to yourself actually
  3060. 1:54:10that's not too complicated
  3061. 1:54:12um
  3062. 1:54:13each one of these layers is like three
  3063. 1:54:15lines of code or something like that and
  3064. 1:54:17most of it is fairly straightforward
  3065. 1:54:18potentially with the notable exception
  3066. 1:54:20of the batch normalization backward pass
  3067. 1:54:22otherwise it's pretty good okay and
  3068. 1:54:25that's everything I wanted to cover for
  3069. 1:54:26this lecture so hopefully you found this
  3070. 1:54:29interesting and what I liked about it
  3071. 1:54:31honestly is that it gave us a very nice
  3072. 1:54:33diversity of layers to back propagate
  3073. 1:54:34through and
  3074. 1:54:36um I think it gives a pretty nice and
  3075. 1:54:38comprehensive sense of how these
  3076. 1:54:39backward passes are implemented and how
  3077. 1:54:41they work and you'd be able to derive
  3078. 1:54:43them yourself but of course in practice
  3079. 1:54:45you probably don't want to and you want
  3080. 1:54:46to use the pythonograd but hopefully you
  3081. 1:54:49have some intuition about how gradients
  3082. 1:54:51flow backwards through the neural net
  3083. 1:54:52starting at the loss and how they flow
  3084. 1:54:55through all the variables and all the
  3085. 1:54:56intermediate results
  3086. 1:54:58and if you understood a good chunk of it
  3087. 1:55:00and if you have a sense of that then you
  3088. 1:55:02can count yourself as one of these buff
  3089. 1:55:03doji's on the left instead of the uh
  3090. 1:55:06those on the right here now in the next
  3091. 1:55:09lecture we're actually going to go to
  3092. 1:55:10recurrent neural nuts lstms and all the
  3093. 1:55:13other variants of RNs and we're going to
  3094. 1:55:16start to complexify the architecture and
  3095. 1:55:17start to achieve better uh log
  3096. 1:55:19likelihoods and so I'm really looking
  3097. 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.