YouTube2Text

Let's reproduce GPT-2 (124M) — Transcript

by Andrej Karpathy · 43,388 words · 6,049 segments · language en · Watch on YouTube

Full transcript

  1. 0:00hi everyone so today we are going to be
  2. 0:02continuing our Zero to Hero series and
  3. 0:04in particular today we are going to
  4. 0:06reproduce the gpt2 model the 124 million
  5. 0:09version of it so when openi released
  6. 0:13gpt2 this was 2019 and they released it
  7. 0:16with this blog post on top of that they
  8. 0:19released this paper and on top of that
  9. 0:21they released this code on GitHub so
  10. 0:23open a/
  11. 0:24gpt2 now when we talk about reproducing
  12. 0:27gpt2 we have to be careful because in
  13. 0:29particular in this video we're going to
  14. 0:30be reproducing the 124 million parameter
  15. 0:33model so the thing to realize is that
  16. 0:35there's always a miniseries when these
  17. 0:37are releases are made so there are the
  18. 0:40gpt2 miniseries made up of models at
  19. 0:42different sizes and usually the biggest
  20. 0:45model is called the
  21. 0:46gpt2 but basically the reason we do that
  22. 0:49is because you can put the model sizes
  23. 0:51on the x-axis of plots like this and on
  24. 0:53the Y AIS you put a lot of uh Downstream
  25. 0:55metrics that you're interested in like
  26. 0:57translation summarization question
  27. 0:58answering and so on and you can chart
  28. 1:00out these scaling laws so basically as
  29. 1:03the model size increases you're getting
  30. 1:05better and better at Downstream metrics
  31. 1:07and so in particular for
  32. 1:09gpt2 if we scroll down in paper there
  33. 1:12are four models in the gpt2 miniseries
  34. 1:15starting at 124 million all the way up
  35. 1:18to 1558 million now the reason my
  36. 1:22numbers the way I say them disagree with
  37. 1:23this table is that this table is wrong
  38. 1:25if you actually go to the uh gpt2 uh
  39. 1:29GitHub repo they sort of say that um
  40. 1:32there was an error in how they added up
  41. 1:33the parameters but basically this is the
  42. 1:35124 million parameter model Etc so the
  43. 1:38124 million parameter had 12 layers in
  44. 1:40the Transformer and it had 768 channels
  45. 1:44in the Transformer 768 dimensions and
  46. 1:47I'm going to be assuming some
  47. 1:48familiarity with what these terms mean
  48. 1:50because I covered all of this in my
  49. 1:51previous video let's build gpt2 uh let's
  50. 1:54build GPT from scratch so I covered that
  51. 1:56in the previous video in this playlist
  52. 1:59now if we do everything correctly and
  53. 2:01everything works out well by the end of
  54. 2:03this video we're going to see something
  55. 2:04like this where we're looking at the
  56. 2:06validation loss which basically um
  57. 2:10measures how good we are at predicting
  58. 2:11the next token in a sequence on some
  59. 2:13validation data that the model has not
  60. 2:15seen during training and we see that we
  61. 2:17go from doing that task not very well
  62. 2:20because we're initializing from scratch
  63. 2:22all the way to doing that task quite
  64. 2:23well um by the end of the training and
  65. 2:26hopefully we're going to beat the gpt2
  66. 2:28uh 124 M model
  67. 2:30now previously when they were working on
  68. 2:32this this is already 5 years ago so this
  69. 2:35was probably a fairly complicated
  70. 2:36optimization at the time and the gpus
  71. 2:38and the compute was a lot smaller today
  72. 2:41you can reproduce this model in roughly
  73. 2:42an hour or probably less even and it
  74. 2:45will cost you about 10 bucks if you want
  75. 2:47to do this on the cloud uh Cloud Compu a
  76. 2:49sort of computer that you can all rent
  77. 2:52and if you pay $10 for that computer you
  78. 2:54wait about an hour or less you can
  79. 2:56actually achieve a model that is as good
  80. 2:58as this model that open ey released and
  81. 3:02uh one more thing to mention is unlike
  82. 3:04many other models open ey did release
  83. 3:06the weights for gpt2 so those weights
  84. 3:08are all available in this repository but
  85. 3:11the gpt2 paper is not always as good
  86. 3:14with all of the details of training so
  87. 3:16in addition to the gpt2 paper we're
  88. 3:18going to be referencing the gpt3 paper
  89. 3:20which is a lot more Concrete in a lot of
  90. 3:22the hyp parameters and optimization
  91. 3:24settings and so on um and it's not a
  92. 3:27huge departure in the architecture from
  93. 3:29the GPT 2 uh version of the model so
  94. 3:31we're going to be referencing both gpt2
  95. 3:33and gpt3 as we try to reproduce gpt2 124
  96. 3:36M uh so let's go so the first thing I
  97. 3:40would like to do is actually start at
  98. 3:41the end or at the Target so in other
  99. 3:43words let's load the GPT to 124 M model
  100. 3:47as it was released by openi and maybe
  101. 3:48take it for a spin let's sample some
  102. 3:50tokens from it now the issue with that
  103. 3:52is when you go into the code base of
  104. 3:54gpt2 and you go into the source and you
  105. 3:56click in on the model. pi you'll realize
  106. 3:58that actually this is using tensorflow
  107. 4:01so the original gpt2 code here was
  108. 4:03written in tensor flow which is
  109. 4:06um you know not let's just say not used
  110. 4:09as much anymore um so we'd like to use
  111. 4:12pytorch uh because it's a lot friendlier
  112. 4:14easier and I just personally like a lot
  113. 4:16more the problem with that is the
  114. 4:17initial code is intenser flow we'd like
  115. 4:19to use pytorch so instead uh to get the
  116. 4:21target we're going to use the hugging
  117. 4:23face Transformers um code which I like a
  118. 4:27lot more so when you go into the
  119. 4:28Transformers source Transformers models
  120. 4:30gpt2 modeling gpt2 Pi you will see that
  121. 4:33they have the gpt2 implementation of
  122. 4:35that Transformer here in this
  123. 4:37file um and it's like medium readable
  124. 4:42but not fully readable um but what it
  125. 4:45does is it did all the work of
  126. 4:47converting all those weights uh from
  127. 4:50tensor flow to pytorch Friendly and so
  128. 4:52it's much easier to load and work with
  129. 4:54so in particular we can look at the
  130. 4:56gpt2 um model here and we can load it
  131. 4:59using hugging face Transformers so
  132. 5:01swinging over this is what that looks
  133. 5:03like from Transformers import the DP GT2
  134. 5:07LM head model and then from pre-train
  135. 5:12gpt2 uh now one awkward thing about this
  136. 5:15is that when you do gpt2 as the model
  137. 5:17that we're loading this actually is the
  138. 5:19124 million parameter model if you want
  139. 5:22the actual the gpt2 the 1.5 billion then
  140. 5:25you actually want to do- XL so this is
  141. 5:28the 12 4 M our Target now what we're
  142. 5:32doing is when we actually get this we're
  143. 5:33initializing the uh pytorch NN module as
  144. 5:37defined here in this
  145. 5:38class from it I want to get just the
  146. 5:41state dict which is just a raw tensors
  147. 5:44so we just have um the tensors of that
  148. 5:46file and by the way here this is a
  149. 5:49jupyter notebook uh but this is jupyter
  150. 5:51notebook running inside vs code uh so I
  151. 5:54like to work with it all in a single
  152. 5:56sort of interface so I like to use vs
  153. 5:57code so this is the jupyter notebook
  154. 6:00extension inside the es
  155. 6:03code so when we get the state dick this
  156. 6:06is just a dict so we can print the key
  157. 6:09and the value which is the tensor and
  158. 6:11let's just look at the shapes so these
  159. 6:13are sort of
  160. 6:14the uh different parameters inside the
  161. 6:17gbt2 model and their shape so the W
  162. 6:22weight for token
  163. 6:25embedding is of size
  164. 6:2750257 by 768 where this is coming from
  165. 6:31is that we have
  166. 6:3250257 tokens in the gpt2 vocabulary um
  167. 6:37and the tokens by the way these are
  168. 6:39exactly the tokens that we spoken about
  169. 6:40in the previous video on my tokenization
  170. 6:43Series so the previous videos just
  171. 6:45before this I go into a ton of detail on
  172. 6:47tokenization gpt2 tokenizer happens to
  173. 6:49have this many tokens for each
  174. 6:53token we have a 768 dimensional
  175. 6:56embedding that is the distributed
  176. 6:58representation that stands in for that
  177. 7:01token so each token is a little string
  178. 7:03piece and then the 768 numbers are the
  179. 7:06vector that represents that
  180. 7:08token and so this is just our lookup
  181. 7:10table for tokens and then here we have
  182. 7:13the lookup table for the positions so
  183. 7:16because gbt2 has a maximum sequence
  184. 7:18length of
  185. 7:191024 we have up to 1,24 positions that
  186. 7:23each token can be attending to in the
  187. 7:25past and every one of those positions in
  188. 7:28gpd2 has a fixed Vector of
  189. 7:31768 that is learned by
  190. 7:33optimization um and so this is the
  191. 7:36position embedding and the token
  192. 7:38embedding um and then everything here is
  193. 7:41just the other weights and biases and
  194. 7:43everything else of this
  195. 7:45Transformer so when you just take for
  196. 7:47example the positional embeddings and
  197. 7:49flatten it out and take just the 20
  198. 7:51elements you can see that these are just
  199. 7:52the parameters these are weights floats
  200. 7:56just we can take and we can plot them so
  201. 7:59these are the position embeddings and we
  202. 8:01get something like this and you can see
  203. 8:03that this has structure and it has
  204. 8:04structure because what we what we have
  205. 8:07here really is every Row in this
  206. 8:10visualization is a different position a
  207. 8:12fixed absolute position in um the range
  208. 8:16from 0 to
  209. 8:171024 and each row here is the
  210. 8:19representation of that position and so
  211. 8:23it has structure because these
  212. 8:24positional embeddings end up learning
  213. 8:26these sinusoids and cosiness um that
  214. 8:29sort of like represent each of these
  215. 8:31positions and uh each row here stands in
  216. 8:35for that position and is processed by
  217. 8:36the Transformer to recover all the
  218. 8:38relative positions and uh sort of
  219. 8:41realize which token is where and um
  220. 8:44attend to them depending on their
  221. 8:45position not just their
  222. 8:47content so when we actually just look
  223. 8:49into an individual column inside these
  224. 8:53and I just grabbed three random columns
  225. 8:55you'll see that for example here we are
  226. 8:57focusing on every every single um
  227. 9:01Channel and we're looking
  228. 9:03at what that channel is doing as a
  229. 9:07function of uh position from one from Z
  230. 9:11to
  231. 9:121223
  232. 9:14really and we can see that some of these
  233. 9:15channels basically like respond more or
  234. 9:17less to different parts of the position
  235. 9:19Spectrum so this green channel uh really
  236. 9:22likes to fire for everything after 200
  237. 9:26uh up to 800 but not less a lot less and
  238. 9:30has a sharp drop off here near zero so
  239. 9:33who knows what these embeddings are
  240. 9:34doing and why they are the way they are
  241. 9:36you can tell for example that because
  242. 9:37they're a bit more Jagged and they're
  243. 9:38kind of noisy you can tell that this
  244. 9:40model was not fully trained and the more
  245. 9:43trained this model was the more you
  246. 9:45would expect to smooth this out and so
  247. 9:47this is telling you that this is a
  248. 9:48little bit of an undertrained model um
  249. 9:51but in principle actually these curves
  250. 9:53don't even have to be smooth this should
  251. 9:55just be totally random noise and in fact
  252. 9:57in the beginning of the optimization it
  253. 9:58is complete random noise because this
  254. 10:01position embedding table is initialized
  255. 10:03completely at random so in the beginning
  256. 10:05you have jaggedness and the fact that
  257. 10:07you end up with something smooth is
  258. 10:09already kind of impressive um that that
  259. 10:11just falls out of the optimization
  260. 10:13because in principle you shouldn't even
  261. 10:14be able to get any single graph out of
  262. 10:16this that makes sense but we actually
  263. 10:18get something that looks a little bit
  264. 10:19noisy but for the most part looks
  265. 10:21sinusoidal like um in the original
  266. 10:24Transformer um in the original
  267. 10:26Transformer paper the attention is all
  268. 10:28you need paper the positional embeddings
  269. 10:30are actually initialized and fixed if I
  270. 10:32remember correctly to sinusoids and
  271. 10:34cosiness of uh different frequencies and
  272. 10:37that's the positional coding and it's
  273. 10:38fixed but in gpt2 these are just
  274. 10:40parameters and they're trained from
  275. 10:41scratch just like any other parameter uh
  276. 10:44and that seems to work about as well and
  277. 10:46so what they do is they kind of like
  278. 10:47recover these sinusoidal like features
  279. 10:50during the
  280. 10:52optimization we can also look at any of
  281. 10:54the other matrices here so here I took
  282. 10:57the first layer of the
  283. 11:00Transformer and looking at like one of
  284. 11:02its weights and just the first block of
  285. 11:05300 by 300 and you see some structure
  286. 11:08but like again like who knows what any
  287. 11:10of this is if you're into mechanistic
  288. 11:12interpretability you might get a real
  289. 11:14kick out of trying to figure out like
  290. 11:16what is going on what is this structure
  291. 11:18and what does this all mean but we're
  292. 11:19not going to be doing that in this video
  293. 11:21but we definitely see that there's some
  294. 11:22interesting structure and that's kind of
  295. 11:24cool what we're mostly interested in is
  296. 11:26we've loaded the weights of this model
  297. 11:28that was released by open Ai and now
  298. 11:30using the hogging face Transformers we
  299. 11:33can not just get all the raw weights but
  300. 11:35we can also get the um what they call
  301. 11:39Pipeline and sample from it so this is
  302. 11:42the prefix hello I'm a language model
  303. 11:44comma and then we're sampling uh 30
  304. 11:47tokens and we getting five sequences and
  305. 11:50I ran this and this is what it produced
  306. 11:53um hell language
  307. 11:55model but what I'm really doing is
  308. 11:57making a human readable document there
  309. 11:59are other languages but those are dot
  310. 12:01dot dot so you can read through these if
  311. 12:03you like but basically these are five
  312. 12:05different completions of the same prefix
  313. 12:07from this uh gbt
  314. 12:092124m now uh if I go here I took this
  315. 12:13example from here and sadly even though
  316. 12:16we are fixing the seed we are getting
  317. 12:18different Generations from the snippet
  318. 12:21than what they got so presumably the
  319. 12:24code changed um but what we see though
  320. 12:28at this stage that's important is that
  321. 12:29we are getting coherent text so we've
  322. 12:32loaded the model successfully we can
  323. 12:34look at all its parameters and the keys
  324. 12:36tell us where in the model these come
  325. 12:39from and we want to actually write our
  326. 12:41own gpt2 class so that we have full
  327. 12:43understanding of what's happening there
  328. 12:44we don't want to be working with
  329. 12:46something like uh the modeling gpt2 Pi
  330. 12:49because it's just too complicated we
  331. 12:50want to write this from scratch
  332. 12:51ourselves so we're going to be
  333. 12:53implementing the GPT model here in
  334. 12:54parallel and as our first task let's
  335. 12:57load the gpt2 124 M into the class that
  336. 13:01we're going to develop here from scratch
  337. 13:04that's going to give us confidence that
  338. 13:06we can load the open ey model and
  339. 13:08therefore there's a setting of Weights
  340. 13:10that exactly is the 124 model but then
  341. 13:13of course what we're going to do is
  342. 13:14we're going to initialize the model from
  343. 13:15scratch instead and try try to train it
  344. 13:18ourselves um on a bunch of documents
  345. 13:20that we're going to get and we're going
  346. 13:22to try to surpass that model so we're
  347. 13:24going to get different weights and
  348. 13:25everything's going to look different
  349. 13:27hopefully better even um
  350. 13:29but uh we're going to have a lot of
  351. 13:31confidence that because we can load the
  352. 13:32openi model we are in the same model
  353. 13:34family and model class and we just have
  354. 13:36to ReDiscover a good setting of the
  355. 13:37weights uh but from scratch so let's now
  356. 13:41write the gbt2 model and let's load the
  357. 13:43weights and make sure that we can also
  358. 13:45generate text that looks coherent okay
  359. 13:48so let's now swing over to the attention
  360. 13:49is all un need paper that started
  361. 13:51everything and let's scroll over to the
  362. 13:53model architecture the original
  363. 13:55Transformer now remember that gpt2 is
  364. 13:57slightly modified from the or or
  365. 13:59Transformer in particular we do not have
  366. 14:02uh the encoder gpt2 is a decoder only
  367. 14:05Transformer as we call it so this entire
  368. 14:07encoder here is missing in addition to
  369. 14:09that this cross attention here that was
  370. 14:12using that encoder is also missing so we
  371. 14:14delete this entire part everything else
  372. 14:18stays almost the same but there are some
  373. 14:20differences that we're going to uh sort
  374. 14:21of look at here so there are two main
  375. 14:26differences when we go to the gb2 page
  376. 14:29under 2.3 model we notice that first
  377. 14:32there's a reshuffling of the layer Norms
  378. 14:34so they change place and second an
  379. 14:38additional layer normalization was added
  380. 14:40here to the final self detention block
  381. 14:43so basically all the layer Norms here
  382. 14:46instead of being after the MLP or after
  383. 14:48the attention they SN before it and an
  384. 14:50additional layer Norm gets added here
  385. 14:52right before the final
  386. 14:54classifier so now let's Implement some
  387. 14:56of the first sort of skeleton NN module
  388. 14:59modules here in our GPT NN module and in
  389. 15:02particular we're going to try to match
  390. 15:04up this schema here that is used by
  391. 15:06hugging face Transformers because that
  392. 15:08will make it much easier to load these
  393. 15:10weights from this state dict so we want
  394. 15:12something that reflects uh this schema
  395. 15:15here so here's what I came up with
  396. 15:19um basically we see that the main
  397. 15:22container here that has all the modules
  398. 15:24is called Transformer so I'm reflecting
  399. 15:26that with an NN module dict and this is
  400. 15:29basically a module that allows you to
  401. 15:30index into the subm modules using keys
  402. 15:34just like a dictionary uh
  403. 15:36strings within it we have the weights of
  404. 15:39the token embeddings WT and that's an N
  405. 15:41embedding and the weights of the
  406. 15:44position embeddings which is also just
  407. 15:45an N embedding and if you remember n
  408. 15:47embedding is really just a fancy little
  409. 15:49wrapper module around just a single um
  410. 15:53single array of numbers a single uh
  411. 15:56block of numbers just like this it's a
  412. 15:58single tensor and an embedding is a
  413. 16:02glorified um wrapper around a tensor
  414. 16:04that allows you to access its elements
  415. 16:07uh by indexing into the
  416. 16:08rows now in addition to that we see here
  417. 16:11that we have a h and then there's a this
  418. 16:14is index using numbers instead of
  419. 16:16indexed using strings so there's a h. 0
  420. 16:191 2 Etc all the way up till h. 11 and
  421. 16:23that's because there are 12 layers here
  422. 16:26in this Transformer so to reflect that
  423. 16:28I'm creating also an H I think that
  424. 16:31probably stands for hidden and instead
  425. 16:33of a module dict this is a model list so
  426. 16:35we can index it using integers exactly
  427. 16:37as we see here 01 2 Etc and the modular
  428. 16:42list has a n layer blocks and the blocks
  429. 16:46are yet to be defined in a module in a
  430. 16:48bit in addition to that following the
  431. 16:50gpt2 paper we have we need an additional
  432. 16:53final layer Norm that we're going to put
  433. 16:56in there and then we have the final
  434. 16:58classifier uh the language model head
  435. 17:01which um projects from 768 the number of
  436. 17:05embedding dimensions in this GPT all the
  437. 17:08way to the vocab size which is
  438. 17:1050257 and gpt2 uses no bias for this
  439. 17:13final uh sort of projection so this is
  440. 17:16the skeleton and you can see that it
  441. 17:19reflects this so the wte is the token
  442. 17:22embeddings here it's called output
  443. 17:24embedding but it's really the token
  444. 17:26embeddings the PE is the positional
  445. 17:29codings uh those two pieces of
  446. 17:31information as we saw previously are
  447. 17:32going to add and then go into the
  448. 17:34Transformer the H is the all the blocks
  449. 17:37in Gray and the LNF is this new layer
  450. 17:40that gets added here by the gpt2 model
  451. 17:43and LM head is this linear part here so
  452. 17:47that's the skeleton of the gpt2 we now
  453. 17:50have to implement the block okay so
  454. 17:53let's now recurse to the block itself so
  455. 17:55we want to define the block um so I'll
  456. 17:59start putting them here so the block I
  457. 18:02like to write out like
  458. 18:04this uh these are some of the
  459. 18:06initializations and then this is the
  460. 18:07actual forward pass of what this block
  461. 18:09computes and notice here that there's a
  462. 18:12change from the Transformer again that
  463. 18:14is mentioned in the gpt2 paper so here
  464. 18:17the layer normalizations are after the
  465. 18:20application of attention or feed forward
  466. 18:22in addition to that note that the
  467. 18:24normalizations are inside the residual
  468. 18:26stream you see how feed forward is
  469. 18:28applied and this arrow goes through and
  470. 18:30through the normalization so that means
  471. 18:33that your residual pathway has
  472. 18:35normalizations inside them and this is
  473. 18:37not very good or desirable uh you
  474. 18:39actually prefer to have a single uh
  475. 18:42clean residual stream all the way from
  476. 18:44supervision all the way down to the
  477. 18:45inputs the tokens and this is very
  478. 18:48desirable and nice because the gradients
  479. 18:51that flow from the top if you remember
  480. 18:54from your microad addition just
  481. 18:56distributes gradients during the
  482. 18:58backwards state to both of its branches
  483. 19:00equally so addition is a branch in the
  484. 19:04gradients and so that means that the
  485. 19:06gradients from the top flows straight to
  486. 19:08the inputs the tokens through the
  487. 19:10residual Pathways unchanged but then in
  488. 19:13addition to that the gradient also flows
  489. 19:14through the blocks and the blocks you
  490. 19:17know contribute their own contribution
  491. 19:18over time and kick in and change the
  492. 19:20optimization over time but basically
  493. 19:22clean residual pathway is desirable from
  494. 19:25an optimization perspective and then the
  495. 19:28this is the pre-normalization version
  496. 19:30where you see that RX first goes through
  497. 19:32the layer normalization and then the
  498. 19:34attention and then goes uh back out to
  499. 19:38go to the L ration number two and the
  500. 19:40multia perceptron sometimes also
  501. 19:43referred to as a feed forward Network or
  502. 19:44an FFN and then that goes into the
  503. 19:47residual stream again and the one more
  504. 19:50thing that is kind of interesting to
  505. 19:51note is that recall that attention is a
  506. 19:53communication operation it is where all
  507. 19:55the tokens and there's 1,24 tokens lined
  508. 19:58up in a sequence and this is where the
  509. 20:00tokens communicate this is where they
  510. 20:02exchange information so attention is a
  511. 20:06um aggregation function it's a pooling
  512. 20:08function it's a weighted sum function it
  513. 20:12is a reduce operation whereas MLP this
  514. 20:16uh MLP here happens at every single
  515. 20:18token individually there's no
  516. 20:20information being collected or exchanged
  517. 20:21between the tokens so the attention is
  518. 20:24the reduce and the MLP is the map and
  519. 20:27what you end up with is that the
  520. 20:28Transformer just ends up just being a
  521. 20:30repeated application of map produce if
  522. 20:33you want to think about it that way so
  523. 20:36um this is where they communicate and
  524. 20:37this is where they think individually
  525. 20:39about the information that they gathered
  526. 20:41and every one of these blocks uh
  527. 20:43iteratively refines the um
  528. 20:46representation is at the residual stream
  529. 20:48so this is our block um slightly
  530. 20:51modified from this picture Okay so let's
  531. 20:53now move on to the MLP so the MLP block
  532. 20:57uh I implemented as follows
  533. 20:59it is relatively straightforward we
  534. 21:00basically have two linear projections
  535. 21:02here that are sandwiched in between the
  536. 21:05G
  537. 21:06nonlinearity so nn. G approximate is 10h
  538. 21:11now when we swing on uh swing over to
  539. 21:13the Pyro documentation this is n.g and
  540. 21:16it has this format and it has two
  541. 21:18versions the original version of G which
  542. 21:20we'll step into into in a bit and the
  543. 21:22approximate version of Galo which we can
  544. 21:24request using
  545. 21:2510 so as you can see just as a preview
  546. 21:28here G is a basically like a reu except
  547. 21:32there's no flat exactly Flat Tail here
  548. 21:35at exactly zero but otherwise it looks
  549. 21:38very much like a slightly smoother reu
  550. 21:41it comes from this paper here Gan error
  551. 21:43linear units and uh you can step through
  552. 21:46this paper and there's some mathematical
  553. 21:48calac reasoning that leads to an
  554. 21:50interpretation that leads to the
  555. 21:51specific formulation it has to do with
  556. 21:53stochastic radial risers and the
  557. 21:56expectation of a modification to
  558. 21:57Adaptive dropout so you can read through
  559. 21:59all of that if you'd like here and
  560. 22:01there's a little bit of history as to
  561. 22:03why there is an an approximate version
  562. 22:05of G and that comes from this issue here
  563. 22:08as far as I can tell and in this issue
  564. 22:11Daniel Hendrix mentions that at the time
  565. 22:14when they developed this nonlinearity
  566. 22:17the Earth function which you need to
  567. 22:19evaluate the exact G was very slow in
  568. 22:21tensor flow so they ended up basically
  569. 22:23developing this approximation and this
  570. 22:25approximation that then ended up being
  571. 22:27picked up by Bert and by GP P2 Etc but
  572. 22:30today there's no real good reason to use
  573. 22:31the approximate version you'd prefer to
  574. 22:33just use the exact version um because I
  575. 22:36my expectation is that there's no big
  576. 22:38difference anymore and this is kind of
  577. 22:40like a historical um kind of Quirk um
  578. 22:43but we are trying to reproduce gpt2
  579. 22:45exactly and gpt2 used the 10h
  580. 22:49approximate version so we prefer to
  581. 22:51stick with
  582. 22:52that um now one other reason to actually
  583. 22:55just intuitively use G instead of veru
  584. 22:57is previously in the in videos in the
  585. 22:59past we've spoken about the dead reu
  586. 23:02neuron problem where in this tale of a
  587. 23:04reu if it's exactly flat at zero any
  588. 23:07activations that fall there will get
  589. 23:09exactly zero gradient there's no change
  590. 23:11there's no adaptation there's no
  591. 23:13development of the network if any of
  592. 23:15these activations end in this flat
  593. 23:17region but the G always contributes a
  594. 23:20local gradient and so there's always
  595. 23:22going to be a change always going to be
  596. 23:23an adaptation and sort of smoothing it
  597. 23:25out ends up empirically working better
  598. 23:27in practice as demonstrated in this
  599. 23:29paper and also as demonstrated by it
  600. 23:31being picked up by the bird paper gbt2
  601. 23:33paper and so on so for that reason we
  602. 23:35adopt this nonlinearity uh here in the
  603. 23:3810 in the gbt2 reproduction now in more
  604. 23:41modern networks also like llama 3 and so
  605. 23:43on this nonlinearity also further
  606. 23:45changes uh to swiglo and other variants
  607. 23:48like that uh but for gpt2 they Ed this
  608. 23:50approximate
  609. 23:51G okay and finally we have the attention
  610. 23:54operation so let me paste in my
  611. 23:57attention
  612. 24:00so I know this is a lot so I'm going to
  613. 24:02go through this a bit quickly a bit
  614. 24:03slowly but not too slowly because we
  615. 24:05have covered this in the previous video
  616. 24:07and I would just point you there um so
  617. 24:10this is the attention operation now in
  618. 24:12the previous video you will remember
  619. 24:13this is not just attention this is um
  620. 24:16multi-headed attention right and so in
  621. 24:19the previous video we had this
  622. 24:20multi-headed attention module and this
  623. 24:23implementation made it obvious that
  624. 24:25these heads are not actually that
  625. 24:26complicated uh there's basically
  626. 24:28in parallel inside every attention block
  627. 24:32there's multiple heads and they're all
  628. 24:33functioning in parallel and uh their
  629. 24:36outputs are just being concatenated and
  630. 24:38that becomes the output of the
  631. 24:40multi-headed attention so the heads are
  632. 24:42just kind of like parallel streams and
  633. 24:45their outputs get
  634. 24:46concatenated and so it was very simple
  635. 24:48and made the head be kind of like U
  636. 24:51fairly straightforward in terms of its
  637. 24:54implementation what happens here is that
  638. 24:56instead of having two separate modules
  639. 24:58and indeed many more modules that get
  640. 24:59concatenated all of that is just put
  641. 25:01into a single uh self attention uh
  642. 25:04module and instead I'm being very
  643. 25:07careful and doing a bunch of transpose
  644. 25:10split um tensor gymnastics to make this
  645. 25:13very efficient in pych but fundamentally
  646. 25:15and algorithmically nothing is different
  647. 25:17from the implementation we saw
  648. 25:19before um in this uh give
  649. 25:22repository so to remind you very briefly
  650. 25:25and I don't want to go in this uh into
  651. 25:27this in too many in too much time but we
  652. 25:30have these tokens lined up in a sequence
  653. 25:32and there's 1,20 of them and then each
  654. 25:35token at this stage of the attention
  655. 25:37emits three vectors the query key and
  656. 25:40the value and first what happens here um
  657. 25:44is that the queries and the keys have to
  658. 25:46multiply each other to get sort of the
  659. 25:49attention um amount like how interesting
  660. 25:52they find each other so they have to
  661. 25:54interact multiplicatively so what we're
  662. 25:56doing here is we're calculating the qkv
  663. 25:58we splitting it and then there's a bunch
  664. 26:00of gymnastics as I mentioned here and
  665. 26:03the way this works is that we're
  666. 26:04basically making the number of heads and
  667. 26:06H into a batch Dimension and so it's a
  668. 26:10batch Dimension just like B so that in
  669. 26:12these operations that follow pytorch
  670. 26:14treats B and NH as batches and it
  671. 26:18applies all the operations on all of
  672. 26:20them in parallel in both the batch and
  673. 26:22the
  674. 26:23heads and the operations that get
  675. 26:25applied are number one the queries and
  676. 26:27the keys intera to give us her attention
  677. 26:30this is the autoaggressive mask that
  678. 26:32makes sure that the tokens only attend
  679. 26:35to tokens before them and never to
  680. 26:37tokens in the
  681. 26:39future the softmax here normalizes the
  682. 26:41attention so it sums to one always and
  683. 26:45then recall from the previous video that
  684. 26:47doing the attention Matrix multiply with
  685. 26:48the values is basically a way to do a
  686. 26:50weighted sum of the values of the tokens
  687. 26:53that we found interesting at every
  688. 26:55single token and then the final
  689. 26:57transpose conf VI and view is just
  690. 26:59reassembling all of that again and this
  691. 27:02actually performs the concatenation
  692. 27:04operation so you can step through this
  693. 27:06uh slowly if you'd like um but it is
  694. 27:08equivalent mathematically to our
  695. 27:10previous implementation is just more
  696. 27:12efficient in P torch so that's why I
  697. 27:14chose this implementation
  698. 27:16instead now in addition to that I'm
  699. 27:18being careful with how I name my
  700. 27:19variables so for example cattin is the
  701. 27:22same as seaten and so actually our keys
  702. 27:25should basically exactly follow the
  703. 27:27schema of the hugging face train
  704. 27:28Transformers code and that will make it
  705. 27:29very easy for us to now Port over all
  706. 27:32the weights from exactly this sort of
  707. 27:34naming conventions because all of our
  708. 27:36variables are named the same thing but
  709. 27:39um at this point we have finished the
  710. 27:41gpt2 implementation and what that allows
  711. 27:44us to do is we don't have to basically
  712. 27:46use uh this file from hugging face which
  713. 27:48is fairly long
  714. 27:50um this
  715. 27:52is uh 2,000 lines of code um instead we
  716. 27:57just have a less than 100 lines of code
  717. 27:59and this is the complete uh gpd2
  718. 28:01implementation so at this stage we
  719. 28:02should just be able to take over all the
  720. 28:04weights set them and then do generation
  721. 28:07so let's see what that looks like okay
  722. 28:09so here I've also changed the GPT config
  723. 28:11so that the numbers here the H
  724. 28:13parameters agree with the gpt2 124 M
  725. 28:15model so the maximum sequence length
  726. 28:17which I call block size here is 124 the
  727. 28:21number of tokens is 50250 257 which if
  728. 28:25you watch my tokenizer video know that
  729. 28:27this is 50,000 m merges BP merges 256
  730. 28:31bite tokens the leaves of the BP tree
  731. 28:35and one special end of text token that
  732. 28:36delimits different documents and can
  733. 28:38start generation as well and there are
  734. 28:4112 layers there are 12 heads in the
  735. 28:43attention and the dimension of the
  736. 28:45Transformers was
  737. 28:46768 so here's how we can now load the
  738. 28:49parameters from hugging face to uh our
  739. 28:52code here and initialize the GPT class
  740. 28:54with those parameters so let me just
  741. 28:56copy paste a bunch of code
  742. 28:59here and I'm not going to go through
  743. 29:00this code too slow too quickly too
  744. 29:03slowly because um honestly it's not that
  745. 29:07interesting it's not that exciting we're
  746. 29:08just loading the weights so it's kind of
  747. 29:10dry but as I mentioned there are four
  748. 29:12models in this miniseries of gpt2 this
  749. 29:15is some of the Jupiter code um code that
  750. 29:18we had here on the right I'm just pting
  751. 29:20it over these are the hyper parameters
  752. 29:22of the gpt2 models uh we're creating the
  753. 29:24config object and creating our own model
  754. 29:27and then what's Happening Here is we're
  755. 29:28creating the state dict both for our
  756. 29:30model and for the hugging face
  757. 29:33model um and then what we're doing here
  758. 29:36is we're going over the hugging face
  759. 29:38model keys and we're copying over those
  760. 29:42tensors and in the process we are kind
  761. 29:45of ignoring a few of the buffers they're
  762. 29:47not parameters they're buffers so for
  763. 29:49example attention dobias uh that's just
  764. 29:51used for the autoaggressive mask and so
  765. 29:53we are ignoring some of those masks and
  766. 29:56uh that's it and then then one
  767. 29:58additional kind of annoyance is that
  768. 30:00this comes from the tensorflow repo and
  769. 30:02I'm not sure how this is a little bit
  770. 30:04annoying but some of the weights are
  771. 30:05transposed from what pytorch would want
  772. 30:08and so manually I hardcoded the weights
  773. 30:10that should be transposed and then we
  774. 30:12transpose them if that is so and then we
  775. 30:15return this model so the from
  776. 30:18pre-trained is a
  777. 30:20Constructor or class method in Python
  778. 30:23that Returns the GPT object if we just
  779. 30:26give it the model type which in our case
  780. 30:28is gpt2 the smallest model that we're
  781. 30:30interested in so this is the code and
  782. 30:33this is how you would use it and um we
  783. 30:35can pop open the terminal here in vs
  784. 30:38code and we can python train gbt2 pi and
  785. 30:44fingers
  786. 30:46crossed okay so we didn't crash and so
  787. 30:50we can load the weights and the biases
  788. 30:52and everything else into our Ann module
  789. 30:55but now let's also get additional
  790. 30:57confidence that this is working and
  791. 30:58let's try to actually generate from this
  792. 31:00model okay now before we can actually
  793. 31:01generate from this model we have to be
  794. 31:03able to forward it we didn't actually
  795. 31:04write that code yet so here's the
  796. 31:06forward
  797. 31:08function so the input to the forward is
  798. 31:11going to be our indices our tokens uh
  799. 31:13token indices and they are always of
  800. 31:16shape B BYT and so we have batch
  801. 31:19dimension of B and then we have the time
  802. 31:22dimension of up to T and the T can't be
  803. 31:26more than the block size the block size
  804. 31:27is is the maximum sequence length so B
  805. 31:30BYT indices arranged is sort of like a
  806. 31:32two-dimensional layout and remember that
  807. 31:35basically every single row of this is of
  808. 31:37size up to uh block size and this is T
  809. 31:41tokens that are in a sequence and then
  810. 31:43we have B independent sequences stacked
  811. 31:46up in a batch so that this is
  812. 31:48efficient now here we are forwarding the
  813. 31:51position embeddings and the token
  814. 31:52embeddings and this code should be very
  815. 31:54recognizable from the previous lecture
  816. 31:56so um we basically use uh a range which
  817. 31:59is kind of like a version of range but
  818. 32:01for pytorch uh and we're iterating from
  819. 32:04Z to T and creating this uh positions uh
  820. 32:07sort of uh indices
  821. 32:10um and then we are making sure that
  822. 32:12they're in the same device as idx
  823. 32:14because we're not going to be training
  824. 32:15on only CPU that's going to be too
  825. 32:16inefficient we want to be training on
  826. 32:18GPU and that's going to come in in a
  827. 32:20bit uh then we have the position
  828. 32:22embeddings and the token embeddings and
  829. 32:24the addition operation of those two now
  830. 32:26notice that the position embed are going
  831. 32:28to be identical for every single row of
  832. 32:31uh of input and so there's broadcasting
  833. 32:33hidden inside this plus where we have to
  834. 32:36create an additional Dimension here and
  835. 32:38then these two add up because the same
  836. 32:40position embeddings apply at every
  837. 32:41single row of our example stacked up in
  838. 32:44a batch then we forward the Transformer
  839. 32:46blocks and finally the last layer norm
  840. 32:49and the LM head so what comes out after
  841. 32:52forward is the logits and if the input
  842. 32:55was B BYT indices then at every single B
  843. 32:58by T we will calculate the uh logits for
  844. 33:02what token comes next in the sequence so
  845. 33:05what is the token B t+1 the one on the
  846. 33:09right of this token and B app size here
  847. 33:12is the number of possible tokens and so
  848. 33:16therefore this is the tensor that we're
  849. 33:17going to obtain and these low jits are
  850. 33:19just a softmax away from becoming
  851. 33:22probabilities so this is the forward
  852. 33:25pass of the network and now we can get
  853. 33:27load and so we're going to be able to
  854. 33:29generate from the model
  855. 33:30imminently okay so now we're going to
  856. 33:32try to set up the identical thing on the
  857. 33:35left here that matches hug and face on
  858. 33:36the right so here we've sampled from the
  859. 33:39pipeline and we sampled five times up to
  860. 33:4230 tokens with the prefix of hello I'm a
  861. 33:45language model and these are the
  862. 33:46completions that we achieved so we're
  863. 33:48going to try to replicate that on the
  864. 33:49left here so number turn sequences is
  865. 33:51five max length is 30 so the first thing
  866. 33:53we do of course is we initialize our
  867. 33:55model then we put it into evaluation
  868. 33:57mode now this is a good practice to put
  869. 33:59the model into eval when you're not
  870. 34:01going to be training it you're just
  871. 34:02going to be using it and I don't
  872. 34:05actually know if this is doing anything
  873. 34:07right now for the following reason our
  874. 34:09model up above here contains no modules
  875. 34:11or layers that actually have a different
  876. 34:14uh Behavior at training or evaluation
  877. 34:16time so for example Dropout batch norm
  878. 34:18and a bunch of other layers have this
  879. 34:20kind of behavior but all of these layers
  880. 34:22that we've used here should be identical
  881. 34:23in both training and evaluation time um
  882. 34:27so so potentially model that eval does
  883. 34:29nothing but then I'm not actually sure
  884. 34:31if this is the case and maybe pytorch
  885. 34:33internals uh do some clever things
  886. 34:35depending on the evaluation mode uh
  887. 34:36inside here the next thing we're doing
  888. 34:39here is we are moving the entire model
  889. 34:41to Cuda so we're moving this all of the
  890. 34:44tensors to GPU so I'm sshed here to a
  891. 34:47cloud box and I have a bunch of gpus on
  892. 34:49this box and here I'm moving the entire
  893. 34:53model and all of its members and all of
  894. 34:54its tensors and everything like that
  895. 34:56everything gets shipped off to basically
  896. 34:59a whole separate computer that is
  897. 35:01sitting on the GPU and the GPU is
  898. 35:03connected to the uh CPU and they can
  899. 35:05communicate but it's basically a whole
  900. 35:06separate computer with its own computer
  901. 35:08architecture and it's really well
  902. 35:09catered to parallel processing tasks
  903. 35:11like those of running neural networks so
  904. 35:14I'm doing this so that the model lives
  905. 35:16on the GPU a whole separate computer and
  906. 35:19it's just going to make our code a lot
  907. 35:20more efficient because all of this stuff
  908. 35:22runs a lot more efficiently on the
  909. 35:25gpus so that's the model
  910. 35:29itself now uh the next thing we want to
  911. 35:31do is we want to start with this as the
  912. 35:34prefix when we do the generation so
  913. 35:37let's actually create those prefix
  914. 35:39tokens so here's the code that I've
  915. 35:41written we're going to import the tich
  916. 35:43token library from open Ai and we're
  917. 35:45going to get the gpt2 encoding so that's
  918. 35:48the tokenizer for gpt2 and then we're
  919. 35:51going to encode this string and get a
  920. 35:54list of integers which are the tokens uh
  921. 35:57now these integers here should actually
  922. 35:59be fairly straightforward because we can
  923. 36:01just copy paste this string and we can
  924. 36:04sort of inspect what it is in tick
  925. 36:05tokenizer so just pasting that in these
  926. 36:08are the tokens that are going to come
  927. 36:09out so this list of integers is what we
  928. 36:12expect tokens to become and as you
  929. 36:15recall if you saw my video of course all
  930. 36:17the tokens they're just little string
  931. 36:19chunks right so these are this is the
  932. 36:21chunc of this string into gpt2
  933. 36:25tokens so once we have those tokens it's
  934. 36:27a list of integers we can create a torch
  935. 36:30tensor out of it in this case it's eight
  936. 36:32tokens and then we're going to replicate
  937. 36:34these eight tokens for five times to get
  938. 36:36five rows of eight tokens and that is
  939. 36:40our initial um input X as I call it here
  940. 36:45and it lives on the GPU as well so X now
  941. 36:48is this idx that we can put into forward
  942. 36:52to get our logits so that we know what
  943. 36:55comes as the sixth token
  944. 36:58uh sorry as the ninth token in every one
  945. 37:01of these five rows okay and we are now
  946. 37:04ready to generate so let me paste in one
  947. 37:05more code block
  948. 37:07here um so what's happening here in this
  949. 37:09code block is we have this x which is of
  950. 37:12size B BYT right so batch by time and
  951. 37:16we're going to be in every iteration of
  952. 37:18this loop we're going to be adding a
  953. 37:19column of new indices into each one of
  954. 37:22these rows right and so these are the
  955. 37:24new indices and we're appending them to
  956. 37:27the the sequence as we're sampling so
  957. 37:29with each Loop iteration we get one more
  958. 37:31column into X and all of the operations
  959. 37:34happen in the context manager of torch.
  960. 37:36nograd this is just telling pytorch that
  961. 37:38we're not going to be calling that
  962. 37:39backward on any of this so it doesn't
  963. 37:41have to cach all the intermediate
  964. 37:43tensors it's not going to have to
  965. 37:44prepare in any way for a potential
  966. 37:46backward later and this saves a lot of
  967. 37:48space and also possibly uh some time so
  968. 37:52we get our low jits we get the loow jits
  969. 37:54at only the last location we throw away
  970. 37:57all the other low jits uh we don't need
  971. 37:59them we only care about the last columns
  972. 38:01low jits so this is being wasteful uh
  973. 38:04but uh this is just kind of like an
  974. 38:06inefficient implementation of
  975. 38:08sampling um so it's correct but
  976. 38:10inefficient so we get the last column of
  977. 38:13loow jits pass it through soft Max to
  978. 38:14get our probabilities then here I'm
  979. 38:16doing top case sampling of 50 and I'm
  980. 38:18doing that because this is the hugging
  981. 38:20face default so just looking at the
  982. 38:23hugging face docks here of a pipeline um
  983. 38:26there's a bunch of
  984. 38:28quarks that go into hugging face and I
  985. 38:32mean it's it's kind of a lot honestly
  986. 38:34but I guess the important one that I
  987. 38:36noticed is that they're using top K by
  988. 38:38default which is 50 and what that does
  989. 38:41is that uh so that's being used here as
  990. 38:43well and what that does is basically we
  991. 38:45want to take our probabilities and we
  992. 38:47only want to keep the top 50
  993. 38:49probabilities and anything that is lower
  994. 38:51than the 50th probability uh we just
  995. 38:54clamp to zero and renormalize and so
  996. 38:56that way we are never sampling very rare
  997. 38:59tokens uh the tokens we're going to be
  998. 39:01sampling are always in the top 50 of
  999. 39:03most likely tokens and this helps keep
  1000. 39:05the model kind of on track and it
  1001. 39:07doesn't blabber on and it doesn't get
  1002. 39:08lost and doesn't go off the rails as
  1003. 39:10easily uh and it kind of like um sticks
  1004. 39:13in the vicinity of likely tokens a lot
  1005. 39:15better so this is the way to do it in
  1006. 39:17pytorch and you can step through it if
  1007. 39:18you like I don't think it's super
  1008. 39:20insightful so I'll speed through it but
  1009. 39:22roughly speaking we get this new column
  1010. 39:24of of tokens we append them on x and
  1011. 39:27basically The Columns of X grow until
  1012. 39:30this y Loop gets tripped up and then
  1013. 39:33finally we have an entire X of size um 5
  1014. 39:38by 30 in this case in this example and
  1015. 39:41we can just basically print all those
  1016. 39:43individual rows so I'm getting all the
  1017. 39:46rows I'm getting all the tokens that
  1018. 39:48were sampled and I'm using the decode
  1019. 39:50function from Tik tokenizer to get back
  1020. 39:52the string which we can print and so
  1021. 39:55terminal new terminal
  1022. 39:59and let me python train
  1023. 40:08gpt2 okay so these are the generations
  1024. 40:11that we're getting hello I'm a language
  1025. 40:13model not a
  1026. 40:15program um new line new line Etc hello
  1027. 40:19I'm a language model and one of the main
  1028. 40:21things that bothers me when they create
  1029. 40:22languages is how easy it becomes to
  1030. 40:23create something that I me so this will
  1031. 40:26just like blabber on right in all these
  1032. 40:27cases now one thing you will notice is
  1033. 40:29that these Generations are not the
  1034. 40:31generations of hugging face here and I
  1035. 40:35can't find the discrepancy to be honest
  1036. 40:37and I didn't fully go through all these
  1037. 40:39options but probably there's something
  1038. 40:40else hiding in on addition to the top P
  1039. 40:43so I'm not able to match it up but just
  1040. 40:45for correctness um down here Below in
  1041. 40:47the juper notebook and using the hugging
  1042. 40:49face model so this is the hugging face
  1043. 40:52model here I was I replicated the code
  1044. 40:56and if I do this and I run that then I
  1045. 40:59am getting the same results so basically
  1046. 41:03the model internals are not wrong it's
  1047. 41:05just I'm not 100% sure what the pipeline
  1048. 41:08does in hugging face and that's why
  1049. 41:09we're not able to match them up but
  1050. 41:11otherwise the code is correct and we've
  1051. 41:13loaded all the um tensors correctly so
  1052. 41:16we're initializing the model correctly
  1053. 41:18and everything here works so long story
  1054. 41:20short uh We've Port it all the weights
  1055. 41:22we initialize the gpt2 this is the exact
  1056. 41:25opening gpt2 and it can generate
  1057. 41:27sequences and they look sensible and now
  1058. 41:30here of course we're initializing with
  1059. 41:32gbt2 model weights but now we want to
  1060. 41:34initialize from scratch from random
  1061. 41:36numbers and we want to actually train a
  1062. 41:38model that will give us sequences as
  1063. 41:40good as or better than these ones in
  1064. 41:44quality and so that's what we turn to
  1065. 41:46next so it turns out that using the
  1066. 41:48random model is actually fairly
  1067. 41:49straightforward because pytorch already
  1068. 41:51initializes our model randomly and by
  1069. 41:53default so when we create the GPT model
  1070. 41:58and the Constructor this is all um all
  1071. 42:00of these layers and modules have random
  1072. 42:03initializers that are there by default
  1073. 42:05so when these linear layers get created
  1074. 42:07and so on there's default Constructors
  1075. 42:10for example using the Javier
  1076. 42:11initialization that we saw in the past
  1077. 42:13uh to construct the weights of these
  1078. 42:15layers and so creating a random model
  1079. 42:18instead of a gpt2 model is actually
  1080. 42:20fairly straightforward and we would just
  1081. 42:22come here and instead we would create
  1082. 42:24model equals GPT and then we want to use
  1083. 42:28the default config GPT config and the
  1084. 42:31default config uses the 124 M parameters
  1085. 42:33so this is the random model
  1086. 42:35initialization and we can run
  1087. 42:42it and we should be able to get uh
  1088. 42:46results now the results here of course
  1089. 42:48are total garbage carbal and that's
  1090. 42:50because this is random model and so
  1091. 42:51we're just getting all these random
  1092. 42:53token string pieces chunked up totally
  1093. 42:55at random so that's what we have right
  1094. 42:57now uh now one more thing I wanted to
  1095. 42:59point out by the way is in case you do
  1096. 43:01not have Cuda available because you
  1097. 43:03don't have a GPU you can still follow
  1098. 43:04along with uh with what we're doing here
  1099. 43:07uh to some extent uh and probably not to
  1100. 43:10the very end because by the end we're
  1101. 43:11going to be using multiple gpus and
  1102. 43:13actually doing a serious training run uh
  1103. 43:15but for now you can actually follow
  1104. 43:16along decently okay uh so one thing that
  1105. 43:19I like to do in pytorch is I like to
  1106. 43:20autod detect the device that is
  1107. 43:22available to you so in particular you
  1108. 43:24could do that like this
  1109. 43:28so here we are trying to detect a device
  1110. 43:30to run on that has the highest compute
  1111. 43:32capability you can think about it that
  1112. 43:33way so by default we start with CPU
  1113. 43:36which of course is available everywhere
  1114. 43:37because every single computer will have
  1115. 43:39a CPU but then we can try to detect do
  1116. 43:42you have a GPU you so use a Cuda and
  1117. 43:44then if you don't have a Cuda uh do you
  1118. 43:47at least have MPS MPS is the back end
  1119. 43:49for Apple silicon so if you have a
  1120. 43:51Macbook that is fairly new you probably
  1121. 43:53have apple silicon on the inside and
  1122. 43:55then that has a GPU that is actually
  1123. 43:57fairly capable uh depending on which
  1124. 43:59MacBook you have and so you can use MPS
  1125. 44:01which will be potentially faster than
  1126. 44:02CPU and so we can print the device here
  1127. 44:05now once we have the device we can
  1128. 44:07actually use it in place of Puda so we
  1129. 44:11just swap it in and notice that here
  1130. 44:14when we call model on X if this x here
  1131. 44:17is on CPU instead of GPU then it will
  1132. 44:21work fine because here in the forward
  1133. 44:23which is where P to will come when we
  1134. 44:26create a pose we were careful to use the
  1135. 44:28device of idx to create this tensor as
  1136. 44:31well and so there won't be any mismatch
  1137. 44:33where one tensor is on CPU one is on GPU
  1138. 44:36and uh that you can't combine those but
  1139. 44:38here we are um carefully initializing on
  1140. 44:41the correct device as indicated by the
  1141. 44:43input to this model so this will autod
  1142. 44:47detect device for me this will be of
  1143. 44:49course
  1144. 44:50GPU so using device
  1145. 44:54Cuda uh but uh you can also run with um
  1146. 44:58as I mentioned another device and it's
  1147. 45:00not going to be too much slower so if I
  1148. 45:01override device here
  1149. 45:03oops if I override device equals
  1150. 45:07CPU
  1151. 45:08then we'll still print Cuda of course
  1152. 45:11but now we're actually using CPU one 2 3
  1153. 45:164 5 6 okay about 6 seconds and actually
  1154. 45:21we're not using torch compile and stuff
  1155. 45:22like that which will speed up everything
  1156. 45:24a lot faster as well but you can follow
  1157. 45:27even on a CPU I think to a decent extent
  1158. 45:30um so that's note on that okay so I do
  1159. 45:32want to loop around eventually into what
  1160. 45:35it means to have different devices in
  1161. 45:36pytorch and what it is exactly that
  1162. 45:38pytorch does in the background for you
  1163. 45:40when you do something like module. 2
  1164. 45:43device or where you take a torch tensor
  1165. 45:45and do A2 device and what exactly
  1166. 45:48happens and how that works but for now
  1167. 45:49I'd like to get to training and I'd like
  1168. 45:51to start training the model and for now
  1169. 45:53let's just say the device makes code go
  1170. 45:55fast um and let's go into how we can
  1171. 45:58actually train the model so to train the
  1172. 46:00model we're going to need some data set
  1173. 46:02and for me the best debugging simplest
  1174. 46:04data set that I like to use is the tiny
  1175. 46:06Shakespeare data set um and it's
  1176. 46:09available at this URL so you can W get
  1177. 46:11it or you can just search tiny
  1178. 46:12Shakespeare data
  1179. 46:13set and so um I have in my file system
  1180. 46:16as just LS input.txt
  1181. 46:18so I already downloaded it and here I'm
  1182. 46:22reading the data set getting the first
  1183. 46:231,000 characters and printing the first
  1184. 46:26100
  1185. 46:27now remember that gpt2 has uh roughly a
  1186. 46:30compression ratio the tokenizer has a
  1187. 46:32compression ratio of rly 3 to1 so th000
  1188. 46:35characters is roughly 300 tokens here uh
  1189. 46:37that will come out of this in the slice
  1190. 46:39that we're currently getting so this is
  1191. 46:42the first few uh
  1192. 46:44characters and uh if you want to get a
  1193. 46:46few more statistics on this we can do
  1194. 46:48work count on input.txt
  1195. 46:50so we can see that this is uh 40,000
  1196. 46:53lines about 200,000 words in this data
  1197. 46:56set and about 1 million bytes in this
  1198. 46:59file and knowing that this file is only
  1199. 47:01asky characters there's no crazy unic
  1200. 47:03code here as far as I know and so every
  1201. 47:05asky character is encoded with one bite
  1202. 47:08and so this is uh the same number
  1203. 47:10roughly a million characters inside this
  1204. 47:12data set so that's the data set size uh
  1205. 47:15by default very small and minimal data
  1206. 47:17set for debugging to get us off the
  1207. 47:19ground in order to tokenize this data
  1208. 47:21set we're going to get Tik token
  1209. 47:23encoding for gbt2 encode the data uh the
  1210. 47:27first um 1,000 characters and then I'm
  1211. 47:30only going to print the first 24 tokens
  1212. 47:33so these are the tokens as a list of
  1213. 47:36integers and if you can read gpt2 tokens
  1214. 47:38you will see that 198 here you'll
  1215. 47:40recognize that as the slashing character
  1216. 47:42so that is a new line and then here for
  1217. 47:45example we have two new lines so that's
  1218. 47:46198 twice here uh so this is just a
  1219. 47:49tokenization of the first 24 tokens so
  1220. 47:52what we want to do now is we want to
  1221. 47:54actually process these token sequences
  1222. 47:56and feed them into a Transformer and in
  1223. 47:59particular we want them we want to
  1224. 48:01rearrange these tokens into this idx
  1225. 48:05variable that we're going to be feeding
  1226. 48:06into the Transformer so we don't want a
  1227. 48:08single very long onedimensional sequence
  1228. 48:10we want an entire batch where each
  1229. 48:12sequence is up to uh is basically T
  1230. 48:16tokens and T cannot be larger than the
  1231. 48:18maximum sequence length and then we have
  1232. 48:21these t uh tlong uh sequences of tokens
  1233. 48:25and we have B independent examples of
  1234. 48:27sequences so how can we create a b BYT
  1235. 48:30tensor that we can feed into the forward
  1236. 48:32out of these onedimensional
  1237. 48:34sequences so here's my favorite way to
  1238. 48:36to achieve this uh so if we take torch
  1239. 48:39and then we create a tensor object out
  1240. 48:41of this list of integers and just the
  1241. 48:42first 24 tokens my favorite way to do
  1242. 48:45this is basically you do a do view of um
  1243. 48:49of uh for example 4x6 which multiply to
  1244. 48:5224 and so it's just a two-dimensional
  1245. 48:54rearrangement of these tokens and you'll
  1246. 48:56is that when you view this
  1247. 48:57onedimensional sequence as
  1248. 48:58two-dimensional 4x6 here the first six
  1249. 49:03uh tokens uh up to here end up being the
  1250. 49:06first row the next six tokens here end
  1251. 49:09up being the second row and so on and so
  1252. 49:12basically it's just going to stack up
  1253. 49:14this the um every six tokens in this
  1254. 49:18case as independent rows and it creates
  1255. 49:20a batch of tokens in this case and so
  1256. 49:23for example if we are token 25 in the
  1257. 49:26Transformer when we feed this in and
  1258. 49:28this becomes the idx this token is going
  1259. 49:30to see these three tokens and it's going
  1260. 49:33to try to predict that 198 comes
  1261. 49:35next so in this way we are able to
  1262. 49:39create this two-dimensional batch that's
  1263. 49:41that's quite nice now in terms of the
  1264. 49:44label that we're going to need for the
  1265. 49:45Target to calculate the loss function
  1266. 49:47how do we get that well we could write
  1267. 49:49some code inside the forward pass
  1268. 49:51because we know that the next uh token
  1269. 49:53in a sequence which is the label is just
  1270. 49:55to the right of us but you'll notice
  1271. 49:57that actually we for this token at the
  1272. 49:59very end 13 we don't actually have the
  1273. 50:02next correct token because we didn't
  1274. 50:03load it so uh we actually didn't get
  1275. 50:07enough information here so I'll show you
  1276. 50:09my favorite way of basically getting
  1277. 50:11these batches and I like to personally
  1278. 50:14have not just the input to the
  1279. 50:15Transformer which I like to call X but I
  1280. 50:18also like to create the labels uh tensor
  1281. 50:21which is of the exact same size as X but
  1282. 50:24contains the targets at every single
  1283. 50:26position
  1284. 50:27and so here's the way that I like to do
  1285. 50:28that I like to make sure that I fetch
  1286. 50:30plus one uh token because we need the
  1287. 50:32ground Truth for the very last token uh
  1288. 50:35for
  1289. 50:3613 and then when we're creating the
  1290. 50:39input we take everything up to the last
  1291. 50:41token not including and view it as 4x6
  1292. 50:44and when we're creating targets we do
  1293. 50:47the buffer but starting at index one not
  1294. 50:50index zero so we're skipping the first
  1295. 50:52element and we view it in the exact same
  1296. 50:54size and then when I print this
  1297. 50:58here's what happens where we see that
  1298. 51:00basically as an example for this token
  1299. 51:0225 its Target was 198 and that's now
  1300. 51:05just stored at the exact same position
  1301. 51:07in the Target tensor which is 198 and
  1302. 51:10also this last token 13 now has its
  1303. 51:13label which is 198 and that's just
  1304. 51:16because we loaded this plus one here so
  1305. 51:19basically this is the way I like to do
  1306. 51:20it you take long sequences you uh view
  1307. 51:24them in two- dimensional terms so that
  1308. 51:26you get batch of time and then we make
  1309. 51:29sure to load one additional token so we
  1310. 51:31basically load a buffer of tokens of B *
  1311. 51:34t+ one and then we sort of offset things
  1312. 51:37and view them and then we have two
  1313. 51:39tensors one of them is the input to the
  1314. 51:41Transformer and the other exactly is the
  1315. 51:43labels and so let's now reorganize this
  1316. 51:46code and um create a very simple data
  1317. 51:50loader object that tries to basically
  1318. 51:52load these tokens and um feed them to
  1319. 51:55the Transformer and calculate the loss
  1320. 51:57okay so I reshuffled the code here uh
  1321. 51:59accordingly so as you can see here I'm
  1322. 52:01temporarily overwriting U to run a CPU
  1323. 52:05and importing TI token and all of this
  1324. 52:06should look familiar we're loading a
  1325. 52:08th000 characters I'm setting BT to just
  1326. 52:10be 4 and 32 right now just because we're
  1327. 52:13debugging we just want to have a single
  1328. 52:15batch that's very small and all of this
  1329. 52:17should now look familiar and follows
  1330. 52:19what we did on the right and then here
  1331. 52:21we get the we create the model and get
  1332. 52:24the lits and so so here as you see I
  1333. 52:28already ran this only runs in a few
  1334. 52:30seconds but because we have a batch of
  1335. 52:32uh 4X 32 our lits are now of size 4X 32x
  1336. 52:3850257 so those are the lit for what
  1337. 52:40comes next at every position and now we
  1338. 52:43have the labels which are stored in y so
  1339. 52:46now is the time to calculate the loss
  1340. 52:48and then do the backward pass and then
  1341. 52:49the optimization so let's first
  1342. 52:51calculate the
  1343. 52:52loss okay so to calculate the loss we're
  1344. 52:55going to adjust the forward function of
  1345. 52:56this NN module in the model and in
  1346. 52:59particular we're not just going to be
  1347. 53:00returning logits but also we're going to
  1348. 53:02return the loss uh and we're going to
  1349. 53:04not just pass in the input in thees but
  1350. 53:06also the targets uh in y and now we will
  1351. 53:12print not Lo just. shape anymore we're
  1352. 53:14actually going to print the loss
  1353. 53:14function and then c. exit of zero so
  1354. 53:17that we skip some of the sampling logic
  1355. 53:20so now let's swing up to the forward
  1356. 53:21function which gets called there because
  1357. 53:25now we also have these optional
  1358. 53:28targets and when we get the targets we
  1359. 53:30can also calculate uh the loss and
  1360. 53:32remember that we want to basically
  1361. 53:34return uh log just loss and loss by
  1362. 53:36default is none
  1363. 53:39but
  1364. 53:40um let's put this here if uh targets is
  1365. 53:45not none then we want to calculate loss
  1366. 53:49and co-pilot is already getting excited
  1367. 53:51here and calculating the what looks to
  1368. 53:53be correct loss it is using the cross
  1369. 53:55entropy loss as is documented here uh so
  1370. 54:00this is a function in pytorch under the
  1371. 54:03functional now what is actually
  1372. 54:05happening here because it looks a little
  1373. 54:06bit scary uh basically uh the F that
  1374. 54:09cross entropy does not like
  1375. 54:10multi-dimensional inputs it can't take a
  1376. 54:12b BYT by vocap size so what's happening
  1377. 54:15here is that we are flattening out this
  1378. 54:17three-dimensional tensor into just two
  1379. 54:19Dimensions the First Dimension is going
  1380. 54:21to be calculated automatically and it's
  1381. 54:23going to be B * T and then the last
  1382. 54:26Dimension is vocap size so basically
  1383. 54:28this is uh flattening out this
  1384. 54:30three-dimensional tensor of logits to
  1385. 54:32just be two- dimensional B * T all
  1386. 54:35individual examples and vocap size on uh
  1387. 54:39in terms of the length of each row and
  1388. 54:41then it's also flattening out the
  1389. 54:42targets which are also two- dimensional
  1390. 54:44at this stage but we're going to just
  1391. 54:46flatten them out so they're just a
  1392. 54:48single tensor of B * T and this can then
  1393. 54:51pass into cross entropy to calculate a
  1394. 54:52loss which we return so this should
  1395. 54:55basically at this point run because this
  1396. 54:57is not too complicated
  1397. 54:59so let's run it and let's see if we
  1398. 55:03should be printing the
  1399. 55:09loss and here we see that we printed 11
  1400. 55:12uh roughly and so
  1401. 55:16um and notice that this is the tensor of
  1402. 55:18a single element which is this number 11
  1403. 55:21now we also want to be able to calculate
  1404. 55:23a reasonable uh kind of starting point
  1405. 55:25for a random rationalized Network so we
  1406. 55:27covered this in previous videos but our
  1407. 55:29vocabulary size is
  1408. 55:3150257 at initialization of the network
  1409. 55:34you would hope that um every vocab
  1410. 55:37element is getting roughly a uniform
  1411. 55:40probability uh so that we're not
  1412. 55:42favoring at initialization any token way
  1413. 55:45too much we're not confidently wrong at
  1414. 55:47initialization so what we're hoping is
  1415. 55:49that the probability of any arbitrary
  1416. 55:51token is roughly 1 over 50,2 57 and now
  1417. 55:55we can sanity check the loss because
  1418. 55:57remember that the cross entropy loss is
  1419. 55:59just basically the negative um log
  1420. 56:01likelihood so if we now take this
  1421. 56:04probability and we take it through the
  1422. 56:06natural logarithm and then we do the
  1423. 56:08negative that is the loss we expect at
  1424. 56:11initialization and we covered this in
  1425. 56:13previous videos so I would expect
  1426. 56:15something around 10.82 and we're seeing
  1427. 56:17something around 11 so it's not way off
  1428. 56:20this is roughly the probability I expect
  1429. 56:21at initialization so that tells me that
  1430. 56:24the at initialization or probability
  1431. 56:26distribtion is roughly diffused it's a
  1432. 56:27good starting point and we can now uh
  1433. 56:30perform the optimization and tell the
  1434. 56:32network which elements you know should
  1435. 56:34follow correctly in what order so at
  1436. 56:37this point we can do a l step backward
  1437. 56:39calculate the gradients and do an
  1438. 56:40optimization so let's get to that okay
  1439. 56:43so let's do the optimization now um so
  1440. 56:46here we
  1441. 56:47have the loss is this is how we get the
  1442. 56:51loss but now basically we want a load
  1443. 56:53for Loop here so 4 I in range let's do
  1444. 56:5550 steps or something like that uh let's
  1445. 56:58create an Optimizer object in
  1446. 57:00pytorch um and so here we are using the
  1447. 57:04atom um Optimizer which is an
  1448. 57:07alternative to the stochastic radian
  1449. 57:08descent Optimizer SGD that we were using
  1450. 57:11so SGD is a lot simpler atom is a bit
  1451. 57:13more involved and I actually
  1452. 57:14specifically like the atom W variation
  1453. 57:17because in my opinion it kind of just
  1454. 57:19like fixes a bug um so adom w is a bug
  1455. 57:22fix of atom is what I would say when we
  1456. 57:25go to the documentation for atom
  1457. 57:27W oh my
  1458. 57:29gosh we see um that it takes a bunch of
  1459. 57:32hyper parameters and it's a little bit
  1460. 57:34more complicated than the SGD we were
  1461. 57:35looking at before uh because in addition
  1462. 57:37to basically updating the parameters
  1463. 57:39with the gradient uh scaled by the
  1464. 57:41Learning rate it keeps these buffers
  1465. 57:43around and it keeps two buffers the m
  1466. 57:46and the V which it calls the first and
  1467. 57:48the second moment so something that
  1468. 57:49looks a bit like momentum and something
  1469. 57:51that looks a bit like RMS prop if you're
  1470. 57:53familiar with it but you don't have to
  1471. 57:55be it's just kind of a normalization
  1472. 57:57that happens on each gradient element
  1473. 57:59individually and speeds up the
  1474. 58:00optimization especially for language
  1475. 58:02models but I'm not going to go into the
  1476. 58:04detail right here we're going to treat
  1477. 58:06it as a bit of a black box and it just
  1478. 58:08optimizes um the objective faster than
  1479. 58:12SGD which is what we've seen in the
  1480. 58:13previous lectures so let's use it as a
  1481. 58:15black box in our case uh create the
  1482. 58:18optimizer object and
  1483. 58:21then go through the optimization
  1484. 58:28the first thing to always make sure the
  1485. 58:30co-pilot did not forget to zero the
  1486. 58:32gradients so um always remember that you
  1487. 58:35have to start with a zero gradient then
  1488. 58:38when you get your loss and you do a DOT
  1489. 58:39backward dot backward adds to gradients
  1490. 58:42so it deposits gradients it it always
  1491. 58:44does a plus equals on whatever the
  1492. 58:46gradients are which is why you must set
  1493. 58:48them to zero so this accumulates the
  1494. 58:50gradient from this loss and then we call
  1495. 58:52the step function on the optimizer to um
  1496. 58:56update the parameters and to um decrease
  1497. 59:00the
  1498. 59:00loss and then we print a step and the
  1499. 59:03loss do item is used here because loss
  1500. 59:06is a tensor with a single element do
  1501. 59:08item will actually uh convert that to a
  1502. 59:11single float and this float will live
  1503. 59:13not will will live on the CPU so this
  1504. 59:16gets to some of the internals again of
  1505. 59:17the devices but loss is a is a tensor
  1506. 59:20with a single element and it lifts on
  1507. 59:22GPU for me because I'm using gpus when
  1508. 59:25you call item P torch behind the scenes
  1509. 59:28will take that one-dimensional tensor
  1510. 59:30ship it back to the CPU uh memory and
  1511. 59:32convert it into a float that we can just
  1512. 59:35print so this is the optimization and
  1513. 59:38this should probably just
  1514. 59:42work let's see what
  1515. 59:45happens actually sorry let me instead of
  1516. 59:47using CPU override let me delete that so
  1517. 59:50this is a bit faster for me and it runs
  1518. 59:52on Cuda
  1519. 59:58oh expected all tensors to be on the
  1520. 1:00:00same device but found at least two
  1521. 1:00:02devices Cuda zero and CPU so Cuda zero
  1522. 1:00:06is the zeroth GPU because I actually
  1523. 1:00:07have eight gpus on this box uh so the
  1524. 1:00:10zeroth GPU in my box and CPU and model
  1525. 1:00:14we have moved to device but when I was
  1526. 1:00:17writing this code I actually introduced
  1527. 1:00:18a bug because buff we never moved to
  1528. 1:00:21device and you have to be careful
  1529. 1:00:23because you can't just do buff dot two
  1530. 1:00:25of
  1531. 1:00:26device um it's not stateful it doesn't
  1532. 1:00:30convert it to be a device it instead uh
  1533. 1:00:33returns pointer to a new memory which is
  1534. 1:00:35on the device so you see how we can just
  1535. 1:00:37do model that two a device that does not
  1536. 1:00:39apply to tensors you have to do buff
  1537. 1:00:42equals
  1538. 1:00:44um b.2 device and then this should work
  1539. 1:00:49okay so what do we expect to see we
  1540. 1:00:52expect to see a reasonable loss in the
  1541. 1:00:53beginning and then we continue to
  1542. 1:00:55optimize just the single batch and so we
  1543. 1:00:57want to see that we can overfit this
  1544. 1:00:58single batch we can we can crush this
  1545. 1:01:01little batch and we can perfectly
  1546. 1:01:02predict the indices on just this little
  1547. 1:01:04batch and indeed that is roughly what
  1548. 1:01:06we're seeing here
  1549. 1:01:08so um we started off at roughly 10.82 11
  1550. 1:01:12in this case and then as we continue
  1551. 1:01:14optimizing on this single batch without
  1552. 1:01:16loading new examples we are making sure
  1553. 1:01:17that we can overfit a single batch and
  1554. 1:01:20we are getting to very very low loss so
  1555. 1:01:21the Transformer is memorizing this
  1556. 1:01:24single individual batch and one more
  1557. 1:01:26thing I didn't mention is uh the
  1558. 1:01:28learning rate here is 3 E4 which is a
  1559. 1:01:30pretty good default for most uh
  1560. 1:01:33optimizations that you want to run at a
  1561. 1:01:35very early debugging stage so this is
  1562. 1:01:38our simple inter Loop and uh we are
  1563. 1:01:41overfitting a single batch and this
  1564. 1:01:42looks good so now what uh what comes
  1565. 1:01:45next is we don't just want to overfit a
  1566. 1:01:46single batch we actually want to do an
  1567. 1:01:48optimization so we actually need to
  1568. 1:01:50iterate these XY batches and create a
  1569. 1:01:52little data loader uh that makes sure
  1570. 1:01:54that we're always getting a fresh batch
  1571. 1:01:56and that we're actually optimizing a
  1572. 1:01:57reasonable objective so let's do that
  1573. 1:01:59next okay so this is what I came up with
  1574. 1:02:01and I wrote a little data loader
  1575. 1:02:03light um so what this data loader does
  1576. 1:02:06is we're importing the token up here
  1577. 1:02:08we're reading the entire text file from
  1578. 1:02:10this single input.txt
  1579. 1:02:12tokenizing it and then we're just
  1580. 1:02:14printing the number of tokens in total
  1581. 1:02:17and the number of batches in a single
  1582. 1:02:19Epoch of iterating over this data set so
  1583. 1:02:22how many unique batches do we output
  1584. 1:02:24before we loop back around the beginning
  1585. 1:02:26of the document and start reading it
  1586. 1:02:28again so we start off at position zero
  1587. 1:02:31and then we simply walk the document in
  1588. 1:02:33batches of B * T so we take chunks of B
  1589. 1:02:36* T and then always Advance by B * T and
  1590. 1:02:40um it's important to note that we're
  1591. 1:02:42always advancing our position by exactly
  1592. 1:02:44B * T but when we're fetching the tokens
  1593. 1:02:47we're actually fetching from current
  1594. 1:02:49position to B * t + 1 and we need that
  1595. 1:02:52plus one because remember uh we need the
  1596. 1:02:55target token
  1597. 1:02:56um for the last token in the current
  1598. 1:02:58batch and so that way we can do um the
  1599. 1:03:02XY exactly as we did it before and if we
  1600. 1:03:07are to um run out of data we'll just
  1601. 1:03:09loop back around to zero so this is one
  1602. 1:03:12way to write a very very simple data
  1603. 1:03:13loader um that simply just goes through
  1604. 1:03:16the file in chunks and is good enough
  1605. 1:03:19for us uh for current purposes and we're
  1606. 1:03:21going to complexify it later and now
  1607. 1:03:24we'd like to come back around here and
  1608. 1:03:26we'd like to actually use our data
  1609. 1:03:27loader so the import Tik token has moved
  1610. 1:03:29up and actually all of this is now
  1611. 1:03:32useless so instead we just want a train
  1612. 1:03:35loader for the training data and we want
  1613. 1:03:38to use the same hyper parameters for
  1614. 1:03:39four so B size was four and time was
  1615. 1:03:4332 and then here we need to get the XY
  1616. 1:03:47for the current batch so let's see if
  1617. 1:03:49copal gets it because this is simple
  1618. 1:03:51enough uh so we call the next batch and
  1619. 1:03:53then we um make sure that we have to
  1620. 1:03:57move our tensors from CPU to the device
  1621. 1:04:02so here when I converted the tokens
  1622. 1:04:05notice that I didn't actually move these
  1623. 1:04:06tokens to the GPU I left them on CPU
  1624. 1:04:10which is the default um and that's just
  1625. 1:04:12because I'm trying not to waste too much
  1626. 1:04:14memory on the GPU in this case this is a
  1627. 1:04:16tiny data set and it would fit uh but
  1628. 1:04:19it's fine to just uh ship it to GPU
  1629. 1:04:21right now for for our purposes right now
  1630. 1:04:24so we get the next batch we keep the
  1631. 1:04:26data loader simple CPU class and then
  1632. 1:04:29here we actually ship it to the GPU and
  1633. 1:04:31do all the computation and uh let's see
  1634. 1:04:34if this runs so python train gbt2 pi and
  1635. 1:04:39what do we expect to see before this
  1636. 1:04:41actually happens what we expect to see
  1637. 1:04:43is now we're actually getting the next
  1638. 1:04:44batch so we expect to not overfit a
  1639. 1:04:47single batch and so I expect our loss to
  1640. 1:04:50come down but not too much and that's
  1641. 1:04:54because I still expect it to come down
  1642. 1:04:55because in the
  1643. 1:04:5750257 tokens many of those tokens never
  1644. 1:05:00occur in our data set so there are some
  1645. 1:05:02very easy gains to be made here in the
  1646. 1:05:04optimization by for example taking the
  1647. 1:05:06biases of all the loits that never occur
  1648. 1:05:08and driving them to negative infinity
  1649. 1:05:11and that would basically just it's just
  1650. 1:05:12that all of these crazy unic codes or
  1651. 1:05:14different languages those tokens never
  1652. 1:05:16occur so their probability should be
  1653. 1:05:17very low and so the gains that we should
  1654. 1:05:19be seeing are along the lines of
  1655. 1:05:22basically deleting the usage of tokens
  1656. 1:05:24that never occur that's probably most of
  1657. 1:05:26the loss gain that we're going to see at
  1658. 1:05:28this scale right now uh but we shouldn't
  1659. 1:05:30come to a zero uh because um we are only
  1660. 1:05:35doing 50 iterations and I don't think
  1661. 1:05:37that's enough to do an eoch right now so
  1662. 1:05:39let's see what we
  1663. 1:05:40got we um we have 338,000
  1664. 1:05:44tokens which makes sense with our 3:1
  1665. 1:05:47compression ratio because there are 1
  1666. 1:05:48million uh characters so one Epoch with
  1667. 1:05:52the current setting of B and T will take
  1668. 1:05:552, 600 batches and we're only doing 50
  1669. 1:05:58batches of optimization in
  1670. 1:06:01here so we start off in a familiar
  1671. 1:06:03territory as expected and then we seem
  1672. 1:06:05to come down to about 6.6 so basically
  1673. 1:06:09things seem to be working okay right now
  1674. 1:06:11with respect to our expectations so
  1675. 1:06:13that's good okay next I want to actually
  1676. 1:06:16fix a bug that we have in our code um
  1677. 1:06:18it's not a major bug but it is a bug
  1678. 1:06:20with respect to how gpt2 training uh
  1679. 1:06:22should
  1680. 1:06:24happen um
  1681. 1:06:26so the buck is the following we were not
  1682. 1:06:28being careful enough when we were
  1683. 1:06:29loading the weights from hugging face
  1684. 1:06:31and we actually missed a little detail
  1685. 1:06:33so if we come
  1686. 1:06:35here notice that um the shape of these
  1687. 1:06:38two tensors is the same so this one here
  1688. 1:06:42is the token embedding at the bottom of
  1689. 1:06:44the
  1690. 1:06:45Transformer right so and this one here
  1691. 1:06:48is the language modeling head at the top
  1692. 1:06:50of the
  1693. 1:06:51Transformer and both of these are
  1694. 1:06:53basically two-dimensional tensors and
  1695. 1:06:55they shape is identical so here the
  1696. 1:06:59first one is the output embedding the
  1697. 1:07:00token embedding and the second one is
  1698. 1:07:02this linear layer at the very top the
  1699. 1:07:04classifier layer both of them are of
  1700. 1:07:07shape
  1701. 1:07:0850257 X
  1702. 1:07:09768 um this one here is giving us our
  1703. 1:07:13token embeddings at the bottom and this
  1704. 1:07:16one here is taking the 768 channels of
  1705. 1:07:18the Transformer and trying to upscale
  1706. 1:07:21that to 50, 257 to get the Lis for the
  1707. 1:07:24next token so they're both the same
  1708. 1:07:27shape but more than that actually if you
  1709. 1:07:29look at um comparing their elements um
  1710. 1:07:33in pytorch this is an element wise
  1711. 1:07:35equality so then we use do all and we
  1712. 1:07:37see that every single element is
  1713. 1:07:39identical and more than that we see that
  1714. 1:07:42if we actually look at the data pointer
  1715. 1:07:44uh this is what this is a way in pytorch
  1716. 1:07:47to get the actual pointer to the uh data
  1717. 1:07:49and the storage we see that actually the
  1718. 1:07:51pointer is identical so not only are
  1719. 1:07:53these two separate tensors that happen
  1720. 1:07:55to have the same shape and elements
  1721. 1:07:57they're actually pointing to the
  1722. 1:07:58identical tensor so what's happening
  1723. 1:08:02here is that this is a common weight
  1724. 1:08:03tying scheme uh that actually comes from
  1725. 1:08:06the original
  1726. 1:08:08um from the original attention is all
  1727. 1:08:10you need paper and actually even the
  1728. 1:08:12reference before it so if we come
  1729. 1:08:16here
  1730. 1:08:19um eddings and softmax in the attention
  1731. 1:08:22is all you need paper they mentioned
  1732. 1:08:24that in our model we shared the same
  1733. 1:08:26weight Matrix between the two embedding
  1734. 1:08:28layers and the pre softmax linear
  1735. 1:08:30transformation similar to 30 um so this
  1736. 1:08:34is an awkward way to phrase that these
  1737. 1:08:36two are shared and they're tied and
  1738. 1:08:38they're the same Matrix and the 30
  1739. 1:08:40reference is this
  1740. 1:08:42paper um so this came out in
  1741. 1:08:452017 and you can read the full paper but
  1742. 1:08:47basically it argues for this weight
  1743. 1:08:49tying scheme and I think intuitively the
  1744. 1:08:53idea for why you might want to do this
  1745. 1:08:54comes from from this paragraph here and
  1746. 1:08:58basically you you can observe
  1747. 1:09:01that um you actually want these two
  1748. 1:09:04matrices to behave similar in the
  1749. 1:09:07following sense if two tokens are very
  1750. 1:09:10similar semantically like maybe one of
  1751. 1:09:12them is all lowercase and the other one
  1752. 1:09:14is all uppercase or it's the same token
  1753. 1:09:16in a different language or something
  1754. 1:09:17like that if you have similarity between
  1755. 1:09:19two tokens presumably you would expect
  1756. 1:09:21that they are uh nearby in the token
  1757. 1:09:23embedding space but in the exact same
  1758. 1:09:26way you'd expect that if you have two
  1759. 1:09:27tokens that are similar semantically
  1760. 1:09:30you'd expect them to get the same
  1761. 1:09:32probabilities at the output of a
  1762. 1:09:33transformer because they are
  1763. 1:09:35semantically similar and so both
  1764. 1:09:39positions in the Transformer at the very
  1765. 1:09:41bottom and at the top have this property
  1766. 1:09:43that similar tokens should have similar
  1767. 1:09:46embeddings or similar weights and so
  1768. 1:09:49this is what motivates their exploration
  1769. 1:09:51here and they they kind of you know I
  1770. 1:09:53don't want to go through the entire
  1771. 1:09:54paper and and uh you can go through it
  1772. 1:09:57but this is what they observe they also
  1773. 1:09:59observe that if you look at the output
  1774. 1:10:00embeddings they also behave like word
  1775. 1:10:02embeddings um if you um if you just kind
  1776. 1:10:06of try to use those weights as word
  1777. 1:10:08embeddings um so they kind of observe
  1778. 1:10:10this similarity they try to tie them and
  1779. 1:10:13they observe that they can get much
  1780. 1:10:14better performance in that way and so
  1781. 1:10:17this was adopted and the attention is
  1782. 1:10:18all need paper and then it was used
  1783. 1:10:20again in gpt2 as well
  1784. 1:10:24so I couldn't find it in the
  1785. 1:10:26Transformers implementation I'm not sure
  1786. 1:10:28where they tie those embeddings but I
  1787. 1:10:30can find it in the original gpt2 code U
  1788. 1:10:34introduced by open aai so this is um
  1789. 1:10:36openai gpt2 Source model and here where
  1790. 1:10:40they are forwarding this model and this
  1791. 1:10:41is in tensorflow but uh that's okay we
  1792. 1:10:44see that they get the wte token
  1793. 1:10:46embeddings and then here is the incoder
  1794. 1:10:50of the token embeddings and the
  1795. 1:10:52position and then here at the bottom
  1796. 1:10:54they Ed the WT again to do the lits so
  1797. 1:10:58when they get the loits it's a math Mo
  1798. 1:11:00of uh this output from the Transformer
  1799. 1:11:02and the wte tensor is
  1800. 1:11:05reused um and so the wte tensor
  1801. 1:11:08basically is used twice on the bottom of
  1802. 1:11:10the Transformer and on the top of the
  1803. 1:11:12Transformer and in the backward pass
  1804. 1:11:14we'll get gradients contributions from
  1805. 1:11:17both branches right and these gradients
  1806. 1:11:19will add up um on the wte tensor um so
  1807. 1:11:23we'll get a contribution from the
  1808. 1:11:24classifier list
  1809. 1:11:25and then at the very end of the
  1810. 1:11:27Transformer we'll get a contribution at
  1811. 1:11:28the at the bottom of it float floating
  1812. 1:11:31again into the wte uh tensor so we want
  1813. 1:11:35to we are currently not sharing WT and
  1814. 1:11:38our code but we want to do
  1815. 1:11:40that um
  1816. 1:11:44so weight sharing scheme um and one way
  1817. 1:11:48to do this let's see if goil gets it oh
  1818. 1:11:50it does okay uh so this is one way to do
  1819. 1:11:54it
  1820. 1:11:56uh
  1821. 1:11:56basically relatively straightforward
  1822. 1:11:59what we're doing here is we're taking
  1823. 1:12:00the wte do weight and we're simply uh
  1824. 1:12:04redirecting it to point to the LM head
  1825. 1:12:08so um this basically copies the data
  1826. 1:12:11pointer right it copies the reference
  1827. 1:12:14and now the wte weight becomes orphaned
  1828. 1:12:17uh the old value of it and uh pytorch
  1829. 1:12:20will clean it up python will clean it up
  1830. 1:12:23and so we are only left with a single
  1831. 1:12:26tensor and it's going to be used twice
  1832. 1:12:28in the forward pass and uh this is to my
  1833. 1:12:31knowledge all that's required so we
  1834. 1:12:34should be able to use this and this
  1835. 1:12:36should probably train uh we're just
  1836. 1:12:39going to basically be using this exact
  1837. 1:12:40same sensor twice and
  1838. 1:12:44um we weren't being careful with
  1839. 1:12:46tracking the likelihoods but uh
  1840. 1:12:48according to the paper and according to
  1841. 1:12:50the results you'd actually expect
  1842. 1:12:51slightly better results doing this and
  1843. 1:12:53in addition to that one other reason
  1844. 1:12:54that this is very very nice for us is
  1845. 1:12:57that this is a ton of parameters right
  1846. 1:12:59uh what is the size here it's 768 *
  1847. 1:13:0350257 so This Is 40 million parameters
  1848. 1:13:07and this is a 124 million parameter
  1849. 1:13:09model so 40 divide 124 so this is like
  1850. 1:13:1230% of the parameters are being saved
  1851. 1:13:15using this weight time scheme and so
  1852. 1:13:18this might be one of the reasons that
  1853. 1:13:20this is working slightly better if
  1854. 1:13:21you're not training the model long
  1855. 1:13:22enough because of the weight tying uh
  1856. 1:13:25you don't have to train as many
  1857. 1:13:26parameters and so you become more
  1858. 1:13:27efficient um in terms of the training
  1859. 1:13:30process uh because you have fewer
  1860. 1:13:32parameters and you're putting in this
  1861. 1:13:34inductive bias that these two embeddings
  1862. 1:13:36should share similarities between tokens
  1863. 1:13:40so this is the way time scheme and we've
  1864. 1:13:42saved a ton of parameters and we expect
  1865. 1:13:44our model to work slightly better
  1866. 1:13:45because of the scheme okay next I would
  1867. 1:13:47like us to be a bit more careful with
  1868. 1:13:49the initialization and to try to follow
  1869. 1:13:50the way gpt2 initialized their model now
  1870. 1:13:54unfortunately the gpt2 paper and the
  1871. 1:13:55gpt3 paper are not very explicit about
  1872. 1:13:58initialization so we kind of have to
  1873. 1:14:00read between the lines uh and instead of
  1874. 1:14:02going to the paper which is quite vague
  1875. 1:14:04um there's a bit of information in the
  1876. 1:14:07code that open I released so when we go
  1877. 1:14:09to the model.py we see that when they
  1878. 1:14:11initialize their weights they are using
  1879. 1:14:13the standard deviation of
  1880. 1:14:150.02 and that's how they they so this is
  1881. 1:14:19a normal distribution for the weights
  1882. 1:14:21and the standard deviation is
  1883. 1:14:230.02 for the bias they initialize that
  1884. 1:14:25with
  1885. 1:14:26zero and then when we scroll down
  1886. 1:14:30here why is this not scrolling
  1887. 1:14:33um the token embeddings are initialized
  1888. 1:14:36at
  1889. 1:14:370.02 and position embeddings at 0.01 for
  1890. 1:14:40some reason so those are the
  1891. 1:14:42initializations and we'd like to mirror
  1892. 1:14:44that in
  1893. 1:14:45gpt2 uh in our module here so here's a
  1894. 1:14:48snippet of code that I sort of came up
  1895. 1:14:50with very
  1896. 1:14:52quickly so what's happening here is at
  1897. 1:14:55the end of our initializer for the GPT
  1898. 1:14:57module we're calling the apply function
  1899. 1:14:59of NN module and that iterates all the
  1900. 1:15:02sub modules of this module and uh
  1901. 1:15:05applies in it weights function on them
  1902. 1:15:08and so what's happening here is that
  1903. 1:15:11we're in we're iterating all the modules
  1904. 1:15:13here and if they are an nn. linear
  1905. 1:15:16module then we're going to make sure to
  1906. 1:15:17initialize the weight using a normal
  1907. 1:15:19with the standard deviation of
  1908. 1:15:210.02 if there's a bias in this layer we
  1909. 1:15:24will make sure to initialize that to
  1910. 1:15:25zero note that zero initialization for
  1911. 1:15:28the bias is not actually the pyto
  1912. 1:15:29default um by default the bias here is
  1913. 1:15:33initialized with a uniform so uh that's
  1914. 1:15:36interesting so we make sure to use zero
  1915. 1:15:38and for the embedding we're just going
  1916. 1:15:40to use 0.02 and um keep it the same um
  1917. 1:15:43so we're not going to change it to 0.01
  1918. 1:15:45for positional because it's about the
  1919. 1:15:47same and then if you look through our
  1920. 1:15:49model the only other layer that requires
  1921. 1:15:51initialization and that has parameters
  1922. 1:15:53is the layer norm and the fighter defer
  1923. 1:15:55initialization sets the scale in the
  1924. 1:15:57layer Norm to be one and the offset in
  1925. 1:16:00the layer Norm to be zero so that's
  1926. 1:16:01exactly what we want and so we're just
  1927. 1:16:03going to uh keep it that way and so this
  1928. 1:16:06is the default initialization if we are
  1929. 1:16:09following the um where is it the uh gpt2
  1930. 1:16:14uh source code that they released I
  1931. 1:16:17would like to point out by the way that
  1932. 1:16:19um typically the standard deviation here
  1933. 1:16:21on this initialization if you follow the
  1934. 1:16:23Javier initialization would be one of
  1935. 1:16:24over the square root of the number of
  1936. 1:16:27features that are incoming into this
  1937. 1:16:28layer but if you'll notice actually 0.02
  1938. 1:16:32is basically consistent with that
  1939. 1:16:34because the the model sizes inside these
  1940. 1:16:36Transformers for gpt2 are roughly 768
  1941. 1:16:391600 Etc so 1 over the square root of
  1942. 1:16:41for example 768 gives us
  1943. 1:16:440.03 if we plug in 600 1,600 we get
  1944. 1:16:490.02 if we plug in three times that
  1945. 1:16:520.014 Etc so basically 0.02 is roughly
  1946. 1:16:56in the vicinity of reasonable values for
  1947. 1:16:59the for um for these initializations
  1948. 1:17:02anyway so so it's not uh completely
  1949. 1:17:05crazy to be hard coding 0.02 here uh but
  1950. 1:17:08you'd like typically uh some something
  1951. 1:17:11that grows with the model size instead
  1952. 1:17:13but we will keep this because that is
  1953. 1:17:15the gpt2 initialization per their source
  1954. 1:17:17code but we are not fully done yet on
  1955. 1:17:19initialization because there's one more
  1956. 1:17:20caveat here so
  1957. 1:17:23here a mod initialization which accounts
  1958. 1:17:26for the accumulation on the residual
  1959. 1:17:27path with model depth is used we scale
  1960. 1:17:30the weight of residual layers of
  1961. 1:17:31initialization by factor of one over squ
  1962. 1:17:33of n where n is the number of residual
  1963. 1:17:35layers so this is what gbt2 paper says
  1964. 1:17:38so we have not implemented that yet and
  1965. 1:17:41uh we can do so now now I'd like to
  1966. 1:17:43actually kind of like motivate a little
  1967. 1:17:44bit what they mean here I think um so
  1968. 1:17:47here's roughly what they
  1969. 1:17:49mean if you start out with zeros in your
  1970. 1:17:52residual stream remember that each
  1971. 1:17:54residual stream is a is of this form
  1972. 1:17:57where we continue adding to it X is X
  1973. 1:18:00plus something some kind of contribution
  1974. 1:18:02so every single block of the residual uh
  1975. 1:18:05Network contributes some uh amount and
  1976. 1:18:09it gets added and so what ends up
  1977. 1:18:11happening is that the variance of the
  1978. 1:18:15activations in the residual stream grows
  1979. 1:18:18so here's a small example if we start at
  1980. 1:18:19zero and then we for 100 times uh we
  1981. 1:18:23have sort of this residual stream of of
  1982. 1:18:25768 uh zeros and then 100 times we add
  1983. 1:18:30um random which is a normal distribution
  1984. 1:18:33zero mean one standard deviation if we
  1985. 1:18:36add to it then by the end the residual
  1986. 1:18:37stream has grown to have standard
  1987. 1:18:39deviation of 10 and that's just because
  1988. 1:18:42um we're always adding um these numbers
  1989. 1:18:47and so this scaling factor that they use
  1990. 1:18:50here exactly compensates for that growth
  1991. 1:18:53so if we take n and we basically um
  1992. 1:18:57scale down every one of these
  1993. 1:18:59contributions into the residual stream
  1994. 1:19:00by one over theare Ro of n so 1 over
  1995. 1:19:03theun of n is n to the 0.5
  1996. 1:19:07right because n the5 is the square root
  1997. 1:19:11and then one over the square root is n.5
  1998. 1:19:14if we scale it in this way then we see
  1999. 1:19:16that we actually get um
  2000. 1:19:20one
  2001. 1:19:21so this is a way to control the growth
  2002. 1:19:24of of activations inside the residual
  2003. 1:19:26stream in the forward pass and so we'd
  2004. 1:19:29like to initialize in the same way where
  2005. 1:19:31these weights that are at the end of
  2006. 1:19:33each block so this C uh layer uh the gbt
  2007. 1:19:38paper proposes to scale down those
  2008. 1:19:40weights by one over the square root of
  2009. 1:19:42the number of residual
  2010. 1:19:43layers so one crude way to implement
  2011. 1:19:46this is the following I don't know if
  2012. 1:19:48this is uh pyro sanctioned but it works
  2013. 1:19:50for me is we'll do in the
  2014. 1:19:53initialization see that s that do
  2015. 1:19:56special nanog
  2016. 1:19:58GPT uh scale in it is one so we're
  2017. 1:20:04setting um kind of like a flag for this
  2018. 1:20:06module there must be a better way in py
  2019. 1:20:08torch right but I don't
  2020. 1:20:11know okay so we're basically attaching
  2021. 1:20:13this flag and trying to make sure that
  2022. 1:20:16it doesn't conflict with anything
  2023. 1:20:17previously and then when we come down
  2024. 1:20:20here this STD should be 0.02 by default
  2025. 1:20:25but then if
  2026. 1:20:27haat um module of this thing
  2027. 1:20:31then STD *
  2028. 1:20:34equals
  2029. 1:20:36um copal is not guessing correctly uh so
  2030. 1:20:39we want one over the square root of the
  2031. 1:20:41number of layers so
  2032. 1:20:44um the number of residual layers here is
  2033. 1:20:47twice
  2034. 1:20:48times Salt out config layers and then
  2035. 1:20:52this times .5 so we want to scale down
  2036. 1:20:57that standard deviation and this should
  2037. 1:20:59be um correct and Implement that I
  2038. 1:21:03should clarify by the way that the two
  2039. 1:21:04times number of layers comes from the
  2040. 1:21:06fact that every single one of our layers
  2041. 1:21:07in the Transformer actually has two
  2042. 1:21:09blocks that add to the ridal pathway
  2043. 1:21:11right we have the attention and then the
  2044. 1:21:13MLP so that's where the two times comes
  2045. 1:21:16from and the other thing to mention is
  2046. 1:21:18that uh what's slightly awkward but
  2047. 1:21:21we're not going to fix it is that um
  2048. 1:21:23because we are weight sharing the wte
  2049. 1:21:26and the LM head in this iteration of our
  2050. 1:21:29old subm modules we're going to actually
  2051. 1:21:31come around to that tensor twice so
  2052. 1:21:33we're going to first initialize it as an
  2053. 1:21:34embedding with 0.02 and then we're going
  2054. 1:21:37to come back around it again in a linear
  2055. 1:21:39and initialize it again using 0.02 and
  2056. 1:21:42it's going to be 0.02 because the LM
  2057. 1:21:44head is of course not not scaled so it's
  2058. 1:21:46not going to come here it's just it's
  2059. 1:21:48going to be basically initialized twice
  2060. 1:21:50using the identical same initialization
  2061. 1:21:52but that's okay and then scrolling over
  2062. 1:21:56here I added uh some code here so that
  2063. 1:21:59we have
  2064. 1:22:00reproducibility um to set the seeds and
  2065. 1:22:03now we should be able to python train
  2066. 1:22:05gpt2 pi and let this running and as far
  2067. 1:22:09as I know this is the gpt2
  2068. 1:22:11initialization uh in the way we've
  2069. 1:22:12implemented it right now so this
  2070. 1:22:16looks uh reasonable to me okay so at
  2071. 1:22:19this point we have the gpt2 model we
  2072. 1:22:21have some confidence that it's correctly
  2073. 1:22:23implemented we've initialized it
  2074. 1:22:24properly and we have a data loader
  2075. 1:22:26that's iterating through data batches
  2076. 1:22:27and we can train so now comes the fun
  2077. 1:22:30part I'd like us to speed up the
  2078. 1:22:31training by a lot so we're getting our
  2079. 1:22:33money's worth with respect to the
  2080. 1:22:34hardware that we are uh using here and
  2081. 1:22:38uh we're going to speed up the training
  2082. 1:22:39by quite a bit uh now you always want to
  2083. 1:22:42start with what Hardware do you have
  2084. 1:22:44what does it offer and are you fully
  2085. 1:22:45utilizing it so in my case if we go to
  2086. 1:22:48Nvidia
  2087. 1:22:49SMI we can see
  2088. 1:22:53that I have eight gpus and each one of
  2089. 1:22:57those gpus is an a100 sxm 80 gb so this
  2090. 1:23:01is the GPU that I have available to me
  2091. 1:23:03in this box now when I look when I use
  2092. 1:23:07um to spin up these kinds of Boxes by
  2093. 1:23:09the way my favorite place to go to is
  2094. 1:23:11Lambda Labs um they do sponsor my
  2095. 1:23:14development and that of my projects uh
  2096. 1:23:17but I this is my favorite place to go
  2097. 1:23:20and this is where you can spin up one of
  2098. 1:23:21these machines and you pay per hour and
  2099. 1:23:23it's very very simple
  2100. 1:23:25so I like to spin them up and then
  2101. 1:23:26connect vsod to it and that's how I
  2102. 1:23:28develop now when we look at the A1 100s
  2103. 1:23:30that are available here a100 80 GB sxm
  2104. 1:23:35is the um GPU that I have here and we
  2105. 1:23:39have a bunch of numbers here for um how
  2106. 1:23:41many calculations you can expect out of
  2107. 1:23:43this GPU so when I come over here
  2108. 1:23:46and I break in right after here so
  2109. 1:23:50python
  2110. 1:23:51trity so I'm breaking in right after we
  2111. 1:23:53calculate the loit and
  2112. 1:23:55laws and the interesting thing I'd like
  2113. 1:23:57you to note is when I do lit. dtype this
  2114. 1:24:02prints a torch. FL 32 so by default iny
  2115. 1:24:06torch when you create tensors um and
  2116. 1:24:08this is the case for all the activations
  2117. 1:24:10and for the parameters of the network
  2118. 1:24:11and so on by default everything is in
  2119. 1:24:13float 32 that means that every single
  2120. 1:24:17number activation or weight and so on is
  2121. 1:24:20using a float representation that has 32
  2122. 1:24:23bits and uh that's actually quite a bit
  2123. 1:24:26of memory and it turns out empirically
  2124. 1:24:27that for deep learning as a
  2125. 1:24:28computational workload this is way too
  2126. 1:24:30much and deep learning and the training
  2127. 1:24:32of these networks can tolerate
  2128. 1:24:34significantly lower precisions um not
  2129. 1:24:37all computational workflows can tolerate
  2130. 1:24:39small Precision so for example um if we
  2131. 1:24:43go back to to the data sheet you'll see
  2132. 1:24:45that actually these gpus support up to
  2133. 1:24:48fp64 and this is quite useful I
  2134. 1:24:50understand for a lot of um scientific
  2135. 1:24:52Computing applications and there really
  2136. 1:24:54need this uh but we don't need that much
  2137. 1:24:56Precision for deep learning training So
  2138. 1:24:59currently we are here
  2139. 1:25:01fp32 and with this code as it is right
  2140. 1:25:04now we expect to get at at most 19.5
  2141. 1:25:08Tera flops of performance that means
  2142. 1:25:10we're doing 19.5 trillion operations
  2143. 1:25:13floating Point operations so this is
  2144. 1:25:15floating Point multiply add most um most
  2145. 1:25:20likely and so these are the floating
  2146. 1:25:23Point operations
  2147. 1:25:25uh now notice that if we are willing to
  2148. 1:25:27go down in Precision so tf32 is a lower
  2149. 1:25:31Precision format we're going to see in a
  2150. 1:25:32second you can actually get an 8X
  2151. 1:25:34Improvement here and if you're willing
  2152. 1:25:36to go down to float 16 or B float 16 you
  2153. 1:25:39can actually get time 16x performance
  2154. 1:25:42all the way to 312 Tera flops you see
  2155. 1:25:45here that Nvidia likes to site numbers
  2156. 1:25:47that have an asterisk here this asterisk
  2157. 1:25:50uh says with sparsity uh but we are not
  2158. 1:25:52going to be using sparsity in R code and
  2159. 1:25:55I don't know that this is very widely
  2160. 1:25:56used in the industry right now so most
  2161. 1:25:58people look at this number here uh
  2162. 1:26:01without sparcity and you'll notice that
  2163. 1:26:03we could have got even more here but
  2164. 1:26:05this is int 8 and int 8 is used for
  2165. 1:26:08inference not for training uh because
  2166. 1:26:11int 8 has a um it basically has um
  2167. 1:26:17uniform
  2168. 1:26:18spacing um and uh we actually require a
  2169. 1:26:21float so that we get a better match to
  2170. 1:26:24the uh normal distributions that occur
  2171. 1:26:28during training of neural networks where
  2172. 1:26:29both activations and weights are
  2173. 1:26:31distributed as a normal distribution and
  2174. 1:26:33so uh floating points are really
  2175. 1:26:35important to to match that uh
  2176. 1:26:38representation so we're not typically
  2177. 1:26:40using int 8 uh for training but we are
  2178. 1:26:42using it for inference and if we bring
  2179. 1:26:45down the Precision we can get a lot more
  2180. 1:26:47Terra flops out of the tensor course
  2181. 1:26:49available in the gpus we'll talk about
  2182. 1:26:51that in a second but in addition to that
  2183. 1:26:53if all of these numbers have fewer bits
  2184. 1:26:56of representation it's going to be much
  2185. 1:26:58easier to move them around and that's
  2186. 1:27:00where we start to get into the memory
  2187. 1:27:02bandwidth and the memory of the model so
  2188. 1:27:04not only do we have a finite capacity of
  2189. 1:27:06the number of bits that our GPU can
  2190. 1:27:08store but in addition to that there's a
  2191. 1:27:11speed with which you can access this
  2192. 1:27:13memory um and you have a certain memory
  2193. 1:27:16bandwidth it's a very precious resource
  2194. 1:27:19and in fact many of the deep learning uh
  2195. 1:27:21work workloads for training are memory
  2196. 1:27:23bound and what that means is actually
  2197. 1:27:25that the tensor cores that do all these
  2198. 1:27:27extremely fast multiplications most of
  2199. 1:27:29the time they're waiting around they're
  2200. 1:27:31idle um because we can't feed them with
  2201. 1:27:34data fast enough we can't load the data
  2202. 1:27:37fast enough from memory so typical
  2203. 1:27:38utilizations of your Hardware if you're
  2204. 1:27:40getting 60% uh utilization you're
  2205. 1:27:43actually doing extremely well um so half
  2206. 1:27:46of the time in a well-tuned application
  2207. 1:27:48your tensor cores are not doing
  2208. 1:27:50multiplies because the data is not
  2209. 1:27:51available so the memory bandwidth here
  2210. 1:27:53is extremely important as well and if we
  2211. 1:27:55come down in the Precision for all the
  2212. 1:27:58floats all the numbers weights and
  2213. 1:28:00activations suddenly require less memory
  2214. 1:28:02so we can store more and we can access
  2215. 1:28:05it faster so everything speeds up and
  2216. 1:28:07it's amazing and now let's reap the
  2217. 1:28:09benefits of it um and let's first look
  2218. 1:28:12at the tensor float 32
  2219. 1:28:14format okay so first of all what are
  2220. 1:28:16tensor cores well tensor course tensor
  2221. 1:28:19core is just an instruction in the a100
  2222. 1:28:22architecture right so so what it does is
  2223. 1:28:25it does basically a little 4x4 Matrix
  2224. 1:28:27multiply so uh this is just matrix
  2225. 1:28:30multiplication here of 4x4 matrices and
  2226. 1:28:35there are multiple configurations as to
  2227. 1:28:38what Precision any of these matrices are
  2228. 1:28:40it in what Precision the internal
  2229. 1:28:42accumulate happens and then what is the
  2230. 1:28:45output Precision input precisions Etc so
  2231. 1:28:47there's a few switches but it's
  2232. 1:28:48basically a 4x4 multiply and then
  2233. 1:28:51anytime we have any operations that
  2234. 1:28:53require Magic multiplication uh they get
  2235. 1:28:55broken up into these into this
  2236. 1:28:58instruction of little 4x4 multiply and
  2237. 1:29:00so everything gets broken up into this
  2238. 1:29:02instruction because it's the fastest way
  2239. 1:29:04to multiply matrices and it turns out
  2240. 1:29:06that most of the computational work that
  2241. 1:29:08we're doing up above uh all of it really
  2242. 1:29:10is matrix multiplication most of the
  2243. 1:29:12work computationally happens in the
  2244. 1:29:14linear layers um linear linear Etc
  2245. 1:29:20there's a few things sandwiched in
  2246. 1:29:21between so there's some additions in
  2247. 1:29:23residuals there's some G nonlinearities
  2248. 1:29:25there's some layer Norms Etc but if you
  2249. 1:29:28just time them you'll see that these are
  2250. 1:29:30nothing like basically the in
  2251. 1:29:32Transformer is just a bunch of Matrix
  2252. 1:29:34multiplications really um and especially
  2253. 1:29:37at this small scale 124 million
  2254. 1:29:39parameter model actually the biggest
  2255. 1:29:42matrix multiplication by far is the
  2256. 1:29:44classifier layer at the top that is a
  2257. 1:29:46massive Matrix multiply of going from
  2258. 1:29:49768 to
  2259. 1:29:5050257 and that Matrix multiply dominates
  2260. 1:29:53anything else that happens in that
  2261. 1:29:55Network roughly speaking so it's Matrix
  2262. 1:29:58multiplies that become a lot faster
  2263. 1:30:00which are hidden inside our linear
  2264. 1:30:02layers and they're accelerated through
  2265. 1:30:05tensor course now the best reference I
  2266. 1:30:07would say for tensor course is basically
  2267. 1:30:09just go to the um a 100 architecture
  2268. 1:30:13white paper and then it's pretty
  2269. 1:30:15detailed and but I think people it's
  2270. 1:30:18like relatively readable mostly if you
  2271. 1:30:20half understand what's happening um so
  2272. 1:30:23figure 9 tensor float
  2273. 1:30:2632 so this is the explanation basically
  2274. 1:30:28for tf32 and what happens here and you
  2275. 1:30:31see that there's many configuration
  2276. 1:30:32options here available so the input
  2277. 1:30:35operands and what precisions are they in
  2278. 1:30:37the accumulator and um what um basically
  2279. 1:30:41the um the internal representation
  2280. 1:30:44within the instruction when you do the
  2281. 1:30:46accumulate of this matrix
  2282. 1:30:48multiplication so the intermediate plus
  2283. 1:30:51equals um of the intermediate little
  2284. 1:30:53vector multiplies here that all happens
  2285. 1:30:55in
  2286. 1:30:57fp32 and then uh this is an aex
  2287. 1:31:00improvement as I mentioned to the Ops
  2288. 1:31:01that we get so tf32 specifically we're
  2289. 1:31:04looking at this row here and the way
  2290. 1:31:06this works
  2291. 1:31:07is
  2292. 1:31:10um normally fp32 has 32 bits
  2293. 1:31:14tf32 is the exact same bits we have one
  2294. 1:31:18sign bit we have eight exponent bits
  2295. 1:31:21except the mantisa bits get cropped in
  2296. 1:31:24the float and so basically um we end up
  2297. 1:31:27with just 19 bits instead of 32 bits
  2298. 1:31:30because the last 133 bits get truncated
  2299. 1:31:33they get dropped um and all this is
  2300. 1:31:36internal to the instruction so none of
  2301. 1:31:38it is visible to anything in our pytorch
  2302. 1:31:41uh none of our pytorch code will change
  2303. 1:31:43all of the numbers will look identical
  2304. 1:31:45it's just that when you call the tensor
  2305. 1:31:47core um instruction internally in the
  2306. 1:31:50hardware it will crop out these 13 bits
  2307. 1:31:54and that allows it to uh calculate this
  2308. 1:31:57little Matrix multiply significantly
  2309. 1:31:59faster 8X faster now of course this
  2310. 1:32:02speed up comes at a cost and the cost is
  2311. 1:32:04that we are reducing the Precision our
  2312. 1:32:07accumulate is still an fp32 our output
  2313. 1:32:09is fp32 our inputs are fp32 but
  2314. 1:32:12internally things get truncated in the
  2315. 1:32:14operand to perform the operation faster
  2316. 1:32:17and so our results are starting to be a
  2317. 1:32:19bit more approximate but empirically
  2318. 1:32:21when you actually train with this you
  2319. 1:32:22basically can't tell the difference
  2320. 1:32:24so the reason I like tf32 is because if
  2321. 1:32:26you can tolerate a little bit of a
  2322. 1:32:28Precision fudge um then this is free
  2323. 1:32:32like none of your codes sees this it's
  2324. 1:32:34fully internal to the operation and the
  2325. 1:32:36operation to you just go 8X faster and
  2326. 1:32:39it's a bit more approximate and so it's
  2327. 1:32:42a pretty sweet spot I would say in
  2328. 1:32:43optimization and uh let's see what that
  2329. 1:32:46looks like first so I've set up our Cod
  2330. 1:32:48to just time the uh iterations so import
  2331. 1:32:51time I changed the hyper parameters so
  2332. 1:32:54that we have something a bit more that
  2333. 1:32:55reflects uh kind of workload that we
  2334. 1:32:57want to run uh because we want to do a
  2335. 1:32:59fairly large run at the end of this so
  2336. 1:33:01let's use batch size 16 and let's now
  2337. 1:33:04use the actual gpt2 um maximum sequence
  2338. 1:33:07length of 10,24
  2339. 1:33:08tokens uh so this is the
  2340. 1:33:11configuration and then for 50 iterations
  2341. 1:33:15I'm just doing something very lazy here
  2342. 1:33:17I'm doing time. time to get the current
  2343. 1:33:19time and then this is the optimization
  2344. 1:33:22Loop and now I want to time how long
  2345. 1:33:24this takes now one issue with working
  2346. 1:33:28with gpus is that as your
  2347. 1:33:32CPU um when your CPU runs it's just
  2348. 1:33:35scheduling work on GPU it's ordering
  2349. 1:33:38some work right and so it send a request
  2350. 1:33:40and then it continues running and so we
  2351. 1:33:43can actually it can happen sometimes
  2352. 1:33:44that we sort of um speed through this
  2353. 1:33:48and we queue up a lot of kernels to run
  2354. 1:33:50on the GPU and then the CPU sort of like
  2355. 1:33:52gets here and takes time at time but
  2356. 1:33:54actually the GPU is still running
  2357. 1:33:56because it takes it time to actually
  2358. 1:33:57work through the work that was scheduled
  2359. 1:34:00to run and so you're just building up a
  2360. 1:34:03queue for the GPU and so actually if you
  2361. 1:34:05need to you want to wait toat data
  2362. 1:34:07synchronize and this will wait for the
  2363. 1:34:10GPU to finish all the work that was
  2364. 1:34:12scheduled to run up above here and then
  2365. 1:34:15we can actually take the time so
  2366. 1:34:17basically we're waiting for the GPU to
  2367. 1:34:19stop this iteration take time and then
  2368. 1:34:22we're going to just print it so
  2369. 1:34:24so here I'm going to run the training
  2370. 1:34:26Loop and here on the right I'm watching
  2371. 1:34:29Nvidia SMI so we start off at zero um
  2372. 1:34:33we're not using the GPU and then by
  2373. 1:34:35default P will use gpu0 so we see that
  2374. 1:34:37it gets filled up and we're using 35 GB
  2375. 1:34:40out of 80 gabt
  2376. 1:34:42available and then here on the left we
  2377. 1:34:45see that because we've cranked up the
  2378. 1:34:47batch
  2379. 1:34:48size now it's only 20 batches to do a
  2380. 1:34:51single Epoch on our tiny Shakespeare
  2381. 1:34:54and we see that we're seeing roughly a
  2382. 1:34:55th000 milliseconds per iteration here
  2383. 1:34:58right
  2384. 1:35:00so the first iteration sometimes is
  2385. 1:35:02slower and that's because pytorch might
  2386. 1:35:04be doing a lot of initializations here
  2387. 1:35:06on the very first iteration and so it's
  2388. 1:35:08probably initializing all these uh
  2389. 1:35:09tensors and buffers to hold all the
  2390. 1:35:11gradients and I'm not 100% sure all the
  2391. 1:35:13work that happens here but uh this could
  2392. 1:35:16be a slower iteration when you're timing
  2393. 1:35:18your logic you always want to be careful
  2394. 1:35:19with that but basically we're seeing a
  2395. 1:35:21th000 milliseconds per iteration
  2396. 1:35:24um and so this will run for roughly 50
  2397. 1:35:26seconds as we have it right now so
  2398. 1:35:29that's our Baseline in flo 32 one more
  2399. 1:35:32thing I wanted to mention is that if
  2400. 1:35:35this doesn't fit into your GPU and
  2401. 1:35:36you're getting out of memory errors then
  2402. 1:35:38start decreasing your batch size until
  2403. 1:35:40things fit so instead of 16 try eight or
  2404. 1:35:42four or whatever you need to fit um the
  2405. 1:35:46batch into your GPU and if you have a
  2406. 1:35:48bigger GPU you can actually potentially
  2407. 1:35:49get away with 32 and so on uh by default
  2408. 1:35:52you want to basically max out has Max
  2409. 1:35:54Max out the batch size that fits on your
  2410. 1:35:56GPU and you want to keep it nice numbers
  2411. 1:35:59so use numbers that have lots of powers
  2412. 1:36:01of two in them so 16 is a good number 8
  2413. 1:36:0524 32 48 These are nice numbers but
  2414. 1:36:09don't use something like 17 uh because
  2415. 1:36:11that will run very inefficiently on a
  2416. 1:36:12GPU uh and we're going to see that a bit
  2417. 1:36:14later as well so for now let's just
  2418. 1:36:17stick with
  2419. 1:36:1816124 and uh the one thing that I added
  2420. 1:36:22also here and I ran it again is I'm
  2421. 1:36:25calculating a tokens per second
  2422. 1:36:27throughput during training
  2423. 1:36:29because we might end up changing the
  2424. 1:36:31backat size around over time but tokens
  2425. 1:36:34per second is the objective measure that
  2426. 1:36:35we actually really care about how many
  2427. 1:36:37tokens of data are we training on and
  2428. 1:36:39what is the throughput of tokens that
  2429. 1:36:41we're getting in our optimization so
  2430. 1:36:43right now we're processing and training
  2431. 1:36:44on 163,000 tokens per second roughly and
  2432. 1:36:48that's a bit more objective
  2433. 1:36:50metric okay so let's now enable tf32 now
  2434. 1:36:53luckily pytorch makes this fairly easy
  2435. 1:36:56for us and uh to enable tf32 you just
  2436. 1:36:59need to do a single line and is this and
  2437. 1:37:02when we go to the py documentation here
  2438. 1:37:04for this function basically this tells
  2439. 1:37:07pych what kind of kernels to run and by
  2440. 1:37:10default I believe it is highest highest
  2441. 1:37:13Precision for mat M and that means that
  2442. 1:37:15everything happens in float 32 just like
  2443. 1:37:18it did before but if we set it to high
  2444. 1:37:20as we do right now Matrix
  2445. 1:37:22multiplications will not use tensor flow
  2446. 1:37:2432 when it's
  2447. 1:37:26available my GPU is a100 so it's an
  2448. 1:37:30ampere series and therefore tf32 is
  2449. 1:37:33available if you have an older GPU this
  2450. 1:37:35might not be available for you but for
  2451. 1:37:38my GPU it's available and so what I
  2452. 1:37:39expect P to do is that every single
  2453. 1:37:41place where we see an nn. linear inside
  2454. 1:37:44there there's a matrix multiplication
  2455. 1:37:46and I expect that matrix multiplication
  2456. 1:37:48now to be um running on tensor course
  2457. 1:37:51utilizing the TF 32%
  2458. 1:37:55so this is the single line of change
  2459. 1:37:58that is I believe necessary and let's
  2460. 1:37:59rerun this now we saw that um in terms
  2461. 1:38:03of the throughput that is promised to us
  2462. 1:38:05we're supposed to be getting 8X roughly
  2463. 1:38:08so let's see what
  2464. 1:38:10happens and that 8X came from here right
  2465. 1:38:15um 8X and it also came from looking at
  2466. 1:38:20it um here 156 T flops instead of of
  2467. 1:38:2419.5 okay so what actually happened uh
  2468. 1:38:27so we're seeing that our throughput
  2469. 1:38:29roughly 3x not aex so we are going we're
  2470. 1:38:35from 1,000 milliseconds we're going down
  2471. 1:38:37to 300 milliseconds and our throughput
  2472. 1:38:39is now about 50,000 tokens per second so
  2473. 1:38:41we have a roughly 3x instead of 8X so
  2474. 1:38:43what happened and basically What's
  2475. 1:38:46Happening Here is again a lot of these
  2476. 1:38:48workloads are memory bound and so even
  2477. 1:38:51though the
  2478. 1:38:52tf32 offers in principle a lot faster
  2479. 1:38:57throughput all of these numbers
  2480. 1:38:59everywhere are still float 32s and it's
  2481. 1:39:01float 32 numbers that are being shipped
  2482. 1:39:03all over the place through the memory
  2483. 1:39:05system and is just costing us way too
  2484. 1:39:07much time to shuttle around all this
  2485. 1:39:08data and so even though we've made the
  2486. 1:39:10multiply itself much faster uh we are
  2487. 1:39:13memory bound and we're not actually
  2488. 1:39:14seeing the full benefit uh that would
  2489. 1:39:16come from uh this napkin math here uh
  2490. 1:39:19that said we are getting one a 3X faster
  2491. 1:39:22throughput and this is free um single
  2492. 1:39:26line of code in P torch all your
  2493. 1:39:28variables are still float 32 everywhere
  2494. 1:39:30it just runs faster and it's slightly
  2495. 1:39:32more approximate but we're not going to
  2496. 1:39:34notice it basically uh so that's
  2497. 1:39:37tf32 okay so let's now continue so we've
  2498. 1:39:41exercised this row and um we saw that we
  2499. 1:39:44can crop out some of the Precision
  2500. 1:39:46inside the operation itself but we saw
  2501. 1:39:49that we're still memory bound we're
  2502. 1:39:50still moving around all these floats
  2503. 1:39:52right otherwise and we're paying that
  2504. 1:39:53cost because of this so let's now
  2505. 1:39:56decrease the amount of stuff that we're
  2506. 1:39:57going to be moving around and we're
  2507. 1:39:59going to do that by dropping down to B
  2508. 1:40:01float 16 so we're only going to be
  2509. 1:40:04maintaining 16 bits per float and we're
  2510. 1:40:07going to use the B flat 16 and I'll
  2511. 1:40:08explain in a bit uh fp16 difference and
  2512. 1:40:12uh we're going to be in this row so when
  2513. 1:40:14we go back to the documentation here for
  2514. 1:40:17the a
  2515. 1:40:18100 um we see here the precisions that
  2516. 1:40:23are are available and this is the
  2517. 1:40:25original fp32 the tf32 crops out the
  2518. 1:40:28Precision and then here in
  2519. 1:40:30bf16 you see that it is very similar to
  2520. 1:40:33tf32 but it's even more aggressive in
  2521. 1:40:36cropping off of the Precision the
  2522. 1:40:38mantisa of this float so the important
  2523. 1:40:40thing with B float 16 is that the
  2524. 1:40:42exponent bits and the sign bit of course
  2525. 1:40:45remain unchanged so if you're familiar
  2526. 1:40:47with your float numbers and I think this
  2527. 1:40:49should should probably be an entire
  2528. 1:40:52video by itself
  2529. 1:40:53the exponent sets the range that you can
  2530. 1:40:56represent of your numbers and the
  2531. 1:40:58Precision is how much Precision you have
  2532. 1:41:00for your numbers and so the range of
  2533. 1:41:04numbers is identical but we can we have
  2534. 1:41:07fewer possibilities within that range
  2535. 1:41:10because we are truncating the Mena so we
  2536. 1:41:12have less Precision in that
  2537. 1:41:14range what that means is that things are
  2538. 1:41:17actually fairly nice because we have the
  2539. 1:41:19original range of numbers that are
  2540. 1:41:21representable in float but we just have
  2541. 1:41:24less Precision for it and the difference
  2542. 1:41:27with fp16 is that they actually touch
  2543. 1:41:29and change the range so fp16 cannot
  2544. 1:41:32represent the full range of fp32 it has
  2545. 1:41:35a reduced range and that's where you
  2546. 1:41:37start to actually run into issues
  2547. 1:41:39because now you need uh these gradient
  2548. 1:41:41scalers and things like that and I'm not
  2549. 1:41:43going to go into the detail of that in
  2550. 1:41:45this video because that's a whole video
  2551. 1:41:48by itself but fb16 actually historically
  2552. 1:41:50came first that was available in the
  2553. 1:41:52Volta series before Amper and so fp16
  2554. 1:41:56came first and everyone started to train
  2555. 1:41:58in fp16 but everyone had to use all
  2556. 1:42:00these gradient scaling operations which
  2557. 1:42:02are kind of annoying and it's an
  2558. 1:42:03additional source of state and
  2559. 1:42:05complexity and the reason for that was
  2560. 1:42:07because the exponent range was reduced
  2561. 1:42:09in fp16 so that's the i e fp16 spec and
  2562. 1:42:13then they came out with bf16 and the
  2563. 1:42:15Ampere and they made it much simpler
  2564. 1:42:18because we're just truncating manessa we
  2565. 1:42:20have the exact same range and we do not
  2566. 1:42:21need gradient scalers so everything is
  2567. 1:42:24much much simpler now when we do use
  2568. 1:42:26bf16 though we are impacting the numbers
  2569. 1:42:30that we might be seeing in our pytorch
  2570. 1:42:32code these this change is not just local
  2571. 1:42:35to the operation itself so let's see how
  2572. 1:42:37that works
  2573. 1:42:39um there's some documentation here that
  2574. 1:42:43so I think this is probably the best
  2575. 1:42:44best page to explain how to use mixed
  2576. 1:42:46Precision in pytorch um because there
  2577. 1:42:49are many other tutorials and so on even
  2578. 1:42:51within pitor documentation that are a
  2579. 1:42:53lot more confusing and so I recommend
  2580. 1:42:55specifically this one because there's
  2581. 1:42:57five other copies that I would not
  2582. 1:42:59recommend and then when we come
  2583. 1:43:02here ignore everything about everything
  2584. 1:43:05ignore everything about gradient
  2585. 1:43:07scalers and only look at torch.
  2586. 1:43:10AutoCast and basically also this comes
  2587. 1:43:13to a single line of code at the end so
  2588. 1:43:15this is the context manager that we
  2589. 1:43:18want and we want to use that in our
  2590. 1:43:21Network when you click into the torch.
  2591. 1:43:25AutoCast autocasting it has a few more
  2592. 1:43:28uh a bit more guideline for you so it's
  2593. 1:43:30telling you do not call B flat 16 on any
  2594. 1:43:34of your tensors just use AutoCast and
  2595. 1:43:36only surround the uh forward pass of the
  2596. 1:43:39model and the loss calculation and
  2597. 1:43:41that's the only two things that you
  2598. 1:43:43should be surrounding leave the backward
  2599. 1:43:45and the optimizer step alone so that's
  2600. 1:43:47the guidance that comes from the P team
  2601. 1:43:49so we're going to follow that guidance
  2602. 1:43:51and for us because the L calculation is
  2603. 1:43:53inside of the model forward pass for us
  2604. 1:43:56we are going to be doing
  2605. 1:43:58this and then we don't want to be using
  2606. 1:44:00torch Flo 16 because if we do that we
  2607. 1:44:02need to start using gradient scalers as
  2608. 1:44:04well so we are going to be using B float
  2609. 1:44:0616 this is only possible to do an ampere
  2610. 1:44:09uh but this means that the changes are
  2611. 1:44:11extremely minimal like basically just
  2612. 1:44:13this one line of
  2613. 1:44:14code um let me first break
  2614. 1:44:19in to here before we actually run this
  2615. 1:44:22so right after logits I'd like to show
  2616. 1:44:25you that different from the tf32 that we
  2617. 1:44:28saw this is actually going to impact our
  2618. 1:44:31tensors
  2619. 1:44:32so this Lis tensor if we now look at
  2620. 1:44:36this and we look at the dtype we
  2621. 1:44:38suddenly see that this is now B float
  2622. 1:44:4016 uh it's not float 32 anymore so our
  2623. 1:44:43activations have been changed the
  2624. 1:44:45activations tensor is now B FL 16 but
  2625. 1:44:48not everything has changed so model.
  2626. 1:44:51Transformer
  2627. 1:44:55wte uh this is the weight uh token
  2628. 1:44:57embedding table it has a weight inside
  2629. 1:45:00it and the dtype of this weight this
  2630. 1:45:02parameter is still torch float 32 so our
  2631. 1:45:06parameters seem to still be in float 32
  2632. 1:45:09but our activations the loits are now in
  2633. 1:45:11P 16 so clearly this is why we get the
  2634. 1:45:14mixed Precision some things pytorch is
  2635. 1:45:16keeping inlow 32 some things pytorch is
  2636. 1:45:19converting to lower Precision um and
  2637. 1:45:23what gets converted at what point is not
  2638. 1:45:26super clear I remember scrolling
  2639. 1:45:30down is it
  2640. 1:45:34here okay I can't find
  2641. 1:45:37it I I thought it was here okay there we
  2642. 1:45:41go so there are a few docks on when
  2643. 1:45:44you're using this AutoCast what gets
  2644. 1:45:46converted to B FL 16 and and when so for
  2645. 1:45:49example only these Matrix multiply like
  2646. 1:45:51operations get converted to float 16 but
  2647. 1:45:54a lot of operations remain in float 32
  2648. 1:45:56so in particular a lot of normalizations
  2649. 1:45:58like layer norms and things like that
  2650. 1:46:00not all of those layers might be
  2651. 1:46:01converted um so only some layers
  2652. 1:46:05selectively would be running B flat 16
  2653. 1:46:07but things like softmax uh layer Norms
  2654. 1:46:10uh log um log soft Max so loss function
  2655. 1:46:14calculations a lot of those things might
  2656. 1:46:15remain in float 32 because they are more
  2657. 1:46:17susceptible to Precision changes major
  2658. 1:46:20multiplies are fairly um
  2659. 1:46:23robust to Precision changes uh so some
  2660. 1:46:26parts of the network are um impacted
  2661. 1:46:29more or less by the Precision
  2662. 1:46:31change um so basically only some parts
  2663. 1:46:34of the of the model are running in
  2664. 1:46:35reduced Precision let's take it for a
  2665. 1:46:38spin and let's actually see what kind of
  2666. 1:46:41improvement we achieve
  2667. 1:46:48here okay so we used to be 333
  2668. 1:46:51milliseconds we're now 300
  2669. 1:46:53and we used to be somewhere around
  2670. 1:46:5450,000 tokens per second we're now at 55
  2671. 1:46:57so we're definitely running faster but
  2672. 1:46:59maybe not a lot faster and that's
  2673. 1:47:02because there are still many many
  2674. 1:47:03bottlenecks in our gbt2 we're just
  2675. 1:47:05getting started but we have dropped down
  2676. 1:47:07the precision as far as we can with my
  2677. 1:47:09current GPU which is a100 we're using
  2678. 1:47:12pytorch AutoCast unfortunately I don't
  2679. 1:47:15actually exactly know what pytorch
  2680. 1:47:17AutoCast do uh does I don't actually
  2681. 1:47:19know exactly what's in B flat 16 what's
  2682. 1:47:22in float 32
  2683. 1:47:23we could go in and we could start to
  2684. 1:47:24scrutinize it um but these are the kinds
  2685. 1:47:27of rules that pytorch has internally and
  2686. 1:47:29unfortunately they don't documented very
  2687. 1:47:31well uh so we're not going to go into
  2688. 1:47:34that into in too much detail but for now
  2689. 1:47:36we are training in B flow 16 we do not
  2690. 1:47:39need a gradient scaler and the reason
  2691. 1:47:40things are running faster is because um
  2692. 1:47:44we are able to run tensor course in B FL
  2693. 1:47:4716 now that means we are in this row but
  2694. 1:47:52uh we are also paying in Precision for
  2695. 1:47:53this uh so um we expect slightly less
  2696. 1:47:57accurate results with respect to the
  2697. 1:47:58original fp32 but empirically in many
  2698. 1:48:01cases this is a worth it uh kind of
  2699. 1:48:04tradeoff because it allows you to run
  2700. 1:48:06faster and you could for example train
  2701. 1:48:07longer and make up for the uh for that
  2702. 1:48:10Precision decrease so um that's b46 for
  2703. 1:48:15now okay so as we can see we are
  2704. 1:48:17currently at about 300 milliseconds uh
  2705. 1:48:19per iteration and we're now going to
  2706. 1:48:21reach for some really heavy weapons in
  2707. 1:48:23the pie torch Arsenal and in particular
  2708. 1:48:25we're going to introduce torch. compile
  2709. 1:48:27so torch. compile is really quite
  2710. 1:48:29incredible infrastructure from the
  2711. 1:48:31pytorch team and it's basically a
  2712. 1:48:32compiler for neural networks like it's
  2713. 1:48:35almost like GCC for CN C++ code this is
  2714. 1:48:38just this GCC of neural nuts so came out
  2715. 1:48:42a while ago and extremely simple to use
  2716. 1:48:46um the way to use torch compile is to do
  2717. 1:48:48this it's a single line of code to
  2718. 1:48:50compile your model and return it now
  2719. 1:48:54this line of code will cost you
  2720. 1:48:55compilation time but as you might guess
  2721. 1:48:57it's going to make the code a lot faster
  2722. 1:48:59so let's actually run that because this
  2723. 1:49:01will take some time to run but currently
  2724. 1:49:03remember we're at 300 milliseconds and
  2725. 1:49:05we'll see what happens now while this is
  2726. 1:49:08running I'd like to explain a little bit
  2727. 1:49:10of what torch. compile does under the
  2728. 1:49:11hood uh so feel free to read this page
  2729. 1:49:15of P torch but basically there's no real
  2730. 1:49:17good reason for you to not use torch
  2731. 1:49:19compile in your pie torch I kind of feel
  2732. 1:49:21like you should be using almost by
  2733. 1:49:23default if you're not uh unless you're
  2734. 1:49:25debugging and you want your code to run
  2735. 1:49:26really fast and there's one line here in
  2736. 1:49:29torch compile that I found that actually
  2737. 1:49:31kind of like gets to why this is faster
  2738. 1:49:33speed up mainly comes from reducing
  2739. 1:49:35python overhead and GPU read wrs so let
  2740. 1:49:38me unpack that a little bit um okay here
  2741. 1:49:41we are okay so we went from 300
  2742. 1:49:43milliseconds we're now running at 129
  2743. 1:49:46milliseconds so this is uh 300 129 about
  2744. 1:49:512.3x Improvement from a single line of
  2745. 1:49:53code in py torch uh so quite incredible
  2746. 1:49:56so what is happening what's happening
  2747. 1:49:57under the hood well when you pass the
  2748. 1:49:59model to torch
  2749. 1:50:01compile what we have here in this NN
  2750. 1:50:04module this is really just the
  2751. 1:50:05algorithmic description of what we'd
  2752. 1:50:08like to happen in our Network and torch
  2753. 1:50:11compile will analyze the entire thing
  2754. 1:50:14and it will look at what operations You'
  2755. 1:50:15like to use and with the benefit of
  2756. 1:50:18knowing exactly what's going to happen
  2757. 1:50:20it doesn't have to run in What's called
  2758. 1:50:22the e mode it doesn't have to just kind
  2759. 1:50:24of like go layer by layer like the
  2760. 1:50:26python interpreter normally would start
  2761. 1:50:29at the
  2762. 1:50:31forward and the python interpreter will
  2763. 1:50:33go okay let's do this operation and then
  2764. 1:50:36let's do that operation and it kind of
  2765. 1:50:38materializes all the operations as it
  2766. 1:50:40goes through uh so these um calculations
  2767. 1:50:43are dispatched and run in this order and
  2768. 1:50:45the python interpreter and this code
  2769. 1:50:47doesn't know what kind of operations are
  2770. 1:50:49going to happen later but torch compile
  2771. 1:50:51sees your entire code at the same time
  2772. 1:50:53and it's able to know what operations
  2773. 1:50:56you intend to run and it will kind of
  2774. 1:50:58optimize that process the first thing it
  2775. 1:51:00will do is will it will take out the
  2776. 1:51:01python interpreter from the forward pass
  2777. 1:51:03entirely and it will kind of compile
  2778. 1:51:05this entire neural net as a single
  2779. 1:51:07object with no python interpreter
  2780. 1:51:09involved so it knows exactly what's
  2781. 1:51:11going to run and we'll just run that and
  2782. 1:51:12it's all going to be running in
  2783. 1:51:14efficient
  2784. 1:51:15code uh the second thing that happens is
  2785. 1:51:18uh this read write that they mentioned
  2786. 1:51:21very briefly so a good example of that I
  2787. 1:51:23think is the G nonlinearity that we've
  2788. 1:51:25been looking at so here we use the n and
  2789. 1:51:28G now this here is me uh basically just
  2790. 1:51:32breaking up the inang Galu uh which you
  2791. 1:51:35remember has this formula so this here
  2792. 1:51:37is the equivalent implementation to
  2793. 1:51:39what's happening inside g algorithmic l
  2794. 1:51:41it's
  2795. 1:51:42identical Now by default if uh we just
  2796. 1:51:46we using this instead of ending. G here
  2797. 1:51:48what would happen without torch compile
  2798. 1:51:51well the python interpreter would make
  2799. 1:51:52its way here and then it would be okay
  2800. 1:51:54well there's an input well let me first
  2801. 1:51:58let me raise this input to the third
  2802. 1:51:59power and it's going to dispatch a
  2803. 1:52:01kernel that takes your input and raises
  2804. 1:52:03it to the third power and that kernel
  2805. 1:52:05will run and when this kernel runs what
  2806. 1:52:08ends up happening is this input is
  2807. 1:52:11stored in the memory of the GPU so
  2808. 1:52:13here's a helpful example of the layout
  2809. 1:52:16of what's happening right you have your
  2810. 1:52:18CPU this is in every single computer
  2811. 1:52:21there's a few cores in there and you
  2812. 1:52:23have your uh Ram uh your memory and the
  2813. 1:52:26CPU can talk to the memory and this is
  2814. 1:52:28all well known but now we've added the
  2815. 1:52:30GPU and the GPU is a slightly different
  2816. 1:52:32architecture of course they can
  2817. 1:52:33communicate and it's different in that
  2818. 1:52:35it's got a lot more course than a CPU
  2819. 1:52:38all of those cores are individually a
  2820. 1:52:40lot simpler too but it also has memory
  2821. 1:52:43right this high bandwidth memory I'm
  2822. 1:52:47sorry if I'm botching it hbm I don't
  2823. 1:52:49even know what that stands for I'm just
  2824. 1:52:51realizing that
  2825. 1:52:53but uh this is the memory and it's very
  2826. 1:52:54equivalent to uh RAM basically in the
  2827. 1:52:58computer and what's happening is that
  2828. 1:53:00input is living in the memory and when
  2829. 1:53:02you do input
  2830. 1:53:05cubed this has to travel to the GPU to
  2831. 1:53:09the course and to all the caches and
  2832. 1:53:12registers on the actual chip of this
  2833. 1:53:15GPU and it has to calculate the all the
  2834. 1:53:17elements to the third and then it saves
  2835. 1:53:19the result back to the memory and it's
  2836. 1:53:22this uh travel time that actually causes
  2837. 1:53:25a lot of issues so here remember this
  2838. 1:53:28memory bandwidth we can communicate
  2839. 1:53:30about 2 terabytes per second which is a
  2840. 1:53:31lot but also we have to Traverse this
  2841. 1:53:35link and it's very slow so here on the
  2842. 1:53:37GPU we're on chip and everything is
  2843. 1:53:39super fast within the chip but going to
  2844. 1:53:41the memory is extremely expensive takes
  2845. 1:53:43extremely long amount of time and so we
  2846. 1:53:46load the input do the calculations and
  2847. 1:53:48load back the output and this round trip
  2848. 1:53:51takes a lot of time
  2849. 1:53:53and now right after we do that we
  2850. 1:53:54multiply by this constant so what
  2851. 1:53:57happens then is we dispatch another
  2852. 1:53:59kernel and then the result travels back
  2853. 1:54:02all the elements get multiplied by a
  2854. 1:54:03constant and then the results travel
  2855. 1:54:06back to the memory and then we take the
  2856. 1:54:09result and we add back input and so this
  2857. 1:54:12entire thing again travels to the GPU
  2858. 1:54:15adds the inputs and gets written back so
  2859. 1:54:18we're making all these round trips from
  2860. 1:54:20the memory to actually where the comput
  2861. 1:54:22happens because all the tensor cores and
  2862. 1:54:24alus and everything like that is all
  2863. 1:54:26stored on the chip in the GPU so we're
  2864. 1:54:28doing a ton of round trips and pytorch
  2865. 1:54:31uh without using torch compile doesn't
  2866. 1:54:33know to optimize this because it doesn't
  2867. 1:54:36know what kind of operations you're
  2868. 1:54:37running later you're just telling it
  2869. 1:54:39raise the power to the third then do
  2870. 1:54:41this then do that and it will just do
  2871. 1:54:43that in that sequence but torch compile
  2872. 1:54:45sees your entire code it will come here
  2873. 1:54:47and it will realize wait all of these
  2874. 1:54:49are elementwise operations and actually
  2875. 1:54:52what I'm going to do is I'm going to do
  2876. 1:54:53a single trip of input to the GPU then
  2877. 1:54:56for every single element I'm going to do
  2878. 1:54:58all of these operations while that
  2879. 1:55:00memory is on the GPU or chunks of it
  2880. 1:55:04rather and then I'm going to write back
  2881. 1:55:06a single time so we're not going to have
  2882. 1:55:07these round trips and that's one example
  2883. 1:55:09of what's called kernel fusion and is a
  2884. 1:55:11major way in which everything is sped up
  2885. 1:55:14so basically if you have your benefit of
  2886. 1:55:15onet and you know exactly what you're
  2887. 1:55:17going to compute you can optimize your
  2888. 1:55:19round trips to the memory and you're not
  2889. 1:55:21going to pay the the memory bandwidth
  2890. 1:55:23cost and that's fundamentally what makes
  2891. 1:55:25some of these operations a lot faster
  2892. 1:55:27and what they mean by read writes
  2893. 1:55:30here so let me erase this because we are
  2894. 1:55:32not using it and yeah we should be using
  2895. 1:55:36torch compile and our code is now
  2896. 1:55:39significantly faster and we're doing
  2897. 1:55:40about
  2898. 1:55:42125,000 tokens per second but we still
  2899. 1:55:45have a long way to go before we move on
  2900. 1:55:47I wanted to supplement the discussion a
  2901. 1:55:49little bit with a few more figures uh
  2902. 1:55:51because this is a complic topic but it's
  2903. 1:55:53worth understanding on a high level uh
  2904. 1:55:55what's happening here and I could
  2905. 1:55:56probably spend an entire video of like
  2906. 1:55:58two hours on this but just the preview
  2907. 1:56:00of that basically so this chip here that
  2908. 1:56:03is uh the GPU this chip is where all the
  2909. 1:56:06calculations happen mostly but this chip
  2910. 1:56:09also does have some memory in it but
  2911. 1:56:12most of the memory by far is here in the
  2912. 1:56:15high bandwidth memory hbm and is
  2913. 1:56:18connected they're connected um but these
  2914. 1:56:20are two separate chips basically
  2915. 1:56:23now here this is a zoom in of kind of
  2916. 1:56:26this cartoon diagram of a GPU and what
  2917. 1:56:30we're seeing here is number one you see
  2918. 1:56:31this hbm I I realize it's probably very
  2919. 1:56:34small for you but on the sides here it
  2920. 1:56:35says hbm and so that that's the links to
  2921. 1:56:38the hbm now the hbm is again off chip on
  2922. 1:56:42the chip there are a large number of
  2923. 1:56:45these streaming
  2924. 1:56:46multiprocessors uh every one of these is
  2925. 1:56:48an SM there's 120 of them in total and
  2926. 1:56:51this is where the a lot of the
  2927. 1:56:52calculations happen and this is a zoom
  2928. 1:56:54in of a single individual as it has
  2929. 1:56:57these four quadrants and see for example
  2930. 1:56:59tensor core this is where a lot of the
  2931. 1:57:00Matrix multiply stuff happens but
  2932. 1:57:02there's all these other units to do all
  2933. 1:57:04different kinds of calculations for fp64
  2934. 1:57:07fp32 and for integers and so on now so
  2935. 1:57:11we have all this uh logic here to do the
  2936. 1:57:13calculations but in addition to that on
  2937. 1:57:15the chip there is memory sprinkled
  2938. 1:57:17throughout the chip so L2 cache is some
  2939. 1:57:21amount of memory that lives on the chip
  2940. 1:57:23and then on the SMS themselves there's
  2941. 1:57:25L1 cache I realized it's probably very
  2942. 1:57:28small for you but this blue bar is L1
  2943. 1:57:31and there's also registers um and so
  2944. 1:57:34there is memory stored here but the way
  2945. 1:57:36this memory is stored is very different
  2946. 1:57:38from the way memory is stored in hbm uh
  2947. 1:57:41this is a very different implementation
  2948. 1:57:44uh using um just in terms of like what
  2949. 1:57:47the Silicon looks like it's a very
  2950. 1:57:48different
  2951. 1:57:49implementation um so here you would
  2952. 1:57:52using transistors and capacitors and
  2953. 1:57:54here it's a very different
  2954. 1:57:55implementation uh with SRAM and what
  2955. 1:57:57that looks like but long story short is
  2956. 1:58:01um there is um memory inside the chip
  2957. 1:58:05but it's not a lot of memory that's the
  2958. 1:58:07critical point so this is some C this is
  2959. 1:58:09a example diagram of a slightly
  2960. 1:58:11different GPU just like here where it
  2961. 1:58:14shows that for example typical numbers
  2962. 1:58:16for CPU Dam memory which is this thing
  2963. 1:58:19here you might have one tab of this
  2964. 1:58:22right but it would be extremely
  2965. 1:58:23expensive to access especially for a GPU
  2966. 1:58:25you have to go through the CPU here now
  2967. 1:58:28next we have the hbm so we have tens of
  2968. 1:58:30gigabytes of hbm memory on a typical GPU
  2969. 1:58:33here but it's as I mentioned very
  2970. 1:58:35expensive to access and then on the chip
  2971. 1:58:38itself everything is extremely fast
  2972. 1:58:40within the chip but we only have couple
  2973. 1:58:4210 megabytes of memory collectively
  2974. 1:58:45throughout the Chip And so there's just
  2975. 1:58:48not enough space because the memory is
  2976. 1:58:50very expensive on the chip and so
  2977. 1:58:52there's not a lot of it but it is
  2978. 1:58:53lightning fast to access in relative
  2979. 1:58:55terms and so basically whenever we have
  2980. 1:58:58these kernels um the more accurate
  2981. 1:59:01picture of what's Happening Here is that
  2982. 1:59:03we take these inputs which live by
  2983. 1:59:05default on the global memory and now we
  2984. 1:59:08need to perform some calculation so we
  2985. 1:59:10start streaming the data from the um
  2986. 1:59:12Global memory to the uh chip we perform
  2987. 1:59:16the calculations on the chip and then
  2988. 1:59:18stream it back and store it back to the
  2989. 1:59:19global memory right and so if we are if
  2990. 1:59:23we don't have torch compile we are
  2991. 1:59:24streaming the data through the chip
  2992. 1:59:26doing the calculations and saving to the
  2993. 1:59:27memory and we're doing those round trips
  2994. 1:59:29many many
  2995. 1:59:30times but uh if it's torch compiled then
  2996. 1:59:33we start streaming the memory as before
  2997. 1:59:35but then while we're on the chip we're
  2998. 1:59:37we're we have a chunk of the uh data
  2999. 1:59:40that we're trying to process so that
  3000. 1:59:42chunk now lives on the chip while it's
  3001. 1:59:44on the chip it's extremely fast to
  3002. 1:59:46operate on so if we have kernel Fusion
  3003. 1:59:48we can do all the operations right there
  3004. 1:59:49in an element-wise fashion and those are
  3005. 1:59:52very cheap and then we do a single round
  3006. 1:59:54trip back to the global memory so
  3007. 1:59:58operator Fusion basically allows you to
  3008. 2:00:00keep your chunk of data on the Chip And
  3009. 2:00:02do lots of calculations on it before you
  3010. 2:00:04write it back and that gives huge
  3011. 2:00:06savings and that's why torch compile
  3012. 2:00:09ends up being a lot faster or that's one
  3013. 2:00:11of the major
  3014. 2:00:12reasons uh so again just a very brief
  3015. 2:00:14intro to the memory hierarchy and
  3016. 2:00:16roughly what torch compile does for you
  3017. 2:00:19now torch compile is amazing but there
  3018. 2:00:21are operations torch compile will not
  3019. 2:00:23find and an amazing example of that is
  3020. 2:00:26Flash attention to which we turn next so
  3021. 2:00:29flash attention comes from this paper
  3022. 2:00:30from uh Stanford in
  3023. 2:00:332022 and it's this incredible algorithm
  3024. 2:00:36for performing attention so um and
  3025. 2:00:39running it a lot faster so flash
  3026. 2:00:41attention will come here and we will
  3027. 2:00:44take out these four
  3028. 2:00:46lines and Flash attention implements
  3029. 2:00:48these four lines really really quickly
  3030. 2:00:51and how does it do that well flash
  3031. 2:00:53attention is a kernel Fusion operation
  3032. 2:00:57so you see here we have um in this
  3033. 2:00:59diagram they're showing P torch and you
  3034. 2:01:02have these four operations uh they're
  3035. 2:01:04including Dropout but we are not using
  3036. 2:01:06Dropout here so we just have these four
  3037. 2:01:08lines of code here and instead of those
  3038. 2:01:11we are fusing them into a single fused
  3039. 2:01:13kernel of flash attention so it's an
  3040. 2:01:16it's a it's a kernel Fusion algorithm
  3041. 2:01:19but it's a kernel Fusion that torch
  3042. 2:01:20compile cannot find
  3043. 2:01:22and the reason that it cannot find it is
  3044. 2:01:24that it um requires an algorithmic
  3045. 2:01:26rewrite of how attention is actually
  3046. 2:01:28implemented here in this case and what's
  3047. 2:01:31remarkable about it is that uh flash
  3048. 2:01:33attention actually if you just count the
  3049. 2:01:35number of flops flash attention does
  3050. 2:01:37more flops than this attention here but
  3051. 2:01:41flash attention is actually
  3052. 2:01:42significantly faster in fact they site
  3053. 2:01:457. six times faster potentially and
  3054. 2:01:48that's because it is very mindful of the
  3055. 2:01:51memory hierarchy as I described it just
  3056. 2:01:53now and so it's very mindful about
  3057. 2:01:55what's in high bandwidth memory what's
  3058. 2:01:57in the shared memory and it is very
  3059. 2:02:00careful with how it orchestrates the
  3060. 2:02:02computation such that we have fewer
  3061. 2:02:04reads and writes to the high bandwidth
  3062. 2:02:06memory and so even though we're doing
  3063. 2:02:08more flops the expensive part is they
  3064. 2:02:10load and store into hbm and that's what
  3065. 2:02:12they avoid and so in particular they do
  3066. 2:02:15not ever materialize this end byend
  3067. 2:02:17attention Matrix this ATT here a flash
  3068. 2:02:21attention is designed such that this
  3069. 2:02:23Matrix never gets materialized at any
  3070. 2:02:25point and it never gets read or written
  3071. 2:02:28to the hbm and this is a very large
  3072. 2:02:30Matrix right so um because this is where
  3073. 2:02:32all the queries and keys interact and
  3074. 2:02:34we're sort of getting
  3075. 2:02:36um for each head for each batch element
  3076. 2:02:40we're getting a t BYT Matrix of
  3077. 2:02:42attention which is a Million numbers
  3078. 2:02:45even for a single head at a single batch
  3079. 2:02:47index at like so so basically this is a
  3080. 2:02:50ton of memory and and this is never
  3081. 2:02:52materialized and the way that this is
  3082. 2:02:54achieved is that basically the
  3083. 2:02:57fundamental algorithmic rewrite here
  3084. 2:02:58relies on this online softmax trick
  3085. 2:03:02which was proposed previously and I'll
  3086. 2:03:03show you the paper in a bit and the
  3087. 2:03:05online softmax trick coming from a
  3088. 2:03:07previous paper um shows how you can
  3089. 2:03:10incrementally evaluate a soft Max
  3090. 2:03:14without having to sort of realize all of
  3091. 2:03:16the inputs to the softmax to do the
  3092. 2:03:18normalization and you do that by having
  3093. 2:03:19these intermediate variables M and L and
  3094. 2:03:22there's an update to them that allows
  3095. 2:03:24you to evaluate the softmax in an online
  3096. 2:03:26manner um now flash attention actually
  3097. 2:03:30so recently flash attention 2 came out
  3098. 2:03:32as well so I have that paper up here as
  3099. 2:03:34well uh that has additional gains to how
  3100. 2:03:36it calculates flash attention and the
  3101. 2:03:38original paper that this is based on
  3102. 2:03:40basically is this online normalizer
  3103. 2:03:42calculation for softmax and remarkably
  3104. 2:03:45it came out of Nvidia and it came out of
  3105. 2:03:46it like really early 2018 so this is 4
  3106. 2:03:50years before flash attention
  3107. 2:03:52and this paper says that we propose a
  3108. 2:03:55way to compute the classical softmax
  3109. 2:03:57with fewer memory accesses and
  3110. 2:03:59hypothesize that this reduction in
  3111. 2:04:00memory accesses should improve softmax
  3112. 2:04:02performance on actual hardware and so
  3113. 2:04:05they are extremely correct in this
  3114. 2:04:08hypothesis but it's really fascinating
  3115. 2:04:10to me that they're from Nvidia and that
  3116. 2:04:12they had this realization but they
  3117. 2:04:13didn't actually take it to the actual
  3118. 2:04:15flash attention that had to come four
  3119. 2:04:18years later from Stanford so I don't
  3120. 2:04:20fully understand the historical how this
  3121. 2:04:22happened historically um but they do
  3122. 2:04:24basically propose this online update to
  3123. 2:04:26the softmax uh right here and this is
  3124. 2:04:29fundamentally what they reuse here to
  3125. 2:04:31calculate the softmax in a streaming
  3126. 2:04:33Manner and then they realize they can
  3127. 2:04:35actually fuse all the other operations
  3128. 2:04:37with the online sofx calculation into a
  3129. 2:04:40single fused kernel flash attention and
  3130. 2:04:42that's what we are about to use so great
  3131. 2:04:45example I think of being aware of um
  3132. 2:04:47memory hierarchy the fact that flops
  3133. 2:04:49don't matter uh the entire memory access
  3134. 2:04:52pattern matters and that torch compile
  3135. 2:04:54is amazing but there are many
  3136. 2:04:55optimizations that are still available
  3137. 2:04:57to us that potentially torch compile
  3138. 2:04:59cannot find maybe maybe one day it could
  3139. 2:05:01but right now it seems like a lot to ask
  3140. 2:05:04so here's what we're going to do we're
  3141. 2:05:05going to use Flash attention and the way
  3142. 2:05:09to do that basically in pytorch is we
  3143. 2:05:11are going to comment out these four
  3144. 2:05:14lines and we're going to replace them
  3145. 2:05:15with a single line and here we are
  3146. 2:05:18calling this compound operation in
  3147. 2:05:20pytorch called scale that product
  3148. 2:05:22attention and uh pytorch will call flash
  3149. 2:05:27attention when you use it in this way
  3150. 2:05:31I'm not actually 100% sure why torch
  3151. 2:05:32compile doesn't realize that these four
  3152. 2:05:34lines should just call flash attention
  3153. 2:05:36in this exact way we have to do it again
  3154. 2:05:38for it which in my opinion is a little
  3155. 2:05:40bit odd but um here we are so you have
  3156. 2:05:46to use this compound up and uh let's
  3157. 2:05:49wait for a few moments before torch comp
  3158. 2:05:51compile gets around to it and then let's
  3159. 2:05:53remember that we achieved 6.05 661 I
  3160. 2:05:58have it here that's the loss we were
  3161. 2:06:00expecting to see and we took 130
  3162. 2:06:03milliseconds uh before this change so
  3163. 2:06:05we're expecting to see the exact same
  3164. 2:06:07result by iteration 49 but we expect to
  3165. 2:06:10see faster runtime because Flash
  3166. 2:06:13attention is just a an algorithmic
  3167. 2:06:14rewrite and it's a faster kernel but it
  3168. 2:06:16doesn't actually change any of the
  3169. 2:06:17computation and we should have the exact
  3170. 2:06:19same optimization so okay so we're a lot
  3171. 2:06:21faster we're at about 95 milliseconds
  3172. 2:06:24and we achiev
  3173. 2:06:286.58 okay so they're basically identical
  3174. 2:06:31up to a floating Point fudge Factor so
  3175. 2:06:34it's the identical computation but it's
  3176. 2:06:36significantly faster going from 130 to
  3177. 2:06:39roughly 90
  3178. 2:06:4096 and so this is um 96 divide
  3179. 2:06:44130ish so this is maybe 27 is%
  3180. 2:06:48Improvement um so uh really interesting
  3181. 2:06:52and that is Flash retention okay we are
  3182. 2:06:54now getting to one of my favorite
  3183. 2:06:57optimizations and it is simultaneously
  3184. 2:06:59the dumbest and the most brilliant
  3185. 2:07:02optimization and it's always a little
  3186. 2:07:03bit surprising to me um anyway so
  3187. 2:07:06basically I mentioned a few minutes ago
  3188. 2:07:08that there are some numbers that are
  3189. 2:07:10nice and some numbers that are ugly so
  3190. 2:07:1364 is a beautiful nice number 128 is
  3191. 2:07:17even nicer 256 is beautiful what makes
  3192. 2:07:20these numbers beautiful is that there
  3193. 2:07:21are many powers of two inside them you
  3194. 2:07:23can divide by two many times and uh
  3195. 2:07:26examples of ugly numbers are like 13 and
  3196. 2:07:2817 and something like that prime numbers
  3197. 2:07:30numbers that are not even and so on and
  3198. 2:07:32so pretty much you always want to use
  3199. 2:07:34nice numbers in all of your code that
  3200. 2:07:36deals with neural networks or Cuda
  3201. 2:07:38because everything in Cuda Works in sort
  3202. 2:07:40of like powers of two and lots of
  3203. 2:07:42kernels are written in terms of powers
  3204. 2:07:45of Two And there are lots of blocks of
  3205. 2:07:47sizes 16 and uh 64 and so on so
  3206. 2:07:50everything is written in those terms and
  3207. 2:07:52you always have special case handling
  3208. 2:07:54for all kinds of uh logic that U when
  3209. 2:07:57your inputs are not made of nice numbers
  3210. 2:08:00so let's see what that looks like
  3211. 2:08:01basically scan your code and look for
  3212. 2:08:03ugly numbers is roughly theistic so
  3213. 2:08:06three times is kind of ugly um I'm not
  3214. 2:08:10100% sure maybe this can be improved but
  3215. 2:08:12this is uh this is ugly and not
  3216. 2:08:15ideal um four times is nice so that's uh
  3217. 2:08:20that's nice
  3218. 2:08:221024 is very nice that's a power of two
  3219. 2:08:2512 is a little bit suspicious um not too
  3220. 2:08:28many powers of two 768 is great 50, 257
  3221. 2:08:32is a really really ugly number um it's
  3222. 2:08:36first of all it's odd so uh and there's
  3223. 2:08:38no not too many powers of two in there
  3224. 2:08:40so this is a very ugly number and it's
  3225. 2:08:43highly suspicious and then when we
  3226. 2:08:45scroll down all these numbers are nice
  3227. 2:08:48and then here we have mostly nice
  3228. 2:08:50numbers except for 25 so in this
  3229. 2:08:53configuration of gpt2 XL a number of
  3230. 2:08:55heads is 25 uh that's a really ugly
  3231. 2:08:57number that's an odd number and um
  3232. 2:09:00actually this did cause a lot of
  3233. 2:09:01headaches for us recently when we're
  3234. 2:09:02trying to optimize some kernels uh to
  3235. 2:09:04run this fast um and required a bunch of
  3236. 2:09:07special case handling so basically these
  3237. 2:09:10numbers are we have some ugly numbers
  3238. 2:09:12and some of them are easier to fix than
  3239. 2:09:13others and in particular the voap size
  3240. 2:09:15being 50257 that's a very ugly number
  3241. 2:09:18very suspicious and we want to fix it
  3242. 2:09:20now when you when you fix these things
  3243. 2:09:23uh one of the easy ways to do that is
  3244. 2:09:24you basically um increase the number
  3245. 2:09:27until it's the nearest power of two that
  3246. 2:09:29you like so here's a much nicer number
  3247. 2:09:32it's
  3248. 2:09:3350304 and why is that because 50304 can
  3249. 2:09:37be divided by 8 or by 16 or by 32
  3250. 2:09:4364 it can even be divided by 128 I think
  3251. 2:09:46yeah so it's a very nice number um so
  3252. 2:09:49what we're going to do here is the GPT
  3253. 2:09:51config and you see that we initialized B
  3254. 2:09:53cap size to
  3255. 2:09:5450257 Let's override just
  3256. 2:09:58that um element to be
  3257. 2:10:0150304 okay so everything else stays the
  3258. 2:10:05same we're just increasing our
  3259. 2:10:06vocabulary size so we're adding it's
  3260. 2:10:09almost like we're adding fake tokens uh
  3261. 2:10:12so that book up size has powers of two
  3262. 2:10:14inside it now actually what I'm doing
  3263. 2:10:16here by the way is I'm increasing the
  3264. 2:10:18amount of computation that our network
  3265. 2:10:19will be doing if you just count the the
  3266. 2:10:21flops on like do the math of how many
  3267. 2:10:23flops we're doing we're going to be
  3268. 2:10:25doing more flops and we still have to
  3269. 2:10:27think through whether this doesn't break
  3270. 2:10:30anything but if I just run this uh let's
  3271. 2:10:33see what we get uh currently this ran in
  3272. 2:10:35maybe
  3273. 2:10:3896.5 milliseconds per step I'm just kind
  3274. 2:10:41of like eyeballing it and let's see what
  3275. 2:10:43kind of a result we're going to
  3276. 2:10:46get uh while this is compiling let's
  3277. 2:10:49think through whether our code actually
  3278. 2:10:51works okay when we increase the vocap
  3279. 2:10:53size like this let's look at where vocap
  3280. 2:10:55size is actually
  3281. 2:10:57used so we swing up to the inet and we
  3282. 2:11:00see that it's used inside the embedding
  3283. 2:11:01table of course so all the way at the
  3284. 2:11:03bottom of the Transformer and it's used
  3285. 2:11:05at the classifier layer all the way at
  3286. 2:11:06the top of the Transformer so in two
  3287. 2:11:08places and let's take a look and we're
  3288. 2:11:11running at 93 so 93 milliseconds instead
  3289. 2:11:14of
  3290. 2:11:1596.5 so we are seeing a roughly yeah 4%
  3291. 2:11:19Improvement here uh by doing more
  3292. 2:11:22calculations and the reason for this is
  3293. 2:11:25we fixed we've made an ugly number into
  3294. 2:11:28a nice number let's I'm going to come
  3295. 2:11:30into the explanation for that a little
  3296. 2:11:32bit again but for now let's just
  3297. 2:11:34convince ourselves that we're not
  3298. 2:11:35breaking anything when we do this so
  3299. 2:11:36first of all we've made the the wte the
  3300. 2:11:39embedding table for the tokens we've
  3301. 2:11:41made it larger it's almost like we
  3302. 2:11:43introduced more tokens at the bottom and
  3303. 2:11:46these tokens are never used because the
  3304. 2:11:48gbt tokenizer only has tokens up to
  3305. 2:11:50$50,000
  3306. 2:11:51256 and so we'll never index into the
  3307. 2:11:55rows that we've added so we're wasting a
  3308. 2:11:57little bit of space here by creating
  3309. 2:11:59memory that's never going to be accessed
  3310. 2:12:01never going to be used Etc now that's
  3311. 2:12:03not fully correct because this wte
  3312. 2:12:06weight ends up being shared and ends up
  3313. 2:12:08being used in the classifier here at the
  3314. 2:12:10end so what is that doing to the
  3315. 2:12:13classifier right here well what what
  3316. 2:12:15that's doing is we're predicting
  3317. 2:12:16additional Dimensions at the classifier
  3318. 2:12:18now and we're predicting probabilities
  3319. 2:12:20for tokens that will of course never be
  3320. 2:12:21present in the training set um and so
  3321. 2:12:25therefore the network has to learn that
  3322. 2:12:27these probabilities uh have to be driven
  3323. 2:12:29to zero and so the logits that the
  3324. 2:12:31network produces have to drive those
  3325. 2:12:33dimensions of the output to negative
  3326. 2:12:35Infinity but it that's no different from
  3327. 2:12:38all the other tokens that are already in
  3328. 2:12:39our data set um or rather that are not
  3329. 2:12:42in our data set so Shakespeare only
  3330. 2:12:45probably uses let's say a th000 tokens
  3331. 2:12:46out of 50,000 to 57 tokens so most of
  3332. 2:12:49the tokens are already being driven to
  3333. 2:12:51zero probability by the optimization we'
  3334. 2:12:53just introduced a few more tokens now
  3335. 2:12:55that in a similar manner will never be
  3336. 2:12:57used and have to be driven to zero in
  3337. 2:12:59probability um so functionally though
  3338. 2:13:02nothing breaks we're using a bit more
  3339. 2:13:05extra um memory but otherwise this is a
  3340. 2:13:08harmless operation as far as I can tell
  3341. 2:13:11but and we're adding calculation but
  3342. 2:13:12it's running faster and it's running
  3343. 2:13:14faster because as I mentioned in Cuda so
  3344. 2:13:17many kernels use uh block tiles and
  3345. 2:13:21these block towels are usually nice
  3346. 2:13:22numbers uh so powers of two so
  3347. 2:13:25calculations are done in like chunks of
  3348. 2:13:2664 or chunks of 32 and when your um when
  3349. 2:13:31your desired calculation doesn't neatly
  3350. 2:13:32fit into those block tiles um there are
  3351. 2:13:36all kinds of boundary kernels that can
  3352. 2:13:38kick in to like do the last part so
  3353. 2:13:42basically in a lot of kernels they will
  3354. 2:13:44chunk at up your input and they will do
  3355. 2:13:46the nice part first and then they have a
  3356. 2:13:47whole second second phase where they
  3357. 2:13:50come back to any that like uh remains uh
  3358. 2:13:54and then they process the remaining part
  3359. 2:13:56and the kernels for that could be very
  3360. 2:13:57inefficient and so you're basically um
  3361. 2:14:00spinning up all this extra compute and
  3362. 2:14:02is extremely inefficient so you might as
  3363. 2:14:04well pad your inputs and um make it fit
  3364. 2:14:07nicely and usually that empiric lens up
  3365. 2:14:10actually running faster um so this is
  3366. 2:14:13another example of a 4% Improvement that
  3367. 2:14:16we've added and this is something that
  3368. 2:14:18also torch compile did not find for us
  3369. 2:14:21you would hope that torch compile at
  3370. 2:14:22some point could figure an optimization
  3371. 2:14:24like this out uh but for now uh this is
  3372. 2:14:27it and I also have to point out that
  3373. 2:14:28we're using pytorch nightly so that's
  3374. 2:14:30why we're only seeing 4% if you're using
  3375. 2:14:33pytorch 2.3.1 or earlier you would
  3376. 2:14:36actually see something like 30%
  3377. 2:14:37Improvement just from this change from
  3378. 2:14:39changing it to from 50,000 to 57 to
  3379. 2:14:4350304 so again one of my favorite
  3380. 2:14:47examples also of having to understand
  3381. 2:14:49the under the hood and how it all works
  3382. 2:14:51and to know what kinds of things to
  3383. 2:14:52Tinker with to push the performance of
  3384. 2:14:54your code okay so at this point we have
  3385. 2:14:56improved the performance by about 11x
  3386. 2:14:58right because we started at about 1,000
  3387. 2:15:00milliseconds per step and we're now down
  3388. 2:15:02to like 93 milliseconds so that's uh
  3389. 2:15:05quite good and we're uh doing a much
  3390. 2:15:08better job of utilizing our GPU
  3391. 2:15:09resources so I'm going to now turn to
  3392. 2:15:12more algorithmic changes uh and
  3393. 2:15:14improvements to the actual optimization
  3394. 2:15:16itself and what we would like to do is
  3395. 2:15:18we would like to follow the hyper
  3396. 2:15:19parameters that are mentioned in the GP
  3397. 2:15:20G2 or gpt2 gpt3 paper now sadly gpt2 is
  3398. 2:15:26uh doesn't actually say too much it's
  3399. 2:15:28very nice of them that they released the
  3400. 2:15:30model weights and the code but the paper
  3401. 2:15:32itself is extremely vague as to the
  3402. 2:15:33optimization details uh the code itself
  3403. 2:15:36that they released as well the code
  3404. 2:15:38we've been looking at this is just the
  3405. 2:15:40inference code so there's no training
  3406. 2:15:41code here and very few hyp parameters so
  3407. 2:15:44this doesn't also tell us too much so
  3408. 2:15:46for that we have to turn to the gpt3
  3409. 2:15:48paper and um in the depending of the
  3410. 2:15:51gpt3 paper um they have a lot more hyper
  3411. 2:15:55parameters here for us to use and the
  3412. 2:15:57gpt3 paper in general is a lot more
  3413. 2:15:59detailed as to uh all of the you know
  3414. 2:16:02small details that go into the model
  3415. 2:16:04training but gpt3 U models were never
  3416. 2:16:07released so gbt2 we have the weights but
  3417. 2:16:10no details and gpt3 we have lots of
  3418. 2:16:11details but no weights so um but roughly
  3419. 2:16:15speaking gpt2 and gpt3 architectures are
  3420. 2:16:17very very similar and um basically there
  3421. 2:16:21are very few changes the context length
  3422. 2:16:23was expanded from 1024 to 2048 and
  3423. 2:16:25that's kind of like the major change uh
  3424. 2:16:28and some of the hyper parameters around
  3425. 2:16:29the Transformer have changed but
  3426. 2:16:31otherwise they're pretty much the same
  3427. 2:16:32model it's just that gpt3 was trained
  3428. 2:16:34for a lot longer on a bigger data set
  3429. 2:16:36and uh has a lot more thorough
  3430. 2:16:38evaluations uh and the gpt3 model is 175
  3431. 2:16:42billion instead of 1.6 billion um in the
  3432. 2:16:46gpt2 so long story short we're going to
  3433. 2:16:49go to gp3 paper to follow along some the
  3434. 2:16:51hyper parameters so to train all the
  3435. 2:16:54versions of gpt3 we use atom with beta 1
  3436. 2:16:56beta 2 of9 and .95 so let's swing over
  3437. 2:17:00here and make sure that the betas
  3438. 2:17:02parameter which you can see here
  3439. 2:17:04defaults to 0.9 and
  3440. 2:17:06999 is actually set to 0.9 and
  3441. 2:17:11.95 and then the Epsilon parameter uh
  3442. 2:17:14you can see is the default is 1 in8 and
  3443. 2:17:17this is also one in8 let's just uh put
  3444. 2:17:19it in so that works
  3445. 2:17:22expit uh now next up they say we clip
  3446. 2:17:25the gra Global Norm of the gradient at
  3447. 2:17:271.0 so what this is referring to is that
  3448. 2:17:30once we calculate the gradients right
  3449. 2:17:32after l. backward um we basically have
  3450. 2:17:35the gradients at all the parameter
  3451. 2:17:37tensors and what people like to do is
  3452. 2:17:40basically uh clip them to have some kind
  3453. 2:17:42of a maximum Norm so in pytor this is
  3454. 2:17:45fairly easy to do uh it's one line of
  3455. 2:17:48code here that we have to insert right
  3456. 2:17:50after we calcul Cal the gradients and
  3457. 2:17:52what this utility function is doing is
  3458. 2:17:55um it's calculating the global Norm of
  3459. 2:17:58the parameters so every single par um
  3460. 2:18:01gradient on all the parameters you
  3461. 2:18:03square it and you add it all up and you
  3462. 2:18:05take a big square root of that and
  3463. 2:18:07that's the norm of the parameter V
  3464. 2:18:10Vector basically it's the it's the
  3465. 2:18:12length of it if you if you'd like to
  3466. 2:18:14look at it that way and we are basically
  3467. 2:18:16making sure that its length is no more
  3468. 2:18:18than 1.0 and we're going to clip it
  3469. 2:18:21and the reason that people like to use
  3470. 2:18:23this is that uh sometimes you can get
  3471. 2:18:25unlucky during your optimization maybe
  3472. 2:18:27it's a bad data batch or something like
  3473. 2:18:28that and if you get very unlucky in the
  3474. 2:18:31batch you might get really high loss and
  3475. 2:18:33really high loss could lead to a really
  3476. 2:18:35high gradient and this could basically
  3477. 2:18:38uh shock your model and shock the
  3478. 2:18:40optimization so people like to use a
  3479. 2:18:42gradient Norm clipping uh to prevent the
  3480. 2:18:45model from um basically getting too big
  3481. 2:18:49of shocks in terms of the gradient
  3482. 2:18:50magnet ude and uh the upper bound it in
  3483. 2:18:53this way it's a bit of a hacky solution
  3484. 2:18:55it's about like a patch on top of like
  3485. 2:18:57deeper issues uh but uh people still do
  3486. 2:19:00it fairly frequently now the clip grad
  3487. 2:19:03Norm Returns the norm of the gradient
  3488. 2:19:05which I like to always visualize uh
  3489. 2:19:08because um it is useful information and
  3490. 2:19:11sometimes you can look at the norm of
  3491. 2:19:13the gradient and if it's well behaved
  3492. 2:19:15things are good if it's climbing things
  3493. 2:19:17are bad and they're destabilizing during
  3494. 2:19:19training sometimes you could get a spike
  3495. 2:19:21in the norm and that means there's some
  3496. 2:19:22kind of an issue or an instability so
  3497. 2:19:25the norm here will be a
  3498. 2:19:28norm uh and let's do a uh 4f or
  3499. 2:19:33something like
  3500. 2:19:34that and I believe this is just a float
  3501. 2:19:37and so we should be able to uh print
  3502. 2:19:40that uh so that's Global gradient
  3503. 2:19:44clipping now they go into the details of
  3504. 2:19:46the learning rate uh scheduler so they
  3505. 2:19:49don't just use a fixed learning rate
  3506. 2:19:51like we do here for 3 E4 but there's
  3507. 2:19:54actually basically a cosine DK learning
  3508. 2:19:57rate schedule um it's got a warm-up and
  3509. 2:20:00it's got a cosine DEC to 10% over some
  3510. 2:20:04Horizon
  3511. 2:20:06um and so we're going to implement uh
  3512. 2:20:09this in a second I just like to see Norm
  3513. 2:20:11printed here okay there we go so what
  3514. 2:20:14happened here is the norm is actually
  3515. 2:20:16really high in the beginning 30 or so
  3516. 2:20:19and you see that as we continue training
  3517. 2:20:21it kind of like
  3518. 2:20:22stabilizes um at values below one um and
  3519. 2:20:27this is not that crazy uncommon for the
  3520. 2:20:30norm to be high in the very first few
  3521. 2:20:31stages basically What's Happening Here
  3522. 2:20:33is the model is completely random and so
  3523. 2:20:35there's a ton of learning happening very
  3524. 2:20:37early in the network but that learning
  3525. 2:20:39is kind of like um you know it's mostly
  3526. 2:20:41learning the biases of the output tokens
  3527. 2:20:44and so it's a bit of an unstable time uh
  3528. 2:20:46but the network usually stabilizes in a
  3529. 2:20:48very few iterations so this looks very
  3530. 2:20:50relatively reasonable to me except
  3531. 2:20:52usually I would expect this looks a
  3532. 2:20:54little bit funky that we go from 28 to 6
  3533. 2:20:56to 2 and then to 10 um it's not
  3534. 2:20:59completely insane but it's just kind of
  3535. 2:21:01a little bit
  3536. 2:21:02funky um okay so let's now get to the
  3537. 2:21:05learning rate schuer so the learning
  3538. 2:21:07rate schedule that's used here in gpt3
  3539. 2:21:09is what's called a cosine Decay learning
  3540. 2:21:12schedule with warmup and the way this
  3541. 2:21:14looks is that the learning rate is
  3542. 2:21:17basically starts right at around zero
  3543. 2:21:19linearly rank s up over some amount of
  3544. 2:21:21time and then comes down with this
  3545. 2:21:24cosine sort of form and comes down to
  3546. 2:21:27some kind of a minimum learning rate
  3547. 2:21:28that's up to you so here the minimum
  3548. 2:21:30learning rate is zero but uh here in the
  3549. 2:21:33paper they said that they use cosine
  3550. 2:21:35Decay for learning rate down to 10% of
  3551. 2:21:37its value over the first 260 billion
  3552. 2:21:40tokens and then training continues 10%
  3553. 2:21:43after and there's a linear warmup over
  3554. 2:21:46the first 375 million tokens so that's
  3555. 2:21:50about the learn R so let's now implement
  3556. 2:21:52this uh so I already implemented it here
  3557. 2:21:55and the way this works is let me scroll
  3558. 2:21:58down first here I changed our training
  3559. 2:22:00Loop a little bit so this was a 4i in
  3560. 2:22:02Max steps I just change it to step now
  3561. 2:22:04so that we have the notion of a step is
  3562. 2:22:07a single optimization step in the in the
  3563. 2:22:09for Loop and then here I get the LR for
  3564. 2:22:13this step of the optimization using a
  3565. 2:22:15new function I call get LR and then in
  3566. 2:22:18pytorch to set the learning rate I think
  3567. 2:22:20this is is the way to set the learning
  3568. 2:22:21rate it's a little bit gnarly um because
  3569. 2:22:24you have to basically there's a notion
  3570. 2:22:25of different par parameter groups that
  3571. 2:22:27could exist in the optimizer and so you
  3572. 2:22:28actually have to iterate over them even
  3573. 2:22:30though we currently have a single param
  3574. 2:22:32group only um and you have to set the LR
  3575. 2:22:34in this for Loop kind of style is is my
  3576. 2:22:37impression right now so we have this
  3577. 2:22:39look of LR we set the learning rate and
  3578. 2:22:42then on the bottom I'm also printing it
  3579. 2:22:45uh so that's all the changes I made to
  3580. 2:22:47this Loop and then of course the get LR
  3581. 2:22:49is my scheduler now it's worth pointing
  3582. 2:22:51out that pytorch actually has learning
  3583. 2:22:53rate schedulers and you can use them and
  3584. 2:22:55I believe there's a cosine learning rate
  3585. 2:22:57schedule in pytorch I just don't really
  3586. 2:22:59love using that code because honestly
  3587. 2:23:02it's like five lines of code and I fully
  3588. 2:23:06understand what's happening inside these
  3589. 2:23:07lines so I don't love to use
  3590. 2:23:09abstractions where they're kind of in
  3591. 2:23:11screwable and then I don't know what
  3592. 2:23:13they're doing so personal style so the
  3593. 2:23:16max learning rate here is let's say 3 E4
  3594. 2:23:19but we're going to see that in gpt3
  3595. 2:23:22here they have a table of what the
  3596. 2:23:25maximum learning rate is for every model
  3597. 2:23:28size so um for for this one basically 12
  3598. 2:23:3412 layer 768 gpt3 so the gpt3 small is
  3599. 2:23:37roughly like a GPT
  3600. 2:23:402124m we see that here they use a
  3601. 2:23:42learning rate of 6 E4 so we could
  3602. 2:23:44actually go higher um in fact we may
  3603. 2:23:46want to try to follow that and just set
  3604. 2:23:48the max LR here at six
  3605. 2:23:51uh then the that's the maximum learning
  3606. 2:23:53rate the minum learning rate is uh 10%
  3607. 2:23:55of that per description in the paper
  3608. 2:23:58some number of steps that we're going to
  3609. 2:24:00warm up over and then the maximum steps
  3610. 2:24:02of the optimization which I now use also
  3611. 2:24:05in the for Loop down here and then you
  3612. 2:24:07can go over this code if you like it's
  3613. 2:24:09not U it's not terribly inside Flor
  3614. 2:24:11interesting I'm just uh modulating based
  3615. 2:24:13on the iteration number which learning
  3616. 2:24:16rate uh there should be so this is the
  3617. 2:24:18warm-up region um
  3618. 2:24:21this is the region after the
  3619. 2:24:22optimization and then this is the region
  3620. 2:24:24sort of in between and this is where I
  3621. 2:24:26calculate the cosine learning rate
  3622. 2:24:28schedule and you can step through this
  3623. 2:24:29in detail if you'd like uh but this is
  3624. 2:24:32basically implementing this
  3625. 2:24:33curve and I ran this already and this is
  3626. 2:24:38what that looks
  3627. 2:24:40like um so when we now run we start at
  3628. 2:24:45um some very low number now note that we
  3629. 2:24:47don't start exactly at zero because that
  3630. 2:24:49would be not useful to update with a
  3631. 2:24:50learning rate of zero that's why there's
  3632. 2:24:52an it+ one so that on the zeroth
  3633. 2:24:54iteration we are not using exactly zero
  3634. 2:24:57we're using something very very low then
  3635. 2:24:59we linearly warm up to maximum learning
  3636. 2:25:02rate which in this case was 34 when I
  3637. 2:25:04ran it but now would be 6 E4 and then it
  3638. 2:25:07starts to decay all the way down to um 3
  3639. 2:25:11E5 which was at the time 10% of the
  3640. 2:25:14original learning rate now one thing we
  3641. 2:25:16are not following exactly is that they
  3642. 2:25:18mentioned that um
  3643. 2:25:21let me see if I can find it
  3644. 2:25:23again we're not exactly following what
  3645. 2:25:26they did
  3646. 2:25:28because uh they mentioned that their
  3647. 2:25:30training Horizon is 300 billion tokens
  3648. 2:25:33and they come down to 10% of the initial
  3649. 2:25:35learning rate of at 260 billion and then
  3650. 2:25:37they train after 260 with 10% so
  3651. 2:25:41basically their Decay time is less than
  3652. 2:25:43the max steps time whereas for us
  3653. 2:25:45they're exactly equal so it's not
  3654. 2:25:47exactly faithful but it's um it's an
  3655. 2:25:51okay um this is okay for us and for our
  3656. 2:25:53purposes right now and um we're just
  3657. 2:25:57going to use this ourselves I don't
  3658. 2:25:58think it makes too too big of a
  3659. 2:26:00difference honestly I should point out
  3660. 2:26:02that what learning rate schedule you use
  3661. 2:26:04is totally up to you there's many
  3662. 2:26:05different types um coign learning rate
  3663. 2:26:08has been popularized a lot by gpt2 and
  3664. 2:26:10gpt3 but people have come up with all
  3665. 2:26:12kinds of uh other learning rate
  3666. 2:26:14schedules um and this is kind of like an
  3667. 2:26:16active area of uh research as to which
  3668. 2:26:18one is the most effective at train these
  3669. 2:26:20networks okay next up the paper talks
  3670. 2:26:23about the gradual batch size increase so
  3671. 2:26:26there's a ramp on the batch size that is
  3672. 2:26:29linear and you start with very small
  3673. 2:26:31batch size and you ramp up to a big
  3674. 2:26:32batch size over time uh we're going to
  3675. 2:26:35actually skip this and we're not going
  3676. 2:26:36to work with it and the reason I don't
  3677. 2:26:38love to use it is that it complicates a
  3678. 2:26:41lot of the arithmetic because you are
  3679. 2:26:42changing the number of tokens that
  3680. 2:26:43you're processing at every single step
  3681. 2:26:45of the optimization and I like to keep
  3682. 2:26:47that math very very simple also my
  3683. 2:26:49understanding is that that this is not
  3684. 2:26:50like a major um Improvement and also my
  3685. 2:26:54understanding is that this is not like
  3686. 2:26:55an algorithmic optimization Improvement
  3687. 2:26:57it's more of a systems and speed
  3688. 2:26:59Improvement and roughly speaking this is
  3689. 2:27:02because uh in the early stages of the
  3690. 2:27:05optimization uh again the model is in a
  3691. 2:27:07very atypical setting and mostly what
  3692. 2:27:10you're learning is that um you're mostly
  3693. 2:27:13learning to ignore the tokens uh that
  3694. 2:27:15don't come up in your training set very
  3695. 2:27:16often you're learning very simple biases
  3696. 2:27:19and and that kind of a thing and so
  3697. 2:27:23every single example that you put
  3698. 2:27:24through your network is basically just
  3699. 2:27:26telling you use these tokens and don't
  3700. 2:27:28use these tokens and so the gradients
  3701. 2:27:30from every single example are actually
  3702. 2:27:31extremely highly correlated they all
  3703. 2:27:33look roughly the same in the in the OR
  3704. 2:27:36original parts of the optimization
  3705. 2:27:38because they're all just telling you
  3706. 2:27:39that these tokens don't appear and these
  3707. 2:27:40tokens do appear and so because the
  3708. 2:27:43gradients are all very similar and
  3709. 2:27:45they're highly correlated then why are
  3710. 2:27:46you doing batch sizes of like Millions
  3711. 2:27:49when if you do a batch size of 32k
  3712. 2:27:51you're basically getting the exact same
  3713. 2:27:53gradient early on in the training and
  3714. 2:27:55then later in the optimization once
  3715. 2:27:57you've learned all the simple stuff
  3716. 2:28:00that's where the actual work starts and
  3717. 2:28:01that's where the gradients become more
  3718. 2:28:02decorrelated per examples and that's
  3719. 2:28:04where they actually offer you sort of
  3720. 2:28:07statistical power in some sense um so
  3721. 2:28:10we're going to skip this just because it
  3722. 2:28:12kind of complicates things and we're
  3723. 2:28:14going to go
  3724. 2:28:15to uh data are sampled without
  3725. 2:28:18replacement during training um so until
  3726. 2:28:21an Epoch boundary is reached so without
  3727. 2:28:23replacement means that they're not
  3728. 2:28:24sampling from some fixed pool and then
  3729. 2:28:27uh take a sequence train on it but then
  3730. 2:28:31also like return the sequence to the
  3731. 2:28:32pool they are exhausting a pool so when
  3732. 2:28:34they draw a sequence it's it's gone
  3733. 2:28:37until the next Epoch of training uh so
  3734. 2:28:39we're already doing that because our
  3735. 2:28:41data loader um iterates over chunks of
  3736. 2:28:44data so there's no replacement they
  3737. 2:28:47don't become eligible to be drawn again
  3738. 2:28:49until the next P so we're basically
  3739. 2:28:51already doing
  3740. 2:28:53that um all models use a weight decay of
  3741. 2:28:560.1 to provide a small amount of
  3742. 2:28:59regularization so let's Implement a
  3743. 2:29:01weight Decay and you see here that I've
  3744. 2:29:03already kind of made the changes and in
  3745. 2:29:04particular instead of creating the
  3746. 2:29:06optimizer right here um I I'm creating a
  3747. 2:29:10new configure optimizers function inside
  3748. 2:29:12the model and I'm passing in some of the
  3749. 2:29:14hyper parameters instead so let's look
  3750. 2:29:17at the configure optimizers which is
  3751. 2:29:18supposed to return the optimizer
  3752. 2:29:24object okay so it looks complicated but
  3753. 2:29:27it's actually really simple and it's
  3754. 2:29:29just um we're just being very careful
  3755. 2:29:31and there's a few settings here to go
  3756. 2:29:32through the most important thing with
  3757. 2:29:34respect to this line is that you see
  3758. 2:29:36there's a weight Decay parameter here
  3759. 2:29:38and I'm passing that
  3760. 2:29:41into um well I'm passing that into
  3761. 2:29:44something called optim groups that
  3762. 2:29:46eventually ends up going into the addom
  3763. 2:29:47W Optimizer um and the weight Decay
  3764. 2:29:50that's by default used in Addam W here
  3765. 2:29:53is 0.01 so it's it's u 10 times lower
  3766. 2:29:57than what's used in gpt3 paper here um
  3767. 2:30:01so the weight dek basically ends up
  3768. 2:30:02making its way into the ADD and W
  3769. 2:30:04through the optimizer groups now what
  3770. 2:30:05else is going on here in this uh
  3771. 2:30:07function so the two things that are
  3772. 2:30:09happening here that are important is
  3773. 2:30:10that I'm splitting up the parameters
  3774. 2:30:12into those that should be weight decayed
  3775. 2:30:14and those that should not be weight
  3776. 2:30:15decayed so in particular it is common to
  3777. 2:30:18not weight decay uh biases and any other
  3778. 2:30:22sort of one-dimensional tensors so the
  3779. 2:30:25one-dimensional tensors are in the no
  3780. 2:30:27Decay prams and these are also things
  3781. 2:30:30like uh layer Norm scales and biases it
  3782. 2:30:33doesn't really make sense to weight
  3783. 2:30:34Decay those you mostly want to weight
  3784. 2:30:36Decay uh the weights that participate in
  3785. 2:30:39Matrix multiplications and you want to
  3786. 2:30:41potentially weight Decay the
  3787. 2:30:43embeddings and uh We've covered in
  3788. 2:30:46previous video why it makes sense to
  3789. 2:30:47Decay the weights because you can sort
  3790. 2:30:49of the it as a regularization because
  3791. 2:30:51when you're pulling down all the weights
  3792. 2:30:53you're forcing the optimization to use
  3793. 2:30:55more of the weights um and you're not
  3794. 2:30:57allowing any one of the weights
  3795. 2:30:59individually to be way too large um
  3796. 2:31:02you're forcing you're forcing the
  3797. 2:31:03network to kind of like distribute the
  3798. 2:31:05work across more channels because
  3799. 2:31:07there's sort of like a pull of gravity
  3800. 2:31:09on the weights
  3801. 2:31:11themselves um so that's why we are
  3802. 2:31:13separating it in those ways here we're
  3803. 2:31:16only decaying the embeddings and the
  3804. 2:31:18mmal participating ways
  3805. 2:31:21uh we're printing the number of uh
  3806. 2:31:22parameters that we decaying and not most
  3807. 2:31:24of the parameters will be decayed and
  3808. 2:31:26then one more thing that we're doing
  3809. 2:31:27here is I'm doing another optimization
  3810. 2:31:31here and previous add and W did not have
  3811. 2:31:34this option but later parts of pytorch
  3812. 2:31:37introduced it and that's why I'm
  3813. 2:31:38guarding it with an inspect do signature
  3814. 2:31:41which is basically checking if this
  3815. 2:31:43fused um quar is present inside atom W
  3816. 2:31:48and then if it is present I'm going to
  3817. 2:31:50end up using it and passing it in here
  3818. 2:31:53because some earlier versions do not
  3819. 2:31:55have fused equals so here's adamw fused
  3820. 2:31:58equals it did not used to exist and it
  3821. 2:32:00was added later and there's some docks
  3822. 2:32:03here for what's happening and basically
  3823. 2:32:05they say that by default they do not use
  3824. 2:32:07fused because it is relatively new and
  3825. 2:32:10we want to give it sufficient big time
  3826. 2:32:12so by default they don't use fused but
  3827. 2:32:13fused is a lot faster when it is
  3828. 2:32:15available and when you're running on
  3829. 2:32:17Cuda and what that does is in instead of
  3830. 2:32:20iterating in a for Loop over all the
  3831. 2:32:22parameter tensors and updating them that
  3832. 2:32:25would launch a lot of kernels right and
  3833. 2:32:27so a fused just means that it's a um all
  3834. 2:32:30those kernels are fused into a single
  3835. 2:32:31kernel you get rid of a lot of overhead
  3836. 2:32:34and you a single time on all the
  3837. 2:32:36parameters call a uh kernel that updates
  3838. 2:32:39them and so it's just basically a kernel
  3839. 2:32:42Fusion for the atom W update instead of
  3840. 2:32:44iterating over all the
  3841. 2:32:47tensors so that's the configure
  3842. 2:32:48optimizers function that I like to use
  3843. 2:32:51and we can rerun and we're not going to
  3844. 2:32:53see any major differences from what we
  3845. 2:32:55saw before but we are going to see some
  3846. 2:32:57prints uh coming from here so let's just
  3847. 2:33:00take a look at what they look
  3848. 2:33:01like so we see that number of Decay
  3849. 2:33:04tensors is 50 and it's most of the
  3850. 2:33:06parameters and number of non- deay
  3851. 2:33:08tensors is 98 and these are the biases
  3852. 2:33:10and the layer Norm parameters mostly and
  3853. 2:33:13that's there's only 100,000 of those so
  3854. 2:33:15most of it is decayed and then we are
  3855. 2:33:18using the fused implementation of ATM W
  3856. 2:33:20which will be a lot faster so if you
  3857. 2:33:22have it available I would advise you to
  3858. 2:33:24use it I'm not actually 100% sure why
  3859. 2:33:26they don't default to it it seems fairly
  3860. 2:33:28benign and
  3861. 2:33:29harmless and also because we are using
  3862. 2:33:31the fused implementation I think this is
  3863. 2:33:34why we have dropped um notice that the
  3864. 2:33:37running time used to be 93 milliseconds
  3865. 2:33:39per step and we're now down to 90
  3866. 2:33:41milliseconds per step because of using
  3867. 2:33:43the fused atom W Optimizer so in a
  3868. 2:33:46single commit here we are introducing
  3869. 2:33:48fused atom getting improvements on the
  3870. 2:33:51time and we're adding or changing the
  3871. 2:33:54weight Decay but we're only weight
  3872. 2:33:56decaying the two dimensional parameters
  3873. 2:33:58the embeddings and the matrices that
  3874. 2:34:00participate in linear so that is this
  3875. 2:34:03and we can take this out and uh yeah
  3876. 2:34:06that is it for this line one more quick
  3877. 2:34:10note before we continue here I just want
  3878. 2:34:11to point out that the relationship
  3879. 2:34:13between weight Decay learning rate batch
  3880. 2:34:15size the atom parameters beta 1 beta 2
  3881. 2:34:18the Epsilon and so on these are very
  3882. 2:34:20complicated uh mathematical
  3883. 2:34:22relationships in the optimization
  3884. 2:34:24literature and um for the most part I'm
  3885. 2:34:27in this video I'm just trying to copy
  3886. 2:34:29paste the settings that open AI used but
  3887. 2:34:31this is a complicated topic uh quite
  3888. 2:34:33deep and um yeah in this video I just
  3889. 2:34:36want to copy the parameters because it's
  3890. 2:34:38a whole different video to really talk
  3891. 2:34:39about that in detail and give it a
  3892. 2:34:41proper Justice instead of just high
  3893. 2:34:42level
  3894. 2:34:43intuitions uh now the next thing that I
  3895. 2:34:45want to move on to is that uh this
  3896. 2:34:48paragraph here by the way we're going to
  3897. 2:34:49turn back around to when we improve our
  3898. 2:34:51data loader for now I want to swing back
  3899. 2:34:54around
  3900. 2:34:56to this
  3901. 2:35:01table where you will notice that um for
  3902. 2:35:04different models we of course have
  3903. 2:35:06different U hyper parameters for the
  3904. 2:35:08Transformer that dictate the size of the
  3905. 2:35:10Transformer Network we also have a
  3906. 2:35:12different learning rate so we're seeing
  3907. 2:35:13the pattern that the bigger networks are
  3908. 2:35:14trained with slightly lower learning
  3909. 2:35:16rates and we also see this batch size
  3910. 2:35:20where in in the small networks they use
  3911. 2:35:22a smaller batch size and in the bigger
  3912. 2:35:23networks they use a bigger batch size
  3913. 2:35:26now the problem with for us is we can't
  3914. 2:35:28just use 0.5 million batch size because
  3915. 2:35:31uh if I just try to come in here and I
  3916. 2:35:33try to set uh this uh B where is my
  3917. 2:35:38b
  3918. 2:35:40um b
  3919. 2:35:44equals where where do I call the DAT
  3920. 2:35:46okay b equal 16 if I try to set um
  3921. 2:35:51well well we have to be careful it's not
  3922. 2:35:520.5 million because this is the badge
  3923. 2:35:54size in the number of tokens every
  3924. 2:35:56single one of our rows is24 tokens so
  3925. 2:36:000.5 E6 1 million divide 1024 this would
  3926. 2:36:04need about a
  3927. 2:36:06488 match size so the problem is I can't
  3928. 2:36:09come in here and set this to 488 uh
  3929. 2:36:12because my GPU would explode um this
  3930. 2:36:15would not fit for sure and so but we
  3931. 2:36:18still want to use this batch size
  3932. 2:36:20because again as I mentioned the batch
  3933. 2:36:22size is correlated with all the other
  3934. 2:36:24optimization hyper parameters and the
  3935. 2:36:26learning rates and so on so we want to
  3936. 2:36:28have a faithful representation of all
  3937. 2:36:29the hyper parameters and therefore we
  3938. 2:36:31need to uh use a bat size of .5 million
  3939. 2:36:34roughly but the question is how do we
  3940. 2:36:37use .5 million if we only have a small
  3941. 2:36:39GPU well for that we need to use what's
  3942. 2:36:41called gradient accumulation uh so we're
  3943. 2:36:44going to turn to that next and it allows
  3944. 2:36:46us to simulate in a Serial way any
  3945. 2:36:48arbitrary batch size that we set and so
  3946. 2:36:51we can do a batch size of .5 million we
  3947. 2:36:54just have to run longer and we have to
  3948. 2:36:56process multiple sequences and basically
  3949. 2:36:59add up all the gradients from them to
  3950. 2:37:02simulate a batch size of .5 million so
  3951. 2:37:04let's turn to that next okay so I
  3952. 2:37:05started the implementation right here
  3953. 2:37:07just by adding these lines of code and
  3954. 2:37:09basically what I did is first I set the
  3955. 2:37:12total batch size that we desire so this
  3956. 2:37:14is exactly .5 million and I used a nice
  3957. 2:37:17number a power of two uh because 2 to
  3958. 2:37:19the 19 is 524 288 so it's roughly .5
  3959. 2:37:23million it's a nice number now our micro
  3960. 2:37:26batch size as we call it now is 16 so
  3961. 2:37:29this is going to be we still have B BYT
  3962. 2:37:32in the SE that go into the Transformer
  3963. 2:37:34and do forward backward but we're not
  3964. 2:37:36going to do an update right we're going
  3965. 2:37:38to do many forward backwards we're going
  3966. 2:37:40to and those gradients are all going to
  3967. 2:37:42plus equals on the parameter gradients
  3968. 2:37:44they're all going to add up so we're
  3969. 2:37:46going to do forward backward grad akum
  3970. 2:37:48steps number of times and then we're
  3971. 2:37:50going to do a single update once all
  3972. 2:37:52that is
  3973. 2:37:53accumulated so in particular our micro
  3974. 2:37:55batch size is just now controlling how
  3975. 2:37:58many tokens how many rows we're
  3976. 2:37:59processing in a single go over a forward
  3977. 2:38:02backward so um here we are doing 16 *
  3978. 2:38:06124 we're doing 16
  3979. 2:38:09384 um tokens per forward backward and
  3980. 2:38:14we are supposed to be doing 2 to the 19
  3981. 2:38:17whoops what am I doing 2 to the
  3982. 2:38:2019 in total so the grat Aon will be
  3983. 2:38:2632 uh so therefore gr AUM here will work
  3984. 2:38:28out to 32 and we have to do 32 forward
  3985. 2:38:32backward um and then a single update now
  3986. 2:38:35we see that we have about 100
  3987. 2:38:37milliseconds for a singer forward
  3988. 2:38:38backward so doing 32 of them will be
  3989. 2:38:41will make every step roughly 3 seconds
  3990. 2:38:44just napkin
  3991. 2:38:46math so that's grum steps but now we
  3992. 2:38:48actually have to Implement that so we're
  3993. 2:38:50going to swing over to our training Loop
  3994. 2:38:54because now this part
  3995. 2:38:56here and this part here the forward and
  3996. 2:38:59the backward we have to now repeat this
  3997. 2:39:0132 times before we do everything else
  3998. 2:39:04that follows so let's uh see how we can
  3999. 2:39:06Implement that so let's come over here
  4000. 2:39:09and actually we do have to load a new
  4001. 2:39:10batch every single time so let me move
  4002. 2:39:12that over here and now this is where we
  4003. 2:39:14have the inner loop so for micro step in
  4004. 2:39:18range graum
  4005. 2:39:20steps we do this and remember that l.
  4006. 2:39:24backward always deposits gradients so
  4007. 2:39:26we're doing inside losta backward
  4008. 2:39:27there's always a plus equals on the
  4009. 2:39:29gradients so in every single L of
  4010. 2:39:31backward gradients will add up on the
  4011. 2:39:33gradient
  4012. 2:39:35tensors um so we lost that backward and
  4013. 2:39:38then we get all the gradients over there
  4014. 2:39:41and then we normalize and everything
  4015. 2:39:43else should just follow um so we're very
  4016. 2:39:47close but actually there's like subtle
  4017. 2:39:50and deep issue here and this is actually
  4018. 2:39:52incorrect so invite I invite you to
  4019. 2:39:54think about why this is not yet
  4020. 2:39:56sufficient um and uh let me fix it then
  4021. 2:39:59okay so I brought back the jupyter
  4022. 2:40:01notebook so we can think about this
  4023. 2:40:02carefully in a simple toy setting and
  4024. 2:40:05see what's happening so let's create a
  4025. 2:40:07very simple neural nut that takes a 16
  4026. 2:40:10Vector of 16 numbers and returns a
  4027. 2:40:11single
  4028. 2:40:12number and then here I'm creating some
  4029. 2:40:15random uh examples X and some targets uh
  4030. 2:40:19y Y and then we are using the mean
  4031. 2:40:21squared loss uh here to calculate the
  4032. 2:40:25loss so basically what this is is four
  4033. 2:40:28individual examples and we're just doing
  4034. 2:40:30Simple regression with the mean squared
  4035. 2:40:31loss over those four
  4036. 2:40:34examples now when we calculate the loss
  4037. 2:40:36and we lost that backward and look at
  4038. 2:40:38the gradient this is the gradient that
  4039. 2:40:40we
  4040. 2:40:41achieve now the loss objective here
  4041. 2:40:44notice that in MSE loss the default for
  4042. 2:40:46the loss function is reduction is mean
  4043. 2:40:49so we're we're calculating the average
  4044. 2:40:52mean loss um the the mean loss here over
  4045. 2:40:56the four examples so this is the exact
  4046. 2:40:59loss objective and this is the average
  4047. 2:41:02the one over four because there are four
  4048. 2:41:03independent examples here and then we
  4049. 2:41:06have the four examples and their mean
  4050. 2:41:08squared error the squared error and then
  4051. 2:41:11this makes it the mean squared error so
  4052. 2:41:14therefore uh we are we calculate the
  4053. 2:41:16squared error and then we normalize it
  4054. 2:41:18to make it the mean over the examples
  4055. 2:41:20and there's four examples here so now
  4056. 2:41:22when we come to the gradient
  4057. 2:41:24accumulation version of it this uh this
  4058. 2:41:28here is the gradient accumulation
  4059. 2:41:30version of it where we have grad acum
  4060. 2:41:32steps of four and I reset the gradient
  4061. 2:41:35we've grum steps of four and now I'm
  4062. 2:41:38evaluating all the examples individually
  4063. 2:41:39instead and calling L that backward on
  4064. 2:41:41them many times and then we're looking
  4065. 2:41:43at the gradient that we achieve from
  4066. 2:41:44that so basically now we forward our
  4067. 2:41:47function calculate the exact same loss
  4068. 2:41:49do a backward and we do that four times
  4069. 2:41:52and when we look at the gradient uh
  4070. 2:41:54you'll notice that the gradients don't
  4071. 2:41:57match so here we uh did a single batch
  4072. 2:42:00of four and here we did uh four gradient
  4073. 2:42:03accumulation steps of batch size one and
  4074. 2:42:06the gradients are not the same and
  4075. 2:42:08basically the the reason that they're
  4076. 2:42:09not the same is exactly because this
  4077. 2:42:11mean squared error gets lost this one
  4078. 2:42:14quarter in this loss gets lost because
  4079. 2:42:16what happens here is the loss of
  4080. 2:42:19objective for every one of the loops is
  4081. 2:42:22just a mean squ error um which in this
  4082. 2:42:25case because there's only a single
  4083. 2:42:26example is just this term here so that
  4084. 2:42:28was the loss in the zeroth eration same
  4085. 2:42:30in the first third and so on and then
  4086. 2:42:33when you do the loss. backward we're
  4087. 2:42:35accumulating gradients and what happens
  4088. 2:42:38is that accumulation in the gradient is
  4089. 2:42:40basically equivalent to doing a sum in
  4090. 2:42:43the
  4091. 2:42:45loss so our loss actually here is this
  4092. 2:42:49without the factor of one quarter
  4093. 2:42:51outside of it so we're missing the
  4094. 2:42:54normalizer and therefore our gradients
  4095. 2:42:56are off and so the way to fix this or
  4096. 2:42:58one of them is basically we can actually
  4097. 2:43:00come here and we can say loss equals
  4098. 2:43:02loss divide
  4099. 2:43:044 and what happens now is that we're
  4100. 2:43:07introducing we're we're scaling our loss
  4101. 2:43:09we're introducing a one quarter in front
  4102. 2:43:11of all of these
  4103. 2:43:14places so all the individual losses are
  4104. 2:43:17now scaled by one quarter and and then
  4105. 2:43:19when we backward all of these accumulate
  4106. 2:43:22with a sum but now there's a one quarter
  4107. 2:43:24inside every one of these components and
  4108. 2:43:26now our losses will be
  4109. 2:43:28equivalent so when I run this you see
  4110. 2:43:32that the U gradients are now identical
  4111. 2:43:35so long story short with this simple
  4112. 2:43:37example uh when you step through it you
  4113. 2:43:39can see that basically the reason that
  4114. 2:43:41this is not correct is because in the
  4115. 2:43:44same way as here in the MSE loss the
  4116. 2:43:46loss that we're calculating here in the
  4117. 2:43:50model is using a reduction of mean as
  4118. 2:43:54well uh so where's the loss after that
  4119. 2:43:57cross
  4120. 2:43:58entropy and by default the reduction uh
  4121. 2:44:01here in Cross entropy is also I don't
  4122. 2:44:03know why they don't show it but it's the
  4123. 2:44:05mean uh the mean uh loss at all the B
  4124. 2:44:08BYT elements
  4125. 2:44:10right so there's a reduction by mean in
  4126. 2:44:13there and if we're just doing this
  4127. 2:44:15gradient accumulation here we're missing
  4128. 2:44:16that and so the way to fix this is to
  4129. 2:44:19simply compensate for the number of
  4130. 2:44:21gradient accumulation steps and we can
  4131. 2:44:23in the same way divide this loss so in
  4132. 2:44:25particular here the number of steps that
  4133. 2:44:26we're doing is loss equals loss divide
  4134. 2:44:31gradient accumulation steps so even uh
  4135. 2:44:33co-pilot s gets the modification but in
  4136. 2:44:36the same way exactly we are scaling down
  4137. 2:44:38the loss so that when we do loss that
  4138. 2:44:40backward which basically corresponds to
  4139. 2:44:42a sum in the objective we are summing up
  4140. 2:44:45the already
  4141. 2:44:46normalized um loss and and therefore
  4142. 2:44:49when we sum up the losses divided by
  4143. 2:44:51grum steps we are recovering the
  4144. 2:44:53additional normalizer uh and so now
  4145. 2:44:56these two will be now this will be
  4146. 2:44:59equivalent to the original uh sort of
  4147. 2:45:01optimization because the gradient will
  4148. 2:45:03come out the same okay so I had to do a
  4149. 2:45:05few more touch-ups and I launched
  4150. 2:45:07launched the optimization here so in
  4151. 2:45:09particular one thing we want to do
  4152. 2:45:10because we want to print things nicely
  4153. 2:45:13is well first of all we need to create
  4154. 2:45:15like an accumulator over the loss we
  4155. 2:45:16can't just print the loss because we'd
  4156. 2:45:18be printing only the final loss at the
  4157. 2:45:20final micro step so instead we have loss
  4158. 2:45:22ofon which I initialize at zero and then
  4159. 2:45:25I accumulate a uh the loss into it and
  4160. 2:45:28I'm using detach so that um uh I'm
  4161. 2:45:31detaching the tensor uh from the graph
  4162. 2:45:35and I'm just trying to keep track of the
  4163. 2:45:36values so I'm making these Leaf nodes
  4164. 2:45:38when I add them so that's lakum and then
  4165. 2:45:42we're printing that here instead of loss
  4166. 2:45:43and then in addition to that I had to
  4167. 2:45:46account for the grum steps inside the
  4168. 2:45:48tokens processed because now the tokens
  4169. 2:45:50processed per step is B * T * gradient
  4170. 2:45:54accumulation so long story short here we
  4171. 2:45:57have the optimization it looks uh
  4172. 2:45:59reasonable right we're starting at a
  4173. 2:46:00good spot we calculated the grum steps
  4174. 2:46:03to be
  4175. 2:46:0432 and uh we're getting about 3 seconds
  4176. 2:46:07here
  4177. 2:46:08right
  4178. 2:46:10um
  4179. 2:46:12and so this looks pretty good now if
  4180. 2:46:14you'd like to verify that uh your
  4181. 2:46:16optimization and the implementation here
  4182. 2:46:18is correct and your working on a side
  4183. 2:46:20well now because we have the total patch
  4184. 2:46:21size and the gradient accumulation steps
  4185. 2:46:24our setting of B is purely a performance
  4186. 2:46:26optimization kind of setting so if you
  4187. 2:46:29have a big GPU you can actually increase
  4188. 2:46:31this to 32 and you'll probably go a bit
  4189. 2:46:33faster if you have a very small GPU you
  4190. 2:46:35can try eight or four but in any case
  4191. 2:46:37you should be getting the exact same
  4192. 2:46:38optimization and the same answers up to
  4193. 2:46:41like a floating Point error because the
  4194. 2:46:43gradient accumulation kicks in and um
  4195. 2:46:46and can um handle everything serially as
  4196. 2:46:48an
  4197. 2:46:49Neary so uh that's it for gradient
  4198. 2:46:51accumulation I think okay so now is the
  4199. 2:46:53time to bring out the heavy weapons uh
  4200. 2:46:56you've noticed that so far we've only
  4201. 2:46:57been using a single GPU for training but
  4202. 2:47:00actually I am paying for eight gpus here
  4203. 2:47:02and so uh we should be putting all of
  4204. 2:47:04them to work and in particular they are
  4205. 2:47:06going to collaborate and uh you know
  4206. 2:47:09optimize over tokens at the same time
  4207. 2:47:12and communicate so that um uh they're
  4208. 2:47:15all kind of collaborating on the
  4209. 2:47:16optimization for this we are going to be
  4210. 2:47:18using the distributed data parallel from
  4211. 2:47:20pytorch there's also a legacy data
  4212. 2:47:22parallel which I recommend you not use
  4213. 2:47:24and that's kind of like you know Legacy
  4214. 2:47:27distributed data parallel Works in a
  4215. 2:47:28very simple way we have eight gpus so
  4216. 2:47:31we're going to uh launch eight processes
  4217. 2:47:35and each process is going to be assigned
  4218. 2:47:36to GPU and for each process the training
  4219. 2:47:40Loop and everything we've worked on so
  4220. 2:47:41far is going to look pretty much the
  4221. 2:47:42same H GPU as far as it's concerned is
  4222. 2:47:45just working on exactly what we've built
  4223. 2:47:47so far but now Secret L there's eight of
  4224. 2:47:49them and they're all going to be
  4225. 2:47:51processing slightly different parts of
  4226. 2:47:52the data and we're going to add one more
  4227. 2:47:56part where once they all calculate their
  4228. 2:47:58gradients there's one more part where we
  4229. 2:48:00do a average of those
  4230. 2:48:03gradients and so that's how they're
  4231. 2:48:05going to be collaborating on uh the
  4232. 2:48:07computational workload here so to use
  4233. 2:48:10all eight of them we're not going to be
  4234. 2:48:12launching our script anymore with just
  4235. 2:48:14um pytorch train
  4236. 2:48:16gbt2 piy we're going to be running it
  4237. 2:48:19with a special command called torrun in
  4238. 2:48:21pytorch we'll see that in a bit and
  4239. 2:48:23torrun uh when it runs our python script
  4240. 2:48:26we'll actually make sure to run eight
  4241. 2:48:28eight of them in parallel and it creates
  4242. 2:48:32these environmental variables where each
  4243. 2:48:34of these processes can look up which uh
  4244. 2:48:37basically which one of the processes it
  4245. 2:48:40is so for example torron will set rank
  4246. 2:48:43local Rank and World size environmental
  4247. 2:48:46variables and so this is a bad way to
  4248. 2:48:48detect whether uh DDP is running so if
  4249. 2:48:51we're using torch run if DDP is
  4250. 2:48:54running then uh we have to make sure
  4251. 2:48:57that K is available because I don't know
  4252. 2:48:58that you can run this on CPU anymore or
  4253. 2:49:01that that makes sense to do um this is
  4254. 2:49:05some um setup code here the important
  4255. 2:49:07part is that there's a world size which
  4256. 2:49:10for us will be eight that's the total
  4257. 2:49:11number of processes running there's a
  4258. 2:49:14rank which is um each process will
  4259. 2:49:17basically run the ex exact same code at
  4260. 2:49:19the exact same time roughly but all the
  4261. 2:49:22process the only difference between
  4262. 2:49:24these processes is that they all have a
  4263. 2:49:26different dtp rank so the um gpu0 will
  4264. 2:49:30have DDP rank of zero GPU 1 will have uh
  4265. 2:49:33rank of one Etc so otherwise they're all
  4266. 2:49:36running the exact same script it's just
  4267. 2:49:38that DDP rank will be a slightly
  4268. 2:49:40different integer and that is the way
  4269. 2:49:42for us to coordinate that they don't for
  4270. 2:49:44example run on the same data we want to
  4271. 2:49:46we want them to run on different parts
  4272. 2:49:47of the data and so on
  4273. 2:49:49now local rank is something that is only
  4274. 2:49:52used in a multi- node setting we only
  4275. 2:49:54have a single node with ag gpus and so
  4276. 2:49:57local rank is the rank of the GPU on a
  4277. 2:50:00single node so from 0 to seven as an
  4278. 2:50:04example but for us we're mostly going to
  4279. 2:50:06be running on a single box so the things
  4280. 2:50:08we care about are Rank and World size
  4281. 2:50:10this is eight and this will be whatever
  4282. 2:50:12it is depending on the GPU uh that uh
  4283. 2:50:15that this particular instantiation of
  4284. 2:50:17the script runs on
  4285. 2:50:19now here we make sure that according to
  4286. 2:50:23the local rank we are setting the device
  4287. 2:50:27to be Cuda colon and colon indicates
  4288. 2:50:30which GPU to use if there are more than
  4289. 2:50:32one gpus so depending on the local rank
  4290. 2:50:36of this process it's going to use just
  4291. 2:50:39the appropriate GPU so there's no
  4292. 2:50:40collisions on which GPU is being used by
  4293. 2:50:42which
  4294. 2:50:43process and finally there's a Boolean
  4295. 2:50:45variable that I like to create which is
  4296. 2:50:47the DDP rank equ equal Z so the master
  4297. 2:50:50process is arbitrarily process number
  4298. 2:50:53zero and it does a lot of the printing
  4299. 2:50:55logging checkpointing Etc and the other
  4300. 2:50:57processes are thought of mostly as a
  4301. 2:50:59compute processes that are assisting and
  4302. 2:51:01so Master process zero will have some
  4303. 2:51:03additional work to do all the other
  4304. 2:51:05processes will uh will mostly just be
  4305. 2:51:06doing forward
  4306. 2:51:08backwards and if we're not using DDP and
  4307. 2:51:10none of these variables are set we
  4308. 2:51:12revert back to single GPU training so
  4309. 2:51:14that means that we only have rank zero
  4310. 2:51:16the world size is just one uh and and we
  4311. 2:51:19are the master process and we try to
  4312. 2:51:21autodetect the device and this is world
  4313. 2:51:24as
  4314. 2:51:25normal so so far all we've done is we've
  4315. 2:51:27initialized
  4316. 2:51:28DDP and uh in the case where we're
  4317. 2:51:31running with torrun which we'll see in a
  4318. 2:51:33bit there's going to be eight copies
  4319. 2:51:35running in parallel each one of them
  4320. 2:51:37will have a different Rank and now we
  4321. 2:51:39have to make sure that everything
  4322. 2:51:41happens uh correctly afterwards so the
  4323. 2:51:44tricky thing with running multiple
  4324. 2:51:45processes is you always have to imagine
  4325. 2:51:48that there's going to be eight processes
  4326. 2:51:50running in parallel so as you read the
  4327. 2:51:52code now you have to imagine there's
  4328. 2:51:54eight you know eight python interpreters
  4329. 2:51:57running down these lines of code and the
  4330. 2:51:59only difference between them is that
  4331. 2:52:01they have a different DDP rank so they
  4332. 2:52:03all come here they all pick the exact
  4333. 2:52:05same seed they all make all of these
  4334. 2:52:08calculations completely unaware of the
  4335. 2:52:10other copies running roughly speaking
  4336. 2:52:12right so they all make the exact same
  4337. 2:52:14calculations and now we have to adjust
  4338. 2:52:16these calculations to take into account
  4339. 2:52:19that there's actually like a certain
  4340. 2:52:21world size and certain ranks so in
  4341. 2:52:24particular these micro batches and
  4342. 2:52:26sequence lengths these are all just per
  4343. 2:52:28GPU right so now there's going to be num
  4344. 2:52:31processes of them running in parallel so
  4345. 2:52:34we have to adjust this right because the
  4346. 2:52:36grum steps now is going to be total B
  4347. 2:52:39size divide B * T time U DDP R
  4348. 2:52:43size because each um process will will
  4349. 2:52:48do B * T and there's this many of
  4350. 2:52:51them and so in addition to that we we
  4351. 2:52:54want to make sure that this fits nicely
  4352. 2:52:56into total batch size which for us it
  4353. 2:52:58will because 16 * 124 * 8 8 gpus is
  4354. 2:53:04131 uh K and so
  4355. 2:53:08524288 this means that our gratum will
  4356. 2:53:10be four with the current settings right
  4357. 2:53:13so there's going to be 16 * 124 process
  4358. 2:53:16on each GPU and then there's a GP pus so
  4359. 2:53:18we're going to be doing
  4360. 2:53:20131,000 tokens in a single forward
  4361. 2:53:23backward on the 8
  4362. 2:53:26gpus so we want to make sure that this
  4363. 2:53:28fits nicely so that we can derive a nice
  4364. 2:53:30gradient accumulation
  4365. 2:53:32steps and uh yeah let's just adjust the
  4366. 2:53:36comments here times uh DDP World size
  4367. 2:53:41okay so each GPU calculates this now
  4368. 2:53:45this is where we start to get run into
  4369. 2:53:46issues right so we are each process is
  4370. 2:53:49going to come by a print and they're all
  4371. 2:53:51going to print so we're going to have
  4372. 2:53:53eight copies of these prints so one way
  4373. 2:53:56to deal with this is exactly this master
  4374. 2:53:58process variable that we have so if
  4375. 2:54:00Master process then guard this and
  4376. 2:54:03that's just so that we just print this a
  4377. 2:54:05single time because otherwise all the
  4378. 2:54:07processes would have computed the exact
  4379. 2:54:08same variables and there's no need to
  4380. 2:54:10print this eight
  4381. 2:54:11times um before getting into the data
  4382. 2:54:14loader and we're going to have to
  4383. 2:54:15refactor it obviously maybe at this
  4384. 2:54:18point is uh we should do some prints and
  4385. 2:54:21uh just take it out for a spin and exit
  4386. 2:54:23at this point so import
  4387. 2:54:26sis and S start exit and print IM
  4388. 2:54:33GPU um DDP
  4389. 2:54:38rank IM GPU DDP Rank and that um
  4390. 2:54:43print
  4391. 2:54:46by so uh so now let's try to run this
  4392. 2:54:49and just see how this works so let's
  4393. 2:54:51take it for a spin just so we see what
  4394. 2:54:52it looks like so normally we use to
  4395. 2:54:54launch python train gpd2 P like this now
  4396. 2:54:57we're going to run with torch run and
  4397. 2:54:59this is what it looks like so torch run
  4398. 2:55:02Standalone number of processes for
  4399. 2:55:04example is eight for us because we have
  4400. 2:55:05eight gpus uh and then change of2 Pi so
  4401. 2:55:09this is what the command would look like
  4402. 2:55:11and torch run again we'll run eight of
  4403. 2:55:13these so let's just see what happens so
  4404. 2:55:16first
  4405. 2:55:18it gets a little busy so there's a lot
  4406. 2:55:20going on here so first of all there's
  4407. 2:55:22some warnings from distributed and I
  4408. 2:55:24don't actually know that these mean
  4409. 2:55:26anything I think this is just like the
  4410. 2:55:28code is setting up and the processes are
  4411. 2:55:29coming online and we're seeing some
  4412. 2:55:31preliminary failure to collect while the
  4413. 2:55:33processes come up I'm not 100% sure
  4414. 2:55:36about that but we start to then get into
  4415. 2:55:39actual prints
  4416. 2:55:41so all the processes went down and then
  4417. 2:55:44the first print actually comes from
  4418. 2:55:46process 5 uh just by chance and then it
  4419. 2:55:50printed so process 5 basically got here
  4420. 2:55:52first it said I'm process on GPU 5 buy
  4421. 2:55:56and then this these prints come from the
  4422. 2:56:00master
  4423. 2:56:01process so process 5 just finished first
  4424. 2:56:04for whatever reason it just depends on
  4425. 2:56:05how the operating system scheduled the
  4426. 2:56:07processes to run uh then gpu0 ended then
  4427. 2:56:10GPU 3 and two and then uh probably
  4428. 2:56:14process 5 or something like that has uh
  4429. 2:56:17exited and and DDP really doesn't like
  4430. 2:56:19that because we didn't properly dispose
  4431. 2:56:21of uh the multi-gpus um setting and so
  4432. 2:56:27process group has not been destroyed
  4433. 2:56:28before we destruct uh so it really
  4434. 2:56:31doesn't like that and in an actual
  4435. 2:56:33application we would want to call
  4436. 2:56:34destroy process group uh so that we
  4437. 2:56:37clean up DDP properly and so it doesn't
  4438. 2:56:40like that too much and then the rest of
  4439. 2:56:41the gpus finish and that's it so
  4440. 2:56:45basically we can't guarantee when these
  4441. 2:56:46processes are running it's totally
  4442. 2:56:48but they are running in parallel we
  4443. 2:56:50don't want them to be printing um and
  4444. 2:56:54next up let's erase
  4445. 2:56:57this next up we want to make sure that
  4446. 2:56:59when we create data loader light we need
  4447. 2:57:01to now make it aware of this
  4448. 2:57:03multi-process um setting because we
  4449. 2:57:06don't want all the processes to be
  4450. 2:57:07loading the exact same data we want
  4451. 2:57:10every process to get its own chunk of
  4452. 2:57:11data so that they're all working on
  4453. 2:57:13different parts of the data set of
  4454. 2:57:14course so let's adjust that so one
  4455. 2:57:17particular particularly simple and a
  4456. 2:57:19naive way to do this is we have to make
  4457. 2:57:21sure that we pass in the rank and the
  4458. 2:57:23size to the data
  4459. 2:57:25loader and then when we come up here we
  4460. 2:57:28see that we now take Rank and processes
  4461. 2:57:29and we save them now the current
  4462. 2:57:32position will not be zero uh because
  4463. 2:57:35what we want is we want to stride out
  4464. 2:57:37all the processes so one way to do this
  4465. 2:57:40is we basically take S.B times salt. T
  4466. 2:57:43and then multiply it by the process
  4467. 2:57:46rank so proc process rank 0 will start
  4468. 2:57:49at zero but process rank one now starts
  4469. 2:57:52at B * T process rank two is starts at 2
  4470. 2:57:55* B * D Etc so that is the
  4471. 2:57:59initialization now we still they still
  4472. 2:58:01do this identically but now when we
  4473. 2:58:04advance we don't Advance by B * T we
  4474. 2:58:06advance by B * T times number of
  4475. 2:58:10processes right so basically um the
  4476. 2:58:14total number of tokens that we're um
  4477. 2:58:16consuming is B * T * number processes
  4478. 2:58:19and they all go off to a different Rank
  4479. 2:58:23and the position has to advance by the
  4480. 2:58:24entire
  4481. 2:58:26chunk and then here B * T time uh s. num
  4482. 2:58:30processes + one would be to exceed
  4483. 2:58:33number of tokens then we're going to
  4484. 2:58:35Loop and when we Loop we want to of
  4485. 2:58:37course Loop in the exact same way so we
  4486. 2:58:39sort of like reset back uh so this is
  4487. 2:58:42the simplest change that I can uh find
  4488. 2:58:45for kind of a very simple distributed
  4489. 2:58:47data Lo light and um you can notice that
  4490. 2:58:50if process rank is zero and non
  4491. 2:58:52processes is one then uh the whole thing
  4492. 2:58:54will be identical to what we had before
  4493. 2:58:56but now we can have actually multiple
  4494. 2:58:58processes uh running and this should
  4495. 2:59:00work
  4496. 2:59:01fine um so that's the data loader okay
  4497. 2:59:05so next up once they've all initialized
  4498. 2:59:07the data loader they come here and they
  4499. 2:59:09all create a GPT model uh so we create
  4500. 2:59:13eight GPT models on eight processes but
  4501. 2:59:15because the seeds are fixed here they
  4502. 2:59:17all create the same identical model they
  4503. 2:59:20all move it to the device of their Rank
  4504. 2:59:22and they all compile the model and
  4505. 2:59:25because the models are identical there
  4506. 2:59:26are eight identical compilations
  4507. 2:59:28happening in parallel but that's okay
  4508. 2:59:31now none of this uh changes because that
  4509. 2:59:33is on a per step basis and we're
  4510. 2:59:34currently working kind of within step
  4511. 2:59:36because we need to um just uh all the
  4512. 2:59:39all the changes we're making are kind of
  4513. 2:59:41like a within step
  4514. 2:59:42changes now the important thing here is
  4515. 2:59:44when we construct the M model we
  4516. 2:59:47actually have a bit of work to to do
  4517. 2:59:48here get loits is deprecated so uh
  4518. 2:59:50create
  4519. 2:59:52model we need to actually wrap the model
  4520. 2:59:55into the distributed data parallel
  4521. 2:59:58container so um this is how we wrap the
  4522. 3:00:01model into the DDP container and these
  4523. 3:00:04are the docs for DDP and they're quite
  4524. 3:00:07extensive and there's a lot of caveats
  4525. 3:00:09and a lot of things to be careful with
  4526. 3:00:10because everything complexifies times 10
  4527. 3:00:12when multiple processes are involved but
  4528. 3:00:15roughly speaking this device IDs I
  4529. 3:00:17believe has to be passed in now
  4530. 3:00:18unfortunately the docs for what device
  4531. 3:00:20IDs is is is extremely unclear uh so
  4532. 3:00:24when you actually like come here this
  4533. 3:00:26comment for what device IDs is is
  4534. 3:00:29roughly
  4535. 3:00:30nonsensical um but I'm pretty sure it's
  4536. 3:00:33supposed to be the DDP local rank so not
  4537. 3:00:35the DDP rank the local rank uh so this
  4538. 3:00:39is what you pass in here this wraps the
  4539. 3:00:41model and in particular what DDP does
  4540. 3:00:43for you is in a forward pass it actually
  4541. 3:00:45behaves identically so um my
  4542. 3:00:48understanding of it is nothing should be
  4543. 3:00:49changed in the forward pass but in the
  4544. 3:00:51backward pass as you are doing the
  4545. 3:00:53backward pass um in the simpl setting
  4546. 3:00:56once the backp passes over on each
  4547. 3:00:59independent GPU each independent GPU has
  4548. 3:01:02the gradient for all the parameters and
  4549. 3:01:05what DDP does for you is once the
  4550. 3:01:06backward pass is over it will call
  4551. 3:01:09what's called all reduce and it
  4552. 3:01:11basically does an average across all the
  4553. 3:01:14uh ranks of their gradients and and then
  4554. 3:01:18it will deposit that average on every
  4555. 3:01:20single rank so every sing Single rank
  4556. 3:01:22will end up with the average on it and
  4557. 3:01:25so basically that's the communication it
  4558. 3:01:27just synchronizes and averages the
  4559. 3:01:28gradients and that's what DDP offers you
  4560. 3:01:31now DDP actually is a little bit more um
  4561. 3:01:34it is a little bit more involved than
  4562. 3:01:35that because as you are doing the
  4563. 3:01:37backward pass through the layers of the
  4564. 3:01:38Transformer it actually can dispatch
  4565. 3:01:41Communications for the gradient while
  4566. 3:01:43the backward pass is still happening so
  4567. 3:01:45there's overlap of the uh communication
  4568. 3:01:47of the gradient and the synchronization
  4569. 3:01:48of them and uh the backward pass and uh
  4570. 3:01:52this is just more efficient and um uh to
  4571. 3:01:55do it that way so that's what DDP does
  4572. 3:01:57for you um forward is unchanged and
  4573. 3:02:00backward is mostly unchanged and we're
  4574. 3:02:02tacking on this average as we'll see in
  4575. 3:02:04a bit okay so now let's go to the uh
  4576. 3:02:08optimization nothing here changes let's
  4577. 3:02:11go to the optimization here the inner
  4578. 3:02:12loop and think through the
  4579. 3:02:13synchronization of uh these gradients in
  4580. 3:02:15the DP so basically by default what
  4581. 3:02:18happens as I mentioned is when you do l.
  4582. 3:02:20backward here it will do the backward
  4583. 3:02:22pass and then it will synchronize the
  4584. 3:02:24gradients um the problem here is because
  4585. 3:02:28of the gradient accumulation steps Loop
  4586. 3:02:30here we don't actually want to do the
  4587. 3:02:33synchronization after every single La
  4588. 3:02:35step backward because we are just
  4589. 3:02:37depositing gradients and we're doing
  4590. 3:02:39that serially and we just want them
  4591. 3:02:40adding up and we don't want to
  4592. 3:02:42synchronize every single time that would
  4593. 3:02:44be extremely wasteful so basically we
  4594. 3:02:46want to add them up and then on the the
  4595. 3:02:48very last uh it's only on the very last
  4596. 3:02:50step when micro when micro step becomes
  4597. 3:02:53gratak steps minus one only at that last
  4598. 3:02:55step do we want to actually do the
  4599. 3:02:58alberu uh to average up the gradients so
  4600. 3:03:02to do that we come here and um the
  4601. 3:03:05official sanctioned way by the way is to
  4602. 3:03:07do this no sync context manager so
  4603. 3:03:10pytorch says this is a context manager
  4604. 3:03:13to disable gradient synchronization
  4605. 3:03:14across DDP processes So within this
  4606. 3:03:17context gradient will be
  4607. 3:03:19accumulated and basically when you do no
  4608. 3:03:21sync there will be no communication so
  4609. 3:03:24they are telling us to do with DDP no
  4610. 3:03:26sync uh do the gradient accumulation
  4611. 3:03:29accumulate grats and then they are
  4612. 3:03:30asking us to do DDP again with another
  4613. 3:03:32input and that backward and I just
  4614. 3:03:35really don't love this I I just really
  4615. 3:03:37don't like it uh the fact that you have
  4616. 3:03:39to copy paste your code here and use a
  4617. 3:03:40context manager and this is just super
  4618. 3:03:42ugly so when I went to this source code
  4619. 3:03:45here you can see that when you enter
  4620. 3:03:48you simply toggle this variable this
  4621. 3:03:51require backward grat sync and this is
  4622. 3:03:54uh being toggled around and changed and
  4623. 3:03:58this is the variable that basically uh
  4624. 3:04:01if you step through it is being toggled
  4625. 3:04:03to determine if the gradient is going to
  4626. 3:04:05be synchronized so I actually just kind
  4627. 3:04:07of like to use that directly uh so
  4628. 3:04:10instead what I like to do is the
  4629. 3:04:13following right here before the L back
  4630. 3:04:15backward if we are using the DDP then um
  4631. 3:04:20then basically we only want to
  4632. 3:04:23synchronize we only want this variable
  4633. 3:04:25to be true when it is the final
  4634. 3:04:28iteration in all the other iterations
  4635. 3:04:31inside the micr steps we want to be
  4636. 3:04:33false so I just toggle it like this so
  4637. 3:04:36required backward graph sync should only
  4638. 3:04:38turn on when the micro step is the last
  4639. 3:04:41step and so I'm toggling this variable
  4640. 3:04:44directly and I hope that that impacts
  4641. 3:04:47last St backwards
  4642. 3:04:48and this is a naughty thing to do
  4643. 3:04:49because you know they could probably
  4644. 3:04:51change the DDP and this variable will go
  4645. 3:04:53away but for now I believe this this
  4646. 3:04:55works and it allows me to avoid the use
  4647. 3:04:57of context managers and code duplication
  4648. 3:05:00I'm just toggling the variable and then
  4649. 3:05:01Lop backward will not synchronize most
  4650. 3:05:03of the steps and it will synchronize the
  4651. 3:05:04very last step and so once this is over
  4652. 3:05:08uh and we come out every single um rank
  4653. 3:05:13will suddenly magically have the average
  4654. 3:05:17of all the gradients that were stored on
  4655. 3:05:20all the ranks so now we have to think
  4656. 3:05:22through whether that is what we want and
  4657. 3:05:24also um if this suffices and whether how
  4658. 3:05:29it works with the loss and what is loss
  4659. 3:05:31AUM so let's think through through that
  4660. 3:05:33now and the problem I'm getting at is
  4661. 3:05:35that we've averaged the gradients which
  4662. 3:05:37is great but the loss AUM has not been
  4663. 3:05:40impacted yet and the and this is outside
  4664. 3:05:43of the DDP container so that is not
  4665. 3:05:45being averaged um and so here when when
  4666. 3:05:47we are printing Los AUM well presumably
  4667. 3:05:49we're only going to be printing on the
  4668. 3:05:51master process uh rank zero and it's
  4669. 3:05:53just going to be printing the losses
  4670. 3:05:55that it saw on its process but instead
  4671. 3:05:57we want it to print the loss over all
  4672. 3:06:00the processes and the average of that
  4673. 3:06:02loss because we did average of gradients
  4674. 3:06:04so we want the average of loss as well
  4675. 3:06:06so simply here after this uh this is the
  4676. 3:06:09code that I've used in the past um and
  4677. 3:06:13instead of LF we want
  4678. 3:06:15Lum so if
  4679. 3:06:18DDP again then this is a p torch
  4680. 3:06:22distributed I import it where do I
  4681. 3:06:24import
  4682. 3:06:26it uh oh gosh so this file is starting
  4683. 3:06:30to get out of control huh so if uh so
  4684. 3:06:33import torch. distributed as dist
  4685. 3:06:36so dist.
  4686. 3:06:38ALU and we're doing the average on Lum
  4687. 3:06:42and so this lakum tensor exists on all
  4688. 3:06:44the ranks when we call all use of
  4689. 3:06:46average it creates the average of those
  4690. 3:06:48numbers and it deposits that average on
  4691. 3:06:51all the ranks so all the ranks after
  4692. 3:06:53this um call will now contain L AUM uh
  4693. 3:06:57averaged up and so when we print here on
  4694. 3:07:00the master process the L AUM is
  4695. 3:07:02identical in all the other ranks as well
  4696. 3:07:04so here if Master process
  4697. 3:07:07oops we want to print like this okay and
  4698. 3:07:10finally we have to be careful because
  4699. 3:07:12we're not processing even more tokens so
  4700. 3:07:15times DDP World size
  4701. 3:07:18that's number of tokens that we've
  4702. 3:07:19processed up
  4703. 3:07:21above
  4704. 3:07:24and everything else should be fine uh
  4705. 3:07:27the only other thing to be careful with
  4706. 3:07:29is as I mentioned you want to destroy
  4707. 3:07:31the process group so that we are nice to
  4708. 3:07:33nickel and it's not going to uh to uh to
  4709. 3:07:35DDP and it's not going to complain to us
  4710. 3:07:38uh when we exit
  4711. 3:07:40here so that should be it let's try to
  4712. 3:07:43take it for a spin okay so I launched
  4713. 3:07:44the script and it should be uh printing
  4714. 3:07:46here imminently we're now training with
  4715. 3:07:488 gpus at the same time so the gradient
  4716. 3:07:51accumulation steps is not 32 it is now
  4717. 3:07:53divide 8 and it's just four uh so um
  4718. 3:07:58otherwise this is what the optimization
  4719. 3:07:59now looks like and wow we're going
  4720. 3:08:01really fast so we're processing 1.5
  4721. 3:08:04million tokens uh per second now so
  4722. 3:08:09these are some serious numbers and the
  4723. 3:08:11tiny shakespare data set is so tiny that
  4724. 3:08:12we're just doing like so many Epoch over
  4725. 3:08:15it most likely but this is roughly what
  4726. 3:08:17looks like um one thing that I had to
  4727. 3:08:20fix by the way is that this was model.
  4728. 3:08:23configure optimizers which Now doesn't
  4729. 3:08:25work because model now is a DDP model so
  4730. 3:08:27instead this has to become raw
  4731. 3:08:29model. configure optimizers where raw
  4732. 3:08:32model is something I create here so
  4733. 3:08:35right after I wrap the model into DDP uh
  4734. 3:08:38I have to create the raw model which in
  4735. 3:08:40the case of DDP is a model. module is
  4736. 3:08:43where it stores the raw and then module
  4737. 3:08:46of gpt2 as we have it which contains the
  4738. 3:08:49uh configure optimizers function that we
  4739. 3:08:51want to call so that's one thing that I
  4740. 3:08:53have to fix otherwise this seems to run
  4741. 3:08:56now one thing you'll notice is that when
  4742. 3:08:57you actually compare this run and the
  4743. 3:08:59numbers in it to the just running a
  4744. 3:09:01single GPU you'll notice that this is
  4745. 3:09:04single GPU run with 32 gratum the
  4746. 3:09:06numbers won't exactly match
  4747. 3:09:09up and uh that's kind of a boring reason
  4748. 3:09:11for why that happens uh the reason for
  4749. 3:09:13that is that in the data loader we're
  4750. 3:09:15basically just iterating through batches
  4751. 3:09:17and slightly different way because now
  4752. 3:09:18we're looking for an entire page of data
  4753. 3:09:21and if that page uh for all the gpus if
  4754. 3:09:24that chunk exceeds the number of tokens
  4755. 3:09:26we just Loop and so actually the single
  4756. 3:09:29GPU and the H GPU process will end up um
  4757. 3:09:33resetting in a slightly different Manner
  4758. 3:09:35and so our batches are slightly
  4759. 3:09:36different and so we get slightly
  4760. 3:09:38different numbers but one way to
  4761. 3:09:39convince yourself that this is okay it
  4762. 3:09:42just make the total batch size much
  4763. 3:09:43smaller and the b and a t and then um
  4764. 3:09:48so I think I used uh 4 * 124 * 8 so I
  4765. 3:09:52used 32768 as a total patch size and
  4766. 3:09:55then um so I made sure that the single
  4767. 3:09:57GPU will do eight creting accumulation
  4768. 3:10:00steps and then the multi-gpu and then
  4769. 3:10:02you're reducing the boundary effects of
  4770. 3:10:04the data loader and you'll see that the
  4771. 3:10:06numbers match up so long story short
  4772. 3:10:08we're now going really really fast the
  4773. 3:10:10optimization is mostly consistent with
  4774. 3:10:12gpt2 and three hyper parameters and uh
  4775. 3:10:16we have outgrown our tiny Shakespeare
  4776. 3:10:18file and we want to upgrade it so let's
  4777. 3:10:20move to next to that next so let's now
  4778. 3:10:22take a look at what data sets were used
  4779. 3:10:23by gpt2 and gpt3 so gbt2 used this web
  4780. 3:10:27Text data set that was never released um
  4781. 3:10:30there's an attempt at reproducing it
  4782. 3:10:32called open web text uh so basically
  4783. 3:10:34roughly speaking what they say here in
  4784. 3:10:35the paper is that they scraped all
  4785. 3:10:37outbound links from Reddit and then uh
  4786. 3:10:41with at least three Karma and that was
  4787. 3:10:43kind of like their starting point and
  4788. 3:10:44they collected all the web P all the web
  4789. 3:10:45pages and all the text in them and so
  4790. 3:10:48this was 45 million links and this ended
  4791. 3:10:50up being 40 GB of text so uh so that's
  4792. 3:10:54roughly what gpt2 says about its data
  4793. 3:10:57set so it's basically outbound links
  4794. 3:10:58from Reddit now when we go over to gpt3
  4795. 3:11:01there's a training data set section and
  4796. 3:11:03that's where they start to talk about um
  4797. 3:11:05common coll which is a lot more uh used
  4798. 3:11:09actually I think even gpt2 talked about
  4799. 3:11:11common coll um but basically it's not a
  4800. 3:11:14very high quality data set all by itself
  4801. 3:11:16because it is extremely noisy this is a
  4802. 3:11:18completely random subset of the internet
  4803. 3:11:20and it's much worse than you think so
  4804. 3:11:22people go into Great Lengths to filter
  4805. 3:11:24common craw because there's good stuff
  4806. 3:11:26in it but most of it is just like ad
  4807. 3:11:27spam random tables and numbers and stock
  4808. 3:11:30tickers and uh it's just total mess
  4809. 3:11:35so that's why people like to train on
  4810. 3:11:38these data mixtures that they curate and
  4811. 3:11:41uh are careful with so a large chunk of
  4812. 3:11:44these data mixtures typically will be
  4813. 3:11:45common C like for example 50% of the
  4814. 3:11:47tokens will be comic but then here in
  4815. 3:11:50gpt3 they're also using web text to from
  4816. 3:11:52before so that's Reddit outbound but
  4817. 3:11:54they're also adding for example books
  4818. 3:11:56and they're adding Wikipedia there's
  4819. 3:11:58many other things you can decide to add
  4820. 3:12:00now this data set for gpt3 was also
  4821. 3:12:02never released so today some of the data
  4822. 3:12:05sets that I'm familiar with that are
  4823. 3:12:06quite good and would be representative
  4824. 3:12:08of something along these lines are
  4825. 3:12:10number one the red pajama data set or
  4826. 3:12:12more specifically for example the slim
  4827. 3:12:14pajama subset of the red pajama data set
  4828. 3:12:17which is a cleaned and D duplicated
  4829. 3:12:19version of it and just to give you a
  4830. 3:12:21sense again it's a bunch of common crawl
  4831. 3:12:24um C4 which is also as far as I know
  4832. 3:12:27more common craw but processed
  4833. 3:12:28differently and then we have GitHub
  4834. 3:12:30books archive Wikipedia stack exchange
  4835. 3:12:33these are the kinds of data sets that
  4836. 3:12:35would go into these data mixtures now
  4837. 3:12:37specifically the one that I like that
  4838. 3:12:38came out recently is called Fine web
  4839. 3:12:41data set uh so this is an attempt to
  4840. 3:12:43basically collect really high quality
  4841. 3:12:45common coll data and filter it in this
  4842. 3:12:48case to 15 trillion tokens and then in
  4843. 3:12:51addition to that more recently
  4844. 3:12:52huggingface released this fine web edu
  4845. 3:12:55subset which is 1.3 trillion of
  4846. 3:12:58educational and 5.4 trillion of high
  4847. 3:13:01educational content so basically they're
  4848. 3:13:03trying to filter common C to very high
  4849. 3:13:06quality educational subsets and uh this
  4850. 3:13:09is the one that we will use there's a
  4851. 3:13:11long uh web page here on fine web and
  4852. 3:13:14they go into a ton of detail about how
  4853. 3:13:16they process the data which is really
  4854. 3:13:17fascinating reading by the way and I
  4855. 3:13:19would definitely recommend if you're
  4856. 3:13:20interested into Data mixtures and so on
  4857. 3:13:22and how data gets processed at these
  4858. 3:13:24scales a look at this uh page and more
  4859. 3:13:27specifically we'll be working with the
  4860. 3:13:28fine web edu I think and it's basically
  4861. 3:13:32educational content from the
  4862. 3:13:34internet uh they show that training on
  4863. 3:13:36educational content in in their metrics
  4864. 3:13:39um uh works really really well and we're
  4865. 3:13:43going to use this sample 10 billion
  4866. 3:13:46tokens subsample of it because we're not
  4867. 3:13:49going to be training on trillions of
  4868. 3:13:50tokens uh we're just going to train on
  4869. 3:13:52uh 10 billion sample of the fine web edu
  4870. 3:13:56because empirically in my previous few
  4871. 3:13:58experiments this actually suffices to
  4872. 3:14:00really get close to gpt2 Performance and
  4873. 3:14:02it's um simple enough to work with and
  4874. 3:14:04so let's work with the sample 10 uh BT
  4875. 3:14:07so our goal will be to download it
  4876. 3:14:10process it and make sure that our data
  4877. 3:14:12loader can work with it so let's get to
  4878. 3:14:15that okay so I introduced another um
  4879. 3:14:18file here that will basically download
  4880. 3:14:21Fine web edu from huging face data sets
  4881. 3:14:24it will pre-process and pre- tokenize
  4882. 3:14:26all of the data and it will save data
  4883. 3:14:28shards to a uh folder on um local disk
  4884. 3:14:34and so while this is running uh just
  4885. 3:14:38wanted to briefly mention that you can
  4886. 3:14:40kind of look through the data set viewer
  4887. 3:14:41here just to get a sense of what's in
  4888. 3:14:43here and it's kind of interesting I mean
  4889. 3:14:45it's a it basically looks like it's
  4890. 3:14:47working fairly well like it's talking
  4891. 3:14:48about nuclear energy in France it's
  4892. 3:14:51talking
  4893. 3:14:52about Mexican
  4894. 3:14:54America some mac PJs Etc so actually it
  4895. 3:14:58seems like their filters are working
  4896. 3:14:59pretty well uh the filters here by the
  4897. 3:15:01way were applied automatically using um
  4898. 3:15:04llama 370b I believe and so uh basically
  4899. 3:15:08llms are judging which content is
  4900. 3:15:10educational and that ends up making it
  4901. 3:15:11through the filter uh so that's pretty
  4902. 3:15:13cool now in terms of the script itself
  4903. 3:15:16I'm not going to go through the full
  4904. 3:15:17script because it's not as interesting
  4905. 3:15:19and not as llm Centric but when you run
  4906. 3:15:22this basically number one we're going to
  4907. 3:15:24load the data set uh which this is all
  4908. 3:15:26huging face code running this you're
  4909. 3:15:28going to need to uh pip install data
  4910. 3:15:31sets um so it's downloading the data set
  4911. 3:15:35then it is tokenizing all of the
  4912. 3:15:37documents inside this data set now when
  4913. 3:15:39we tokenize the documents you'll notice
  4914. 3:15:42that um to tokenize a single document uh
  4915. 3:15:46we first
  4916. 3:15:47start the tokens with the end of text
  4917. 3:15:49token and this is a special token in the
  4918. 3:15:51gpt2 tokenizer as you know so
  4919. 3:15:5450256 is the ID of the end of text and
  4920. 3:15:57this is what begins a document even
  4921. 3:15:59though it's called end of text but this
  4922. 3:16:01is uh the first token that begins a
  4923. 3:16:03document then we extend with all of the
  4924. 3:16:06tokens of that document then we create a
  4925. 3:16:08numpy array out of that we make sure
  4926. 3:16:11that all the tokens are between
  4927. 3:16:14oh okay let me debug this
  4928. 3:16:17okay so apologies for that uh it just
  4929. 3:16:19had to do with me using a float division
  4930. 3:16:21in Python it must be integer division so
  4931. 3:16:23that this is an INT and everything is
  4932. 3:16:25nice um okay but basically the
  4933. 3:16:28tokenization here is relatively
  4934. 3:16:29straightforward returns tokens in mp.
  4935. 3:16:32un6 uh we're using .16 to save a little
  4936. 3:16:35bit of space because 2 to the 16us 1 is
  4937. 3:16:3965,000 so the gpt2 max token ID is well
  4938. 3:16:43below that and then here there's a bunch
  4939. 3:16:45of multiprocessing code and it's
  4940. 3:16:47honestly not that exciting so I'm not
  4941. 3:16:48going to step through it but we're
  4942. 3:16:50loading the data set we're tokenizing it
  4943. 3:16:52and we're saving everything to shards
  4944. 3:16:55and the shards are numpy files uh so
  4945. 3:16:58just storing a numpy array and uh which
  4946. 3:17:01is very very similar to torch
  4947. 3:17:03tensors and the first Shard 0000 is a
  4948. 3:17:07Val a validation Shard and all the other
  4949. 3:17:09shards are uh training shards and as I
  4950. 3:17:12mentioned they all have 100 million
  4951. 3:17:14tokens in them exactly um and and that
  4952. 3:17:17just makes it easier to work with as to
  4953. 3:17:20Shard the files because if we just have
  4954. 3:17:22a single massive file sometimes they can
  4955. 3:17:24be hard to work with on the disk and so
  4956. 3:17:26sharting it is just kind of um nicer
  4957. 3:17:28from that
  4958. 3:17:30perspective and uh yeah so we'll just
  4959. 3:17:32let this run this will be probably um
  4960. 3:17:3630ish minutes or so and then we're going
  4961. 3:17:38to come back to actually train on this
  4962. 3:17:39data and we're going to be actually
  4963. 3:17:41doing some legit pre-training in this
  4964. 3:17:42case this is a good data set we're doing
  4965. 3:17:45lots of tokens per second we have 8 gpus
  4966. 3:17:48the code is ready and so we're actually
  4967. 3:17:50going to be doing a serious training run
  4968. 3:17:52so let's get P it back in a bit okay so
  4969. 3:17:54we're back so uh if we LS edu fine web
  4970. 3:17:58we see that there's now 100 charts in it
  4971. 3:18:02um and that makes sense because each
  4972. 3:18:03chart is 100 million tokens so 100
  4973. 3:18:06charts of that is 10 billion tokens in
  4974. 3:18:08total now swinging over to the main file
  4975. 3:18:11I made some adjustments to our data
  4976. 3:18:12loader again and that's because we're
  4977. 3:18:14not running with uh Shakespeare anymore
  4978. 3:18:17we want to use the fine web shards and
  4979. 3:18:20so you'll see some code here that
  4980. 3:18:21additionally basically can load these
  4981. 3:18:23shards uh we load the um un6 numpy file
  4982. 3:18:28we convert it to a torch. long tensor
  4983. 3:18:30which is what a lot of the layers up top
  4984. 3:18:32expect by default and then here we're
  4985. 3:18:35just enumerating all the shards I also
  4986. 3:18:38added a split to data load of light so
  4987. 3:18:40we can uh load the split train but also
  4988. 3:18:42the split Val uh the zero
  4989. 3:18:44split and then we can load the shards
  4990. 3:18:47and then here we also have not just the
  4991. 3:18:49current position now but also the
  4992. 3:18:51current Shard so we have a position
  4993. 3:18:53inside A Shard and then when we uh run
  4994. 3:18:55out of tokens in A Single Shard we first
  4995. 3:18:58Advance The Shard and loop if we need to
  4996. 3:19:01and then we get the tokens and readjust
  4997. 3:19:03the position so this data loader will
  4998. 3:19:06now iterate all the shards as well so I
  4999. 3:19:09Chang that and then the other thing that
  5000. 3:19:11I did while uh the data was processing
  5001. 3:19:14is our train loader now has split train
  5002. 3:19:17of course and down here I set up some I
  5003. 3:19:20set up some numbers
  5004. 3:19:21so we are doing 2 to the
  5005. 3:19:249 uh tokens per uh per um per step and
  5006. 3:19:31we want to do roughly 10 billion tokens
  5007. 3:19:35um because that's how many unique tokens
  5008. 3:19:36we have so if we did 10 billion tokens
  5009. 3:19:39then divide that by 29 we see that this
  5010. 3:19:41is 1973 steps so that's where that's
  5011. 3:19:44from and then the GPT three paper says
  5012. 3:19:47that they warm up the learning rate over
  5013. 3:19:49375 million tokens so I came here and
  5014. 3:19:53375 E6 tokens divide uh 2 to the
  5015. 3:19:5719 is 715 steps so that's why warm-up
  5016. 3:20:01steps is set to 715 so this will exactly
  5017. 3:20:04match um the warm-up schedule that gpt3
  5018. 3:20:07used and I think 715 by the way is very
  5019. 3:20:10uh mild and this could be made
  5020. 3:20:12significantly more aggressive probably
  5021. 3:20:13even like 100 is good enough um
  5022. 3:20:17but it's okay let's leave it for now so
  5023. 3:20:18that we have the exact hyper parameters
  5024. 3:20:20of gpt3 so I fix that and then um that's
  5025. 3:20:25pretty much it we can we can run so we
  5026. 3:20:28have our script
  5027. 3:20:29here and we can
  5028. 3:20:32launch and actually sorry let me do one
  5029. 3:20:34more
  5030. 3:20:38thing excuse
  5031. 3:20:40me for my GPU I can actually fit more
  5032. 3:20:43batch size and I believe I can fat I can
  5033. 3:20:45fit 60 4 on my GPU as a micro bash size
  5034. 3:20:50so let me try
  5035. 3:20:54that I could be misremembering but that
  5036. 3:20:57means 64 * 124 per GPU and then we have
  5037. 3:21:00a gpus so that means we would not even
  5038. 3:21:02be doing gradient accumulation if this
  5039. 3:21:04fits because uh this just multi
  5040. 3:21:06multiplies out to uh the full total bat
  5041. 3:21:09size so no gradient
  5042. 3:21:12accumulation and that would run pretty
  5043. 3:21:14quickly if that fits
  5044. 3:21:26let's go let's go I mean if this works
  5045. 3:21:29then this is basically a serious
  5046. 3:21:31pre-training run um we're not logging
  5047. 3:21:34we're not evaluating the validation
  5048. 3:21:35split we're not running any evaluations
  5049. 3:21:37yet so it's not we haven't crossed our
  5050. 3:21:39te's and dotted our eyes but uh if we
  5051. 3:21:42let this run for a while we're going to
  5052. 3:21:44actually get a pretty good model and the
  5053. 3:21:46model that might even be on par with or
  5054. 3:21:49better than gpt2 124 M okay so it looks
  5055. 3:21:54like everything is going great we're
  5056. 3:21:55processing 1.5 million tokens per
  5057. 3:21:58second uh everything here looks good
  5058. 3:22:03we're doing 330 milliseconds per
  5059. 3:22:06iteration and we have to do a total
  5060. 3:22:09of uh where are we printing that 1973 so
  5061. 3:22:1319073 times 0.33
  5062. 3:22:17is this many seconds this many minutes
  5063. 3:22:20so this will run for 1.7
  5064. 3:22:24hours uh so one and a half hour run uh
  5065. 3:22:28like this and uh we don't even have to
  5066. 3:22:30use gradient accumulation which is nice
  5067. 3:22:31and you might not have that luxury in
  5068. 3:22:33your GPU in that case just start
  5069. 3:22:35decreasing the batch size until things
  5070. 3:22:37fit but keep it to nice
  5071. 3:22:39numbers um so that's pretty exciting
  5072. 3:22:42we're currently warming up the learning
  5073. 3:22:43rate so you see that it's still very low
  5074. 3:22:45one4 so this will ramp up over the next
  5075. 3:22:48few steps all the way to 6 e
  5076. 3:22:50Nega uh 4
  5077. 3:22:53here very cool so now what I'd like to
  5078. 3:22:56do is uh let's cross the T and do our
  5079. 3:22:58eyes let's evaluate on the validation
  5080. 3:23:00split and let's try to figure out how we
  5081. 3:23:02can run evals how we can do logging how
  5082. 3:23:05we can visualize our losses and all the
  5083. 3:23:07good stuff so let's get to that before
  5084. 3:23:09we actually do the run okay so I've
  5085. 3:23:11adjusted the code so that we're
  5086. 3:23:13evaluating on the validation split so
  5087. 3:23:15creating the Val loader just by passing
  5088. 3:23:17in Split equals Val that will basically
  5089. 3:23:19create a data loader just for the uh
  5090. 3:23:21validation
  5091. 3:23:22Shard um the other thing I did is in the
  5092. 3:23:25data loader I introduced a new function
  5093. 3:23:27reset which is called at init and it
  5094. 3:23:29basically resets the data loader and
  5095. 3:23:31that is very useful because when we come
  5096. 3:23:34to the main training Loop now so this is
  5097. 3:23:37the code that I've added and basically
  5098. 3:23:39every 100th iteration including the
  5099. 3:23:41zeroth iteration we put the model into
  5100. 3:23:44evaluation mode we reset the Val loader
  5101. 3:23:47and then um no gradients involved we're
  5102. 3:23:50going to
  5103. 3:23:52basically accumulate the gradients over
  5104. 3:23:54say 20 steps and then average it all up
  5105. 3:23:58and print out the validation loss and so
  5106. 3:24:01that basically is the exact same logic
  5107. 3:24:03as the training Loop roughly but there's
  5108. 3:24:06no loss that backward it's only
  5109. 3:24:07inference we're just measuring the loss
  5110. 3:24:09we're adding it up everything else
  5111. 3:24:11otherwise applies and is exactly as
  5112. 3:24:13we've seen it before and so this will
  5113. 3:24:15print the validation laws
  5114. 3:24:16um every 100th iteration including on
  5115. 3:24:19the very first
  5116. 3:24:20iteration uh so that's nice that will
  5117. 3:24:23tell us some amount some a little bit
  5118. 3:24:25about how much we're overfitting that
  5119. 3:24:27said like uh we have roughly Infinity
  5120. 3:24:29data so we're mostly expecting our train
  5121. 3:24:31and Val loss to be about the same but
  5122. 3:24:33the other reason I'm kind of interested
  5123. 3:24:35in this is because we can take the GPT
  5124. 3:24:362124m as openi released it we can
  5125. 3:24:39initialize from it and we can basically
  5126. 3:24:41see what kind of loss it achieves on the
  5127. 3:24:43validation loss as well and that gives
  5128. 3:24:45us kind of an indication as to uh how
  5129. 3:24:47much that model would generalize to 124
  5130. 3:24:49M but it's not an sorry to fine web edu
  5131. 3:24:52validation split that said it's not a
  5132. 3:24:55super fair comparison to gpt2 because it
  5133. 3:24:57was trained on a very different data
  5134. 3:24:58distribution but it's still kind of like
  5135. 3:25:00an interesting data point and in any
  5136. 3:25:02case you would always want to have a
  5137. 3:25:03validation split in a training run like
  5138. 3:25:06this so that you can make sure that you
  5139. 3:25:08are not um overfitting and this is
  5140. 3:25:11especially a concern if we were to make
  5141. 3:25:13more Epoch in our training data um so
  5142. 3:25:16for example right now we're just doing a
  5143. 3:25:18single Epoch but if we get to a point
  5144. 3:25:20where we want to train on 10 epochs or
  5145. 3:25:21something like that we would be really
  5146. 3:25:23careful with maybe we are memorizing
  5147. 3:25:26that data too much if we have a big
  5148. 3:25:28enough model and our validation split
  5149. 3:25:30would be one way to tell whether that is
  5150. 3:25:32happening okay and in addition to that
  5151. 3:25:34if you remember at bottom of our script
  5152. 3:25:36we had all of this orphaned code for
  5153. 3:25:37sampling from way back when so I deleted
  5154. 3:25:40that code and I moved it up um to here
  5155. 3:25:43so once in a while we simply value
  5156. 3:25:45validation
  5157. 3:25:46once in a while we sample we generate
  5158. 3:25:49samples and then uh we do that only
  5159. 3:25:52every 100 steps and we train on every
  5160. 3:25:55single step so that's how I have a
  5161. 3:25:56structure right now and I've been
  5162. 3:25:58running this for 10,000 iterations so
  5163. 3:26:00here are some samples on neration
  5164. 3:26:021,000
  5165. 3:26:05um hello I'm a language model and I'm
  5166. 3:26:07not able to get more
  5167. 3:26:09creative I'm a language model and
  5168. 3:26:10languages file you're learning about
  5169. 3:26:12here is or is the beginning of a
  5170. 3:26:14computer
  5171. 3:26:16okay so this is all like pretty uh this
  5172. 3:26:19is still a garble uh but we're only at
  5173. 3:26:21ration 1,000 and we've only just barely
  5174. 3:26:24reached maximum learning rate uh so this
  5175. 3:26:26is still learning uh we're about to get
  5176. 3:26:28some more samples coming up in
  5177. 3:26:321,00 okay
  5178. 3:26:35um okay this is you know the model is
  5179. 3:26:38still is still a young baby okay so uh
  5180. 3:26:42basically all of this sampling code that
  5181. 3:26:44I've put here everything should be
  5182. 3:26:45familiar with to you and came from
  5183. 3:26:47before the only thing that I did is I
  5184. 3:26:49created a generator object in pytorch so
  5185. 3:26:52that I have a direct control over the
  5186. 3:26:54sampling of the random numbers don't
  5187. 3:26:56because I don't want to impact the RNG
  5188. 3:26:58state of the random number generator
  5189. 3:27:00that is the global one used for training
  5190. 3:27:02I want this to be completely outside of
  5191. 3:27:04the training Loop and so I'm using a
  5192. 3:27:07special sampling RNG and then I make
  5193. 3:27:09sure to seed it that every single rank
  5194. 3:27:12has a different seed and then I pass in
  5195. 3:27:14here where we sort of consumer in the
  5196. 3:27:17numbers in multinomial where the
  5197. 3:27:18sampling happens I make sure to pass in
  5198. 3:27:20the generator object there otherwise
  5199. 3:27:22this is identical uh now the other thing
  5200. 3:27:25is um you'll notice that we're running a
  5201. 3:27:27bit slower that's because I actually had
  5202. 3:27:29to disable torch. compile to get this to
  5203. 3:27:32sample and um so we're running a bit
  5204. 3:27:34slower so for some reason it works with
  5205. 3:27:36no torch compile but when I torch
  5206. 3:27:37compile my model I get a really scary
  5207. 3:27:39error from pytorch and I have no idea
  5208. 3:27:41how to resolve it right now so probably
  5209. 3:27:43by the time you see this code released
  5210. 3:27:45or something like that maybe it's fixed
  5211. 3:27:47but for now I'm just going to do end
  5212. 3:27:49false um and I'm going to bring back
  5213. 3:27:51toor compile and you're not going to get
  5214. 3:27:54samples and I I think I'll fix this
  5215. 3:27:56later uh by the way um I will be
  5216. 3:27:59releasing all this code and actually
  5217. 3:28:01I've been very careful about making get
  5218. 3:28:03commits every time we add something and
  5219. 3:28:05so I'm going to release the entire repo
  5220. 3:28:07that starts completely from scratch all
  5221. 3:28:09the way to uh now and after this as well
  5222. 3:28:12and so everything should be exactly
  5223. 3:28:13documented in the git commit history um
  5224. 3:28:16um and so I think that will be nice so
  5225. 3:28:19hopefully by the time you go to GitHub
  5226. 3:28:20uh this is removed and it's working and
  5227. 3:28:22I will have fixed the bug okay so I have
  5228. 3:28:24the optimization running here and it's
  5229. 3:28:26stepping and we're on step 6,000 or so
  5230. 3:28:28so we're about 30% through training now
  5231. 3:28:31while this is training I would like to
  5232. 3:28:32introduce one evaluation that we're
  5233. 3:28:34going to use to supplement the
  5234. 3:28:35validation set and that is the H swag
  5235. 3:28:38eval so hos swag comes from this paper
  5236. 3:28:42back in 2019 so it's a 5-year-old eval
  5237. 3:28:44now and the way H swag works is there is
  5238. 3:28:47basically a sentence completion data set
  5239. 3:28:50so it's a multiple choice for every one
  5240. 3:28:52of these questions we have uh basically
  5241. 3:28:54a shared context like a woman is outside
  5242. 3:28:57with a bucket and a dog the dog is
  5243. 3:28:59running around trying to avoid bath she
  5244. 3:29:02a Rises the bucket off with soap and
  5245. 3:29:04blow dry the dog's head B uses a hose to
  5246. 3:29:08keep it from getting soapy C gets the
  5247. 3:29:11dog wet and it runs away again or D gets
  5248. 3:29:14into a bathtub with the dog
  5249. 3:29:16and so basically the idea is that these
  5250. 3:29:19multiple choice are constructed so that
  5251. 3:29:22one of them is a natural continuation of
  5252. 3:29:25the um sentence and the others are
  5253. 3:29:30not and uh the others might not make
  5254. 3:29:32sense like uses the host to keep it from
  5255. 3:29:34getting soaped that makes no sense and
  5256. 3:29:36so what happens is that models that are
  5257. 3:29:38not trained very well are not able to
  5258. 3:29:40tell these apart but models that have a
  5259. 3:29:43lot of World Knowledge and can tell uh
  5260. 3:29:45which um and can tell a lot about the
  5261. 3:29:48world will be able to create these
  5262. 3:29:50completions and these sentences are
  5263. 3:29:52sourced from activity net and from Wiki
  5264. 3:29:55how and at the bottom of the uh
  5265. 3:30:00paper there's kind of like a cool chart
  5266. 3:30:03of the kinds of domains in Wiki house so
  5267. 3:30:05there's a lot of sentences from
  5268. 3:30:07computers and electronics and Homes and
  5269. 3:30:09Garden and it has kind of a broad
  5270. 3:30:11coverage of the kinds of things you need
  5271. 3:30:13to know about the world in order to find
  5272. 3:30:15the most likely completion and um the
  5273. 3:30:19identity of that of that completion one
  5274. 3:30:22more thing that's kind of interesting
  5275. 3:30:23about H swag is the way it was
  5276. 3:30:25constructed is that the incorrect um
  5277. 3:30:28options are deliberately um
  5278. 3:30:32adversarially sourced so they're not
  5279. 3:30:34just random sentences they're actually
  5280. 3:30:37sentences generated by language models
  5281. 3:30:39and they're generated in such a way that
  5282. 3:30:41language models basically find them
  5283. 3:30:42difficult but humans find them easy and
  5284. 3:30:45so they mentioned that humans have a 95%
  5285. 3:30:47accuracy on this set but at the time the
  5286. 3:30:49state-of-the-art language models had
  5287. 3:30:51only 48% and so at the time this was a
  5288. 3:30:54good Benchmark now you can read the
  5289. 3:30:57details of this paper to to learn more
  5290. 3:30:59um the thing to point out though is that
  5291. 3:31:01this is 5 years ago and since then what
  5292. 3:31:03happened to H swag is that it's been
  5293. 3:31:05totally just uh
  5294. 3:31:08um solved and so now the language models
  5295. 3:31:11here are 96% so basically the 4% the
  5296. 3:31:14last 4% is probably errors in the data
  5297. 3:31:16set or the questions are really really
  5298. 3:31:18hard and so basically this data set is
  5299. 3:31:20kind of crushed with respect to language
  5300. 3:31:22models but back then the best language
  5301. 3:31:23model was only at about 50% uh but this
  5302. 3:31:27is how far things got but still the the
  5303. 3:31:30reason people like H swag and it's not
  5304. 3:31:33used by the way in gpt2 but in gpt3
  5305. 3:31:37there is H swag eval and lots of people
  5306. 3:31:39use H
  5307. 3:31:41swag and so for gpt3 we have results
  5308. 3:31:45here
  5309. 3:31:46that are cited so we know what percent
  5310. 3:31:48accuracies gpt3 um attains at all these
  5311. 3:31:51different model checkpoints for H swag
  5312. 3:31:54eval and the reason people like it is
  5313. 3:31:56because H swag is a smooth eval and it
  5314. 3:31:59is an eval that offers quote unquote
  5315. 3:32:01early signal uh so early signal means
  5316. 3:32:04that even small language models are
  5317. 3:32:06going to start at the random chance of
  5318. 3:32:0825% but they're going to slowly improve
  5319. 3:32:11and you're going to see 25 26 27 Etc and
  5320. 3:32:15uh you can see slow Improvement even
  5321. 3:32:17when the models are very small and it's
  5322. 3:32:19very early so it's smooth it has early
  5323. 3:32:23signal and um it's been around for a
  5324. 3:32:26long time so that's why people kind of
  5325. 3:32:28like this
  5326. 3:32:29eval uh now the way that we're going to
  5327. 3:32:32evaluate this is as
  5328. 3:32:34follows as I mentioned we have a shared
  5329. 3:32:37context and this is kind of like a
  5330. 3:32:39multiple choice task but instead of
  5331. 3:32:41giving the model a multiple choice
  5332. 3:32:42question and asking it for A B C or D uh
  5333. 3:32:46we can't do that because these models
  5334. 3:32:47when they are so small as we are seeing
  5335. 3:32:49here the models can't actually do
  5336. 3:32:51multiple choice they don't understand
  5337. 3:32:53the concept of associating a label to
  5338. 3:32:55one of the options of multiple choice uh
  5339. 3:32:58they don't understand that so we have to
  5340. 3:32:59give it to them in a native form and the
  5341. 3:33:01native form is a token completion so
  5342. 3:33:05here's what we do we construct a batch
  5343. 3:33:06of four rows and uh T tokens whatever
  5344. 3:33:10that t happens to be then the shared
  5345. 3:33:13context that is basically the context
  5346. 3:33:15for the for choices the tokens of that
  5347. 3:33:17are shared across all of the rows and
  5348. 3:33:20then we have the four options so we kind
  5349. 3:33:22of like lay them out and then only one
  5350. 3:33:25of the options is correct in this case
  5351. 3:33:26label three option three and so um this
  5352. 3:33:30is the correct option and option one two
  5353. 3:33:32and for are
  5354. 3:33:33incorrect now these options might be of
  5355. 3:33:36different lengths so what we do is we
  5356. 3:33:38sort of like take the longest length and
  5357. 3:33:40that's the size of the batch B BYT and
  5358. 3:33:42then some of these uh here are going to
  5359. 3:33:45be pded Dimensions so they're going to
  5360. 3:33:47be unused and so we need the tokens we
  5361. 3:33:51need the correct label and we need a
  5362. 3:33:53mask that tells us which tokens are
  5363. 3:33:55active and the mask is then zero for
  5364. 3:33:58these uh padded areas so that's how we
  5365. 3:34:01construct these batches and then in
  5366. 3:34:04order to get the language model to
  5367. 3:34:05predict A B C or D the way this works is
  5368. 3:34:08basically we're just going to look at
  5369. 3:34:10the tokens their probabilities and we're
  5370. 3:34:12going to pick the option that gets the
  5371. 3:34:15lowest or the highest average
  5372. 3:34:18probability for the token so for the
  5373. 3:34:22tokens because that is the most likely
  5374. 3:34:25completion according to the language
  5375. 3:34:27model so we're just going to look at the
  5376. 3:34:29um probabilities here and average them
  5377. 3:34:33up across the options and pick the one
  5378. 3:34:35with the highest probability roughly
  5379. 3:34:38speaking so this is how we're going to
  5380. 3:34:40do H swag
  5381. 3:34:42um and this is I believe also how uh
  5382. 3:34:46gpt3 did it um this is how gpt3 did it
  5383. 3:34:50as far as I know but you should note
  5384. 3:34:52that some of the other evals where you
  5385. 3:34:54might see H swag may not do it this way
  5386. 3:34:57they may do it in a multiple choice
  5387. 3:34:58format where you sort of uh give the the
  5388. 3:35:00context a single time and then the four
  5389. 3:35:02completions and so the model is able to
  5390. 3:35:05see all the four options before it picks
  5391. 3:35:07the best possible option and that's
  5392. 3:35:08actually an easier task for a model
  5393. 3:35:11because you get to see the other options
  5394. 3:35:12when you're picking your choice um but
  5395. 3:35:15unfortunately models at our size can't
  5396. 3:35:17do that only models at a bigger size are
  5397. 3:35:20able to do that and so our models are
  5398. 3:35:22actually slightly handicapped in this
  5399. 3:35:23way that they are not going to see the
  5400. 3:35:25other options they're only going to see
  5401. 3:35:27one option at a time and they just have
  5402. 3:35:29to assign probabilities and the correct
  5403. 3:35:31option has to win out in this metric all
  5404. 3:35:34right so let's now implement this very
  5405. 3:35:36briefly and incorporate it into our
  5406. 3:35:38script okay so what I've done here is
  5407. 3:35:40I've introduced a new file called hell
  5408. 3:35:42swag. py that you can take a look into
  5409. 3:35:45and I'm not going to to step through all
  5410. 3:35:46of it because uh this is not exactly
  5411. 3:35:48like deep code deep code it's kind of
  5412. 3:35:51like a little bit tedious honestly
  5413. 3:35:53because what's happening is I'm
  5414. 3:35:54downloading hsac from GitHub and I'm
  5415. 3:35:56rendering all of its examples and there
  5416. 3:35:58are a total of 10,000 examples I am
  5417. 3:36:00rendering them into this format um and
  5418. 3:36:04so here at the end of this render
  5419. 3:36:07example function you can see that I'm
  5420. 3:36:09returning the
  5421. 3:36:11tokens uh the tokens of this um 4xt
  5422. 3:36:16uh array of Tokens The Mask which tells
  5423. 3:36:19us which parts are the options and
  5424. 3:36:21everything else is zero and the label
  5425. 3:36:24that is the correct label and so that
  5426. 3:36:26allows us to then iterate the examples
  5427. 3:36:28and render them and I have an evaluate
  5428. 3:36:30function here which can load a um gpt2
  5429. 3:36:33from huging face and it runs the eval
  5430. 3:36:36here um and it basically just calculates
  5431. 3:36:40uh just as I described it predicts the
  5432. 3:36:42option that has the lowest or the
  5433. 3:36:45highest prob ility and the way to do
  5434. 3:36:47that actually is we can basically
  5435. 3:36:48evaluate the cross entropy loss so we're
  5436. 3:36:51basically evaluating the loss of
  5437. 3:36:53predicting the next token in a sequence
  5438. 3:36:55and then we're looking at the row that
  5439. 3:36:57has the lowest average loss and that's
  5440. 3:37:01the uh option that we pick as the
  5441. 3:37:04prediction and then we do some stats and
  5442. 3:37:06prints and stuff like that so that is a
  5443. 3:37:08way to evaluate L swag now if you go up
  5444. 3:37:11here I'm showing that for GPT 2124m if
  5445. 3:37:14you run this script you're going to see
  5446. 3:37:16that H swag gets
  5447. 3:37:1929.5% um so that's the performance we
  5448. 3:37:22get here now remember that random Chan
  5449. 3:37:23is 25% so we haven't gone too far and
  5450. 3:37:27gpt2 XL which is the biggest the gpt2
  5451. 3:37:31gets all the way up to 49% roughly so uh
  5452. 3:37:34these are pretty low values considering
  5453. 3:37:36that today's state-ofthe-art is more
  5454. 3:37:37like 95% uh so these are definitely
  5455. 3:37:40older models by now and then there's one
  5456. 3:37:42more thing called Uther harness which is
  5457. 3:37:44a very piece of infrastructure for
  5458. 3:37:46running evals for language models and
  5459. 3:37:48they get slightly different numbers and
  5460. 3:37:50I'm not 100% sure what the discrepancy
  5461. 3:37:52is for these um it could be that they
  5462. 3:37:54actually do the multiple choice uh
  5463. 3:37:57instead of just the completions and that
  5464. 3:37:59could be the um uh the discrepancy but
  5465. 3:38:02I'm not 100% sure about that i' have to
  5466. 3:38:04take a look but for now our script
  5467. 3:38:06reports 2955 and so that is the number
  5468. 3:38:08that we'd like to beat if we are
  5469. 3:38:10training a GPD 2124m from scratch and
  5470. 3:38:13ourselves um
  5471. 3:38:16so now I'm going to go into actually
  5472. 3:38:19incorporating this eval into our main
  5473. 3:38:22training script and um and basically
  5474. 3:38:26because we want to evaluate it in a
  5475. 3:38:28periodic manner so that we can track H
  5476. 3:38:30swag and how it evolves over time and
  5477. 3:38:32see when when and if we cross uh this
  5478. 3:38:362955 um sort of region so let's now walk
  5479. 3:38:41through some of the changes to train
  5480. 3:38:42gpt2 thatp the first thing I did here is
  5481. 3:38:45I actually made use compile optional
  5482. 3:38:47kind of and I disabled it by default and
  5483. 3:38:51the problem with that is the problem
  5484. 3:38:53with compile is that unfortunately it
  5485. 3:38:55does make our code faster but it
  5486. 3:38:56actually breaks the evaluation code and
  5487. 3:38:58the sampling code it gives me a very
  5488. 3:39:00gnarly message and I don't know why so
  5489. 3:39:02hopefully by the time you get to the
  5490. 3:39:04codebase when I put it up on GitHub uh
  5491. 3:39:06we're going to fix that by then but for
  5492. 3:39:07now I'm running without torch compile
  5493. 3:39:09which is why you see this be a bit
  5494. 3:39:11slower so we're running without torch
  5495. 3:39:13compile I also create cre a log
  5496. 3:39:15directory log where we can place our
  5497. 3:39:18log.txt which will record the train loss
  5498. 3:39:22validation loss and the H swag
  5499. 3:39:23accuracies so a very simple text file
  5500. 3:39:25and we're going to uh open for writing
  5501. 3:39:28so that it sort of starts empty and then
  5502. 3:39:30we're going to append to
  5503. 3:39:32it I created a simple variable that um
  5504. 3:39:36helps tell us when we have a last step
  5505. 3:39:39and then basically periodically inside
  5506. 3:39:40this Loop every 250th iteration or at
  5507. 3:39:44the last step we're going to evaluate
  5508. 3:39:46the validation loss and then every 250th
  5509. 3:39:50iteration um we are going to evaluate H
  5510. 3:39:53swag but only if we are not using
  5511. 3:39:56compile because compile breaks it so I'm
  5512. 3:39:59going to come back to this code for
  5513. 3:40:01evaluating H swag in a second and then
  5514. 3:40:04every 250th iteration as well we're also
  5515. 3:40:06going to sample from the model and so
  5516. 3:40:08you should recognize this as our ancient
  5517. 3:40:10code from way back when we started the
  5518. 3:40:12video and we're just sampling from the
  5519. 3:40:13model
  5520. 3:40:15and then finally here um these are if
  5521. 3:40:18we're not after we validate sample and
  5522. 3:40:21evaluate hell swag we actually do a
  5523. 3:40:23training step here and so this is one
  5524. 3:40:26step of uh training and you should be
  5525. 3:40:28pretty familiar with all of what this
  5526. 3:40:30does and at the end here once we get our
  5527. 3:40:32training laws we write it to the file so
  5528. 3:40:35the only thing that changed that I
  5529. 3:40:37really added is this entire section for
  5530. 3:40:38H swag eval and the way this works is
  5531. 3:40:41I'm trying to get all the gpus to
  5532. 3:40:43collaborate on the H swag and so we're
  5533. 3:40:45iterating all the examples and then each
  5534. 3:40:48process only picks the examples that
  5535. 3:40:52assigned to it so we sort of take I and
  5536. 3:40:54moded by the world size and we have to
  5537. 3:40:56make it equal to rank otherwise we
  5538. 3:40:58continue and then we render an example
  5539. 3:41:01put it on the GPU we get the low jits
  5540. 3:41:04then I create a helper function that
  5541. 3:41:05helps us basically predict the option
  5542. 3:41:08with the lowest loss so this comes here
  5543. 3:41:10the prediction and then if it's correct
  5544. 3:41:12we sort of keep count and then if
  5545. 3:41:15multiple processes were collaborating on
  5546. 3:41:17all this then we need to synchronize
  5547. 3:41:18their stats and so the way one way to do
  5548. 3:41:21that is to package up our statistics
  5549. 3:41:23here into tensors which we can then call
  5550. 3:41:26this. alberon and
  5551. 3:41:29sum and then here we sort of um unwrap
  5552. 3:41:33them from tensors so that we just have
  5553. 3:41:35ins and then here the master process
  5554. 3:41:37will print and log the hellis swag
  5555. 3:41:40accuracy
  5556. 3:41:41so that's kind of the that's kind of it
  5557. 3:41:45and that's what I'm running right here
  5558. 3:41:47so you see this optimization here and uh
  5559. 3:41:50we just had a generation and this is
  5560. 3:41:52Step 10,000 out of about 20,000 right so
  5561. 3:41:55we are halfway done and these are the
  5562. 3:41:58kinds of samples that uh we are getting
  5563. 3:41:59at this stage so let's take a look hello
  5564. 3:42:02I'm a language model so I'd like to use
  5565. 3:42:04it to generate some kinds of output
  5566. 3:42:07hello I'm a language model and I'm a
  5567. 3:42:08developer for a lot of
  5568. 3:42:10companies Al language
  5569. 3:42:12model uh let's see if I can find fun
  5570. 3:42:17one
  5571. 3:42:28um I don't know you can go through this
  5572. 3:42:30yourself but certainly the predictions
  5573. 3:42:32are getting less and less random uh it
  5574. 3:42:34seems like the model is a little bit
  5575. 3:42:35more self-aware and using language uh
  5576. 3:42:38that is a bit
  5577. 3:42:39more uh specific to it being language
  5578. 3:42:43model hello I'm a language model and
  5579. 3:42:45like how the language is used to
  5580. 3:42:46communicate I'm a language model and I'm
  5581. 3:42:48going to be speaking English and German
  5582. 3:42:52okay I don't know so let's just wait
  5583. 3:42:53until this optimization finishes and uh
  5584. 3:42:56we'll see what kind of samples we get
  5585. 3:42:57and we're also going to look at the
  5586. 3:42:59train Val and the hway accuracy and see
  5587. 3:43:03how we're doing with respect to
  5588. 3:43:06gpt2 okay good morning so focusing For a
  5589. 3:43:09Moment On The jupyter Notebook here on
  5590. 3:43:11the right I created a new cell that
  5591. 3:43:13basically allows us to visualize the the
  5592. 3:43:15train Val and Hela and um the hel score
  5593. 3:43:19and you can step through this it
  5594. 3:43:21basically like parses the log file that
  5595. 3:43:22we are writing and um a lot of this is
  5596. 3:43:25just like boring ma plot lip code but
  5597. 3:43:28basically this is what our optimization
  5598. 3:43:30looks like
  5599. 3:43:32so we ran for
  5600. 3:43:3819,731 billion tokens which is whoops oh
  5601. 3:43:41my gosh which is one Epoch of the sample
  5602. 3:43:4410B of webd on the left we have the loss
  5603. 3:43:48and the in blue we have the training
  5604. 3:43:50loss in Orange we have the validation
  5605. 3:43:52loss and red as a horizontal line we
  5606. 3:43:55have the opening IG gpt2 124 M model
  5607. 3:43:58checkpoint when it's just evaluated on
  5608. 3:44:00the validation set of um of this fine
  5609. 3:44:04web edu uh so you can see that we are
  5610. 3:44:06surpassing this orange is below the red
  5611. 3:44:09so we're surpassing the validation set
  5612. 3:44:11of this data set and like I mentioned
  5613. 3:44:13the data set distribution is very
  5614. 3:44:15different from what gpt2 trained on so
  5615. 3:44:16this is not an exactly fair comparison
  5616. 3:44:19but it's a good cross check uh to uh to
  5617. 3:44:22look at now we would ideally like
  5618. 3:44:25something that is withheld and
  5619. 3:44:27comparable and somewhat standard um and
  5620. 3:44:30so for us that is helis swag and so on
  5621. 3:44:33here we see the H swag progress we made
  5622. 3:44:35from 25% all the way here in red we see
  5623. 3:44:39the open gpt2 124 M model in red so it
  5624. 3:44:44achieves this h bag here and the the
  5625. 3:44:47gpt3 model 124 M which was trained on
  5626. 3:44:50300 billion tokens achieves green so
  5627. 3:44:54that's over here so you see that we
  5628. 3:44:56basically surpassed the gbt2 24m uh
  5629. 3:45:00model right here uh which is uh really
  5630. 3:45:03nice now interestingly we were able to
  5631. 3:45:07do so with only training on 10 billion
  5632. 3:45:08tokens while gpt2 was trained on 100
  5633. 3:45:11billion tokens so uh for some reason we
  5634. 3:45:14were able to get away with significantly
  5635. 3:45:16fewer tokens for training there are many
  5636. 3:45:18possibilities to as to why we could
  5637. 3:45:21match or surpass this accuracy um with
  5638. 3:45:24only 10 million training so number one
  5639. 3:45:27um it could be that opening gbt2 was
  5640. 3:45:30trained on a much wider data
  5641. 3:45:32distribution so in particular fine web
  5642. 3:45:34edu is all English it's not multilingual
  5643. 3:45:38and there's not that much math and code
  5644. 3:45:40um and so math and code and multilingual
  5645. 3:45:43could have been stealing capacity from
  5646. 3:45:45the original gpt2 model and um basically
  5647. 3:45:50that could be partially the reason why
  5648. 3:45:52uh this is not working out there's many
  5649. 3:45:54other reasons um so for example the H
  5650. 3:45:57swag eval is fairly old uh maybe 5 years
  5651. 3:45:59or so it is possible that aspects of H
  5652. 3:46:02swag in some way or even identically
  5653. 3:46:04have made it into the training Set uh of
  5654. 3:46:07fine web we don't know for sure but if
  5655. 3:46:10that was the case then we are basically
  5656. 3:46:11looking at the training curve instead of
  5657. 3:46:12the validation curve so long story short
  5658. 3:46:15this is not a perfect eval and there's
  5659. 3:46:16some caveats here uh but at least we
  5660. 3:46:19have some confidence that that we're not
  5661. 3:46:20doing something completely wrong and
  5662. 3:46:23um and uh it's probably the case that
  5663. 3:46:26when people try to create these data
  5664. 3:46:27sets they try to make sure that test
  5665. 3:46:29sets that are very common are not part
  5666. 3:46:31of the training set for example uh when
  5667. 3:46:33hugging face created the fine web BDU
  5668. 3:46:35they use H swag as an eval so I would
  5669. 3:46:37hope that they make sure that they D
  5670. 3:46:39duplicate and that there's no hella swag
  5671. 3:46:41in the training set but we can't be sure
  5672. 3:46:45uh the other thing I wanted to address
  5673. 3:46:46briefly is look at this loss curve this
  5674. 3:46:48looks really this looks really wrong
  5675. 3:46:50here I don't actually know 100% what
  5676. 3:46:52this is and I suspect it's because the
  5677. 3:46:55uh 10 billion sample of fine web edu was
  5678. 3:46:58not properly shuffled um and there's
  5679. 3:47:01some issue here uh with the data that I
  5680. 3:47:04don't fully understand yet and there's
  5681. 3:47:06some weird periodicity to it um and
  5682. 3:47:08because we are in a very lazy way sort
  5683. 3:47:10of serializing all the tokens and just
  5684. 3:47:12iterating all them from scratch without
  5685. 3:47:14doing any permutation or any random
  5686. 3:47:16sampling ourselves I think we're
  5687. 3:47:18inheriting some of the ordering that
  5688. 3:47:21they have in the data set so uh this is
  5689. 3:47:24not ideal but hopefully by the time you
  5690. 3:47:26get to this repo uh some of these things
  5691. 3:47:28by the way will hopefully be fixed and I
  5692. 3:47:32will release this build n GPT repo and
  5693. 3:47:35right now it looks a little ugly and
  5694. 3:47:37preliminary uh so hopefully by the time
  5695. 3:47:39you get here it's nicer but down here
  5696. 3:47:41I'm going to show aada and I'm going to
  5697. 3:47:44talk about about some of the things that
  5698. 3:47:45happened after the video and I expect
  5699. 3:47:48that we will have fixed uh the small
  5700. 3:47:50issue uh but for now basically this
  5701. 3:47:52shows that uh our training is not uh
  5702. 3:47:55completely wrong and it shows that uh
  5703. 3:47:58we're able to surpass the accuracy with
  5704. 3:48:00only 10x the token budget um and
  5705. 3:48:03possibly it could be also that the data
  5706. 3:48:05set may have improved so uh the original
  5707. 3:48:08uh gpt2 data set was web text it's
  5708. 3:48:11possible that not a lot of care and
  5709. 3:48:12attention went into the data set this
  5710. 3:48:14was very early in llms whereas now
  5711. 3:48:17there's a lot more scrutiny on good
  5712. 3:48:18practices around uh D duplication
  5713. 3:48:20filtering uh quality filtering and so on
  5714. 3:48:23and it's possible that the data that
  5715. 3:48:24we're training on is just of higher
  5716. 3:48:25quality per token and that could be
  5717. 3:48:27giving us a boost as well so a number of
  5718. 3:48:30cave has to think about but for now uh
  5719. 3:48:32we're pretty happy with this um and yeah
  5720. 3:48:36now the next thing I was interested in
  5721. 3:48:37is as you see it's a morning now so
  5722. 3:48:39there was an overnight and I wanted to
  5723. 3:48:41basically see how far I could push the
  5724. 3:48:43result so uh to do an overnight run I
  5725. 3:48:46basically did instead of one Epoch which
  5726. 3:48:48took roughly two hours I just did a
  5727. 3:48:50times four so that that would take eight
  5728. 3:48:52hours while I was sleeping and so we did
  5729. 3:48:54four Epoch or roughly 40 billion uh
  5730. 3:48:56tokens of training and I was trying to
  5731. 3:48:59see how far we could get um and so this
  5732. 3:49:01was the only change and I reran the
  5733. 3:49:03script and when I point uh and read the
  5734. 3:49:05log file at uh at the 40b uh this is
  5735. 3:49:08what the curve look
  5736. 3:49:10like okay so to narrate this number one
  5737. 3:49:13we are seeing this issue here here with
  5738. 3:49:15the periodicity through the different
  5739. 3:49:17Epoch and something really weird with
  5740. 3:49:19the fine web edu data set and that is to
  5741. 3:49:22be determined uh but otherwise we are
  5742. 3:49:25seeing that the H swag actually went up
  5743. 3:49:27by a lot and we almost we almost made it
  5744. 3:49:31uh to the GPT 324m accuracy uh up here
  5745. 3:49:35uh but not quite so uh it's too bad that
  5746. 3:49:37I didn't sleep slightly longer um and uh
  5747. 3:49:41I think if this was an uh five Epoch run
  5748. 3:49:44we may have gotten here now one thing to
  5749. 3:49:47point out is that if you're doing multi
  5750. 3:49:49Epoch runs uh we're not actually being
  5751. 3:49:51very careful in our data loader and
  5752. 3:49:53we're not um I this data loader goes
  5753. 3:49:56through the data in exactly the same
  5754. 3:49:59format and exactly the same order and
  5755. 3:50:01this is kind of suboptimal and you would
  5756. 3:50:03want to look into extensions where you
  5757. 3:50:05actually permute the data uh randomly
  5758. 3:50:08you permute the documents around in
  5759. 3:50:10Every Single Shard on every single new
  5760. 3:50:12Epoch um and po even permute the
  5761. 3:50:16shards and that would go a long way into
  5762. 3:50:18decreasing the pricity and it's also
  5763. 3:50:20better for the optimization so that
  5764. 3:50:22you're not seeing things ident in the
  5765. 3:50:23identical format and you're introducing
  5766. 3:50:25some of the some uh Randomness in how
  5767. 3:50:28the documents follow each other because
  5768. 3:50:29you have to remember that in every
  5769. 3:50:31single row these documents follow each
  5770. 3:50:33other and then there's the end of text
  5771. 3:50:34token and then the next document so the
  5772. 3:50:36documents are currently glued together
  5773. 3:50:39in the exact same identical manner but
  5774. 3:50:41we actually want to break break up the
  5775. 3:50:43documents and shuffle them around
  5776. 3:50:45because the order of the documents
  5777. 3:50:46shouldn't matter and they shouldn't um
  5778. 3:50:49basically we want to break up that
  5779. 3:50:50dependence because it's a kind of a
  5780. 3:50:51spous correlation and so our data lad is
  5781. 3:50:54not currently doing that and that's one
  5782. 3:50:56Improvement uh you could think of
  5783. 3:50:58making um the other thing to point out
  5784. 3:51:01is we're almost matching gpt3 accuracy
  5785. 3:51:03with only 40 billion tokens gpt3 trained
  5786. 3:51:06on 300 billion tokens so again we're
  5787. 3:51:08seeing about a 10x um Improvement here
  5788. 3:51:11with respect to learning efficiency uh
  5789. 3:51:14the other thing I wanted to and I don't
  5790. 3:51:16actually know exactly what to attribute
  5791. 3:51:18this to other than some of the things
  5792. 3:51:19that I already mentioned previously for
  5793. 3:51:21the previous run uh the other thing I
  5794. 3:51:23wanted to briefly mention is uh the max
  5795. 3:51:26LR here I saw some people already play
  5796. 3:51:29with this a little bit in a previous
  5797. 3:51:31related repository um and it turns out
  5798. 3:51:33that you can actually almost like three
  5799. 3:51:35xas so it's possible that the maximum
  5800. 3:51:37learning rate can be a lot higher and
  5801. 3:51:39for some reason the gpt3 hyper
  5802. 3:51:40parameters that we are inheriting are
  5803. 3:51:42actually extremely conservative and you
  5804. 3:51:44can actually get away with a Higher
  5805. 3:51:45Learning rate and it would train faster
  5806. 3:51:47so a lot of these hyper parameters um
  5807. 3:51:50are quite tunable and feel free to play
  5808. 3:51:52with them and they're probably not set
  5809. 3:51:54precisely correctly and um it's possible
  5810. 3:51:59that you can get away with doing this
  5811. 3:52:01basically and if you wanted to exactly
  5812. 3:52:03be faithful to gpt3 you would also want
  5813. 3:52:07to make the following difference you'd
  5814. 3:52:10want to come here and the sequence
  5815. 3:52:11length of gpt3 is 2x it's 20 48 instead
  5816. 3:52:15of 1,24 so you would come here change
  5817. 3:52:17this to 248 for T and then if you want
  5818. 3:52:20the exact same number of tokens uh half
  5819. 3:52:22a million per iteration or per step you
  5820. 3:52:25want to then decrease this to 32 so they
  5821. 3:52:28still multiply to half a mil so that
  5822. 3:52:31would give your model sequence length
  5823. 3:52:33equal to that of gpt3 and in that case
  5824. 3:52:36basically the
  5825. 3:52:37um the models would be roughly identical
  5826. 3:52:40as far as I'm as far as I'm aware
  5827. 3:52:42because again gpt2 and gpt3 are very
  5828. 3:52:44very similar models now we can also look
  5829. 3:52:47at some of the samples here from the
  5830. 3:52:48model that was trained overnight so this
  5831. 3:52:51is
  5832. 3:52:52the optimization and you see that here
  5833. 3:52:55we stepped all the way to
  5834. 3:52:5776290 also or so and these are the hos
  5835. 3:53:02mag we achieved was 33.2 4 and these are
  5836. 3:53:06some of the samples from the model and
  5837. 3:53:08you can see that if you read through
  5838. 3:53:10this and pause the video briefly you can
  5839. 3:53:11see that they are a lot more coherent uh
  5840. 3:53:14so
  5841. 3:53:15um and they're actually addressing the
  5842. 3:53:17fact that it's a language model almost
  5843. 3:53:21so uh hello I'm a language model and I
  5844. 3:53:24try to be as accurate as
  5845. 3:53:27possible um I'm a language model not a
  5846. 3:53:29programming
  5847. 3:53:31language I know how to communicate uh I
  5848. 3:53:34use
  5849. 3:53:35Python
  5850. 3:53:37um I don't know if you pause this and
  5851. 3:53:40look at it and then compare it to the
  5852. 3:53:41one to the model that was only trained
  5853. 3:53:43for 10 billion uh you will see that
  5854. 3:53:45these are a lot more coherent and you
  5855. 3:53:47can play with this uh
  5856. 3:53:48yourself one more thing I added to The
  5857. 3:53:50Code by the way is this chunk of code
  5858. 3:53:52here so basically right after we
  5859. 3:53:54evaluate the validation loss if we are
  5860. 3:53:56the master process in addition to
  5861. 3:53:58logging the validation loss every 5,000
  5862. 3:54:01steps we're also going to save the
  5863. 3:54:02checkpoint which is really just the
  5864. 3:54:04state dictionary of the model and so
  5865. 3:54:07checkpointing is nice just because uh
  5866. 3:54:09you can save the model and later you can
  5867. 3:54:11uh use it in some way if you wanted to
  5868. 3:54:13resume the optimiz ation then in
  5869. 3:54:15addition to saving the model we have to
  5870. 3:54:17also save the optimizer State dict
  5871. 3:54:20because remember that the optimizer has
  5872. 3:54:21a few additional buffers because of adom
  5873. 3:54:24so it's got the m and V and uh you need
  5874. 3:54:28to also resume the optimizer properly
  5875. 3:54:30you have to be careful with your RNG
  5876. 3:54:31seeds uh random number generators and so
  5877. 3:54:33on so if you wanted to exactly be able
  5878. 3:54:35to resume optimization you have to think
  5879. 3:54:37through the state of the of the training
  5880. 3:54:40process but if you just want to save the
  5881. 3:54:41model this is how you would do it and
  5882. 3:54:43one one nice reason why you might want
  5883. 3:54:45to do this is because you may want to
  5884. 3:54:47evaluate the model a lot more carefully
  5885. 3:54:50so here we are only kind of like winging
  5886. 3:54:52the hell swag eval but you may want to
  5887. 3:54:54use something um nicer like for example
  5888. 3:54:57the Luther uh Luther evaluation hardness
  5889. 3:55:01evaluation hardness hardness um so this
  5890. 3:55:06is a way to also evaluate language
  5891. 3:55:08models and um so it's possible that um
  5892. 3:55:13you may want to use basically different
  5893. 3:55:15infrastructure to more thoroughly
  5894. 3:55:17evaluate the models on different um
  5895. 3:55:20evaluations and compare it to the
  5896. 3:55:21opening gbt2 model on many other um
  5897. 3:55:25tasks like for example that involve math
  5898. 3:55:26code or different languages and so on so
  5899. 3:55:29this is a nice functionality to have as
  5900. 3:55:30well
  5901. 3:55:32um and then the other thing I wanted to
  5902. 3:55:34mention is that everything we've built
  5903. 3:55:36here this is only the pre-training step
  5904. 3:55:39so um the GPT here is a it dreams
  5905. 3:55:42documents it just predicts the next to
  5906. 3:55:44you can't talk to it like you can talk
  5907. 3:55:46to chat GPT uh chat GPT if you wanted to
  5908. 3:55:49talk to the model we have to fine-tune
  5909. 3:55:51it into the chat format and it's not
  5910. 3:55:54actually like that complicated if you're
  5911. 3:55:55looking at supervised fine-tuning or sft
  5912. 3:55:58really what that means is we're just
  5913. 3:55:59swapping out a data set into a data set
  5914. 3:56:01that is a lot more conversational and
  5915. 3:56:03there's a user assistant user assistant
  5916. 3:56:04kind of structure and we just fine-tune
  5917. 3:56:06on it and then we um we basically fill
  5918. 3:56:09in the user tokens and we sample the
  5919. 3:56:11assistant tokens it's not a lot more
  5920. 3:56:13deeper than that uh but basically we
  5921. 3:56:15swap out the data set and continue
  5922. 3:56:17training uh but for now we're going to
  5923. 3:56:19stop at uh pre-training one more thing
  5924. 3:56:21that I wanted to briefly show you is
  5925. 3:56:23that of course what we've built up today
  5926. 3:56:25was building towards nanog GPT which is
  5927. 3:56:27this repository from earlier uh but also
  5928. 3:56:30there's actually another nanog GPT
  5929. 3:56:32implementation and it's hiding in a more
  5930. 3:56:34recent project that I've been working on
  5931. 3:56:36called llm Doc and lm. C is a pure Cuda
  5932. 3:56:41implementation of gpt2 or gpt3 training
  5933. 3:56:44and it just directly uses uh Cuda and is
  5934. 3:56:47written as Cuda now the nanog gbt here
  5935. 3:56:51acts as reference code in pytorch to the
  5936. 3:56:53C implementation so we're trying to
  5937. 3:56:55exactly match up the two but we're
  5938. 3:56:57hoping that the C Cuda is faster and of
  5939. 3:56:59course currently that seems to be the
  5940. 3:57:01case um because it is a direct optimized
  5941. 3:57:04implementation so train gpt2 Pi in LL
  5942. 3:57:06M.C is basically the nanog GPT and when
  5943. 3:57:10you scroll through this file you'll find
  5944. 3:57:12a lot of things that very much look like
  5945. 3:57:16um things that we've built up in this
  5946. 3:57:19lecture and then when you look at train
  5947. 3:57:21gpt2 docu uh this is the C Cuda
  5948. 3:57:25implementation so there's a lot of MPI
  5949. 3:57:27nickel GPU Cuda
  5950. 3:57:30cc++ and you have to be familiar with
  5951. 3:57:32that but uh um when this is built up we
  5952. 3:57:37can actually run the two side by side
  5953. 3:57:39and they're going to produce the exact
  5954. 3:57:40same results but lm. C actually runs
  5955. 3:57:43faster so let's see that so on the left
  5956. 3:57:45I have pytorch a nanog GPT looking thing
  5957. 3:57:49on the right I have the llmc call and
  5958. 3:57:52here I'm going to launch the
  5959. 3:57:54two both of these are going to be
  5960. 3:57:55running on a single GPU and here I'm
  5961. 3:57:57putting the lm. C on GPU 1 and this one
  5962. 3:58:00will grab uh gpu0 by default and
  5963. 3:58:05then we can see here that lm. c
  5964. 3:58:08compiled and then allocate space and
  5965. 3:58:11it's
  5966. 3:58:12stepping so
  5967. 3:58:15basically uh meanwhile P torch is still
  5968. 3:58:17compiling because torch compile is a bit
  5969. 3:58:19slower here than the lm. C nbcc Cuda
  5970. 3:58:24compile and so this program has already
  5971. 3:58:26started running and uh we're still
  5972. 3:58:28waiting here for torch compile now of
  5973. 3:58:30course uh this is a very specific
  5974. 3:58:33implementation to gpt2 and 3 a pytorch
  5975. 3:58:35is a very general neural network
  5976. 3:58:37framework so they're not exactly
  5977. 3:58:38comparable but if you're only interested
  5978. 3:58:39in training gpt2 and 3 lm. C is very
  5979. 3:58:43fast it takes less space it's faster to
  5980. 3:58:46start and it's faster per
  5981. 3:58:49step and so P started to Stepping here
  5982. 3:58:53and as you can see we're running at
  5983. 3:58:54about 223,000 tokens per second here and
  5984. 3:58:57about 185,000 tokens per second here um
  5985. 3:59:03so quite a bit slower but I don't have
  5986. 3:59:05full confidence that I exactly squeezed
  5987. 3:59:08out all the juice from the pytorch
  5988. 3:59:09implementation but the important thing
  5989. 3:59:11here is notice that if I Aline up the
  5990. 3:59:14steps you will see that the losses and
  5991. 3:59:16Norms that are printed between these two
  5992. 3:59:18are
  5993. 3:59:19identical so on the left we have the pie
  5994. 3:59:21torch and on the right this C
  5995. 3:59:24implementation and they're the same
  5996. 3:59:25except this one runs faster uh so that's
  5997. 3:59:28kind of I wanted to show you also
  5998. 3:59:30briefly lm. C and this is a parallel
  5999. 3:59:33implementation and it's also something
  6000. 3:59:35that you may want to uh play with or
  6001. 3:59:37look at and um it's kind of interesting
  6002. 3:59:39okay so at this point I should probably
  6003. 3:59:40start wrapping up the video because I
  6004. 3:59:42think it's getting way longer than I
  6005. 3:59:44anticipated uh but we did Cover a lot of
  6006. 3:59:46ground and we built everything from
  6007. 3:59:48scratch so as a brief summary we were
  6008. 3:59:50looking at the gpt2 and GPT 3
  6009. 3:59:55papers we were looking at how you set up
  6010. 3:59:57these training runs uh and all the
  6011. 3:59:59considerations involved we wrote
  6012. 4:00:01everything from scratch and then we saw
  6013. 4:00:03that over the duration of either a
  6014. 4:00:042-hour training run or an overnight run
  6015. 4:00:07we can actually match the 124 million
  6016. 4:00:09parameter checkpoints of gbt2 and gpt3
  6017. 4:00:12uh to a very large extent
  6018. 4:00:14um in principle the code that we wrote
  6019. 4:00:16would be able to train even bigger
  6020. 4:00:18models if you have the patients or the
  6021. 4:00:19Computing resources uh and so you could
  6022. 4:00:21potentially think about training some of
  6023. 4:00:23the bigger checkpoints as well um there
  6024. 4:00:26are a few remaining issues to address
  6025. 4:00:28what's happening with the loss here
  6026. 4:00:30which I suspect has to do with the fine
  6027. 4:00:31web edu data sampling uh why can't we
  6028. 4:00:34turn on Torch compile uh it currently
  6029. 4:00:36breaks generation and H swag what's up
  6030. 4:00:39with that in the data loader we should
  6031. 4:00:41probably be permuting our data when we
  6032. 4:00:43reach boundaries so there's a few more
  6033. 4:00:45issues like that and I expect to be
  6034. 4:00:47documenting some of those over time in
  6035. 4:00:49the uh build n GPT repository here which
  6036. 4:00:53I'm going to be releasing with this
  6037. 4:00:55video if you have any questions or like
  6038. 4:00:57to talk about anything that we covered
  6039. 4:00:59please go to discussions tab uh so we
  6040. 4:01:02can talk here uh or please go to issues
  6041. 4:01:04or pull request pull requests um
  6042. 4:01:07depending on what you'd like to
  6043. 4:01:08contribute or also have a look at the uh
  6044. 4:01:11Zero to Hero Discord and uh I'm going to
  6045. 4:01:14be hanging out here on N GPT
  6046. 4:01:17um otherwise for now I'm pretty happy
  6047. 4:01:20about where we got um and I hope you
  6048. 4:01:23enjoyed the video and I will see you
  6049. 4:01:25later

About this transcript

This page contains the full transcript of Let's reproduce GPT-2 (124M) by Andrej Karpathy, generated from the public captions YouTube serves with the video. The transcript has 43,388 words across 6,049 segments, with the original timestamps preserved so you can click any line to jump to that moment in the embedded player.

What you can do with it

Use the transcript to take notes, quote the speaker, build a study guide, generate a summary with ChatGPT or Claude via the YouTube Summary tool, or export it as a timed subtitle file with YouTube to SRT. You can also re-open it in the transcriber to translate the transcript into 100+ languages.

Free YouTube transcript tool

YouTube2Text is a free YouTube transcript generator — no signup, no daily limit. Paste any YouTube link and get the full transcript instantly, with timestamps, click-to-jump, translation to 100+ languages, AI prompts for ChatGPT, Claude, and Gemini, and exports to TXT, SRT, VTT, or Markdown.