# Trace - lecture_02_recording

https://cs336.stanford.edu/lectures/?trace=lecture_02_recording

lecture_02.py☀️⚪️🅴⬛⬅️➡️↖️↗️⤴️
1import torch.nn.functional as F
2import timeit
3import torch
4from typing import Iterable
5from torch import nn
6from torch.utils.checkpoint import checkpoint
7import numpy as np
8from edtrace import text, image, link
9from lecture_util import article_link
10from einops import rearrange, einsum, reduce
11from references import deepseek_v3_2_2025, adagrad_2011, nemotron_3_2025
12from gpu_util import cuda_if_available
13from facts import h100_flop_per_sec, h100_bytes_per_sec
14
15
16def main():
17
18
19
20
21
22
23
24
![](https://cs336.stanford.edu/lectures/var/files/image-75df5937a21eff96383d1c7c2ff05132-https_pbs_twimg_com_media_HE1P1HmaUAAjLXF_format_jpg_name_medium)
25
26
27
28
29
30
31
32
33 motivating_questions()
34
35
36
37
38
39
40 # Memory accounting
41 tensors_basics()
42 tensors_memory()
43 tensors_on_gpus()
44
45 # Compute accounting
46 tensor_einops()
47 tensor_operations_flops()
48
49 arithmetic_intensity()
50
51 # Full example
52 deep_linear_network()
53 gradients_basics()
54 gradients_flops()
55 optimizer()
56 train_loop()
57
58 # More memory optimizations
59 gradient_accumulation()
60 activation_checkpointing()
61
62
63
64
65
66
67
68
69
70
71def motivating_questions():
72
73 total_flops = 6 * 70e9 * 15e12
74 h100_flop_per_sec = 1979e12 / 2
75 mfu = 0.5
76 flops_per_day = h100_flop_per_sec * mfu * 1024 * 60 * 60 * 24
77 days = total_flops / flops_per_day 
78
79
80 h100_bytes = 80e9
81 bytes_per_parameter = 2 + 2 + (4 + 4) # parameters, gradients, optimizer state 
82 num_parameters = (h100_bytes * 8) / bytes_per_parameter 
83
84
85
86
87
88def tensors_basics():
89
90
91
92
[[DeepSeek-AI+ 2025]](https://arxiv.org/abs/2512.02556)
DeepSeek-V3.2: Pushing the Frontier of Open Large Language Models
DeepSeek-AI, Aixin Liu, Aoxue Mei, Bangcai Lin, Bing Xue ... (254 more) ... Yukun Zha, Zekai Zhang, Zhe Ju, Zhen Zhang, Zihua Qu
2025-12-02
We introduce DeepSeek-V3.2, a model that harmonizes high computational efficiency with superior reasoning and agent performance. The key technical breakthroughs of DeepSeek-V3.2 are as follows: (1) DeepSeek Sparse Attention (DSA): We introduce DSA, an efficient attention mechanism that substantially reduces computational complexity while preserving model performance in long-context scenarios. (2) Scalable Reinforcement Learning Framework: By implementing a robust reinforcement learning protocol and scaling post-training compute, DeepSeek-V3.2 performs comparably to GPT-5. Notably, our high-compute variant, DeepSeek-V3.2-Speciale, surpasses GPT-5 and exhibits reasoning proficiency on par with Gemini-3.0-Pro, achieving gold-medal performance in both the 2025 International Mathematical Olympiad (IMO) and the International Olympiad in Informatics (IOI). (3) Large-Scale Agentic Task Synthesis Pipeline: To integrate reasoning into tool-use scenarios, we developed a novel synthesis pipeline that systematically generates training data at scale. This methodology facilitates scalable agentic post-training, yielding substantial improvements in generalization and instruction-following robustness within complex, interactive environments.
93
[[DeepSeek v3.2 model on Hugging Face]](https://huggingface.co/deepseek-ai/DeepSeek-V3.2?show_file_info=model.safetensors.index.json)
DeepSeek v3.2 model on Hugging Face
94
95 x = torch.zeros(4) # vector 
96 x = torch.zeros(4, 8) # matrix 
97 x = torch.zeros(4, 8) # rank-3 tensor 
98
99
100def tensors_memory():
101
102
103
104
[[Wikipedia]](https://en.wikipedia.org/wiki/Single-precision_floating-point_format)
Wikipedia
105
![](https://cs336.stanford.edu/lectures/images/fp32.png)
106
107
108
109
110
111
112 x = torch.zeros(4, 8) 
113 assert x.dtype == torch.float32 # Default type
114 assert x.numel() == 4 * 8
115 assert x.element_size() == 4 # Float is 4 bytes
116 assert get_memory_usage(x) == 4 * 8 * 4 # 128 bytes
117
118
119 assert get_memory_usage(torch.empty(12288 * 4, 12288)) == 2304 * 1024 * 1024 # 2.3 GB 
120
121
122
[[Wikipedia]](https://en.wikipedia.org/wiki/Half-precision_floating-point_format)
Wikipedia
123
![](https://cs336.stanford.edu/lectures/images/fp16.png)
124
125 x = torch.zeros(4, 8, dtype=torch.float16) 
126 assert x.element_size() == 2
127
128 x = torch.tensor([1e-8], dtype=torch.float16) 
129 assert x == 0 # Underflow!
130
131
132
133
[[Wikipedia]](https://en.wikipedia.org/wiki/Bfloat16_floating-point_format)
Wikipedia
134
![](https://cs336.stanford.edu/lectures/images/bf16.png)
135
136
137
138 x = torch.tensor([1e-8], dtype=torch.bfloat16) 
139 assert x != 0 # No underflow!
140
141
142
143
144
145
[[Micikevicius+ 2017]](https://arxiv.org/pdf/1710.03740.pdf)
Mixed Precision Training
Paulius Micikevicius, Sharan Narang, Jonah Alben, Gregory Diamos, Erich Elsen ... (1 more) ... Boris Ginsburg, Michael Houston, Oleksii Kuchaiev, Ganesh Venkatesh, Hao Wu
2017-10-10
Deep neural networks have enabled progress in a wide variety of applications. Growing the size of the neural network typically results in improved accuracy. As model sizes grow, the memory and compute requirements for training these models also increases. We introduce a technique to train deep neural networks using half precision floating point numbers. In our technique, weights, activations and gradients are stored in IEEE half-precision format. Half-precision floating numbers have limited numerical range compared to single-precision numbers. We propose two techniques to handle this loss of information. Firstly, we recommend maintaining a single-precision copy of the weights that accumulates the gradients after each optimizer step. This single-precision copy is rounded to half-precision format during training. Secondly, we propose scaling the loss appropriately to handle the loss of information with half-precision gradients. We demonstrate that this approach works for a wide variety of models including convolution neural networks, recurrent neural networks and generative adversarial networks. This technique works for large scale models with more than 100 million parameters trained on large datasets. Using this approach, we can reduce the memory consumption of deep learning models by nearly 2x. In future processors, we can also expect a significant computation speedup using half-precision hardware units.
146
147
148
149
[[docs]](https://pytorch.org/docs/stable/amp.html)
docs
150
151
152
153
154
![](https://cs336.stanford.edu/lectures/var/files/image-df6d7649a3bdb77cfdc38092d8387a99-https_docs_nvidia_com_deeplearning_transformer-engine_user-guide__images_fp8_formats_png)
155
156
[[Micikevicius+ 2022]](https://arxiv.org/pdf/2209.05433.pdf)
FP8 Formats for Deep Learning
Paulius Micikevicius, Dusan Stosic, Neil Burgess, Marius Cornea, Pradeep Dubey ... (5 more) ... Naveen Mellempudi, Stuart Oberman, Mohammad Shoeybi, Michael Siu, Hao Wu
2022-09-12
FP8 is a natural progression for accelerating deep learning training inference beyond the 16-bit formats common in modern processors. In this paper we propose an 8-bit floating point (FP8) binary interchange format consisting of two encodings - E4M3 (4-bit exponent and 3-bit mantissa) and E5M2 (5-bit exponent and 2-bit mantissa). While E5M2 follows IEEE 754 conventions for representatio of special values, E4M3's dynamic range is extended by not representing infinities and having only one mantissa bit-pattern for NaNs. We demonstrate the efficacy of the FP8 format on a variety of image and language tasks, effectively matching the result quality achieved by 16-bit training sessions. Our study covers the main modern neural network architectures - CNNs, RNNs, and Transformer-based models, leaving all the hyperparameters unchanged from the 16-bit baseline training sessions. Our training experiments include large, up to 175B parameter, language models. We also examine FP8 post-training-quantization of language models trained using 16-bit formats that resisted fixed point int8 quantization.
157
158
159
160
161
162
163
[[NVIDIA+ 2025]](https://arxiv.org/abs/2512.20856)
NVIDIA Nemotron 3: Efficient and Open Intelligence
NVIDIA, Aaron Blakeman, Aaron Grattafiori, Aarti Basant, Abhibha Gupta ... (348 more) ... Zhen Dong, Zhongbo Zhu, Zihan Liu, Zijia Chen, Zijie Yan
2025-12-24
We introduce the Nemotron 3 family of models - Nano, Super, and Ultra. These models deliver strong agentic, reasoning, and conversational capabilities. The Nemotron 3 family uses a Mixture-of-Experts hybrid Mamba-Transformer architecture to provide best-in-class throughput and context lengths of up to 1M tokens. Super and Ultra models are trained with NVFP4 and incorporate LatentMoE, a novel approach that improves model quality. The two larger models also include MTP layers for faster text generation. All Nemotron 3 models are post-trained using multi-environment reinforcement learning enabling reasoning, multi-step tool use, and support granular reasoning budget control. Nano, the smallest model, outperforms comparable models in accuracy while remaining extremely cost-efficient for inference. Super is optimized for collaborative agents and high-volume workloads such as IT ticket automation. Ultra, the largest model, provides state-of-the-art accuracy and reasoning performance. Nano is released together with its technical report and this white paper, while Super and Ultra will follow in the coming months. We will openly release the model weights, pre- and post-training software, recipes, and all data for which we hold redistribution rights.
164
165
166
167
168def tensors_on_gpus():
169
170 x = torch.zeros(32, 32)
171 assert x.device == torch.device("cpu")
172
173
174
![](https://cs336.stanford.edu/lectures/images/cpu-gpu.png)
175
176
177 if not torch.cuda.is_available():
178 return
179
180 num_gpus = torch.cuda.device_count() 
181 for i in range(num_gpus):
182 properties = torch.cuda.get_device_properties(i) 
183
184 memory_allocated = torch.cuda.memory_allocated() 
185
186 text("Move the tensor to GPU memory (device 0).")
187 y = x.to("cuda:0")
188 assert y.device == torch.device("cuda", 0)
189
190 text("Or create a tensor directly on the GPU:")
191 z = torch.zeros(32, 32, device="cuda:0")
192
193 new_memory_allocated = torch.cuda.memory_allocated() 
194 memory_used = new_memory_allocated - memory_allocated 
195 assert memory_used == 2 * (32 * 32 * 4) # 2 32x32 matrices of 4-byte floats
196
197
198
199def tensor_einops():
200 einops_motivation()
201
202
203
204
[[Einops tutorial]](https://einops.rocks/1-einops-basics/)
Einops tutorial
205
206 einops_einsum()
207 einops_reduce()
208 einops_rearrange()
209
210
211def einops_motivation():
212
213
214 x = torch.ones(2, 2, 3) # batch seq hidden 
215 y = torch.ones(2, 2, 3) # batch seq hidden 
216 z = x @ y.transpose(-2, -1) # batch seq seq 
217
218
219
220def einops_einsum():
221
222
223 x = torch.ones(3, 4) # seq1 hidden 
224 y = torch.ones(4, 3) # hidden seq2 
225
226 # Old way
227 z = x @ y # seq1 seq2 
228
229 # New (einops) way
230 z = einsum(x, y, "seq1 hidden, hidden seq2 -> seq1 seq2") 
231
232
233
234 x = torch.ones(2, 3, 4) # batch seq1 hidden 
235 y = torch.ones(2, 3, 4) # batch seq2 hidden 
236
237 # Old way
238 z = x @ y.transpose(-2, -1) # batch seq1 seq2 
239
240 # New (einops) way
241 z = einsum(x, y, "batch seq1 hidden, batch seq2 hidden -> batch seq1 seq2") 
242
243
244 # Or can use `...` to represent broadcasting over any number of dimensions
245 z = einsum(x, y, "... seq1 hidden, ... seq2 hidden -> ... seq1 seq2") 
246
247
248def einops_reduce():
249
250 x = torch.ones(2, 3, 4) # batch seq hidden 
251
252 # Old way
253 y = x.sum(dim=-1) 
254
255 # New (einops) way
256 y = reduce(x, "... hidden -> ...", "sum") 
257
258
259def einops_rearrange():
260
261
262
263 x = torch.ones(3, 8) # seq total_hidden 
264
265 w = torch.ones(4, 4) # hidden1 hidden2 
266
267 # Break up `total_hidden` into two dimensions (`heads` and `hidden1`
268 x = rearrange(x, "... (heads hidden1) -> ... heads hidden1", heads=2) 
269
270 # Perform the transformation by `w`
271 x = einsum(x, w, "... hidden1, hidden1 hidden2 -> ... hidden2") 
272
273 # Combine `heads` and `hidden2` back together
274 x = rearrange(x, "... heads hidden2 -> ... (heads hidden2)") 
275
276
277def tensor_operations_flops():
278
279
280
281
282
283
284
285
286
287
[[article]](https://lambdalabs.com/blog/demystifying-gpt-3)
article
288
[[article]](https://patmcguinness.substack.com/p/gpt-4-details-revealed)
article
289
290
[[spec]](https://resources.nvidia.com/en-us-tensor-core/nvidia-tensor-core-gpu-datasheet)
spec
291 h100_flop_per_sec = 1979e12 / 2
292
293
294 total_flops = 8 * (60 * 60 * 24 * 7) * h100_flop_per_sec 
295
296
297
298
299
300
301
302 if torch.cuda.is_available():
303 B = 16384 # Number of points
304 D = 32768 # Dimension
305 K = 8192 # Number of outputs
306 else:
307 B = 1024
308 D = 256
309 K = 64
310
311 x = torch.ones(B, D, device=cuda_if_available())
312 w = torch.randn(D, K, device=cuda_if_available())
313 y = x @ w
314
315 actual_num_flops = 2 * B * D * K 
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330 actual_time = time_matmul(x, w) 
331 actual_flop_per_sec = actual_num_flops / actual_time 
332
333
334
[[H100 spec]](https://resources.nvidia.com/en-us-tensor-core/nvidia-tensor-core-gpu-datasheet)
H100 spec
335
336 promised_flop_per_sec = get_promised_flop_per_sec(x.dtype) 
337
338
339
340
341 mfu = actual_flop_per_sec / promised_flop_per_sec 
342
343
344
345
346
347
348
349
350def gradients_basics():
351
352
353
354
355
356
357
358 x = torch.tensor([1., 2, 3])
359 w = torch.tensor([1., 1, 1], requires_grad=True) # Want gradient
360 pred_y = x @ w
361 loss = 0.5 * (pred_y - 5).pow(2)
362
363
364 loss.backward()
365 assert loss.grad is None
366 assert pred_y.grad is None
367 assert x.grad is None
368 assert torch.equal(w.grad, torch.tensor([1, 2, 3])) 
369
370
371def gradients_flops():
372
373
374
![](https://cs336.stanford.edu/lectures/images/deep-network.png)
375
376 B = 1024 # Number of points
377 D = 256 # Dimension
378
379
380 x = torch.ones(B, D, device=cuda_if_available())
381 w1 = torch.randn(D, D, device=cuda_if_available(), requires_grad=True)
382 w2 = torch.randn(D, D, device=cuda_if_available(), requires_grad=True)
383
384 # Forward pass
385 h1 = einsum(x, w1, "batch in, in out -> batch out") # x 
386 h2 = einsum(h1, w2, "batch in, in out -> batch out") # h1 
387 loss = (h2.mean() - 0)**2 # Regress everything to 0 (arbitrary)
388
389 # Backward pass
390 h1.retain_grad() # For debugging
391 h2.retain_grad() # For debugging
392 loss.backward()
393
394
395
396
397
398 num_forward_flops = 2 * B * D * D 
399
400
401
402
403
404
405
406 h1_grad = einsum(h2.grad, w2, "batch out, in out -> batch in")
407 assert torch.allclose(h1.grad, h1_grad)
408
409 w2_grad = einsum(h2.grad, h1, "batch out, batch in -> in out")
410 assert torch.allclose(w2.grad, w2_grad)
411
412 num_backward_flops = (2 * B * D * D) + (2 * B * D * D) 
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428def arithmetic_intensity():
429
![](https://cs336.stanford.edu/lectures/images/compute-memory.png)
430
431
432
433
434
435
436
437
438
439
440 assert h100_flop_per_sec == 1979e12 / 2 # Half without sparsity
441 assert h100_bytes_per_sec == 3.35e12
442
443 arithmetic_intensity_relu()
444 arithmetic_intensity_gelu()
445 arithmetic_intensity_dot_product()
446 arithmetic_intensity_matrix_vector_product()
447 arithmetic_intensity_matmul()
448
449 # Let's visualize it
450 roofline_plots()
451
452
453def arithmetic_intensity_relu():
454 n = 1024 * 1024
455 x = torch.ones(n, dtype=torch.bfloat16, device=cuda_if_available())
456 y = torch.relu(x)
457
458 bytes = (2 * n) + (2 * n) # Read x, write y (bf16 is 2 bytes/float)
459 flops = n # n comparisons
460
461 communication_time = bytes / h100_bytes_per_sec 
462 computation_time = flops / h100_flop_per_sec 
463
464
465 total_time = max(communication_time, computation_time) 
466
467
468
469
470
471
472
473
474
475 h100_accelerator_intensity = h100_flop_per_sec / h100_bytes_per_sec 
476
477
478 arithmetic_intensity = flops / bytes # ~1/2 
479
480
481
482
483
484 assert arithmetic_intensity < h100_accelerator_intensity
485
486
487
488
489def arithmetic_intensity_gelu():
490 n = 1024
491 x = torch.ones(n, dtype=torch.bfloat16, device=cuda_if_available())
492 y = F.gelu(x) # GELU(x) = 0.5 x (1 + tanh(sqrt(2/pi) (x + 0.044715 x^3)))
493
494 bytes = (2 * n) + (2 * n) # Read x, write y (bf16 is 2 bytes/float)
495 flops = 20 * n # tanh can approximated in various ways (e.g., polynomial)
496
497 arithmetic_intensity = flops / bytes
498
499 h100_accelerator_intensity = h100_flop_per_sec / h100_bytes_per_sec 
500 assert arithmetic_intensity < h100_accelerator_intensity
501
502
503
504
505
506
507def arithmetic_intensity_dot_product():
508 n = 1024
509 x = torch.ones(n, dtype=torch.bfloat16, device=cuda_if_available())
510 w = torch.ones(n, dtype=torch.bfloat16, device=cuda_if_available())
511 y = x @ w
512
513 bytes = (2 * n) + (2 * n) + 2 # Read x, read w, write y
514 flops = 2 * n - 1 # n multiplications, n-1 additions
515
516 arithmetic_intensity = flops / bytes # ~1/2 
517
518 h100_accelerator_intensity = h100_flop_per_sec / h100_bytes_per_sec 
519 assert arithmetic_intensity < h100_accelerator_intensity
520
521
522
523def arithmetic_intensity_matrix_vector_product():
524 n = 1024
525 x = torch.ones(n, dtype=torch.bfloat16, device=cuda_if_available())
526 w = torch.ones(n, n, dtype=torch.bfloat16, device=cuda_if_available())
527 y = x @ w
528
529 bytes = (2 * n) + (2 * n * n) + (2 * n) # Read x, read w, write y
530 flops = n * (2 * n - 1) # n dot-products
531
532 arithmetic_intensity = flops / bytes # ~1 
533
534 h100_accelerator_intensity = h100_flop_per_sec / h100_bytes_per_sec 
535 assert arithmetic_intensity < h100_accelerator_intensity
536
537
538def arithmetic_intensity_matmul():
539 n = 1024
540 x = torch.ones(n, n, dtype=torch.bfloat16, device=cuda_if_available())
541 w = torch.ones(n, n, dtype=torch.bfloat16, device=cuda_if_available())
542 y = x @ w
543
544 bytes = (2 * n * n) + (2 * n * n) + (2 * n * n) # Read x, read w, write y
545 flops = n * n * (2 * n - 1) # n^2 dot products
546
547 arithmetic_intensity = flops / bytes # ~n/3 
548
549 h100_accelerator_intensity = h100_flop_per_sec / h100_bytes_per_sec 
550 assert arithmetic_intensity > h100_accelerator_intensity
551
552
553
554
555
556
557
558
559
560def roofline_plots():
561
562
![](https://cs336.stanford.edu/lectures/var/files/image-42d32b9c87939fe9a4b0a268d6d02ea7-https_jax-ml_github_io_scaling-book_assets_img_roofline-improved-1400_webp)
563
564
565
566
567
[[https://jax-ml.github.io/scaling-book/roofline/]](https://jax-ml.github.io/scaling-book/roofline/)
568
569
570def deep_linear_network():
571
![](https://cs336.stanford.edu/lectures/images/deep-network.png)
572
573
574 # Define the network
575 D = 8 # Dimensionality of input, activations, and output
576 L = 3 # Number of layers
577 model = DeepNetwork(dim=D, num_layers=L).to(cuda_if_available())
578
579 num_parameters = get_num_parameters(model) 
580 assert num_parameters == (D * D) * L
581
582 # Run the model on a batch of data
583 B = 4 # Batch size
584 x = torch.randn(B, D, device=cuda_if_available()) 
585 y = model(x) 
586
587
588class Block(nn.Module):
589 """Simple block that applies a linear transformation followed by a ReLU nonlinearity."""
590 def __init__(self, dim: int):
591 super().__init__()
592 self.weight = nn.Parameter(torch.randn(dim, dim) / np.sqrt(dim))
593
594 def forward(self, x: torch.Tensor) -> torch.Tensor:
595 x = x @ self.weight # Linear
596 x = F.relu(x) # Activation
597 return x
598
599
600class DeepNetwork(nn.Module):
601 """Map `dim`-vector to a `dim`-vector."""
602 def __init__(self, dim: int, num_layers: int):
603 super().__init__()
604 self.layers = nn.ModuleList([Block(dim) for i in range(num_layers)])
605
606 def forward(self, x: torch.Tensor) -> torch.Tensor:
607 # Apply all the layers sequentially
608 for layer in self.layers:
609 x = layer(x) 
610 return x
611
612
613def optimizer():
614
615 B = 2 # Batch size
616 D = 4 # Dimensionality of input, activations, and output
617 L = 3 # Number of layers
618 model = DeepNetwork(dim=D, num_layers=L).to(cuda_if_available()) 
619
620
621
622
623
624
625
626
[[Duchi+ 2011]](https://www.jmlr.org/papers/volume12/duchi11a/duchi11a.pdf)
627 optimizer = AdaGrad(model.parameters(), lr=0.01) 
628 state = model.state_dict() 
629
630
631 x = torch.randn(B, D, device=cuda_if_available())
632 y = torch.tensor([4., 5.], device=cuda_if_available())
633 pred_y = model(x).mean() 
634 loss = F.mse_loss(input=pred_y, target=y)
635 loss.backward()
636
637 # Take a step
638 optimizer.step()
639 optimizer_state = {i: dict(p_state) for i, (p, p_state) in enumerate(optimizer.state.items())} 
640
641 # Free up the memory
642 optimizer.zero_grad(set_to_none=True)
643
644
645
646 # Parameters
647 parameter_memory = 2 * (D * D * L) # (2 bytes for bf16) 
648
649 # Activations
650 activation_memory = 2 * B * D * L # (2 bytes for bf16) 
651
652 # Gradients
653 gradient_memory = 2 * parameter_memory # (2 bytes for bf16) 
654
655 # Optimizer states
656 optimizer_state_memory = 4 * parameter_memory # (4 bytes for fp32) 
657
658
659
660 # Putting it all together
661 total_memory = parameter_memory + activation_memory + gradient_memory + optimizer_state_memory 
662
663
664 num_parameters = D * D * L
665 flops = 6 * B * num_parameters 
666
667
668
669
670
671
672
[[article]](https://erees.dev/transformer-memory/)
article
673
[[article]](https://www.adamcasson.com/posts/transformer-flops)
article
674
675
676class AdaGrad(torch.optim.Optimizer):
677 def __init__(self, params: Iterable[nn.Parameter], lr: float = 0.01):
678 super(AdaGrad, self).__init__(params, dict(lr=lr))
679
680 def step(self):
681 for group in self.param_groups:
682 lr = group["lr"]
683 for p in group["params"]:
684 # Optimizer state
685 state = self.state[p]
686 grad = p.grad.data
687
688 # Get squared gradients g2 = sum_{i<t} g_i^2
689 g2 = state.get("g2", torch.zeros_like(grad))
690
691 # Update optimizer state
692 g2 += torch.square(grad)
693 state["g2"] = g2
694
695 # Update parameters
696 p.data -= lr * grad / torch.sqrt(g2 + 1e-5)
697
698
699def train_loop():
700 # True linear function with weights (0, 1, 2, ..., D-1)
701 D = 16 # Dimensionality
702 true_w = torch.arange(D, dtype=torch.float32, device=cuda_if_available())
703
704 # Data loader that generates (x, y) pairs
705 B = 4 # Batch size
706 def get_batch() -> tuple[torch.Tensor, torch.Tensor]:
707 x = torch.randn(B, D).to(cuda_if_available())
708 true_y = x @ true_w
709 return (x, true_y)
710
711 # Define the model and optimizer
712 L = 2 # Number of layers
713 model = DeepNetwork(dim=D, num_layers=L).to(cuda_if_available())
714 optimizer = AdaGrad(model.parameters(), lr=0.01)
715
716 # Train!
717 num_train_steps = 10
718 for t in range(num_train_steps):
719 # Get data
720 x, y = get_batch()
721
722 # Forward (compute loss)
723 pred_y = model(x).mean()
724 loss = F.mse_loss(pred_y, y)
725
726 # Backward (compute gradients)
727 loss.backward()
728
729 # Update parameters
730 optimizer.step()
731 optimizer.zero_grad(set_to_none=True)
732
733
734def gradient_accumulation():
735
736
737 B = 1024
738 L = 16
739 D = 1024
740 activation_memory = B * L * D 
741
742
743
744
745 micro_batch_size = 256
746 activation_memory = micro_batch_size * L * D 
747
748
749def activation_checkpointing():
750
751
752
753
![](https://cs336.stanford.edu/lectures/images/deep-network.png)
754
755 B = 64 # Batch size
756 L = 16 # Number of layers
757 D = 1024 # Dimensionality
758
759 x = torch.randn(B, D, device=cuda_if_available())
760 activation_memory = B * D * 2 * L 
761
762 model = DeepNetwork(dim=D, num_layers=L).to(cuda_if_available()) 
763 memory = get_max_memory_usage(lambda: model(x).backward()) 
764
765
766
767
768
769
770
771
772
773 # Store all activations: x g1 h1 g2 h2 g3 h3 g4 h4
774 # Activation checkpointing: x h1 h2 h3 h4
775
776 # Define the model with checkpointing
777 model = DeepNetworkCheckpointed(dim=D, num_layers=L).to(cuda_if_available()) 
778 checkpointed_memory = get_max_memory_usage(lambda: model(x).backward()) 
779
780
781
782 # Store all layers: | h1 h2 h3 h4 h5 h6 h7 h8 h9 |
783 # Store no layers: | |
784 # Store some layers: | h3 h6 h9 |
785
786
787
788
789
790
791
792class DeepNetworkCheckpointed(nn.Module):
793 """Same as DeepNetwork, but with activation checkpointing."""
794 def __init__(self, dim: int, num_layers: int):
795 super().__init__()
796 self.layers = nn.ModuleList([Block(dim) for i in range(num_layers)])
797
798 def forward(self, x: torch.Tensor) -> torch.Tensor:
799 # Apply all the layers sequentially
800 for layer in self.layers:
801 # KEY: line: we only store activations 
802 x = torch.utils.checkpoint.checkpoint(layer, x) 
803 return x
804
805def get_max_memory_usage(func):
806 """Measure how much memmory calling `func` uses."""
807 if not torch.cuda.is_available():
808 return 0 # Can't measure it without GPUs!
809
810 torch.cuda.empty_cache()
811 torch.cuda.reset_peak_memory_stats()
812 func()
813 return torch.cuda.max_memory_allocated()
814
815
816############################################################
817
818def get_memory_usage(x: torch.Tensor):
819 return x.numel() * x.element_size()
820
821
822def get_promised_flop_per_sec(dtype: torch.dtype) -> float:
823 """Return the peak FLOP/s for `device` operating on `dtype`."""
824 if not torch.cuda.is_available():
825 # No CUDA device available, so can't get FLOP/s
826 return 1
827 properties = torch.cuda.get_device_properties(cuda_if_available())
828
829 if "A100" in properties.name:
830 # https://www.nvidia.com/content/dam/en-zz/Solutions/Data-Center/a100/pdf/nvidia-a100-datasheet-us-nvidia-1758950-r4-web.pdf
831 if dtype == torch.float32:
832 return 19.5e12
833 if dtype in (torch.bfloat16, torch.float16):
834 return 312e12
835 raise ValueError(f"Unknown dtype: {dtype}")
836
837 if "H100" in properties.name:
838 # https://www.nvidia.com/en-us/data-center/h100/
839 if dtype == torch.float32:
840 return 67.5e12
841 if dtype in (torch.bfloat16, torch.float16):
842 return 1979e12 / 2 # 1979 is for sparse, dense is half of that
843 raise ValueError(f"Unknown dtype: {dtype}")
844
845 if "B200" in properties.name:
846 # https://www.primeline-solutions.com/media/categories/server/nach-gpu/nvidia-hgx-h200/nvidia-blackwell-b200-datasheet.pdf
847 if dtype == torch.float32:
848 return 75e12
849 if dtype in (torch.bfloat16, torch.float16):
850 return 4.5e15 / 2 # 4.5e15 is for sparse, dense is half of that
851 raise ValueError(f"Unknown dtype: {dtype}")
852
853
854def time_matmul(a: torch.Tensor, b: torch.Tensor) -> float:
855 """Return the number of seconds required to perform `a @ b`."""
856
857 # Wait until previous CUDA threads are done
858 if torch.cuda.is_available():
859 torch.cuda.synchronize()
860
861 def run():
862 # Perform the operation
863 a @ b
864
865 # Wait until CUDA threads are done
866 if torch.cuda.is_available():
867 torch.cuda.synchronize()
868
869 # Time the operation `num_trials` times
870 num_trials = 5
871 total_time = timeit.timeit(run, number=num_trials)
872
873 return total_time / num_trials
874
875
876def get_num_parameters(model: nn.Module) -> int:
877 return sum(param.numel() for param in model.parameters())
878
879
880if __name__ == "__main__":
881 main()
