YouTube2Text

Let's build GPT: from scratch, in code, spelled out. — Transcript

by Andrej Karpathy · 21,030 words · 2,955 segments · language en · Watch on YouTube

Full transcript

  1. 0:00hi everyone so by now you have probably
  2. 0:02heard of chat GPT it has taken the world
  3. 0:04and AI Community by storm and it is a
  4. 0:07system that allows you to interact with
  5. 0:09an AI and give it text based tasks so
  6. 0:12for example we can ask chat GPT to write
  7. 0:15us a small Hau about how important it is
  8. 0:16that people understand Ai and then they
  9. 0:18can use it to improve the world and make
  10. 0:20it more prosperous so when we run this
  11. 0:23AI knowledge brings prosperity for all
  12. 0:25to see Embrace its
  13. 0:27power okay not bad and so you could see
  14. 0:29that chpt went from left to right and
  15. 0:32generated all these words SE sort of
  16. 0:35sequentially now I asked it already the
  17. 0:37exact same prompt a little bit earlier
  18. 0:39and it generated a slightly different
  19. 0:41outcome ai's power to grow ignorance
  20. 0:44holds us back learn Prosperity weights
  21. 0:47so uh pretty good in both cases and
  22. 0:49slightly different so you can see that
  23. 0:50chat GPT is a probabilistic system and
  24. 0:52for any one prompt it can give us
  25. 0:54multiple answers sort of uh replying to
  26. 0:57it now this is just one example of a
  27. 0:59problem people have come up with many
  28. 1:01many examples and there are entire
  29. 1:03websites that index interactions with
  30. 1:06chpt and so many of them are quite
  31. 1:08humorous explain HTML to me like I'm a
  32. 1:10dog uh write release notes for chess 2
  33. 1:14write a note about Elon Musk buying a
  34. 1:16Twitter and so on so as an example uh
  35. 1:20please write a breaking news article
  36. 1:21about a leaf falling from a
  37. 1:23tree uh and a shocking turn of events a
  38. 1:26leaf has fallen from a tree in the local
  39. 1:28park Witnesses report that the leaf
  40. 1:30which was previously attached to a
  41. 1:31branch of a tree attached itself and
  42. 1:33fell to the ground very dramatic so you
  43. 1:36can see that this is a pretty remarkable
  44. 1:37system and it is what we call a language
  45. 1:40model uh because it um it models the
  46. 1:43sequence of words or characters or
  47. 1:46tokens more generally and it knows how
  48. 1:49sort of words follow each other in
  49. 1:50English language and so from its
  50. 1:52perspective what it is doing is it is
  51. 1:55completing the sequence so I give it the
  52. 1:57start of a sequence and it completes the
  53. 2:00sequence with the outcome and so it's a
  54. 2:02language model in that sense now I would
  55. 2:05like to focus on the under the hood of
  56. 2:07um under the hood components of what
  57. 2:09makes CH GPT work so what is the neural
  58. 2:12network under the hood that models the
  59. 2:14sequence of these words and that comes
  60. 2:17from this paper called attention is all
  61. 2:19you need in 2017 a landmark paper a
  62. 2:23landmark paper in AI that produced and
  63. 2:25proposed the Transformer
  64. 2:27architecture so GPT is uh short for
  65. 2:31generally generatively pre-trained
  66. 2:33Transformer so Transformer is the neuron
  67. 2:35nut that actually does all the heavy
  68. 2:36lifting under the hood it comes from
  69. 2:39this paper in 2017 now if you read this
  70. 2:41paper this uh reads like a pretty random
  71. 2:44machine translation paper and that's
  72. 2:46because I think the authors didn't fully
  73. 2:47anticipate the impact that the
  74. 2:49Transformer would have on the field and
  75. 2:51this architecture that they produced in
  76. 2:52the context of machine translation in
  77. 2:54their case actually ended up taking over
  78. 2:57uh the rest of AI in the next 5 years
  79. 3:00after and so this architecture with
  80. 3:02minor changes was copy pasted into a
  81. 3:05huge amount of applications in AI in
  82. 3:07more recent years and that includes at
  83. 3:10the core of chat GPT now we are not
  84. 3:13going to what I'd like to do now is I'd
  85. 3:15like to build out something like chat
  86. 3:17GPT but uh we're not going to be able to
  87. 3:19of course reproduce chat GPT this is a
  88. 3:21very serious production grade system it
  89. 3:23is trained on uh a good chunk of
  90. 3:26internet and then there's a lot of uh
  91. 3:29pre-training and fine-tuning stages to
  92. 3:31it and so it's very complicated what I'd
  93. 3:33like to focus on is just to train a
  94. 3:36Transformer based language model and in
  95. 3:38our case it's going to be a character
  96. 3:40level language model I still think that
  97. 3:43is uh very educational with respect to
  98. 3:45how these systems work so I don't want
  99. 3:47to train on the chunk of Internet we
  100. 3:48need a smaller data set in this case I
  101. 3:51propose that we work with uh my favorite
  102. 3:53toy data set it's called tiny
  103. 3:55Shakespeare and um what it is is
  104. 3:57basically it's a concatenation of all of
  105. 3:59the works of sh Shakespeare in my
  106. 4:00understanding and so this is all of
  107. 4:02Shakespeare in a single file uh this
  108. 4:05file is about 1 megab and it's just all
  109. 4:07of
  110. 4:08Shakespeare and what we are going to do
  111. 4:10now is we're going to basically model
  112. 4:12how these characters uh follow each
  113. 4:14other so for example given a chunk of
  114. 4:16these characters like this uh given some
  115. 4:19context of characters in the past the
  116. 4:22Transformer neural network will look at
  117. 4:24the characters that I've highlighted and
  118. 4:26is going to predict that g is likely to
  119. 4:28come next in the sequence and it's going
  120. 4:30to do that because we're going to train
  121. 4:31that Transformer on Shakespeare and it's
  122. 4:34just going to try to produce uh
  123. 4:36character sequences that look like this
  124. 4:39and in that process is going to model
  125. 4:40all the patterns inside this data so
  126. 4:43once we've trained the system i' just
  127. 4:45like to give you a preview we can
  128. 4:47generate infinite Shakespeare and of
  129. 4:49course it's a fake thing that looks kind
  130. 4:51of like
  131. 4:53Shakespeare
  132. 4:55um apologies for there's some Jank that
  133. 4:59I'm not able to resolve in in here but
  134. 5:02um you can see how this is going
  135. 5:05character by character and it's kind of
  136. 5:07like predicting Shakespeare like
  137. 5:09language so verily my Lord the sites
  138. 5:12have left the again the king coming with
  139. 5:15my curses with precious pale and then
  140. 5:19tranos say something else Etc and this
  141. 5:21is just coming out of the Transformer in
  142. 5:23a very similar manner as it would come
  143. 5:25out in chat GPT in our case character by
  144. 5:27character in chat GPT uh it's coming out
  145. 5:31on the token by token level and tokens
  146. 5:33are these sort of like little subword
  147. 5:35pieces so they're not Word level they're
  148. 5:36kind of like word chunk
  149. 5:38level um and now I've already written
  150. 5:43this entire code uh to train these
  151. 5:45Transformers um and it is in a GitHub
  152. 5:48repository that you can find and it's
  153. 5:50called nanog
  154. 5:51GPT so nanog GPT is a repository that
  155. 5:54you can find in my GitHub and it's a
  156. 5:56repository for training Transformers um
  157. 5:59on any given text and what I think is
  158. 6:02interesting about it because there's
  159. 6:03many ways to train Transformers but this
  160. 6:05is a very simple implementation so it's
  161. 6:06just two files of 300 lines of code each
  162. 6:10one file defines the GPT model the
  163. 6:12Transformer and one file trains it on
  164. 6:14some given Text data set and here I'm
  165. 6:17showing that if you train it on a open
  166. 6:18web Text data set which is a fairly
  167. 6:20large data set of web pages then I
  168. 6:22reproduce the the performance of
  169. 6:25gpt2 so gpt2 is an early version of open
  170. 6:29AI GPT uh from 2017 if I recall
  171. 6:32correctly and I've only so far
  172. 6:34reproduced the the smallest 124 million
  173. 6:36parameter model uh but basically this is
  174. 6:38just proving that the codebase is
  175. 6:39correctly arranged and I'm able to load
  176. 6:42the uh neural network weights that openi
  177. 6:45has released later so you can take a
  178. 6:48look at the finished code here in N GPT
  179. 6:50but what I would like to do in this
  180. 6:51lecture is I would like to basically uh
  181. 6:55write this repository from scratch so
  182. 6:57we're going to begin with an empty file
  183. 6:59and we're we're going to define a
  184. 7:00Transformer piece by piece we're going
  185. 7:03to train it on the tiny Shakespeare data
  186. 7:05set and we'll see how we can then uh
  187. 7:08generate infinite Shakespeare and of
  188. 7:10course this can copy paste to any
  189. 7:12arbitrary Text data set uh that you like
  190. 7:14uh but my goal really here is to just
  191. 7:16make you understand and appreciate uh
  192. 7:18how under the hood chat GPT works and um
  193. 7:22really all that's required is a
  194. 7:24Proficiency in Python and uh some basic
  195. 7:27understanding of um calculus and
  196. 7:29statistics
  197. 7:30and it would help if you also see my
  198. 7:32previous videos on the same YouTube
  199. 7:34channel in particular my make more
  200. 7:35series where I um Define smaller and
  201. 7:40simpler neural network language models
  202. 7:42uh so multi perceptrons and so on it
  203. 7:45really introduces the language modeling
  204. 7:46framework and then uh here in this video
  205. 7:49we're going to focus on the Transformer
  206. 7:50neural network itself okay so I created
  207. 7:53a new Google collab uh jup notebook here
  208. 7:57and this will allow me to later easily
  209. 7:58share this code that we're going to
  210. 8:00develop together uh with you so you can
  211. 8:01follow along so this will be in a video
  212. 8:03description uh later now here I've just
  213. 8:07done some preliminaries I downloaded the
  214. 8:09data set the tiny Shakespeare data set
  215. 8:10at this URL and you can see that it's
  216. 8:12about a 1 Megabyte file then here I open
  217. 8:15the input.txt file and just read in all
  218. 8:17the text of the string and we see that
  219. 8:20we are working with 1 million characters
  220. 8:22roughly and the first 1,000 characters
  221. 8:24if we just print them out are basically
  222. 8:26what you would expect this is the first
  223. 8:281,000 characters of the tiny Shakespeare
  224. 8:30data set roughly up to here so so far so
  225. 8:34good next we're going to take this text
  226. 8:37and the text is a sequence of characters
  227. 8:39in Python so when I call the set
  228. 8:41Constructor on it I'm just going to get
  229. 8:44the set of all the characters that occur
  230. 8:46in this text and then I call list on
  231. 8:49that to create a list of those
  232. 8:51characters instead of just a set so that
  233. 8:53I have an ordering an arbitrary ordering
  234. 8:56and then I sort that so basically we get
  235. 8:59just all the characters that occur in
  236. 9:00the entire data set and they're sorted
  237. 9:02now the number of them is going to be
  238. 9:04our vocabulary size these are the
  239. 9:06possible elements of our sequences and
  240. 9:09we see that when I print here the
  241. 9:11characters there's 65 of them in total
  242. 9:14there's a space character and then all
  243. 9:16kinds of special characters and then U
  244. 9:19capitals and lowercase letters so that's
  245. 9:21our vocabulary and that's the sort of
  246. 9:23like possible uh characters that the
  247. 9:25model can see or emit okay so next we
  248. 9:29will would like to develop some strategy
  249. 9:31to tokenize the input text now when
  250. 9:35people say tokenize they mean convert
  251. 9:36the raw text as a string to some
  252. 9:39sequence of integers According to some
  253. 9:41uh notebook According to some vocabulary
  254. 9:43of possible elements so as an example
  255. 9:46here we are going to be building a
  256. 9:48character level language model so we're
  257. 9:49simply going to be translating
  258. 9:50individual characters into integers so
  259. 9:53let me show you uh a chunk of code that
  260. 9:55sort of does that for us so we're
  261. 9:57building both the encoder and the
  262. 9:58decoder
  263. 10:00and let me just talk through what's
  264. 10:01happening
  265. 10:02here when we encode an arbitrary text
  266. 10:05like hi there we're going to receive a
  267. 10:08list of integers that represents that
  268. 10:10string so for example 46 47 Etc and then
  269. 10:14we also have the reverse mapping so we
  270. 10:17can take this list and decode it to get
  271. 10:20back the exact same string so it's
  272. 10:22really just like a translation to
  273. 10:24integers and back for arbitrary string
  274. 10:26and for us it is done on a character
  275. 10:28level
  276. 10:30now the way this was achieved is we just
  277. 10:31iterate over all the characters here and
  278. 10:34create a lookup table from the character
  279. 10:35to the integer and vice versa and then
  280. 10:38to encode some string we simply
  281. 10:40translate all the characters
  282. 10:41individually and to decode it back we
  283. 10:44use the reverse mapping and concatenate
  284. 10:46all of it now this is only one of many
  285. 10:49possible encodings or many possible sort
  286. 10:51of tokenizers and it's a very simple one
  287. 10:54but there's many other schemas that
  288. 10:55people have come up with in practice so
  289. 10:57for example Google uses a sentence
  290. 10:59piece uh so sentence piece will also
  291. 11:02encode text into um integers but in a
  292. 11:05different schema and using a different
  293. 11:08vocabulary and sentence piece is a
  294. 11:10subword uh sort of tokenizer and what
  295. 11:13that means is that um you're not
  296. 11:15encoding entire words but you're not
  297. 11:17also encoding individual characters it's
  298. 11:19it's a subword unit level and that's
  299. 11:22usually what's adopted in practice for
  300. 11:24example also openai has this Library
  301. 11:26called tick token that uses a bite pair
  302. 11:28encode
  303. 11:29tokenizer um and that's what GPT uses
  304. 11:33and you can also just encode words into
  305. 11:35like hell world into a list of integers
  306. 11:38so as an example I'm using the Tik token
  307. 11:40Library here I'm getting the encoding
  308. 11:43for gpt2 or that was used for gpt2
  309. 11:46instead of just having 65 possible
  310. 11:48characters or tokens they have 50,000
  311. 11:51tokens and so when they encode the exact
  312. 11:54same string High there we only get a
  313. 11:57list of three integers but those
  314. 11:59integers are not between 0 and 64 they
  315. 12:01are between Z and 5,
  316. 12:055,256 so basically you can trade off the
  317. 12:09code book size and the sequence lengths
  318. 12:12so you can have very long sequences of
  319. 12:13integers with very small vocabularies or
  320. 12:16we can have short um sequences of
  321. 12:20integers with very large vocabularies
  322. 12:23and so typically people use in practice
  323. 12:25these subword encodings but I'd like to
  324. 12:28keep our token ier very simple so we're
  325. 12:30using character level tokenizer and that
  326. 12:33means that we have very small code books
  327. 12:35we have very simple encode and decode
  328. 12:37functions uh but we do get very long
  329. 12:40sequences as a result but that's the
  330. 12:42level at which we're going to stick with
  331. 12:43this lecture because it's the simplest
  332. 12:45thing okay so now that we have an
  333. 12:46encoder and a decoder effectively a
  334. 12:49tokenizer we can tokenize the entire
  335. 12:51training set of Shakespeare so here's a
  336. 12:53chunk of code that does that and I'm
  337. 12:55going to start to use the pytorch
  338. 12:56library and specifically the torch.
  339. 12:58tensor from the pytorch library so we're
  340. 13:01going to take all of the text in tiny
  341. 13:03Shakespeare encode it and then wrap it
  342. 13:05into a torch. tensor to get the data
  343. 13:08tensor so here's what the data tensor
  344. 13:10looks like when I look at just the first
  345. 13:121,000 characters or the 1,000 elements
  346. 13:14of it so we see that we have a massive
  347. 13:16sequence of integers and this sequence
  348. 13:18of integers here is basically an
  349. 13:20identical translation of the first
  350. 13:2210,000 characters
  351. 13:24here so I believe for example that zero
  352. 13:27is a new line character and maybe one
  353. 13:29one is a space not 100% sure but from
  354. 13:32now on the entire data set of text is
  355. 13:34re-represented as just it's just
  356. 13:35stretched out as a single very large uh
  357. 13:38sequence of
  358. 13:39integers let me do one more thing before
  359. 13:41we move on here I'd like to separate out
  360. 13:43our data set into a train and a
  361. 13:45validation split so in particular we're
  362. 13:48going to take the first 90% of the data
  363. 13:51set and consider that to be the training
  364. 13:52data for the Transformer and we're going
  365. 13:54to withhold the last 10% at the end of
  366. 13:56it to be the validation data and this
  367. 13:59will help us understand to what extent
  368. 14:01our model is overfitting so we're going
  369. 14:03to basically hide and keep the
  370. 14:04validation data on the side because we
  371. 14:06don't want just a perfect memorization
  372. 14:08of this exact Shakespeare we want a
  373. 14:11neural network that sort of creates
  374. 14:12Shakespeare like uh text and so it
  375. 14:15should be fairly likely for it to
  376. 14:17produce the actual like stowed away uh
  377. 14:21true Shakespeare text um and so we're
  378. 14:24going to use this to uh get a sense of
  379. 14:26the overfitting okay so now we would
  380. 14:28like to start plugging these text
  381. 14:30sequences or integer sequences into the
  382. 14:32Transformer so that it can train and
  383. 14:34learn those patterns now the important
  384. 14:36thing to realize is we're never going to
  385. 14:38actually feed entire text into a
  386. 14:40Transformer all at once that would be
  387. 14:42computationally very expensive and
  388. 14:44prohibitive so when we actually train a
  389. 14:46Transformer on a lot of these data sets
  390. 14:48we only work with chunks of the data set
  391. 14:50and when we train the Transformer we
  392. 14:52basically sample random little chunks
  393. 14:53out of the training set and train on
  394. 14:55just chunks at a time and these chunks
  395. 14:58have basically some kind of a length and
  396. 15:01some maximum length now the maximum
  397. 15:04length typically at least in the code I
  398. 15:06usually write is called block size you
  399. 15:08can you can uh find it under different
  400. 15:10names like context length or something
  401. 15:12like that let's start with the block
  402. 15:14size of just eight and let me look at
  403. 15:16the first train data characters the
  404. 15:18first block size plus one characters
  405. 15:20I'll explain why plus one in a
  406. 15:22second so this is the first nine
  407. 15:24characters in the sequence in the
  408. 15:27training set now what I'd like to point
  409. 15:30out is that when you sample a chunk of
  410. 15:31data like this so say the these nine
  411. 15:34characters out of the training set this
  412. 15:36actually has multiple examples packed
  413. 15:38into it and uh that's because all of
  414. 15:41these characters follow each other and
  415. 15:43so what this thing is going to say when
  416. 15:47we plug it into a Transformer is we're
  417. 15:49going to actually simultaneously train
  418. 15:50it to make prediction at every one of
  419. 15:52these
  420. 15:53positions now in the in a chunk of nine
  421. 15:56characters there's actually eight indiv
  422. 15:58ual examples packed in there so there's
  423. 16:01the example that when 18 when in the
  424. 16:04context of 18 47 likely comes next in a
  425. 16:08context of 18 and 47 56 comes next in a
  426. 16:12context of 18 47 56 57 can come next and
  427. 16:16so on so that's the eight individual
  428. 16:18examples let me actually spell it out
  429. 16:20with
  430. 16:21code so here's a chunk of code to
  431. 16:24illustrate X are the inputs to the
  432. 16:26Transformer it will just be the first
  433. 16:28block size characters y will be the uh
  434. 16:32next block size characters so it's
  435. 16:34offset by one and that's because y are
  436. 16:37the targets for each position in the
  437. 16:40input and then here I'm iterating over
  438. 16:42all the block size of eight and the
  439. 16:45context is always all the characters in
  440. 16:47x uh up to T and including T and the
  441. 16:51target is always the teth character but
  442. 16:53in the targets array y so let me just
  443. 16:56run
  444. 16:57this and basically it spells out what I
  445. 16:59said in words uh these are the eight
  446. 17:02examples hidden in a chunk of nine
  447. 17:04characters that we uh sampled from the
  448. 17:08training set I want to mention one more
  449. 17:11thing we train on all the eight examples
  450. 17:14here with context between one all the
  451. 17:16way up to context of block size and we
  452. 17:19train on that not just for computational
  453. 17:20reasons because we happen to have the
  454. 17:22sequence already or something like that
  455. 17:23it's not just done for efficiency it's
  456. 17:26also done um to make the Transformer
  457. 17:28Network be used to seeing contexts all
  458. 17:32the way from as little as one all the
  459. 17:33way to block size and we'd like the
  460. 17:36transform to be used to seeing
  461. 17:38everything in between and that's going
  462. 17:39to be useful later during inference
  463. 17:41because while we're sampling we can
  464. 17:43start the sampling generation with as
  465. 17:45little as one character of context and
  466. 17:47the Transformer knows how to predict the
  467. 17:49next character with all the way up to
  468. 17:51just context of one and so then it can
  469. 17:53predict everything up to block size and
  470. 17:55after block size we have to start
  471. 17:56truncating because the Transformer will
  472. 17:58will never um receive more than block
  473. 18:01size inputs when it's predicting the
  474. 18:03next
  475. 18:03character Okay so we've looked at the
  476. 18:06time dimension of the tensors that are
  477. 18:07going to be feeding into the Transformer
  478. 18:09there's one more Dimension to care about
  479. 18:11and that is the batch Dimension and so
  480. 18:13as we're sampling these chunks of text
  481. 18:15we're going to be actually every time
  482. 18:17we're going to feed them into a
  483. 18:18Transformer we're going to have many
  484. 18:20batches of multiple chunks of text that
  485. 18:22are all like stacked up in a single
  486. 18:23tensor and that's just done for
  487. 18:25efficiency just so that we can keep the
  488. 18:27gpus busy uh because they are very good
  489. 18:29at parallel processing of um of data and
  490. 18:33so we just want to process multiple
  491. 18:35chunks all at the same time but those
  492. 18:37chunks are processed completely
  493. 18:38independently they don't talk to each
  494. 18:39other and so on so let me basically just
  495. 18:42generalize this and introduce a batch
  496. 18:44Dimension here's a chunk of
  497. 18:46code let me just run it and then I'm
  498. 18:48going to explain what it
  499. 18:50does so here because we're going to
  500. 18:52start sampling random locations in the
  501. 18:54data set to pull chunks from I am
  502. 18:57setting the seed so that um in the
  503. 19:00random number generator so that the
  504. 19:01numbers I see here are going to be the
  505. 19:02same numbers you see later if you try to
  506. 19:04reproduce this now the batch size here
  507. 19:07is how many independent sequences we are
  508. 19:09processing every forward backward pass
  509. 19:11of the
  510. 19:12Transformer the block size as I
  511. 19:14explained is the maximum context length
  512. 19:16to make those predictions so let's say B
  513. 19:19size four block size eight and then
  514. 19:21here's how we get batch for any
  515. 19:23arbitrary split if the split is a
  516. 19:25training split then we're going to look
  517. 19:26at train data otherwise at valid data
  518. 19:30that gives us the data array and then
  519. 19:33when I Generate random positions to grab
  520. 19:35a chunk out of I actually grab I
  521. 19:38actually generate batch size number of
  522. 19:41Random offsets so because this is four
  523. 19:44we are ex is going to be a uh four
  524. 19:47numbers that are randomly generated
  525. 19:49between zero and Len of data minus block
  526. 19:51size so it's just random offsets into
  527. 19:53the training
  528. 19:54set and then X's as I explained are the
  529. 19:58first first block size characters
  530. 20:00starting at I the Y's are the offset by
  531. 20:05one of that so just add plus one and
  532. 20:08then we're going to get those chunks for
  533. 20:10every one of integers I INX and use a
  534. 20:14torch. stack to take all those uh uh
  535. 20:17one-dimensional tensors as we saw here
  536. 20:20and we're going to um stack them up at
  537. 20:24rows and so they all become a row in a
  538. 20:274x8 tensor
  539. 20:29so here's where I'm printing then when I
  540. 20:32sample a batch XB and YB the inputs to
  541. 20:35the Transformer now are the input X is
  542. 20:39the 4x8 tensor four uh rows of eight
  543. 20:44columns and each one of these is a chunk
  544. 20:47of the training
  545. 20:48set and then the targets here are in the
  546. 20:52associated array Y and they will come in
  547. 20:54to the Transformer all the way at the
  548. 20:55end uh to um create the loss function
  549. 20:59uh so they will give us the correct
  550. 21:01answer for every single position inside
  551. 21:03X and then these are the four
  552. 21:06independent
  553. 21:07rows so spelled out as we did
  554. 21:11before uh this 4x8 array contains a
  555. 21:14total of 32 examples and they're
  556. 21:17completely independent as far as the
  557. 21:19Transformer is
  558. 21:20concerned uh so when the input is 24 the
  559. 21:25target is 43 or rather 43 here in the Y
  560. 21:28array
  561. 21:29when the input is 2443 the target is
  562. 21:3158 uh when the input is 24 43 58 the
  563. 21:34target is 5 Etc or like when it is a 52
  564. 21:38581 the target is
  565. 21:4058 right so you can sort of see this
  566. 21:43spelled out these are the 32 independent
  567. 21:45examples packed in to a single batch of
  568. 21:48the input X and then the desired targets
  569. 21:51are in y and so now this integer tensor
  570. 21:57of um X is going to feed into the
  571. 22:00Transformer and that Transformer is
  572. 22:02going to simultaneously process all
  573. 22:04these examples and then look up the
  574. 22:06correct um integers to predict in every
  575. 22:08one of these positions in the tensor y
  576. 22:11okay so now that we have our batch of
  577. 22:13input that we'd like to feed into a
  578. 22:15Transformer let's start basically
  579. 22:16feeding this into neural networks now
  580. 22:19we're going to start off with the
  581. 22:20simplest possible neural network which
  582. 22:22in the case of language modeling in my
  583. 22:23opinion is the Byram language model and
  584. 22:25we've covered the Byram language model
  585. 22:26in my make more series in a lot of depth
  586. 22:29and so here I'm going to sort of go
  587. 22:31faster and let's just Implement pytorch
  588. 22:33module directly that implements the byr
  589. 22:36language
  590. 22:36model so I'm importing the pytorch um NN
  591. 22:41module uh for
  592. 22:43reproducibility and then here I'm
  593. 22:44constructing a Byram language model
  594. 22:46which is a subass of NN
  595. 22:48module and then I'm calling it and I'm
  596. 22:51passing it the inputs and the targets
  597. 22:53and I'm just printing now when the
  598. 22:55inputs on targets come here you see that
  599. 22:57I'm just taking the index uh the inputs
  600. 23:00X here which I rename to idx and I'm
  601. 23:03just passing them into this token
  602. 23:04embedding table so it's going on here is
  603. 23:07that here in the Constructor we are
  604. 23:09creating a token embedding table and it
  605. 23:12is of size vocap size by vocap
  606. 23:15size and we're using an. embedding which
  607. 23:18is a very thin wrapper around basically
  608. 23:20a tensor of shape voap size by vocab
  609. 23:23size and what's happening here is that
  610. 23:25when we pass idx here every single
  611. 23:28integer in our input is going to refer
  612. 23:30to this embedding table and it's going
  613. 23:32to pluck out a row of that embedding
  614. 23:34table corresponding to its index so 24
  615. 23:37here will go into the embedding table
  616. 23:39and we'll pluck out the 24th row and
  617. 23:42then 43 will go here and pluck out the
  618. 23:4443d row Etc and then pytorch is going to
  619. 23:47arrange all of this into a batch by Time
  620. 23:50by channel uh tensor in this case batch
  621. 23:53is four time is eight and C which is the
  622. 23:57channels is vocab size or 65 and so
  623. 24:01we're just going to pluck out all those
  624. 24:02rows arrange them in a b by T by C and
  625. 24:05now we're going to interpret this as the
  626. 24:07logits which are basically the scores
  627. 24:10for the next character in the sequence
  628. 24:12and so what's happening here is we are
  629. 24:14predicting what comes next based on just
  630. 24:17the individual identity of a single
  631. 24:19token and you can do that because um I
  632. 24:22mean currently the tokens are not
  633. 24:23talking to each other and they're not
  634. 24:25seeing any context except for they're
  635. 24:26just seeing themselves so I'm a f I'm a
  636. 24:29token number five and then I can
  637. 24:32actually make pretty decent predictions
  638. 24:33about what comes next just by knowing
  639. 24:35that I'm token five because some
  640. 24:37characters uh know um C follow other
  641. 24:39characters in in typical scenarios so we
  642. 24:42saw a lot of this in a lot more depth in
  643. 24:44the make more series and here if I just
  644. 24:46run this then we currently get the
  645. 24:49predictions the scores the lits for
  646. 24:53every one of the 4x8 positions now that
  647. 24:55we've made predictions about what comes
  648. 24:57next we'd like to evaluate the loss
  649. 24:58function and so in make more series we
  650. 25:00saw that a good way to measure a loss or
  651. 25:03like a quality of the predictions is to
  652. 25:05use the negative log likelihood loss
  653. 25:07which is also implemented in pytorch
  654. 25:09under the name cross entropy so what we'
  655. 25:12like to do here is loss is the cross
  656. 25:15entropy on the predictions and the
  657. 25:17targets and so this measures the quality
  658. 25:20of the logits with respect to the
  659. 25:21Targets in other words we have the
  660. 25:24identity of the next character so how
  661. 25:26well are we predicting the next
  662. 25:28character based on the lits and
  663. 25:30intuitively the correct um the correct
  664. 25:33dimension of low jits uh depending on
  665. 25:36whatever the target is should have a
  666. 25:38very high number and all the other
  667. 25:39dimensions should be very low number
  668. 25:41right now the issue is that this won't
  669. 25:44actually this is what we want we want to
  670. 25:46basically output the logits and the
  671. 25:50loss this is what we want but
  672. 25:52unfortunately uh this won't actually run
  673. 25:55we get an error message but intuitively
  674. 25:57we want to uh measure this now when we
  675. 26:01go to the pytorch um cross entropy
  676. 26:04documentation here um we're trying to
  677. 26:08call the cross entropy in its functional
  678. 26:10form uh so that means we don't have to
  679. 26:11create like a module for it but here
  680. 26:14when we go to the documentation you have
  681. 26:16to look into the details of how pitor
  682. 26:18expects these inputs and basically the
  683. 26:20issue here is ptor expects if you have
  684. 26:24multi-dimensional input which we do
  685. 26:25because we have a b BYT by C tensor then
  686. 26:28it actually really wants the channels to
  687. 26:31be the second uh Dimension here so if
  688. 26:35you um so basically it wants a b by C
  689. 26:38BYT instead of a b by T by C and so it's
  690. 26:42just the details of how P torch treats
  691. 26:45um these kinds of inputs and so we don't
  692. 26:49actually want to deal with that so what
  693. 26:51we're going to do instead is we need to
  694. 26:52basically reshape our logits so here's
  695. 26:54what I like to do I like to take
  696. 26:56basically give names to the dimensions
  697. 26:58so lit. shape is B BYT by C and unpack
  698. 27:01those numbers and then let's uh say that
  699. 27:04logits equals lit. View and we want it
  700. 27:07to be a b * c b * T by C so just a two-
  701. 27:11dimensional
  702. 27:12array right so we're going to take all
  703. 27:15the we're going to take all of these um
  704. 27:18positions here and we're going to uh
  705. 27:20stretch them out in a onedimensional
  706. 27:22sequence and uh preserve the channel
  707. 27:25Dimension as the second
  708. 27:26dimension so we're just kind of like
  709. 27:28stretching out the array so it's two-
  710. 27:29dimensional and in that case it's going
  711. 27:31to better conform to what pytorch uh
  712. 27:33sort of expects in its Dimensions now we
  713. 27:36have to do the same to targets because
  714. 27:38currently targets are um of shape B by T
  715. 27:44and we want it to be just B * T so
  716. 27:47onedimensional now alternatively you
  717. 27:49could always still just do minus one
  718. 27:51because pytor will guess what this
  719. 27:53should be if you want to lay it out uh
  720. 27:55but let me just be explicit and say p *
  721. 27:57t once we've reshaped this it will match
  722. 28:00the cross entropy case and then we
  723. 28:03should be able to evaluate our
  724. 28:06loss okay so that R now and we can do
  725. 28:10loss and So currently we see that the
  726. 28:12loss is
  727. 28:134.87 now because our uh we have 65
  728. 28:17possible vocabulary elements we can
  729. 28:19actually guess at what the loss should
  730. 28:20be and in
  731. 28:22particular we covered negative log
  732. 28:24likelihood in a lot of detail we are
  733. 28:26expecting log or lawn of um 1 over 65
  734. 28:32and negative of that so we're expecting
  735. 28:34the loss to be about 4.1 17 but we're
  736. 28:37getting 4.87 and so that's telling us
  737. 28:40that the initial predictions are not uh
  738. 28:42super diffuse they've got a little bit
  739. 28:43of entropy and so we're guessing wrong
  740. 28:47uh so uh yes but actually we're I a we
  741. 28:50are able to evaluate the loss okay so
  742. 28:53now that we can evaluate the quality of
  743. 28:54the model on some data we'd like to also
  744. 28:57be able to generate from the model so
  745. 28:59let's do the generation now I'm going to
  746. 29:01go again a little bit faster here
  747. 29:03because I covered all this already in
  748. 29:04previous
  749. 29:05videos
  750. 29:07so here's a generate function for the
  751. 29:11model so we take some uh we take the the
  752. 29:15same kind of input idx here and
  753. 29:18basically this is the current uh context
  754. 29:22of some characters in a batch in some
  755. 29:24batch so it's also B BYT and the job of
  756. 29:28generate is to basically take this B BYT
  757. 29:30and extend it to be B BYT + 1 plus 2
  758. 29:32plus 3 and so it's just basically it
  759. 29:34continues the generation in all the
  760. 29:36batch dimensions in the time Dimension
  761. 29:39So that's its job and it will do that
  762. 29:41for Max new tokens so you can see here
  763. 29:43on the bottom there's going to be some
  764. 29:45stuff here but on the bottom whatever is
  765. 29:47predicted is concatenated on top of the
  766. 29:50previous idx along the First Dimension
  767. 29:53which is the time Dimension to create a
  768. 29:54b BYT + one so that becomes a new idx so
  769. 29:58the job of generate is to take a b BYT
  770. 30:00and make it a b BYT plus 1 plus 2 plus
  771. 30:02three as many as we want Max new tokens
  772. 30:05so this is the generation from the model
  773. 30:08now inside the generation what what are
  774. 30:10we doing we're taking the current
  775. 30:11indices we're getting the predictions so
  776. 30:15we get uh those are in the low jits and
  777. 30:18then the loss here is going to be
  778. 30:19ignored because um we're not we're not
  779. 30:21using that and we have no targets that
  780. 30:23are sort of ground truth targets that
  781. 30:25we're going to be comparing with
  782. 30:28then once we get the logits we are only
  783. 30:30focusing on the last step so instead of
  784. 30:33a b by T by C we're going to pluck out
  785. 30:36the negative-1 the last element in the
  786. 30:38time Dimension because those are the
  787. 30:40predictions for what comes next so that
  788. 30:42gives us the logits which we then
  789. 30:44convert to probabilities via softmax and
  790. 30:47then we use tor. multinomial to sample
  791. 30:49from those probabilities and we ask
  792. 30:51pytorch to give us one sample and so idx
  793. 30:54next will become a b by one because in
  794. 30:57each uh one of the batch Dimensions
  795. 31:00we're going to have a single prediction
  796. 31:01for what comes next so this num samples
  797. 31:03equals one will make this be a
  798. 31:06one and then we're going to take those
  799. 31:08integers that come from the sampling
  800. 31:10process according to the probability
  801. 31:11distribution given here and those
  802. 31:13integers got just concatenated on top of
  803. 31:15the current sort of like running stream
  804. 31:17of integers and this gives us a b BYT +
  805. 31:20one and then we can return that now one
  806. 31:24thing here is you see how I'm calling
  807. 31:26self of idx which will end up going to
  808. 31:29the forward function I'm not providing
  809. 31:31any Targets So currently this would give
  810. 31:33an error because targets is uh is uh
  811. 31:36sort of like not given so targets has to
  812. 31:39be optional so targets is none by
  813. 31:41default and then if targets is none then
  814. 31:44there's no loss to create so it's just
  815. 31:47loss is none but else all of this
  816. 31:50happens and we can create a loss so this
  817. 31:53will make it so um if we have the
  818. 31:56targets we provide them and get a loss
  819. 31:57if we have no targets it will'll just
  820. 31:59get the
  821. 32:00loits so this here will generate from
  822. 32:02the model um and let's take that for a
  823. 32:06ride
  824. 32:08now oops so I have another code chunk
  825. 32:11here which will generate for the model
  826. 32:13from the model and okay this is kind of
  827. 32:15crazy so maybe let me let me break this
  828. 32:18down so these are the idx
  829. 32:23right I'm creating a batch will be just
  830. 32:26one time will be just one so I'm
  831. 32:30creating a little one by one tensor and
  832. 32:32it's holding a zero and the D type the
  833. 32:35data type is uh integer so zero is going
  834. 32:38to be how we kick off the generation and
  835. 32:40remember that zero is uh is the element
  836. 32:44standing for a new line character so
  837. 32:45it's kind of like a reasonable thing to
  838. 32:47to feed in as the very first character
  839. 32:49in a sequence to be the new
  840. 32:51line um so it's going to be idx which
  841. 32:54we're going to feed in here then we're
  842. 32:56going to ask for 100 tokens
  843. 32:58and then. generate will continue that
  844. 33:01now because uh generate works on the
  845. 33:05level of batches we we then have to
  846. 33:07index into the zero throw to basically
  847. 33:09unplug the um the single batch Dimension
  848. 33:13that exists and then that gives us a um
  849. 33:18time steps just a onedimensional array
  850. 33:20of all the indices which we will convert
  851. 33:23to simple python list from pytorch
  852. 33:26tensor so that that can feed into our
  853. 33:28decode function and uh convert those
  854. 33:32integers into text so let me bring this
  855. 33:34back and we're generating 100 tokens
  856. 33:37let's
  857. 33:37run and uh here's the generation that we
  858. 33:40achieved so obviously it's garbage and
  859. 33:43the reason it's garbage is because this
  860. 33:44is a totally random model so next up
  861. 33:47we're going to want to train this model
  862. 33:49now one more thing I wanted to point out
  863. 33:50here is this function is written to be
  864. 33:53General but it's kind of like ridiculous
  865. 33:55right now because
  866. 33:58we're feeding in all this we're building
  867. 33:59out this context and we're concatenating
  868. 34:02it all and we're always feeding it all
  869. 34:05into the model but that's kind of
  870. 34:07ridiculous because this is just a simple
  871. 34:09Byram model so to make for example this
  872. 34:11prediction about K we only needed this W
  873. 34:14but actually what we fed into the model
  874. 34:15is we fed the entire sequence and then
  875. 34:18we only looked at the very last piece
  876. 34:20and predicted K so the only reason I'm
  877. 34:23writing it in this way is because right
  878. 34:25now this is a byr model but I'd like to
  879. 34:27keep keep this function fixed and I'd
  880. 34:29like it to work um later when our
  881. 34:32characters actually um basically look
  882. 34:35further in the history and so right now
  883. 34:37the history is not used so this looks
  884. 34:39silly uh but eventually the history will
  885. 34:42be used and so that's why we want to uh
  886. 34:44do it this way so just a quick comment
  887. 34:46on that so now we see that this is um
  888. 34:49random so let's train the model so it
  889. 34:51becomes a bit less random okay let's Now
  890. 34:53train the model so first what I'm going
  891. 34:55to do is I'm going to create a pyour
  892. 34:57optimization object so here we are using
  893. 35:00the optimizer ATM W um now in a make
  894. 35:05more series we've only ever use tastic
  895. 35:06gradi in descent the simplest possible
  896. 35:08Optimizer which you can get using the
  897. 35:10SGD instead but I want to use Adam which
  898. 35:12is a much more advanced and popular
  899. 35:14Optimizer and it works extremely well
  900. 35:16for uh typical good setting for the
  901. 35:19learning rate is roughly 3 E4 uh but for
  902. 35:22very very small networks like is the
  903. 35:23case here you can get away with much
  904. 35:25much higher learning rates R3 or even
  905. 35:28higher probably but let me create the
  906. 35:30optimizer object which will basically
  907. 35:33take the gradients and uh update the
  908. 35:35parameters using the
  909. 35:36gradients and then here our batch size
  910. 35:40up above was only four so let me
  911. 35:41actually use something bigger let's say
  912. 35:4332 and then for some number of steps um
  913. 35:46we are sampling a new batch of data
  914. 35:48we're evaluating the loss uh we're
  915. 35:51zeroing out all the gradients from the
  916. 35:52previous step getting the gradients for
  917. 35:54all the parameters and then using those
  918. 35:56gradients to up update our parameters so
  919. 35:58typical training loop as we saw in the
  920. 36:00make more series so let me now uh run
  921. 36:04this for say 100 iterations and let's
  922. 36:07see what kind of losses we're going to
  923. 36:09get so we started around
  924. 36:124.7 and now we're getting to down to
  925. 36:14like 4.6 4.5 Etc so the optimization is
  926. 36:18definitely happening but um let's uh
  927. 36:22sort of try to increase number of
  928. 36:23iterations and only print at the
  929. 36:25end because we probably want train for
  930. 36:29longer okay so we're down to 3.6
  931. 36:34roughly roughly down to
  932. 36:40three this is the most janky
  933. 36:46optimization okay it's working let's
  934. 36:48just do
  935. 36:5010,000 and then from here we want to
  936. 36:53copy this and hopefully that we're going
  937. 36:56to get something reason and of course
  938. 36:58it's not going to be Shakespeare from a
  939. 37:00byr model but at least we see that the
  940. 37:01loss is improving and uh hopefully we're
  941. 37:05expecting something a bit more
  942. 37:06reasonable okay so we're down at about
  943. 37:082.5 is let's see what we get okay
  944. 37:12dramatic improvements certainly on what
  945. 37:14we had here so let me just increase the
  946. 37:17number of tokens okay so we see that
  947. 37:19we're starting to get something at least
  948. 37:21like reasonable is
  949. 37:25um certainly not shakes spear but uh the
  950. 37:29model is making progress so that is the
  951. 37:31simplest possible
  952. 37:33model so now what I'd like to do
  953. 37:36is obviously this is a very simple model
  954. 37:39because the tokens are not talking to
  955. 37:41each other so given the previous context
  956. 37:43of whatever was generated we're only
  957. 37:45looking at the very last character to
  958. 37:46make the predictions about what comes
  959. 37:48next so now these uh now these tokens
  960. 37:50have to start talking to each other and
  961. 37:53figuring out what is in the context so
  962. 37:55that they can make better predictions
  963. 37:56for what comes next and this is how
  964. 37:57we're going to kick off the uh
  965. 37:59Transformer okay so next I took the code
  966. 38:02that we developed in this juper notebook
  967. 38:03and I converted it to be a script and
  968. 38:05I'm doing this because I just want to
  969. 38:08simplify our intermediate work into just
  970. 38:10the final product that we have at this
  971. 38:12point so in the top here I put all the
  972. 38:15hyp parameters that we to find I
  973. 38:16introduced a few and I'm going to speak
  974. 38:18to that in a little bit otherwise a lot
  975. 38:20of this should be recognizable uh
  976. 38:23reproducibility read data get the
  977. 38:25encoder and the decoder create the train
  978. 38:27into splits uh use the uh kind of like
  979. 38:30data loader um that gets a batch of the
  980. 38:34inputs and Targets this is new and I'll
  981. 38:36talk about it in a second now this is
  982. 38:39the Byram language model that we
  983. 38:40developed and it can forward and give us
  984. 38:43a logits and loss and it can
  985. 38:45generate and then here we are creating
  986. 38:48the optimizer and this is the training
  987. 38:51Loop so everything here should look
  988. 38:53pretty familiar now some of the small
  989. 38:55things that I added number one I added
  990. 38:57the ability to run on a GPU if you have
  991. 39:00it so if you have a GPU then you can
  992. 39:02this will use Cuda instead of just CPU
  993. 39:04and everything will be a lot more faster
  994. 39:07now when device becomes Cuda then we
  995. 39:09need to make sure that when we load the
  996. 39:11data we move it to
  997. 39:13device when we create the model we want
  998. 39:15to move uh the model parameters to
  999. 39:18device so as an example here we have the
  1000. 39:21N an embedding table and it's got a
  1001. 39:23weight inside it which stores the uh
  1002. 39:26sort of lookup table so so that would be
  1003. 39:27moved to the GPU so that all the
  1004. 39:29calculations here happen on the GPU and
  1005. 39:32they can be a lot faster and then
  1006. 39:34finally here when I'm creating the
  1007. 39:35context that feeds in to generate I have
  1008. 39:37to make sure that I create it on the
  1009. 39:39device number two what I introduced is
  1010. 39:43uh the fact that here in the training
  1011. 39:46Loop here I was just printing the um l.
  1012. 39:50item inside the training Loop but this
  1013. 39:53is a very noisy measurement of the
  1014. 39:54current loss because every batch will be
  1015. 39:56more or less lucky and so what I want to
  1016. 39:59do usually um is uh I have an estimate
  1017. 40:02loss function and the estimate loss
  1018. 40:05basically then um goes up here and it
  1019. 40:10averages up the loss over multiple
  1020. 40:12batches so in particular we're going to
  1021. 40:15iterate eval iter times and we're going
  1022. 40:17to basically get our loss and then we're
  1023. 40:19going to get the average loss for both
  1024. 40:21splits and so this will be a lot less
  1025. 40:24noisy so here when we call the estimate
  1026. 40:26loss we're we're going to report the uh
  1027. 40:28pretty accurate train and validation
  1028. 40:31loss now when we come back up you'll
  1029. 40:33notice a few things here I'm setting the
  1030. 40:35model to evaluation phase and down here
  1031. 40:38I'm resetting it back to training phase
  1032. 40:40now right now for our model as is this
  1033. 40:42doesn't actually do anything because the
  1034. 40:44only thing inside this model is this uh
  1035. 40:46nn. embedding and um this this um
  1036. 40:51Network would behave both would behave
  1037. 40:53the same in both evaluation mode and
  1038. 40:55training mode we have no drop off layers
  1039. 40:57we have no batm layers Etc but it is a
  1040. 41:00good practice to Think Through what mode
  1041. 41:02your neural network is in because some
  1042. 41:04layers will have different Behavior Uh
  1043. 41:07at inference time or training time and
  1044. 41:11there's also this context manager torch
  1045. 41:12up nograd and this is just telling
  1046. 41:14pytorch that everything that happens
  1047. 41:16inside this function we will not call do
  1048. 41:18backward on and so pytorch can be a lot
  1049. 41:21more efficient with its memory use
  1050. 41:23because it doesn't have to store all the
  1051. 41:25intermediate variables uh because we're
  1052. 41:27never going to call backward and so it
  1053. 41:29can it can be a lot more memory
  1054. 41:30efficient in that way so also a good
  1055. 41:32practice to tpy torch when we don't
  1056. 41:35intend to do back
  1057. 41:36propagation so right now this script is
  1058. 41:39about 120 lines of code of and that's
  1059. 41:43kind of our starter code I'm calling it
  1060. 41:45b.p and I'm going to release it later
  1061. 41:48now running this
  1062. 41:50script gives us output in the terminal
  1063. 41:52and it looks something like this it
  1064. 41:54basically as I ran this code uh it was
  1065. 41:57giving me the train loss and Val loss
  1066. 41:59and we see that we convert to somewhere
  1067. 42:01around
  1068. 42:012.5 with the pyr model and then here's
  1069. 42:04the sample that we produced at the
  1070. 42:07end and so we have everything packaged
  1071. 42:09up in the script and we're in a good
  1072. 42:11position now to iterate on this okay so
  1073. 42:13we are almost ready to start writing our
  1074. 42:15very first self attention block for
  1075. 42:18processing these uh tokens now before we
  1076. 42:22actually get there I want to get you
  1077. 42:24used to a mathematical trick that is
  1078. 42:26used in the self attention inside a
  1079. 42:28Transformer and is really just like at
  1080. 42:30the heart of an an efficient
  1081. 42:32implementation of self attention and so
  1082. 42:34I want to work with this toy example to
  1083. 42:36just get you used to this operation and
  1084. 42:38then it's going to make it much more
  1085. 42:39clear once we actually get to um to it
  1086. 42:43uh in the script
  1087. 42:44again so let's create a b BYT by C where
  1088. 42:47BT and C are just 48 and two in the toy
  1089. 42:50example and these are basically channels
  1090. 42:53and we have uh batches and we have the
  1091. 42:55time component and we have information
  1092. 42:58at each point in the sequence so
  1093. 43:01see now what we would like to do is we
  1094. 43:03would like these um tokens so we have up
  1095. 43:06to eight tokens here in a batch and
  1096. 43:08these eight tokens are currently not
  1097. 43:10talking to each other and we would like
  1098. 43:11them to talk to each other we'd like to
  1099. 43:13couple them and in particular we don't
  1100. 43:17we we want to couple them in a very
  1101. 43:18specific way so the token for example at
  1102. 43:21the fifth location it should not
  1103. 43:23communicate with tokens in the sixth
  1104. 43:25seventh and eighth location
  1105. 43:27because uh those are future tokens in
  1106. 43:29the sequence the token on the fifth
  1107. 43:31location should only talk to the one in
  1108. 43:33the fourth third second and first so
  1109. 43:36it's only so information only flows from
  1110. 43:38previous context to the current time
  1111. 43:40step and we cannot get any information
  1112. 43:42from the future because we are about to
  1113. 43:44try to predict the
  1114. 43:45future so what is the easiest way for
  1115. 43:49tokens to communicate okay the easiest
  1116. 43:52way I would say is okay if we're up to
  1117. 43:54if we're a fifth token and I'd like to
  1118. 43:56communicate with my past the simplest
  1119. 43:58way we can do that is to just do a
  1120. 44:00weight is to just do an average of all
  1121. 44:03the um of all the preceding elements so
  1122. 44:06for example if I'm the fif token I would
  1123. 44:08like to take the channels uh that make
  1124. 44:10up that are information at my step but
  1125. 44:13then also the channels from the fourth
  1126. 44:15step third step second step and the
  1127. 44:17first step I'd like to average those up
  1128. 44:19and then that would become sort of like
  1129. 44:21a feature Vector that summarizes me in
  1130. 44:23the context of my history now of course
  1131. 44:26just doing a sum or like an average is
  1132. 44:28an extremely weak form of interaction
  1133. 44:30like this communication is uh extremely
  1134. 44:32lossy we've lost a ton of information
  1135. 44:34about the spatial Arrangements of all
  1136. 44:35those tokens uh but that's okay for now
  1137. 44:38we'll see how we can bring that
  1138. 44:39information back later for now what we
  1139. 44:41would like to do is for every single
  1140. 44:43batch element independently for every
  1141. 44:46teeth token in that sequence we'd like
  1142. 44:49to now calculate the average of all the
  1143. 44:53vectors in all the previous tokens and
  1144. 44:55also at this token so let's write that
  1145. 44:58out um I have a small snippet here and
  1146. 45:01instead of just fumbling around let me
  1147. 45:03just copy paste it and talk to
  1148. 45:05it so in other words we're going to
  1149. 45:08create X and B is short for bag of words
  1150. 45:12because bag of words is um is kind of
  1151. 45:15like um a term that people use when you
  1152. 45:17are just averaging up things so this is
  1153. 45:19just a bag of words basically there's a
  1154. 45:21word stored on every one of these eight
  1155. 45:23locations and we're doing a bag of words
  1156. 45:25we're just averaging
  1157. 45:27so in the beginning we're going to say
  1158. 45:28that it's just initialized at Zero and
  1159. 45:30then I'm doing a for Loop here so we're
  1160. 45:32not being efficient yet that's coming
  1161. 45:34but for now we're just iterating over
  1162. 45:36all the batch Dimensions independently
  1163. 45:38iterating over time and then the
  1164. 45:40previous uh tokens are at this uh batch
  1165. 45:45Dimension and then everything up to and
  1166. 45:47including the teeth token okay so when
  1167. 45:51we slice out X in this way X prev
  1168. 45:54Becomes of shape um how many T elements
  1169. 45:58there were in the past and then of
  1170. 46:00course C so all the two-dimensional
  1171. 46:02information from these little tokens so
  1172. 46:05that's the previous uh sort of chunk of
  1173. 46:08um tokens from my current sequence and
  1174. 46:12then I'm just doing the average or the
  1175. 46:13mean over the zero Dimension so I'm
  1176. 46:15averaging out the time here and I'm just
  1177. 46:19going to get a little c one dimensional
  1178. 46:21Vector which I'm going to store in X bag
  1179. 46:23of words so I can run this and and uh
  1180. 46:27this is not going to be very informative
  1181. 46:30because let's see so this is X of Zer so
  1182. 46:32this is the zeroth batch element and
  1183. 46:35then expo at zero now you see how the at
  1184. 46:40the first location here you see that the
  1185. 46:42two are equal and that's because it's
  1186. 46:45we're just doing an average of this one
  1187. 46:46token but here this one is now an
  1188. 46:49average of these two and now this one is
  1189. 46:53an average of these
  1190. 46:54three and so on
  1191. 46:57so uh and this last one is the average
  1192. 47:01of all of these elements so vertical
  1193. 47:03average just averaging up all the tokens
  1194. 47:05now gives this outcome
  1195. 47:07here so this is all well and good uh but
  1196. 47:10this is very inefficient now the trick
  1197. 47:12is that we can be very very efficient
  1198. 47:14about doing this using matrix
  1199. 47:16multiplication so that's the
  1200. 47:18mathematical trick and let me show you
  1201. 47:19what I mean let's work with the toy
  1202. 47:21example here let me run it and I'll
  1203. 47:24explain I have a simple Matrix here that
  1204. 47:27is a 3X3 of all ones a matrix B of just
  1205. 47:31random numbers and it's a 3x2 and a
  1206. 47:33matrix C which will be 3x3 multip 3x2
  1207. 47:36which will give out a 3x2 so here we're
  1208. 47:39just using um matrix multiplication so a
  1209. 47:43multiply B gives us
  1210. 47:46C okay so how are these numbers in C um
  1211. 47:51achieved right so this number in the top
  1212. 47:54left is the first row of a dot product
  1213. 47:57with the First Column of B and since all
  1214. 48:00the the row of a right now is all just
  1215. 48:02ones then the do product here with with
  1216. 48:05this column of B is just going to do a
  1217. 48:07sum of these of this column so 2 + 6 + 6
  1218. 48:11is
  1219. 48:1214 the element here in the output of C
  1220. 48:15is also the first column here the first
  1221. 48:17row of a multiplied now with the second
  1222. 48:20column of B so 7 + 4 + 5 is 16 now you
  1223. 48:25see that there's repeating elements here
  1224. 48:26so this 14 again is because this row is
  1225. 48:28again all ones and it's multiplying the
  1226. 48:30First Column of B so we get 14 and this
  1227. 48:33one is and so on so this last number
  1228. 48:35here is the last row do product last
  1229. 48:39column now the trick here is uh the
  1230. 48:42following this is just a boring number
  1231. 48:44of um it's just a boring array of all
  1232. 48:48ones but torch has this function called
  1233. 48:50Trail which is short for a
  1234. 48:54triangular uh something like that and
  1235. 48:56you can wrap it in torch up once and it
  1236. 48:58will just return the lower triangular
  1237. 49:00portion of this
  1238. 49:03okay so now it will basically zero out
  1239. 49:06uh these guys here so we just get the
  1240. 49:08lower triangular part well what happens
  1241. 49:10if we do
  1242. 49:14that so now we'll have a like this and B
  1243. 49:17like this and now what are we getting
  1244. 49:18here in C well what is this number well
  1245. 49:22this is the first row times the First
  1246. 49:24Column and because this is zeros
  1247. 49:28uh these elements here are now ignored
  1248. 49:30so we just get a two and then this
  1249. 49:32number here is the first row times the
  1250. 49:35second column and because these are
  1251. 49:37zeros they get ignored and it's just
  1252. 49:39seven this seven multiplies this one but
  1253. 49:42look what happened here because this is
  1254. 49:43one and then zeros we what ended up
  1255. 49:46happening is we're just plucking out the
  1256. 49:48row of this row of B and that's what we
  1257. 49:51got now here we have one 1 Z so here 110
  1258. 49:57do product with these two columns will
  1259. 49:59now give us 2 + 6 which is 8 and 7 + 4
  1260. 50:02which is 11 and because this is 111 we
  1261. 50:05ended up with the addition of all of
  1262. 50:07them and so basically depending on how
  1263. 50:10many ones and zeros we have here we are
  1264. 50:12basically doing a sum currently of a
  1265. 50:16variable number of these rows and that
  1266. 50:18gets deposited into
  1267. 50:20C So currently we're doing sums because
  1268. 50:23these are ones but we can also do
  1269. 50:25average right and you can start to see
  1270. 50:27how we could do average uh of the rows
  1271. 50:29of B uh sort of in an incremental
  1272. 50:32fashion because we don't have to we can
  1273. 50:35basically normalize these rows so that
  1274. 50:37they sum to one and then we're going to
  1275. 50:39get an average so if we took a and then
  1276. 50:41we did aals
  1277. 50:43aide torch. sum in the um of a in the um
  1278. 50:51oneth Dimension and then let's keep them
  1279. 50:55as true so so therefore the broadcasting
  1280. 50:57will work out so if I rerun this you see
  1281. 51:00now that these rows now sum to one so
  1282. 51:04this row is one this row is 0. 5.5 Z and
  1283. 51:07here we get 1/3 and now when we do a
  1284. 51:09multiply B what are we getting here we
  1285. 51:12are just getting the first row first row
  1286. 51:15here now we are getting the average of
  1287. 51:18the first two
  1288. 51:20rows okay so 2 and six average is four
  1289. 51:23and four and seven average is
  1290. 51:255.5 and on the bottom here we are now
  1291. 51:27getting the average of these three rows
  1292. 51:31so the average of all of elements of B
  1293. 51:33are now deposited here and so you can
  1294. 51:36see that by manipulating these uh
  1295. 51:40elements of this multiplying Matrix and
  1296. 51:42then multiplying it with any given
  1297. 51:44Matrix we can do these averages in this
  1298. 51:47incremental fashion because we just get
  1299. 51:50um and we can manipulate that based on
  1300. 51:53the elements of a okay so that's very
  1301. 51:55convenient so let's let's swing back up
  1302. 51:57here and see how we can vectorize this
  1303. 51:59and make it much more efficient using
  1304. 52:00what we've learned so in
  1305. 52:03particular we are going to produce an
  1306. 52:05array a but here I'm going to call it we
  1307. 52:08short for weights but this is our
  1308. 52:11a and this is how much of every row we
  1309. 52:14want to average up and it's going to be
  1310. 52:17an average because you can see that
  1311. 52:18these rows sum to
  1312. 52:20one so this is our a and then our B in
  1313. 52:23this example of course is X
  1314. 52:27so what's going to happen here now is
  1315. 52:29that we are going to have an expo
  1316. 52:312 and this Expo 2 is going to be way
  1317. 52:36multiplying
  1318. 52:38RX so let's think this true way is T BYT
  1319. 52:42and this is Matrix multiplying in
  1320. 52:44pytorch a b by T by
  1321. 52:47C and it's giving us uh different what
  1322. 52:50shape so pytorch will come here and it
  1323. 52:52will see that these shapes are not the
  1324. 52:54same so it will create a batch Dimension
  1325. 52:57here and this is a batched matrix
  1326. 53:00multiply and so it will apply this
  1327. 53:02matrix multiplication in all the batch
  1328. 53:04elements um in parallel and individually
  1329. 53:08and then for each batch element there
  1330. 53:09will be a t BYT multiplying T by C
  1331. 53:12exactly as we had
  1332. 53:15below so this will now create B by T by
  1333. 53:20C and Expo 2 will now become identical
  1334. 53:24to Expo
  1335. 53:28so we can see that torch. all close of
  1336. 53:32xbo and xbo 2 should be true
  1337. 53:36now so this kind of like convinces us
  1338. 53:38that uh these are in fact um the same so
  1339. 53:43xbo and xbo 2 if I just print
  1340. 53:47them uh okay we're not going to be able
  1341. 53:49to okay we're not going to be able to
  1342. 53:51just stare it down but
  1343. 53:54um well let me try Expo basically just
  1344. 53:56at the zeroth element and Expo two at
  1345. 53:58the zeroth element so just the first
  1346. 53:59batch and we should see that this and
  1347. 54:02that should be identical which they
  1348. 54:04are right so what happened here the
  1349. 54:07trick is we were able to use batched
  1350. 54:09Matrix multiply to do this uh
  1351. 54:12aggregation really and it's a weighted
  1352. 54:15aggregation and the weights are
  1353. 54:17specified in this um T BYT array and
  1354. 54:21we're basically doing weighted sums and
  1355. 54:24uh these weighted sums are are U
  1356. 54:26according to uh the weights inside here
  1357. 54:28they take on sort of this triangular
  1358. 54:31form and so that means that a token at
  1359. 54:33the teth dimension will only get uh sort
  1360. 54:36of um information from the um tokens
  1361. 54:39perceiving it so that's exactly what we
  1362. 54:41want and finally I would like to rewrite
  1363. 54:43it in one more way and we're going to
  1364. 54:46see why that's useful so this is the
  1365. 54:48third version and it's also identical to
  1366. 54:50the first and second but let me talk
  1367. 54:53through it it uses
  1368. 54:54softmax so Trill here is this Matrix
  1369. 55:00lower triangular
  1370. 55:01ones way begins as all
  1371. 55:05zero okay so if I just print way in the
  1372. 55:07beginning it's all zero then I
  1373. 55:11used masked fill so what this is doing
  1374. 55:15is we. masked fill it's all zeros and
  1375. 55:18I'm saying for all the elements where
  1376. 55:20Trill is equal equal Z make them be
  1377. 55:23negative Infinity so all the elements
  1378. 55:26where Trill is zero will become negative
  1379. 55:28Infinity now so this is what we get and
  1380. 55:32then the final line here is
  1381. 55:36softmax so if I take a softmax along
  1382. 55:38every single so dim is negative one so
  1383. 55:40along every single row if I do softmax
  1384. 55:44what is that going to
  1385. 55:46do well softmax is um is also like a
  1386. 55:51normalization operation right and so
  1387. 55:54spoiler alert you get the exact same
  1388. 55:58Matrix let me bring back to
  1389. 56:00softmax and recall that in softmax we're
  1390. 56:02going to exponentiate every single one
  1391. 56:04of these and then we're going to divide
  1392. 56:06by the sum and so if we exponentiate
  1393. 56:10every single element here we're going to
  1394. 56:11get a one and here we're going to get uh
  1395. 56:14basically zero 0 z0 Z everywhere else
  1396. 56:17and then when we normalize we just get
  1397. 56:19one here we're going to get one one and
  1398. 56:21then zeros and then softmax will again
  1399. 56:24divide and this will give us 5.5 and so
  1400. 56:27on and so this is also the uh the same
  1401. 56:30way to produce uh this mask now the
  1402. 56:33reason that this is a bit more
  1403. 56:34interesting and the reason we're going
  1404. 56:36to end up using it in self
  1405. 56:37attention is that these weights here
  1406. 56:41begin uh with zero and you can think of
  1407. 56:44this as like an interaction strength or
  1408. 56:46like an affinity so basically it's
  1409. 56:49telling us how much of each uh token
  1410. 56:52from the past do we want to Aggregate
  1411. 56:54and average up
  1412. 56:57and then this line is saying tokens from
  1413. 56:59the past cannot communicate by setting
  1414. 57:02them to negative Infinity we're saying
  1415. 57:04that we will not aggregate anything from
  1416. 57:06those
  1417. 57:07tokens and so basically this then goes
  1418. 57:09through softmax and through the weighted
  1419. 57:11and this is the aggregation through
  1420. 57:12matrix
  1421. 57:14multiplication and so what this is now
  1422. 57:16is you can think of these as um these
  1423. 57:19zeros are currently just set by us to be
  1424. 57:21zero but a quick preview is that these
  1425. 57:25affinities between the tokens are not
  1426. 57:27going to be just constant at zero
  1427. 57:29they're going to be data dependent these
  1428. 57:31tokens are going to start looking at
  1429. 57:32each other and some tokens will find
  1430. 57:34other tokens more or less interesting
  1431. 57:37and depending on what their values are
  1432. 57:39they're going to find each other
  1433. 57:41interesting to different amounts and I'm
  1434. 57:42going to call those affinities I think
  1435. 57:45and then here we are saying the future
  1436. 57:47cannot communicate with the past we're
  1437. 57:49we're going to clamp them and then when
  1438. 57:51we normalize and sum we're going to
  1439. 57:53aggregate uh sort of their values
  1440. 57:56depending on how interesting they find
  1441. 57:57each other and so that's the preview for
  1442. 57:59self attention and basically long story
  1443. 58:03short from this entire section is that
  1444. 58:05you can do weighted aggregations of your
  1445. 58:07past
  1446. 58:08Elements by having by using matrix
  1447. 58:12multiplication of a lower triangular
  1448. 58:14fashion and then the elements here in
  1449. 58:17the lower triangular part are telling
  1450. 58:18you how much of each element uh fuses
  1451. 58:21into this position so we're going to use
  1452. 58:24this trick now to develop the self
  1453. 58:25attention block block so first let's get
  1454. 58:27some quick preliminaries out of the way
  1455. 58:30first the thing I'm kind of bothered by
  1456. 58:31is that you see how we're passing in
  1457. 58:33vocap size into the Constructor there's
  1458. 58:35no need to do that because vocap size is
  1459. 58:36already defined uh up top as a global
  1460. 58:38variable so there's no need to pass this
  1461. 58:40stuff
  1462. 58:41around next what I want to do is I don't
  1463. 58:44want to actually create I want to create
  1464. 58:46like a level of indirection here where
  1465. 58:47we don't directly go to the embedding
  1466. 58:49for the um logits but instead we go
  1467. 58:52through this intermediate phase because
  1468. 58:54we're going to start making that bigger
  1469. 58:57so let me introduce a new variable n
  1470. 58:59embed it shorted for number of embedding
  1471. 59:02Dimensions so
  1472. 59:04nbed here will be say 32 that was a
  1473. 59:09suggestion from GitHub co-pilot by the
  1474. 59:11way um it also suest 32 which is a good
  1475. 59:14number so this is an embedding table and
  1476. 59:16only 32 dimensional
  1477. 59:18embeddings so then here this is not
  1478. 59:21going to give us logits directly instead
  1479. 59:23this is going to give us token
  1480. 59:24embeddings that's I'm going to call it
  1481. 59:27and then to go from the token Tings to
  1482. 59:29the logits we're going to need a linear
  1483. 59:30layer so self. LM head let's call it
  1484. 59:34short for language modeling head is n
  1485. 59:36and linear from n ined up to vocap size
  1486. 59:39and then when we swing over here we're
  1487. 59:41actually going to get the loits by
  1488. 59:43exactly what the co-pilot says now we
  1489. 59:46have to be careful here because this C
  1490. 59:48and this C are not equal um this is nmed
  1491. 59:52C and this is vocap size so let's just
  1492. 59:55say that n ined is equal to
  1493. 59:57C and then this just creates one spous
  1494. 1:00:01layer of interaction through a linear
  1495. 1:00:02layer but uh this should basically
  1496. 1:00:11run so we see that this runs and uh this
  1497. 1:00:15currently looks kind of spous but uh
  1498. 1:00:17we're going to build on top of this now
  1499. 1:00:19next up so far we've taken these indices
  1500. 1:00:22and we've encoded them based on the
  1501. 1:00:23identity of the uh tokens in inside idx
  1502. 1:00:28the next thing that people very often do
  1503. 1:00:30is that we're not just encoding the
  1504. 1:00:31identity of these tokens but also their
  1505. 1:00:33position so we're going to have a second
  1506. 1:00:35position uh embedding table here so
  1507. 1:00:38self. position embedding table is an an
  1508. 1:00:41embedding of block size by an embed and
  1509. 1:00:44so each position from zero to block size
  1510. 1:00:46minus one will also get its own
  1511. 1:00:47embedding vector and then here first let
  1512. 1:00:50me decode B BYT from idx do
  1513. 1:00:54shape and then here we're also going to
  1514. 1:00:56have a pause embedding which is the
  1515. 1:00:58positional embedding and these are this
  1516. 1:01:00is to arrange so this will be basically
  1517. 1:01:03just integers from Z to T minus one and
  1518. 1:01:06all of those integers from 0 to T minus
  1519. 1:01:08one get embedded through the table to
  1520. 1:01:09create a t by
  1521. 1:01:11C and then here this gets renamed to
  1522. 1:01:14just say x and x will be the addition of
  1523. 1:01:18the token embeddings with the positional
  1524. 1:01:20embeddings and here the broadcasting
  1525. 1:01:22note will work out so B by T by C plus T
  1526. 1:01:25by C
  1527. 1:01:26this gets right aligned a new dimension
  1528. 1:01:28of one gets added and it gets
  1529. 1:01:30broadcasted across
  1530. 1:01:31batch so at this point x holds not just
  1531. 1:01:34the token identities but the positions
  1532. 1:01:37at which these tokens occur and this is
  1533. 1:01:39currently not that useful because of
  1534. 1:01:41course we just have a simple byr model
  1535. 1:01:43so it doesn't matter if you're in the
  1536. 1:01:44fifth position the second position or
  1537. 1:01:46wherever it's all translation invariant
  1538. 1:01:48at this stage uh so this information
  1539. 1:01:50currently wouldn't help uh but as we
  1540. 1:01:52work on the self attention block we'll
  1541. 1:01:54see that this starts to matter
  1542. 1:01:59okay so now we get the Crux of self
  1543. 1:02:01attention so this is probably the most
  1544. 1:02:03important part of this video to
  1545. 1:02:05understand we're going to implement a
  1546. 1:02:07small self attention for a single
  1547. 1:02:08individual head as they're called so we
  1548. 1:02:11start off with where we were so all of
  1549. 1:02:13this code is familiar so right now I'm
  1550. 1:02:16working with an example where I Chang
  1551. 1:02:17the number of channels from 2 to 32 so
  1552. 1:02:20we have a 4x8 arrangement of tokens and
  1553. 1:02:24each to and the information each token
  1554. 1:02:26is currently 32 dimensional but we just
  1555. 1:02:28are working with random
  1556. 1:02:30numbers now we saw here that the code as
  1557. 1:02:34we had it before does a uh simple weight
  1558. 1:02:37simple average of all the past tokens
  1559. 1:02:41and the current token so it's just the
  1560. 1:02:43previous information and current
  1561. 1:02:44information is just being mixed together
  1562. 1:02:45in an average and that's what this code
  1563. 1:02:48currently achieves and it Doo by
  1564. 1:02:50creating this lower triangular structure
  1565. 1:02:52which allows us to mask out this uh we
  1566. 1:02:55uh Matrix that we create so we mask it
  1567. 1:02:59out and then we normalize it and
  1568. 1:03:01currently when we initialize the
  1569. 1:03:03affinities between all the different
  1570. 1:03:05sort of tokens or nodes I'm going to use
  1571. 1:03:08those terms
  1572. 1:03:09interchangeably so when we initialize
  1573. 1:03:11the affinities between all the different
  1574. 1:03:13tokens to be zero then we see that way
  1575. 1:03:16gives us this um structure where every
  1576. 1:03:18single row has these um uniform numbers
  1577. 1:03:22and so that's what that's what then uh
  1578. 1:03:25in this Matrix multiply makes it so that
  1579. 1:03:27we're doing a simple
  1580. 1:03:28average now we don't actually want this
  1581. 1:03:32to be all uniform because different uh
  1582. 1:03:36tokens will find different other tokens
  1583. 1:03:38more or less interesting and we want
  1584. 1:03:40that to be data dependent so for example
  1585. 1:03:42if I'm a vowel then maybe I'm looking
  1586. 1:03:44for consonants in my past and maybe I
  1587. 1:03:46want to know what those consonants are
  1588. 1:03:48and I want that information to flow to
  1589. 1:03:50me and so I want to now gather
  1590. 1:03:52information from the past but I want to
  1591. 1:03:54do it in the data dependent way and this
  1592. 1:03:56is the problem that self attention
  1593. 1:03:58solves now the way self attention solves
  1594. 1:04:00this is the following every single node
  1595. 1:04:03or every single token at each position
  1596. 1:04:06will emit two vectors it will emit a
  1597. 1:04:09query and it will emit a
  1598. 1:04:12key now the query Vector roughly
  1599. 1:04:15speaking is what am I looking for and
  1600. 1:04:18the key Vector roughly speaking is what
  1601. 1:04:20do I
  1602. 1:04:21contain and then the way we get
  1603. 1:04:24affinities between these uh tokens now
  1604. 1:04:27in a sequence is we basically just do a
  1605. 1:04:29do product between the keys and the
  1606. 1:04:31queries so my query dot products with
  1607. 1:04:35all the keys of all the other tokens and
  1608. 1:04:37that dot product now becomes
  1609. 1:04:41wayy and so um if the key and the query
  1610. 1:04:45are sort of aligned they will interact
  1611. 1:04:47to a very high amount and then I will
  1612. 1:04:50get to learn more about that specific
  1613. 1:04:52token as opposed to any other token in
  1614. 1:04:55the sequence
  1615. 1:04:56so let's implement this
  1616. 1:05:00now we're going to implement a
  1617. 1:05:03single what's called head of self
  1618. 1:05:07attention so this is just one head
  1619. 1:05:09there's a hyper parameter involved with
  1620. 1:05:10these heads which is the head size and
  1621. 1:05:13then here I'm initializing linear
  1622. 1:05:15modules and I'm using bias equals false
  1623. 1:05:18so these are just going to apply a
  1624. 1:05:19matrix multiply with some fixed
  1625. 1:05:21weights and now let me produce a key and
  1626. 1:05:26q k and Q by forwarding these modules on
  1627. 1:05:29X so the size of this will now
  1628. 1:05:32become B by T by 16 because that is the
  1629. 1:05:36head size and the same here B by T by
  1630. 1:05:4416 so this being the head size so you
  1631. 1:05:47see here that when I forward this linear
  1632. 1:05:49on top of my X all the tokens in all the
  1633. 1:05:52positions in the B BYT Arrangement all
  1634. 1:05:55of them them in parallel and
  1635. 1:05:57independently produce a key and a query
  1636. 1:05:59so no communication has happened
  1637. 1:06:01yet but the communication comes now all
  1638. 1:06:04the queries will do product with all the
  1639. 1:06:07keys so basically what we want is we
  1640. 1:06:09want way now or the affinities between
  1641. 1:06:12these to be query multiplying key but we
  1642. 1:06:16have to be careful with uh we can't
  1643. 1:06:18Matrix multiply this we actually need to
  1644. 1:06:20transpose uh K but we have to be also
  1645. 1:06:23careful because these are when you have
  1646. 1:06:25The Bash Dimension so in particular we
  1647. 1:06:27want to transpose uh the last two
  1648. 1:06:30dimensions dimension1 and dimension -2
  1649. 1:06:33so
  1650. 1:06:36-21 and so this Matrix multiply now will
  1651. 1:06:40basically do the following B by T by
  1652. 1:06:4416 Matrix multiplies B by 16 by T to
  1653. 1:06:49give us B by T by
  1654. 1:06:53T right
  1655. 1:06:56so for every row of B we're now going to
  1656. 1:06:58have a t Square Matrix giving us the
  1657. 1:07:01affinities and these are now the way so
  1658. 1:07:04they're not zeros they are now coming
  1659. 1:07:06from this dot product between the keys
  1660. 1:07:08and the queries so this can now run I
  1661. 1:07:11can I can run this and the weighted
  1662. 1:07:13aggregation now is a function in a data
  1663. 1:07:16Bandon manner between the keys and
  1664. 1:07:18queries of these nodes so just
  1665. 1:07:20inspecting what happened
  1666. 1:07:22here the way takes on this form
  1667. 1:07:26and you see that before way was uh just
  1668. 1:07:29a constant so it was applied in the same
  1669. 1:07:31way to all the batch elements but now
  1670. 1:07:33every single batch elements will have
  1671. 1:07:34different sort of we because uh every
  1672. 1:07:37single batch element contains different
  1673. 1:07:39uh tokens at different positions and so
  1674. 1:07:41this is not data dependent so when we
  1675. 1:07:44look at just the zeroth uh Row for
  1676. 1:07:47example in the input these are the
  1677. 1:07:49weights that came out and so you can see
  1678. 1:07:51now that they're not just exactly
  1679. 1:07:53uniform um and in particular as an
  1680. 1:07:55example here for the last row this was
  1681. 1:07:58the eighth token and the eighth token
  1682. 1:08:00knows what content it has and it knows
  1683. 1:08:02at what position it's in and now the E
  1684. 1:08:04token based on that uh creates a query
  1685. 1:08:08hey I'm looking for this kind of stuff
  1686. 1:08:10um I'm a vowel I'm on the E position I'm
  1687. 1:08:12looking for any consonant at positions
  1688. 1:08:14up to four and then all the nodes get to
  1689. 1:08:18emit keys and maybe one of the channels
  1690. 1:08:20could be I am a I am a consonant and I
  1691. 1:08:23am in a position up to four and that
  1692. 1:08:25that key would have a high number in
  1693. 1:08:27that specific Channel and that's how the
  1694. 1:08:29query and the key when they do product
  1695. 1:08:31they can find each other and create a
  1696. 1:08:33high affinity and when they have a high
  1697. 1:08:35Affinity like say uh this token was
  1698. 1:08:38pretty interesting to uh to this eighth
  1699. 1:08:41token when they have a high Affinity
  1700. 1:08:43then through the softmax I will end up
  1701. 1:08:45aggregating a lot of its information
  1702. 1:08:47into my position and so I'll get to
  1703. 1:08:49learn a lot about
  1704. 1:08:51it now just this we're looking at way
  1705. 1:08:55after this has already happened um let
  1706. 1:08:59me erase this operation as well so let
  1707. 1:09:01me erase the masking and the softmax
  1708. 1:09:03just to show you the under the hood
  1709. 1:09:04internals and how that works so without
  1710. 1:09:07the masking in the softmax Whey comes
  1711. 1:09:09out like this right this is the outputs
  1712. 1:09:11of the do products um and these are the
  1713. 1:09:14raw outputs and they take on values from
  1714. 1:09:15negative you know two to positive two
  1715. 1:09:18Etc so that's the raw interactions and
  1716. 1:09:21raw affinities between all the nodes but
  1717. 1:09:24now if I'm going if I'm a fifth node I
  1718. 1:09:26will not want to aggregate anything from
  1719. 1:09:28the sixth node seventh node and the
  1720. 1:09:30eighth node so actually we use the upper
  1721. 1:09:32triangular masking so those are not
  1722. 1:09:35allowed to
  1723. 1:09:37communicate and now we actually want to
  1724. 1:09:40have a nice uh distribution uh so we
  1725. 1:09:42don't want to aggregate negative .11 of
  1726. 1:09:45this node that's crazy so instead we
  1727. 1:09:47exponentiate and normalize and now we
  1728. 1:09:49get a nice distribution that sums to one
  1729. 1:09:51and this is telling us now in the data
  1730. 1:09:52dependent manner how much of information
  1731. 1:09:54to aggregate from any of these tokens in
  1732. 1:09:56the
  1733. 1:09:58past so that's way and it's not zeros
  1734. 1:10:01anymore but but it's calculated in this
  1735. 1:10:04way now there's one more uh part to a
  1736. 1:10:08single self attention head and that is
  1737. 1:10:10that when we do the aggregation we don't
  1738. 1:10:12actually aggregate the tokens exactly we
  1739. 1:10:15aggregate we produce one more value here
  1740. 1:10:17and we call that the
  1741. 1:10:20value so in the same way that we
  1742. 1:10:22produced p and query we're also going to
  1743. 1:10:23create a value
  1744. 1:10:26and
  1745. 1:10:26then here we don't
  1746. 1:10:30aggregate X we calculate a v which is
  1747. 1:10:34just achieved by uh propagating this
  1748. 1:10:37linear on top of X again and then we
  1749. 1:10:40output way multiplied by V so V is the
  1750. 1:10:44elements that we aggregate or the the
  1751. 1:10:46vectors that we aggregate instead of the
  1752. 1:10:47raw
  1753. 1:10:48X and now of course uh this will make it
  1754. 1:10:51so that the output here of this single
  1755. 1:10:53head will be 16 dimensional because that
  1756. 1:10:55is the head
  1757. 1:10:57size so you can think of X as kind of
  1758. 1:10:59like private information to this token
  1759. 1:11:01if you if you think about it that way so
  1760. 1:11:03X is kind of private to this token so
  1761. 1:11:06I'm a fifth token at some and I have
  1762. 1:11:08some identity and uh my information is
  1763. 1:11:11kept in Vector X and now for the
  1764. 1:11:14purposes of the single head here's what
  1765. 1:11:16I'm interested in here's what I have and
  1766. 1:11:20if you find me interesting here's what I
  1767. 1:11:21will communicate to you and that's
  1768. 1:11:23stored in v and so V is the thing that
  1769. 1:11:26gets aggregated for the purposes of this
  1770. 1:11:28single head between the different
  1771. 1:11:30notes and that's uh basically the self
  1772. 1:11:34attention mechanism this is this is what
  1773. 1:11:36it does there are a few notes that I
  1774. 1:11:39would make like to make about attention
  1775. 1:11:41number one attention is a communication
  1776. 1:11:44mechanism you can really think about it
  1777. 1:11:46as a communication mechanism where you
  1778. 1:11:48have a number of nodes in a directed
  1779. 1:11:50graph where basically you have edges
  1780. 1:11:52pointed between noes like
  1781. 1:11:53this and what happens is every node has
  1782. 1:11:56some Vector of information and it gets
  1783. 1:11:58to aggregate information via a weighted
  1784. 1:12:01sum from all of the nodes that point to
  1785. 1:12:03it and this is done in a data dependent
  1786. 1:12:06manner so depending on whatever data is
  1787. 1:12:08actually stored that you should not at
  1788. 1:12:09any point in time now our graph doesn't
  1789. 1:12:13look like this our graph has a different
  1790. 1:12:15structure we have eight nodes because
  1791. 1:12:17the block size is eight and there's
  1792. 1:12:18always eight to
  1793. 1:12:20tokens and uh the first node is only
  1794. 1:12:23pointed to by itself the second node is
  1795. 1:12:25pointed to by the first node and itself
  1796. 1:12:27all the way up to the eighth node which
  1797. 1:12:29is pointed to by all the previous nodes
  1798. 1:12:32and itself and so that's the structure
  1799. 1:12:34that our directed graph has or happens
  1800. 1:12:37happens to have in Auto regressive sort
  1801. 1:12:38of scenario like language modeling but
  1802. 1:12:41in principle attention can be applied to
  1803. 1:12:42any arbitrary directed graph and it's
  1804. 1:12:44just a communication mechanism between
  1805. 1:12:46the nodes the second note is that notice
  1806. 1:12:48that there is no notion of space so
  1807. 1:12:51attention simply acts over like a set of
  1808. 1:12:53vectors in this graph and so by default
  1809. 1:12:56these nodes have no idea where they are
  1810. 1:12:58positioned in the space and that's why
  1811. 1:12:59we need to encode them positionally and
  1812. 1:13:02sort of give them some information that
  1813. 1:13:03is anchored to a specific position so
  1814. 1:13:05that they sort of know where they are
  1815. 1:13:08and this is different than for example
  1816. 1:13:09from convolution because if you're run
  1817. 1:13:11for example a convolution operation over
  1818. 1:13:13some input there's a very specific sort
  1819. 1:13:15of layout of the information in space
  1820. 1:13:18and the convolutional filters sort of
  1821. 1:13:20act in space and so it's it's not like
  1822. 1:13:23an attention in ATT ention is just a set
  1823. 1:13:26of vectors out there in space they
  1824. 1:13:27communicate and if you want them to have
  1825. 1:13:29a notion of space you need to
  1826. 1:13:31specifically add it which is what we've
  1827. 1:13:33done when we calculated the um relative
  1828. 1:13:36the positional encode encodings and
  1829. 1:13:38added that information to the vectors
  1830. 1:13:40the next thing that I hope is very clear
  1831. 1:13:41is that the elements across the batch
  1832. 1:13:43Dimension which are independent examples
  1833. 1:13:45never talk to each other they're always
  1834. 1:13:47processed independently and this is a
  1835. 1:13:49batched matrix multiply that applies
  1836. 1:13:51basically a matrix multiplication uh
  1837. 1:13:53kind of in parallel across the batch
  1838. 1:13:54dimension so maybe it would be more
  1839. 1:13:56accurate to say that in this analogy of
  1840. 1:13:58a directed graph we really have because
  1841. 1:14:00the back size is four we really have
  1842. 1:14:03four separate pools of eight nodes and
  1843. 1:14:05those eight nodes only talk to each
  1844. 1:14:07other but in total there's like 32 nodes
  1845. 1:14:08that are being processed uh but there's
  1846. 1:14:11um sort of four separate pools of eight
  1847. 1:14:13you can look at it that way the next
  1848. 1:14:15note is that here in the case of
  1849. 1:14:18language modeling uh we have this
  1850. 1:14:20specific uh structure of directed graph
  1851. 1:14:22where the future tokens will not
  1852. 1:14:24communicate to the Past tokens but this
  1853. 1:14:27doesn't necessarily have to be the
  1854. 1:14:28constraint in the general case and in
  1855. 1:14:30fact in many cases you may want to have
  1856. 1:14:32all of the uh noes talk to each other uh
  1857. 1:14:35fully so as an example if you're doing
  1858. 1:14:37sentiment analysis or something like
  1859. 1:14:38that with a Transformer you might have a
  1860. 1:14:40number of tokens and you may want to
  1861. 1:14:42have them all talk to each other fully
  1862. 1:14:45because later you are predicting for
  1863. 1:14:46example the sentiment of the sentence
  1864. 1:14:49and so it's okay for these NOS to talk
  1865. 1:14:50to each other and so in those cases you
  1866. 1:14:53will use an encoder block of self
  1867. 1:14:55attention and uh all it means that it's
  1868. 1:14:58an encoder block is that you will delete
  1869. 1:15:00this line of code allowing all the noes
  1870. 1:15:02to completely talk to each other what
  1871. 1:15:04we're implementing here is sometimes
  1872. 1:15:06called a decoder block and it's called a
  1873. 1:15:09decoder because it is sort of like a
  1874. 1:15:12decoding language and it's got this
  1875. 1:15:15autor regressive format where you have
  1876. 1:15:17to mask with the Triangular Matrix so
  1877. 1:15:19that uh nodes from the future never talk
  1878. 1:15:22to the Past because they would give away
  1879. 1:15:24the answer
  1880. 1:15:25and so basically in encoder blocks you
  1881. 1:15:27would delete this allow all the noes to
  1882. 1:15:29talk in decoder blocks this will always
  1883. 1:15:31be present so that you have this
  1884. 1:15:33triangular structure uh but both are
  1885. 1:15:35allowed and attention doesn't care
  1886. 1:15:36attention supports arbitrary
  1887. 1:15:38connectivity between nodes the next
  1888. 1:15:40thing I wanted to comment on is you keep
  1889. 1:15:41me you keep hearing me say attention
  1890. 1:15:43self attention Etc there's actually also
  1891. 1:15:45something called cross attention what is
  1892. 1:15:47the
  1893. 1:15:47difference
  1894. 1:15:49so basically the reason this attention
  1895. 1:15:52is self attention is because because the
  1896. 1:15:55keys queries and the values are all
  1897. 1:15:57coming from the same Source from X so
  1898. 1:16:01the same Source X produces Keys queries
  1899. 1:16:03and values so these nodes are self
  1900. 1:16:05attending but in principle attention is
  1901. 1:16:08much more General than that so for
  1902. 1:16:10example an encoder decoder Transformers
  1903. 1:16:12uh you can have a case where the queries
  1904. 1:16:15are produced from X but the keys and the
  1905. 1:16:17values come from a whole separate
  1906. 1:16:18external source and sometimes from uh
  1907. 1:16:21encoder blocks that encode some context
  1908. 1:16:23that we'd like to condition on
  1909. 1:16:25and so the keys and the values will
  1910. 1:16:26actually come from a whole separate
  1911. 1:16:28Source those are nodes on the side and
  1912. 1:16:31here we're just producing queries and
  1913. 1:16:32we're reading off information from the
  1914. 1:16:34side so cross attention is used when
  1915. 1:16:37there's a separate source of nodes we'd
  1916. 1:16:40like to pull information from into our
  1917. 1:16:42nodes and it's self attention if we just
  1918. 1:16:45have nodes that would like to look at
  1919. 1:16:46each other and talk to each other so
  1920. 1:16:48this attention here happens to be self
  1921. 1:16:51attention but in principle um attention
  1922. 1:16:55is a lot more General okay and the last
  1923. 1:16:57note at this stage is if we come to the
  1924. 1:16:59attention is all need paper here we've
  1925. 1:17:01already implemented attention so given
  1926. 1:17:03query key and value we've U multiplied
  1927. 1:17:06the query and a key we've soft maxed it
  1928. 1:17:09and then we are aggregating the values
  1929. 1:17:11there's one more thing that we're
  1930. 1:17:12missing here which is the dividing by
  1931. 1:17:13one / square root of the head size the
  1932. 1:17:16DK here is the head size why are they
  1933. 1:17:18doing this finds this important so they
  1934. 1:17:21call it the scaled attention and it's
  1935. 1:17:24kind of like an important normalization
  1936. 1:17:25to basically
  1937. 1:17:26have the problem is if you have unit gsh
  1938. 1:17:29and inputs so zero mean unit variance K
  1939. 1:17:32and Q are unit gashin then if you just
  1940. 1:17:34do we naively then you see that your we
  1941. 1:17:37actually will be uh the variance will be
  1942. 1:17:38on the order of head size which in our
  1943. 1:17:40case is 16 but if you multiply by one
  1944. 1:17:43over head size square root so this is
  1945. 1:17:45square root and this is one
  1946. 1:17:47over then the variance of we will be one
  1947. 1:17:50so it will be
  1948. 1:17:52preserved now why is this important
  1949. 1:17:54you'll not notice that way
  1950. 1:17:56here will feed into
  1951. 1:17:58softmax and so it's really important
  1952. 1:18:00especially at initialization that we be
  1953. 1:18:03fairly diffuse so in our case here we
  1954. 1:18:06sort of locked out here and we had a
  1955. 1:18:10fairly diffuse numbers here so um like
  1956. 1:18:13this now the problem is that because of
  1957. 1:18:15softmax if weight takes on very positive
  1958. 1:18:18and very negative numbers inside it
  1959. 1:18:20softmax will actually converge towards
  1960. 1:18:22one hot vectors and so I can illustrate
  1961. 1:18:25that here um say we are applying softmax
  1962. 1:18:29to a tensor of values that are very
  1963. 1:18:31close to zero then we're going to get a
  1964. 1:18:33diffuse thing out of
  1965. 1:18:34softmax but the moment I take the exact
  1966. 1:18:36same thing and I start sharpening it
  1967. 1:18:38making it bigger by multiplying these
  1968. 1:18:40numbers by eight for example you'll see
  1969. 1:18:42that the softmax will start to sharpen
  1970. 1:18:44and in fact it will sharpen towards the
  1971. 1:18:46max so it will sharpen towards whatever
  1972. 1:18:48number here is the highest and so um
  1973. 1:18:51basically we don't want these values to
  1974. 1:18:52be too extreme especially at
  1975. 1:18:53initialization otherwise softmax will be
  1976. 1:18:55way too peaky and um you're basically
  1977. 1:18:58aggregating um information from like a
  1978. 1:19:01single node every node just agregates
  1979. 1:19:03information from a single other node
  1980. 1:19:04that's not what we want especially at
  1981. 1:19:06initialization and so the scaling is
  1982. 1:19:08used just to control the variance at
  1983. 1:19:11initialization okay so having said all
  1984. 1:19:13that let's now take our self attention
  1985. 1:19:15knowledge and let's uh take it for a
  1986. 1:19:17spin so here in the code I created this
  1987. 1:19:19head module and it implements a single
  1988. 1:19:22head of self attention so you give it a
  1989. 1:19:24head size and then here it creates the
  1990. 1:19:26key query and the value linear layers
  1991. 1:19:29typically people don't use biases in
  1992. 1:19:31these uh so those are the linear
  1993. 1:19:33projections that we're going to apply to
  1994. 1:19:34all of our nodes now here I'm creating
  1995. 1:19:37this Trill variable Trill is not a
  1996. 1:19:40parameter of the module so in sort of
  1997. 1:19:41pytorch naming conventions uh this is
  1998. 1:19:43called a buffer it's not a parameter and
  1999. 1:19:46you have to call it you have to assign
  2000. 1:19:47it to the module using a register buffer
  2001. 1:19:49so that creates the trill uh the triang
  2002. 1:19:52lower triangular Matrix and we're given
  2003. 1:19:55the input X this should look very
  2004. 1:19:56familiar now we calculate the keys the
  2005. 1:19:58queries we C calculate the attention
  2006. 1:20:00scores inside way uh we normalize it so
  2007. 1:20:03we're using scaled attention here then
  2008. 1:20:06we make sure that uh future doesn't
  2009. 1:20:08communicate with the past so this makes
  2010. 1:20:10it a decoder block and then softmax and
  2011. 1:20:13then aggregate the value and
  2012. 1:20:15output then here in the language model
  2013. 1:20:17I'm creating a head in the Constructor
  2014. 1:20:20and I'm calling it self attention head
  2015. 1:20:22and the head size I'm going to keep as
  2016. 1:20:24the same and embed just for
  2017. 1:20:27now and then here once we've encoded the
  2018. 1:20:31information with the token embeddings
  2019. 1:20:32and the position embeddings we're simply
  2020. 1:20:34going to feed it into the self attention
  2021. 1:20:36head and then the output of that is
  2022. 1:20:38going to go into uh the decoder language
  2023. 1:20:42modeling head and create the logits so
  2024. 1:20:44this the sort of the simplest way to
  2025. 1:20:46plug in a self attention component uh
  2026. 1:20:49into our Network right now I had to make
  2027. 1:20:51one more change which is that here in
  2028. 1:20:55the generate uh we have to make sure
  2029. 1:20:57that our idx that we feed into the model
  2030. 1:21:01because now we're using positional
  2031. 1:21:02embeddings we can never have more than
  2032. 1:21:04block size coming in because if idx is
  2033. 1:21:07more than block size then our position
  2034. 1:21:09embedding table is going to run out of
  2035. 1:21:11scope because it only has embeddings for
  2036. 1:21:12up to block size and so therefore I
  2037. 1:21:15added some uh code here to crop the
  2038. 1:21:17context that we're going to feed into
  2039. 1:21:20self um so that uh we never pass in more
  2040. 1:21:23than block siiz elements
  2041. 1:21:25so those are the changes and let's Now
  2042. 1:21:27train the network okay so I also came up
  2043. 1:21:29to the script here and I decreased the
  2044. 1:21:30learning rate because uh the self
  2045. 1:21:32attention can't tolerate very very high
  2046. 1:21:34learning rates and then I also increased
  2047. 1:21:36number of iterations because the
  2048. 1:21:37learning rate is lower and then I
  2049. 1:21:39trained it and previously we were only
  2050. 1:21:41able to get to up to 2.5 and now we are
  2051. 1:21:43down to 2.4 so we definitely see a
  2052. 1:21:46little bit of an improvement from 2.5 to
  2053. 1:21:482.4 roughly uh but the text is still not
  2054. 1:21:51amazing so clearly the self attention
  2055. 1:21:53head is doing some useful communication
  2056. 1:21:56but um we still have a long way to go
  2057. 1:21:59okay so now we've implemented the scale.
  2058. 1:22:01product attention now next up and the
  2059. 1:22:02attention is all you need paper there's
  2060. 1:22:05something called multi-head attention
  2061. 1:22:07and what is multi-head attention it's
  2062. 1:22:09just applying multiple attentions in
  2063. 1:22:11parallel and concatenating their results
  2064. 1:22:13so they have a little bit of diagram
  2065. 1:22:15here I don't know if this is super clear
  2066. 1:22:18it's really just multiple attentions in
  2067. 1:22:20parallel so let's Implement that fairly
  2068. 1:22:23straightforward
  2069. 1:22:25if we want a multi-head attention then
  2070. 1:22:27we want multiple heads of self attention
  2071. 1:22:28running in parallel so in pytorch we can
  2072. 1:22:32do this by simply creating multiple
  2073. 1:22:35heads so however heads how however many
  2074. 1:22:38heads you want and then what is the head
  2075. 1:22:39size of each and then we run all of them
  2076. 1:22:43in parallel into a list and simply
  2077. 1:22:46concatenate all of the outputs and we're
  2078. 1:22:48concatenating over the channel
  2079. 1:22:50Dimension so the way this looks now is
  2080. 1:22:53we don't have just a single ATT
  2081. 1:22:56that uh has a hit size of 32 because
  2082. 1:22:59remember n Ed is
  2083. 1:23:0032 instead of having one Communication
  2084. 1:23:03channel we now have four communication
  2085. 1:23:06channels in parallel and each one of
  2086. 1:23:08these communication channels typically
  2087. 1:23:10will be uh smaller uh correspondingly so
  2088. 1:23:14because we have four communication
  2089. 1:23:15channels we want eight dimensional self
  2090. 1:23:18attention and so from each Communication
  2091. 1:23:20channel we're going to together eight
  2092. 1:23:22dimensional vectors and then we have
  2093. 1:23:23four of them and that concatenates to
  2094. 1:23:25give us 32 which is the original and
  2095. 1:23:28embed and so this is kind of similar to
  2096. 1:23:30um if you're familiar with convolutions
  2097. 1:23:32this is kind of like a group convolution
  2098. 1:23:34uh because basically instead of having
  2099. 1:23:36one large convolution we do convolution
  2100. 1:23:38in groups and uh that's multi-headed
  2101. 1:23:40self
  2102. 1:23:41attention and so then here we just use
  2103. 1:23:44essay heads self attention heads instead
  2104. 1:23:47now I actually ran it and uh scrolling
  2105. 1:23:51down I ran the same thing and then we
  2106. 1:23:53now get this down to 2.28 roughly and
  2107. 1:23:57the output is still the generation is
  2108. 1:23:58still not amazing but clearly the
  2109. 1:24:00validation loss is improving because we
  2110. 1:24:02were at 2.4 just now and so it helps to
  2111. 1:24:05have multiple communication channels
  2112. 1:24:07because obviously these tokens have a
  2113. 1:24:09lot to talk about they want to find the
  2114. 1:24:11consonants the vowels they want to find
  2115. 1:24:13the vowels just from certain positions
  2116. 1:24:15uh they want to find any kinds of
  2117. 1:24:17different things and so it helps to
  2118. 1:24:19create multiple independent channels of
  2119. 1:24:20communication gather lots of different
  2120. 1:24:22types of data and then uh decode the
  2121. 1:24:25output now going back to the paper for a
  2122. 1:24:27second of course I didn't explain this
  2123. 1:24:28figure in full detail but we are
  2124. 1:24:30starting to see some components of what
  2125. 1:24:32we've already implemented we have the
  2126. 1:24:33positional encodings the token encodings
  2127. 1:24:35that add we have the masked multi-headed
  2128. 1:24:37attention implemented now here's another
  2129. 1:24:41multi-headed attention which is a cross
  2130. 1:24:42attention to an encoder which we haven't
  2131. 1:24:45we're not going to implement in this
  2132. 1:24:46case I'm going to come back to that
  2133. 1:24:48later but I want you to notice that
  2134. 1:24:50there's a feed forward part here and
  2135. 1:24:52then this is grouped into a block that
  2136. 1:24:53gets repeat it again and again now the
  2137. 1:24:56feedforward part here is just a simple
  2138. 1:24:57uh multi-layer perceptron
  2139. 1:25:00um so the multi-headed so here position
  2140. 1:25:04wise feed forward networks is just a
  2141. 1:25:06simple little MLP so I want to start
  2142. 1:25:08basically in a similar fashion also
  2143. 1:25:10adding computation into the network and
  2144. 1:25:13this computation is on a per node level
  2145. 1:25:16so I've already implemented it and you
  2146. 1:25:18can see the diff highlighted on the left
  2147. 1:25:20here when I've added or changed things
  2148. 1:25:22now before we had the self multi-headed
  2149. 1:25:25self attention that did the
  2150. 1:25:26communication but we went way too fast
  2151. 1:25:28to calculate the logits so the tokens
  2152. 1:25:31looked at each other but didn't really
  2153. 1:25:32have a lot of time to think on what they
  2154. 1:25:35found from the other tokens and so what
  2155. 1:25:38I've implemented here is a little feet
  2156. 1:25:40forward single layer and this little
  2157. 1:25:42layer is just a linear followed by a Rel
  2158. 1:25:45nonlinearity and that's that's it so
  2159. 1:25:48it's just a little layer and then I call
  2160. 1:25:50it feed
  2161. 1:25:52forward um and embed
  2162. 1:25:54and then this feed forward is just
  2163. 1:25:56called sequentially right after the self
  2164. 1:25:58attention so we self attend then we feed
  2165. 1:26:01forward and you'll notice that the feet
  2166. 1:26:02forward here when it's applying linear
  2167. 1:26:04this is on a per token level all the
  2168. 1:26:06tokens do this independently so the self
  2169. 1:26:09attention is the communication and then
  2170. 1:26:11once they've gathered all the data now
  2171. 1:26:13they need to think on that data
  2172. 1:26:15individually and so that's what feed
  2173. 1:26:16forward is doing and that's why I've
  2174. 1:26:18added it here now when I train this the
  2175. 1:26:21validation LW actually continues to go
  2176. 1:26:23down now to 2. 24 which is down from
  2177. 1:26:262.28 uh the output still look kind of
  2178. 1:26:28terrible but at least we've improved the
  2179. 1:26:31situation and so as a preview we're
  2180. 1:26:34going to now start to intersperse the
  2181. 1:26:37communication with the computation and
  2182. 1:26:39that's also what the Transformer does
  2183. 1:26:42when it has blocks that communicate and
  2184. 1:26:44then compute and it groups them and
  2185. 1:26:46replicates them okay so let me show you
  2186. 1:26:49what we'd like to do we'd like to do
  2187. 1:26:51something like this we have a block and
  2188. 1:26:53this block is is basically this part
  2189. 1:26:55here except for the cross
  2190. 1:26:57attention now the block basically
  2191. 1:26:59intersperses communication and then
  2192. 1:27:01computation the computation the
  2193. 1:27:03communication is done using multi-headed
  2194. 1:27:05selfelf attention and then the
  2195. 1:27:07computation is done using a feed forward
  2196. 1:27:08Network on all the tokens
  2197. 1:27:11independently now what I've added here
  2198. 1:27:14also is you'll
  2199. 1:27:16notice this takes the number of
  2200. 1:27:18embeddings in the embedding Dimension
  2201. 1:27:19and number of heads that we would like
  2202. 1:27:21which is kind of like group size in
  2203. 1:27:22group convolution and and I'm saying
  2204. 1:27:24that number of heads we'd like is four
  2205. 1:27:26and so because this is 32 we calculate
  2206. 1:27:29that because this is 32 the number of
  2207. 1:27:31heads should be four um the head size
  2208. 1:27:34should be eight so that everything sort
  2209. 1:27:36of works out Channel wise um so this is
  2210. 1:27:39how the Transformer structures uh sort
  2211. 1:27:41of the uh the sizes typically so the
  2212. 1:27:44head size will become eight and then
  2213. 1:27:45this is how we want to intersperse them
  2214. 1:27:47and then here I'm trying to create
  2215. 1:27:49blocks which is just a sequential
  2216. 1:27:51application of block block block so that
  2217. 1:27:53we're interspersing communication feed
  2218. 1:27:55forward many many times and then finally
  2219. 1:27:57we decode now I actually tried to run
  2220. 1:28:01this and the problem is this doesn't
  2221. 1:28:02actually give a very good uh answer and
  2222. 1:28:05very good result and the reason for that
  2223. 1:28:07is we're start starting to actually get
  2224. 1:28:09like a pretty deep neural net and deep
  2225. 1:28:11neural Nets uh suffer from optimization
  2226. 1:28:13issues and I think that's what we're
  2227. 1:28:14kind of like slightly starting to run
  2228. 1:28:16into so we need one more idea that we
  2229. 1:28:18can borrow from the um Transformer paper
  2230. 1:28:21to resolve those difficulties now there
  2231. 1:28:23are two optimizations that dramatically
  2232. 1:28:25help with the depth of these networks
  2233. 1:28:27and make sure that the networks remain
  2234. 1:28:29optimizable let's talk about the first
  2235. 1:28:31one the first one in this diagram is you
  2236. 1:28:33see this Arrow here and then this arrow
  2237. 1:28:36and this Arrow those are skip
  2238. 1:28:38connections or sometimes called residual
  2239. 1:28:40connections they come from this paper uh
  2240. 1:28:43the presidual learning for image
  2241. 1:28:44recognition from about
  2242. 1:28:462015 uh that introduced the concept now
  2243. 1:28:51these are basically what it means is you
  2244. 1:28:53transform data but then you have a skip
  2245. 1:28:55connection with addition from the
  2246. 1:28:57previous features now the way I like to
  2247. 1:29:00visualize it uh that I prefer is the
  2248. 1:29:03following here the computation happens
  2249. 1:29:05from the top to bottom and basically you
  2250. 1:29:08have this uh residual pathway and you
  2251. 1:29:11are free to Fork off from the residual
  2252. 1:29:13pathway perform some computation and
  2253. 1:29:15then project back to the residual
  2254. 1:29:16pathway via addition and so you go from
  2255. 1:29:19the the uh inputs to the targets only
  2256. 1:29:22via plus and plus plus and the reason
  2257. 1:29:25this is useful is because during back
  2258. 1:29:27propagation remember from our microG
  2259. 1:29:29grad video earlier addition distributes
  2260. 1:29:32gradients equally to both of its
  2261. 1:29:34branches that that fed as the input and
  2262. 1:29:37so the supervision or the gradients from
  2263. 1:29:40the loss basically hop through every
  2264. 1:29:43addition node all the way to the input
  2265. 1:29:46and then also Fork off into the residual
  2266. 1:29:50blocks but basically you have this
  2267. 1:29:52gradient Super Highway that goes
  2268. 1:29:53directly from the supervision all the
  2269. 1:29:55way to the input unimpeded and then
  2270. 1:29:58these viral blocks are usually
  2271. 1:29:59initialized in the beginning so they
  2272. 1:30:01contribute very very little if anything
  2273. 1:30:03to the residual pathway they they are
  2274. 1:30:05initialized that way so in the beginning
  2275. 1:30:07they are sort of almost kind of like not
  2276. 1:30:09there but then during the optimization
  2277. 1:30:11they come online over time and they uh
  2278. 1:30:14start to contribute but at least at the
  2279. 1:30:17initialization you can go from directly
  2280. 1:30:19supervision to the input gradient is
  2281. 1:30:21unimpeded and just flows and then the
  2282. 1:30:23blocks over time
  2283. 1:30:24kick in and so that dramatically helps
  2284. 1:30:27with the optimization so let's implement
  2285. 1:30:29this so coming back to our block here
  2286. 1:30:31basically what we want to do is we want
  2287. 1:30:33to do xal
  2288. 1:30:35X+ self attention and xal X+ self. feed
  2289. 1:30:39forward so this is X and then we Fork
  2290. 1:30:43off and do some communication and come
  2291. 1:30:45back and we Fork off and we do some
  2292. 1:30:46computation and come back so those are
  2293. 1:30:49residual connections and then swinging
  2294. 1:30:51back up here we also have to introd use
  2295. 1:30:54this projection so nn.
  2296. 1:30:57linear and uh this is going to be
  2297. 1:31:00from after we concatenate this this is
  2298. 1:31:03the prze and embed so this is the output
  2299. 1:31:05of the self tension itself but then we
  2300. 1:31:08actually want the uh to apply the
  2301. 1:31:11projection and that's the
  2302. 1:31:13result so the projection is just a
  2303. 1:31:15linear transformation of the outcome of
  2304. 1:31:16this
  2305. 1:31:17layer so that's the projection back into
  2306. 1:31:20the virual pathway and then here in a
  2307. 1:31:22feet forward it's going to be the same
  2308. 1:31:23same thing I could have a a self doot
  2309. 1:31:26projection here as well but let me just
  2310. 1:31:28simplify it and let me uh couple it
  2311. 1:31:32inside the same sequential container and
  2312. 1:31:34so this is the projection layer going
  2313. 1:31:36back into the residual
  2314. 1:31:38pathway and
  2315. 1:31:40so that's uh well that's it so now we
  2316. 1:31:43can train this so I implemented one more
  2317. 1:31:44small change when you look into the
  2318. 1:31:47paper again you see that the
  2319. 1:31:49dimensionality of input and output is
  2320. 1:31:51512 for them and they're saying that the
  2321. 1:31:53inner layer here in the feet forward has
  2322. 1:31:55dimensionality of 248 so there's a
  2323. 1:31:57multiplier of four and so the inner
  2324. 1:32:00layer of the feet forward Network should
  2325. 1:32:02be multiplied by four in terms of
  2326. 1:32:04Channel sizes so I came here and I
  2327. 1:32:06multiplied four times embed here for the
  2328. 1:32:08feed forward and then from four times
  2329. 1:32:10nmed coming back down to nmed when we go
  2330. 1:32:13back to the pro uh to the projection so
  2331. 1:32:15adding a bit of computation here and
  2332. 1:32:17growing that layer that is in the
  2333. 1:32:19residual block on the side of the
  2334. 1:32:21residual
  2335. 1:32:22pathway and then I train this and we
  2336. 1:32:24actually get down all the way to uh 2.08
  2337. 1:32:27validation loss and we also see that
  2338. 1:32:29network is starting to get big enough
  2339. 1:32:30that our train loss is getting ahead of
  2340. 1:32:32validation loss so we're starting to see
  2341. 1:32:33like a little bit of
  2342. 1:32:35overfitting and um our our
  2343. 1:32:38um uh Generations here are still not
  2344. 1:32:41amazing but at least you see that we can
  2345. 1:32:42see like is here this now grief syn like
  2346. 1:32:46this starts to almost look like English
  2347. 1:32:48so um yeah we're starting to really get
  2348. 1:32:50there okay and the second Innovation
  2349. 1:32:52that is very helpful for optimizing very
  2350. 1:32:54deep neural networks is right here so we
  2351. 1:32:57have this addition now that's the
  2352. 1:32:58residual part but this Norm is referring
  2353. 1:33:00to something called layer Norm so layer
  2354. 1:33:03Norm is implemented in pytorch it's a
  2355. 1:33:04paper that came out a while back here
  2356. 1:33:09um and layer Norm is very very similar
  2357. 1:33:11to bash Norm so remember back to our
  2358. 1:33:14make more series part three we
  2359. 1:33:16implemented bash
  2360. 1:33:17normalization and uh bash normalization
  2361. 1:33:19basically just made sure that um Across
  2362. 1:33:22The Bash dimension any individual neuron
  2363. 1:33:25had unit uh Gan um distribution so it
  2364. 1:33:30was zero mean and unit standard
  2365. 1:33:32deviation one standard deviation output
  2366. 1:33:35so what I did here is I'm copy pasting
  2367. 1:33:37the bashor 1D that we developed in our
  2368. 1:33:39make more series and see here we can
  2369. 1:33:42initialize for example this module and
  2370. 1:33:44we can have a batch of 32 100
  2371. 1:33:47dimensional vectors feeding through the
  2372. 1:33:48bachor layer so what this does is it
  2373. 1:33:52guarantees that when we look at just the
  2374. 1:33:54zeroth column it's a zero mean one
  2375. 1:33:58standard deviation so it's normalizing
  2376. 1:34:00every single column of this uh input now
  2377. 1:34:04the rows are not uh going to be
  2378. 1:34:06normalized by default because we're just
  2379. 1:34:08normalizing columns so let's now
  2380. 1:34:10Implement layer Norm uh it's very
  2381. 1:34:12complicated look we come here we change
  2382. 1:34:15this from zero to one so we don't
  2383. 1:34:18normalize The Columns we normalize the
  2384. 1:34:20rows and now we've implemented layer
  2385. 1:34:23Norm
  2386. 1:34:25so now the columns are not going to be
  2387. 1:34:28normalized um but the rows are going to
  2388. 1:34:31be normalized for every individual
  2389. 1:34:33example it's 100 dimensional Vector is
  2390. 1:34:35normalized uh in this way and because
  2391. 1:34:38our computation Now does not span across
  2392. 1:34:40examples we can delete all of this
  2393. 1:34:43buffers stuff uh because uh we can
  2394. 1:34:45always apply this operation and don't
  2395. 1:34:48need to maintain any running buffers so
  2396. 1:34:50we don't need the
  2397. 1:34:52buffers uh we
  2398. 1:34:54don't There's no distinction between
  2399. 1:34:56training and test
  2400. 1:34:58time uh and we don't need these running
  2401. 1:35:00buffers we do keep gamma and beta we
  2402. 1:35:03don't need the momentum we don't care if
  2403. 1:35:05it's training or not and this is now a
  2404. 1:35:08layer
  2405. 1:35:09norm and it normalizes the rows instead
  2406. 1:35:12of the columns and this here is
  2407. 1:35:15identical to basically this here so
  2408. 1:35:19let's now Implement layer Norm in our
  2409. 1:35:21Transformer before I incorporate the
  2410. 1:35:23layer Norm I just wanted to note that as
  2411. 1:35:25I said very few details about the
  2412. 1:35:27Transformer have changed in the last 5
  2413. 1:35:28years but this is actually something
  2414. 1:35:30that slightly departs from the original
  2415. 1:35:31paper you see that the ADD and Norm is
  2416. 1:35:34applied after the
  2417. 1:35:36transformation but um in now it is a bit
  2418. 1:35:40more uh basically common to apply the
  2419. 1:35:42layer Norm before the transformation so
  2420. 1:35:44there's a reshuffling of the layer Norms
  2421. 1:35:46uh so this is called the prorm
  2422. 1:35:48formulation and that's the one that
  2423. 1:35:49we're going to implement as well so
  2424. 1:35:50select deviation from the original paper
  2425. 1:35:53basically we need two layer Norms layer
  2426. 1:35:55Norm one is uh NN do layer norm and we
  2427. 1:35:59tell it how many um what is the
  2428. 1:36:01embedding Dimension and we need the
  2429. 1:36:03second layer norm and then here the
  2430. 1:36:06layer Norms are applied immediately on X
  2431. 1:36:09so self. layer Norm one applied on X and
  2432. 1:36:13self. layer Norm two applied on X before
  2433. 1:36:15it goes into self attention and feed
  2434. 1:36:18forward and uh the size of the layer
  2435. 1:36:20Norm here is an ed so 32 so when the
  2436. 1:36:23layer Norm is normalizing our features
  2437. 1:36:26it is uh the normalization here uh
  2438. 1:36:30happens the mean and the variance are
  2439. 1:36:32taken over 32 numbers so the batch and
  2440. 1:36:34the time act as batch Dimensions both of
  2441. 1:36:37them so this is kind of like a per token
  2442. 1:36:40um transformation that just normalizes
  2443. 1:36:42the features and makes them a unit mean
  2444. 1:36:46uh unit Gan at
  2445. 1:36:48initialization but of course because
  2446. 1:36:50these layer Norms inside it have these
  2447. 1:36:52gamma and beta training
  2448. 1:36:54parameters uh the layer Norm will U
  2449. 1:36:57eventually create outputs that might not
  2450. 1:36:59be unit gion but the optimization will
  2451. 1:37:01determine that so for now this is the uh
  2452. 1:37:05this is incorporating the layer norms
  2453. 1:37:06and let's train them on okay so I let it
  2454. 1:37:09run and we see that we get down to 2.06
  2455. 1:37:12which is better than the previous 2.08
  2456. 1:37:14so a slight Improvement by adding the
  2457. 1:37:15layer norms and I'd expect that they
  2458. 1:37:17help uh even more if we had bigger and
  2459. 1:37:19deeper Network one more thing I forgot
  2460. 1:37:21to add is that there should be a layer
  2461. 1:37:23Norm here also typically as at the end
  2462. 1:37:26of the Transformer and right before the
  2463. 1:37:28final uh linear layer that decodes into
  2464. 1:37:31vocabulary so I added that as well so at
  2465. 1:37:35this stage we actually have a pretty
  2466. 1:37:36complete uh Transformer according to the
  2467. 1:37:38original paper and it's a decoder only
  2468. 1:37:40Transformer I'll I'll talk about that in
  2469. 1:37:42a second uh but at this stage uh the
  2470. 1:37:44major pieces are in place so we can try
  2471. 1:37:46to scale this up and see how well we can
  2472. 1:37:47push this number now in order to scale
  2473. 1:37:50out the model I had to perform some
  2474. 1:37:51cosmetic changes here to make it nicer
  2475. 1:37:54so I introduced this variable called n
  2476. 1:37:56layer which just specifies how many
  2477. 1:37:57layers of the blocks we're going to have
  2478. 1:38:01I created a bunch of blocks and we have
  2479. 1:38:02a new variable number of heads as well I
  2480. 1:38:05pulled out the layer Norm here and uh so
  2481. 1:38:07this is identical now one thing that I
  2482. 1:38:10did briefly change is I added a Dropout
  2483. 1:38:13so Dropout is something that you can add
  2484. 1:38:15right before the residual connection
  2485. 1:38:17back right before the connection back
  2486. 1:38:19into the residual pathway so we can drop
  2487. 1:38:22out that as l layer here we can drop out
  2488. 1:38:26uh here at the end of the multi-headed
  2489. 1:38:27exension as well and we can also drop
  2490. 1:38:30out here uh when we calculate the um
  2491. 1:38:34basically affinities and after the
  2492. 1:38:36softmax we can drop out some of those so
  2493. 1:38:38we can randomly prevent some of the
  2494. 1:38:40nodes from
  2495. 1:38:41communicating and so Dropout uh comes
  2496. 1:38:43from this paper from 2014 or so and
  2497. 1:38:49basically it takes your neural
  2498. 1:38:50nut and it randomly every forward
  2499. 1:38:53backward pass shuts off some subset of
  2500. 1:38:56uh neurons so randomly drops them to
  2501. 1:38:59zero and trains without them and what
  2502. 1:39:02this does effectively is because the
  2503. 1:39:04mask of what's being dropped out is
  2504. 1:39:06changed every single forward backward
  2505. 1:39:07pass it ends up kind of uh training an
  2506. 1:39:11ensemble of sub networks and then at
  2507. 1:39:13test time everything is fully enabled
  2508. 1:39:15and kind of all of those sub networks
  2509. 1:39:16are merged into a single Ensemble if you
  2510. 1:39:18can if you want to think about it that
  2511. 1:39:20way so I would read the paper to get the
  2512. 1:39:22full detail for now we're just going to
  2513. 1:39:24stay on the level of this is a
  2514. 1:39:25regularization technique and I added it
  2515. 1:39:28because I'm about to scale up the model
  2516. 1:39:30quite a bit and I was concerned about
  2517. 1:39:32overfitting so now when we scroll up to
  2518. 1:39:34the top uh we'll see that I changed a
  2519. 1:39:36number of hyper parameters here about
  2520. 1:39:38our neural nut so I made the batch size
  2521. 1:39:40be much larger now it's 64 I changed the
  2522. 1:39:43block size to be 256 so previously it
  2523. 1:39:46was just eight eight characters of
  2524. 1:39:47context now it is 256 characters of
  2525. 1:39:50context to predict the 257th
  2526. 1:39:54uh I brought down the learning rate a
  2527. 1:39:55little bit because the neural net is now
  2528. 1:39:57much bigger so I brought down the
  2529. 1:39:58learning rate the embedding Dimension is
  2530. 1:40:01now 384 and there are six heads so 384
  2531. 1:40:05divide 6 means that every head is 64
  2532. 1:40:08dimensional as it as a standard and then
  2533. 1:40:11there's going to be six layers of that
  2534. 1:40:13and the Dropout will be at 02 so every
  2535. 1:40:15forward backward pass 20% of all of
  2536. 1:40:18these um intermediate calculations are
  2537. 1:40:21disabled and dropped to zero
  2538. 1:40:24and then I already trained this and I
  2539. 1:40:25ran it so uh drum roll how well does it
  2540. 1:40:28perform so let me just scroll up
  2541. 1:40:31here we get a validation loss of
  2542. 1:40:341.48 which is actually quite a bit of an
  2543. 1:40:37improvement on what we had before which
  2544. 1:40:38I think was 2.07 so it went from 2.07
  2545. 1:40:41all the way down to 1.48 just by scaling
  2546. 1:40:43up this neural nut with the code that we
  2547. 1:40:45have and this of course ran for a lot
  2548. 1:40:47longer this maybe trained for I want to
  2549. 1:40:49say about 15 minutes on my a100 GPU so
  2550. 1:40:52that's a pretty a GPU and if you don't
  2551. 1:40:54have a GPU you're not going to be able
  2552. 1:40:56to reproduce this uh on a CPU this would
  2553. 1:40:59be um I would not run this on a CPU or
  2554. 1:41:01MacBook or something like that you'll
  2555. 1:41:03have to Brak down the number of uh
  2556. 1:41:04layers and the embedding Dimension and
  2557. 1:41:06so on uh but in about 15 minutes we can
  2558. 1:41:09get this kind of a result and um I'm
  2559. 1:41:12printing some of the Shakespeare here
  2560. 1:41:15but what I did also is I printed 10,000
  2561. 1:41:17characters so a lot more and I wrote
  2562. 1:41:18them to a file and so here we see some
  2563. 1:41:21of the outputs
  2564. 1:41:24so it's a lot more recognizable as the
  2565. 1:41:26input text file so the input text file
  2566. 1:41:29just for reference looked like this so
  2567. 1:41:31there's always like someone speaking in
  2568. 1:41:33this manner and uh our predictions now
  2569. 1:41:37take on that form except of course
  2570. 1:41:40they're they're nonsensical when you
  2571. 1:41:41actually read them
  2572. 1:41:43so it is every crimp tap be a house oh
  2573. 1:41:47those
  2574. 1:41:48prepation we give
  2575. 1:41:51heed um you know
  2576. 1:41:56Oho sent me you mighty
  2577. 1:41:59Lord anyway so you can read through this
  2578. 1:42:02um it's nonsensical of course but this
  2579. 1:42:04is just a Transformer trained on a
  2580. 1:42:06character level for 1 million characters
  2581. 1:42:09that come from Shakespeare so there's
  2582. 1:42:10sort of like blabbers on in Shakespeare
  2583. 1:42:12like manner but it doesn't of course
  2584. 1:42:14make sense at this scale uh but I think
  2585. 1:42:18I think still a pretty good
  2586. 1:42:19demonstration of what's
  2587. 1:42:20possible so now
  2588. 1:42:24I think uh that kind of like concludes
  2589. 1:42:26the programming section of this video we
  2590. 1:42:28basically kind of uh did a pretty good
  2591. 1:42:30job and um of implementing this
  2592. 1:42:32Transformer uh but the picture doesn't
  2593. 1:42:35exactly match up to what we've done so
  2594. 1:42:37what's going on with all these digital
  2595. 1:42:38Parts here so let me finish explaining
  2596. 1:42:41this architecture and why it looks so
  2597. 1:42:43funky basically what's happening here is
  2598. 1:42:45what we implemented here is a decoder
  2599. 1:42:47only Transformer so there's no component
  2600. 1:42:50here this part is called the encoder and
  2601. 1:42:52there's no cross attention block here
  2602. 1:42:55our block only has a self attention and
  2603. 1:42:58the feet forward so it is missing this
  2604. 1:43:00third in between piece here this piece
  2605. 1:43:03does cross attention so we don't have it
  2606. 1:43:05and we don't have the encoder we just
  2607. 1:43:07have the decoder and the reason we have
  2608. 1:43:08a decoder only uh is because we are just
  2609. 1:43:12uh generating text and it's
  2610. 1:43:13unconditioned on anything we're just
  2611. 1:43:15we're just blabbering on according to a
  2612. 1:43:16given data set what makes it a decoder
  2613. 1:43:19is that we are using the Triangular mask
  2614. 1:43:21in our uh trans former so it has this
  2615. 1:43:24Auto regressive property where we can
  2616. 1:43:26just uh go and sample from it so the
  2617. 1:43:28fact that it's using the Triangular
  2618. 1:43:30triangular mask to mask out the
  2619. 1:43:32attention makes it a decoder and it can
  2620. 1:43:34be used for language modeling now the
  2621. 1:43:37reason that the original paper had an
  2622. 1:43:39incoder decoder architecture is because
  2623. 1:43:41it is a machine translation paper so it
  2624. 1:43:43is concerned with a different setting in
  2625. 1:43:45particular it expects some uh tokens
  2626. 1:43:49that encode say for example French and
  2627. 1:43:52then it is expecting to decode the
  2628. 1:43:54translation in English so so you
  2629. 1:43:56typically these here are special tokens
  2630. 1:43:59so you are expected to read in this and
  2631. 1:44:02condition on it and then you start off
  2632. 1:44:04the generation with a special token
  2633. 1:44:05called start so this is a special new
  2634. 1:44:08token um that you introduce and always
  2635. 1:44:10place in the beginning and then the
  2636. 1:44:12network is expected to Output neural
  2637. 1:44:15networks are awesome and then a special
  2638. 1:44:17end token to finish the
  2639. 1:44:20generation so this part here will be
  2640. 1:44:23decoded exactly as we we've done it
  2641. 1:44:25neural networks are awesome will be
  2642. 1:44:27identical to what we did but unlike what
  2643. 1:44:29we did they wanton to condition the
  2644. 1:44:32generation on some additional
  2645. 1:44:34information and in that case this
  2646. 1:44:36additional information is the French
  2647. 1:44:38sentence that they should be
  2648. 1:44:39translating so what they do now is they
  2649. 1:44:42bring in the encoder now the encoder
  2650. 1:44:45reads this part here so we're only going
  2651. 1:44:48to take the part of French and we're
  2652. 1:44:50going to uh create tokens from it
  2653. 1:44:52exactly as we've seen in our video and
  2654. 1:44:54we're going to put a Transformer on it
  2655. 1:44:57but there's going to be no triangular
  2656. 1:44:58mask and so all the tokens are allowed
  2657. 1:45:00to talk to each other as much as they
  2658. 1:45:02want and they're just encoding
  2659. 1:45:04whatever's the content of this French uh
  2660. 1:45:07sentence once they've encoded it they
  2661. 1:45:10they basically come out in the top here
  2662. 1:45:13and then what happens here is in our
  2663. 1:45:14decoder which does the uh language
  2664. 1:45:17modeling there's an additional
  2665. 1:45:20connection here to the outputs of the
  2666. 1:45:22encoder
  2667. 1:45:23and that is brought in through a cross
  2668. 1:45:26attention so the queries are still
  2669. 1:45:28generated from X but now the keys and
  2670. 1:45:30the values are coming from the side the
  2671. 1:45:32keys and the values are coming from the
  2672. 1:45:34top generated by the nodes that came
  2673. 1:45:36outside of the de the encoder and those
  2674. 1:45:40tops the keys and the values there the
  2675. 1:45:42top of it feed in on a side into every
  2676. 1:45:45single block of the decoder and so
  2677. 1:45:47that's why there's an additional cross
  2678. 1:45:49attention and really what it's doing is
  2679. 1:45:51it's conditioning the decoding
  2680. 1:45:53not just on the past of this current
  2681. 1:45:55decoding but also on having seen the
  2682. 1:45:59full fully encoded French um prompt sort
  2683. 1:46:04of and so it's an encoder decoder model
  2684. 1:46:06which is why we have those two
  2685. 1:46:07Transformers an additional block and so
  2686. 1:46:09on so we did not do this because we have
  2687. 1:46:12no we have nothing to encode there's no
  2688. 1:46:13conditioning we just have a text file
  2689. 1:46:15and we just want to imitate it and
  2690. 1:46:16that's why we are using a decoder only
  2691. 1:46:19Transformer exactly as done in
  2692. 1:46:21GPT okay okay so now I wanted to do a
  2693. 1:46:24very brief walkthrough of nanog GPT
  2694. 1:46:26which you can find in my GitHub and uh
  2695. 1:46:28nanog GPT is basically two files of
  2696. 1:46:30Interest there's train.py and model.py
  2697. 1:46:33train.py is all the boilerplate code for
  2698. 1:46:35training the network it is basically all
  2699. 1:46:38the stuff that we had here it's the
  2700. 1:46:40training loop it's just that it's a lot
  2701. 1:46:42more complicated because we're saving
  2702. 1:46:44and loading checkpoints and pre-trained
  2703. 1:46:46weights and we are uh decaying the
  2704. 1:46:48learning rate and compiling the model
  2705. 1:46:50and using distributed training across
  2706. 1:46:51multiple nodes or GP use so the training
  2707. 1:46:54Pi gets a little bit more hairy
  2708. 1:46:56complicated uh there's more options Etc
  2709. 1:46:59but the model.py should look very very
  2710. 1:47:01um similar to what we've done here in
  2711. 1:47:04fact the model is is almost identical so
  2712. 1:47:08first here we have the causal self
  2713. 1:47:09attention block and all of this should
  2714. 1:47:11look very very recognizable to you we're
  2715. 1:47:13producing queries Keys values we're
  2716. 1:47:16doing Dot products we're masking
  2717. 1:47:18applying soft Maxs optionally dropping
  2718. 1:47:20out and here we are pulling the wi the
  2719. 1:47:23values what is different here is that in
  2720. 1:47:25our code I have separated out the
  2721. 1:47:30multi-headed detention into just a
  2722. 1:47:31single individual head and then here I
  2723. 1:47:34have multiple heads and I explicitly
  2724. 1:47:36concatenate them whereas here uh all of
  2725. 1:47:39it is implemented in a batched manner
  2726. 1:47:41inside a single causal self attention
  2727. 1:47:43and so we don't just have a b and a T
  2728. 1:47:45and A C Dimension we also end up with a
  2729. 1:47:47fourth dimension which is the heads and
  2730. 1:47:50so it just gets a lot more sort of hairy
  2731. 1:47:52because we have four dimensional array
  2732. 1:47:54um tensors now but it is um equivalent
  2733. 1:47:57mathematically so the exact same thing
  2734. 1:47:59is happening as what we have it's just
  2735. 1:48:01it's a bit more efficient because all
  2736. 1:48:02the heads are now treated as a batch
  2737. 1:48:04Dimension as
  2738. 1:48:05well then we have the multier perceptron
  2739. 1:48:08it's using the Galu nonlinearity which
  2740. 1:48:10is defined here except instead of Ru and
  2741. 1:48:13this is done just because opening I used
  2742. 1:48:14it and I want to be able to load their
  2743. 1:48:17checkpoints uh the blocks of the
  2744. 1:48:19Transformer are identical to communicate
  2745. 1:48:21in the compute phase as we saw and then
  2746. 1:48:23the GPT will be identical we have the
  2747. 1:48:25position encodings token encodings the
  2748. 1:48:27blocks the layer Norm at the end uh the
  2749. 1:48:30final linear layer and this should look
  2750. 1:48:33all very recognizable and there's a bit
  2751. 1:48:35more here because I'm loading
  2752. 1:48:36checkpoints and stuff like that I'm
  2753. 1:48:38separating out the parameters into those
  2754. 1:48:40that should be weight decayed and those
  2755. 1:48:42that
  2756. 1:48:42shouldn't um but the generate function
  2757. 1:48:44should also be very very similar so a
  2758. 1:48:47few details are different but you should
  2759. 1:48:48definitely be able to look at this uh
  2760. 1:48:51file and be able to understand little
  2761. 1:48:52the pieces now so let's now bring things
  2762. 1:48:55back to chat GPT what would it look like
  2763. 1:48:57if we wanted to train chat GPT ourselves
  2764. 1:48:59and how does it relate to what we
  2765. 1:49:00learned today well to train in chat GPT
  2766. 1:49:03there are roughly two stages first is
  2767. 1:49:05the pre-training stage and then the
  2768. 1:49:07fine-tuning stage in the pre-training
  2769. 1:49:09stage uh we are training on a large
  2770. 1:49:12chunk of internet and just trying to get
  2771. 1:49:14a first decoder only Transformer to
  2772. 1:49:17babble text so it's very very similar to
  2773. 1:49:20what we've done ourselves except we've
  2774. 1:49:23done like a tiny little baby
  2775. 1:49:24pre-training step um and so in our case
  2776. 1:49:28uh this is how you print a number of
  2777. 1:49:30parameters I printed it and it's about
  2778. 1:49:3210 million so this Transformer that I
  2779. 1:49:35created here to create little
  2780. 1:49:37Shakespeare um Transformer was about 10
  2781. 1:49:40million parameters our data set is
  2782. 1:49:42roughly 1 million uh characters so
  2783. 1:49:45roughly 1 million tokens but you have to
  2784. 1:49:47remember that opening I is different
  2785. 1:49:48vocabulary they're not on the Character
  2786. 1:49:50level they use these um subword chunks
  2787. 1:49:53of words and so they have a vocabulary
  2788. 1:49:55of 50,000 roughly elements and so their
  2789. 1:49:58sequences are a bit more condensed so
  2790. 1:50:01our data set the Shakespeare data set
  2791. 1:50:03would be probably around 300,000 uh
  2792. 1:50:05tokens in the open AI vocabulary roughly
  2793. 1:50:09so we trained about 10 million parameter
  2794. 1:50:11model on roughly 300,000 tokens now when
  2795. 1:50:14you go to the gpt3
  2796. 1:50:16paper and you look at the Transformers
  2797. 1:50:20that they trained they trained a number
  2798. 1:50:22of trans Transformers of different sizes
  2799. 1:50:24but the biggest Transformer here has 175
  2800. 1:50:27billion parameters uh so ours is again
  2801. 1:50:2910 million they used this number of
  2802. 1:50:31layers in the Transformer this is the
  2803. 1:50:34nmed this is the number of heads and
  2804. 1:50:36this is the head size and then this is
  2805. 1:50:39the batch size uh so ours was
  2806. 1:50:4365 and the learning rate is similar now
  2807. 1:50:46when they train this Transformer they
  2808. 1:50:47trained on 300 billion tokens so again
  2809. 1:50:51remember ours is about 300,000
  2810. 1:50:53so this is uh about a millionfold
  2811. 1:50:56increase and this number would not be
  2812. 1:50:57even that large by today's standards
  2813. 1:50:59you'd be going up uh 1 trillion and
  2814. 1:51:01above so they are training a
  2815. 1:51:04significantly larger
  2816. 1:51:06model on uh a good chunk of the internet
  2817. 1:51:10and that is the pre-training stage but
  2818. 1:51:12otherwise these hyper parameters should
  2819. 1:51:13be fairly recognizable to you and the
  2820. 1:51:15architecture is actually like nearly
  2821. 1:51:17identical to what we implemented
  2822. 1:51:18ourselves but of course it's a massive
  2823. 1:51:20infrastructure challenge to train this
  2824. 1:51:22you're talking about typically thousands
  2825. 1:51:24of gpus having to you know talk to each
  2826. 1:51:27other to train models of this size so
  2827. 1:51:29that's just a pre-training stage now
  2828. 1:51:32after you complete the pre-training
  2829. 1:51:33stage uh you don't get something that
  2830. 1:51:35responds to your questions with answers
  2831. 1:51:38and is not helpful and Etc you get a
  2832. 1:51:40document
  2833. 1:51:41completer right so it babbles but it
  2834. 1:51:44doesn't Babble Shakespeare it babbles
  2835. 1:51:46internet it will create arbitrary news
  2836. 1:51:48articles and documents and it will try
  2837. 1:51:50to complete documents because that's
  2838. 1:51:51what it's trained for it's trying to
  2839. 1:51:52complete the sequence so when you give
  2840. 1:51:54it a question it would just uh
  2841. 1:51:56potentially just give you more questions
  2842. 1:51:58it would follow with more questions it
  2843. 1:52:00will do whatever it looks like the some
  2844. 1:52:02close document would do in the training
  2845. 1:52:05data on the internet and so who knows
  2846. 1:52:07you're getting kind of like undefined
  2847. 1:52:08Behavior it might basically answer with
  2848. 1:52:11to questions with other questions it
  2849. 1:52:13might ignore your question it might just
  2850. 1:52:15try to complete some news article it's
  2851. 1:52:17totally unineed as we say so the second
  2852. 1:52:20fine-tuning stage is to actually align
  2853. 1:52:22it to be an assistant and uh this is the
  2854. 1:52:25second stage and so this chat GPT block
  2855. 1:52:28post from openi talks a little bit about
  2856. 1:52:30how the stage is achieved we basically
  2857. 1:52:34um there's roughly three steps to to
  2858. 1:52:36this stage uh so what they do here is
  2859. 1:52:39they start to collect training data that
  2860. 1:52:41looks specifically like what an
  2861. 1:52:42assistant would do so these are
  2862. 1:52:44documents that have to format where the
  2863. 1:52:46question is on top and then an answer is
  2864. 1:52:47below and they have a large number of
  2865. 1:52:50these but probably not on the order of
  2866. 1:52:51the internet uh this is probably on the
  2867. 1:52:53of maybe thousands of examples and so
  2868. 1:52:58they they then fine-tune the model to
  2869. 1:53:00basically only focus on documents that
  2870. 1:53:03look like that and so you're starting to
  2871. 1:53:05slowly align it so it's going to expect
  2872. 1:53:07a question at the top and it's going to
  2873. 1:53:08expect to complete the answer and uh
  2874. 1:53:11these very very large models are very
  2875. 1:53:13sample efficient during their
  2876. 1:53:14fine-tuning so this actually somehow
  2877. 1:53:16works but that's just step one that's
  2878. 1:53:19just fine tuning so then they actually
  2879. 1:53:20have more steps where okay the second
  2880. 1:53:23step is you let the model respond and
  2881. 1:53:25then different Raiders look at the
  2882. 1:53:27different responses and rank them for
  2883. 1:53:29their preference as to which one is
  2884. 1:53:30better than the other they use that to
  2885. 1:53:32train a reward model so they can predict
  2886. 1:53:35uh basically using a different network
  2887. 1:53:37how much of any candidate
  2888. 1:53:39response would be desirable and then
  2889. 1:53:43once they have a reward model they run
  2890. 1:53:45po which is a form of polic policy
  2891. 1:53:47gradient um reinforcement learning
  2892. 1:53:49Optimizer to uh fine-tune this sampling
  2893. 1:53:53policy uh so that the answers that the
  2894. 1:53:55GP chat GPT now generates are expected
  2895. 1:53:59to score a high reward according to the
  2896. 1:54:02reward model and so basically there's a
  2897. 1:54:04whole aligning stage here or fine-tuning
  2898. 1:54:07stage it's got multiple steps in between
  2899. 1:54:09there as well and it takes the model
  2900. 1:54:11from being a document completer to a
  2901. 1:54:14question answerer and that's like a
  2902. 1:54:16whole separate stage a lot of this data
  2903. 1:54:19is not available publicly it is internal
  2904. 1:54:21to open AI and uh it's much harder to
  2905. 1:54:24replicate this stage um and so that's
  2906. 1:54:27roughly what would give you a chat GPT
  2907. 1:54:29and nanog GPT focuses on the
  2908. 1:54:31pre-training stage okay and that's
  2909. 1:54:32everything that I wanted to cover today
  2910. 1:54:35so we trained to summarize a decoder
  2911. 1:54:38only Transformer following this famous
  2912. 1:54:41paper attention is all you need from
  2913. 1:54:432017 and so that's basically a GPT we
  2914. 1:54:47trained it on Tiny Shakespeare and got
  2915. 1:54:50sensible results
  2916. 1:54:52all of the training code is
  2917. 1:54:54roughly 200 lines of code I will be
  2918. 1:54:57releasing this um code base so also it
  2919. 1:55:01comes with all the git log commits along
  2920. 1:55:04the way as we built it
  2921. 1:55:05up in addition to this code I'm going to
  2922. 1:55:08release the um notebook of course the
  2923. 1:55:10Google collab and I hope that gave you a
  2924. 1:55:13sense for how you can train um these
  2925. 1:55:16models like say gpt3 that will be um
  2926. 1:55:19architecturally basically identical to
  2927. 1:55:20what we have but they are somewhere
  2928. 1:55:22between 10,000 and 1 million times
  2929. 1:55:24bigger depending on how you count and so
  2930. 1:55:27uh that's all I have for now uh we did
  2931. 1:55:30not talk about any of the fine-tuning
  2932. 1:55:32stages that would typically go on top of
  2933. 1:55:33this so if you're interested in
  2934. 1:55:35something that's not just language
  2935. 1:55:36modeling but you actually want to you
  2936. 1:55:38know say perform tasks um or you want
  2937. 1:55:40them to be aligned in a specific way or
  2938. 1:55:43you want um to detect sentiment or
  2939. 1:55:45anything like that basically anytime you
  2940. 1:55:47don't want something that's just a
  2941. 1:55:48document completer you have to complete
  2942. 1:55:50further stages of fine tuning which did
  2943. 1:55:52not cover uh and that could be simple
  2944. 1:55:55supervised fine tuning or it can be
  2945. 1:55:57something more fancy like we see in chat
  2946. 1:55:58jpt where we actually train a reward
  2947. 1:56:00model and then do rounds of Po to uh
  2948. 1:56:03align it with respect to the reward
  2949. 1:56:04model so there's a lot more that can be
  2950. 1:56:06done on top of it I think for now we're
  2951. 1:56:08starting to get to about two hours Mark
  2952. 1:56:10uh so I'm going to um kind of finish
  2953. 1:56:13here uh I hope you enjoyed the lecture
  2954. 1:56:15uh and uh yeah go forth and transform
  2955. 1:56:18see you later

About this transcript

This page contains the full transcript of Let's build GPT: from scratch, in code, spelled out. by Andrej Karpathy, generated from the public captions YouTube serves with the video. The transcript has 21,030 words across 2,955 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.