# Trace - lecture_02

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

lecture_02.py☀️⚪️🅴⬛⬅️➡️↖️↗️⤴️
1import math
2import torch.nn.functional as F
3import timeit
4from typing import Iterable
5import torch
6from torch import nn
7from einops import rearrange, einsum, reduce
8
9from edtrace import text, image, link
10from lecture_util import article_link
11from gpu_util import cuda_if_available, get_max_memory_usage
12from facts import h100_flop_per_sec, h100_bytes_per_sec
13from references import deepseek_v3_2_2025, adagrad_2011, nemotron_3_super_2026
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 # Memory and compute accounting for training
52 deep_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 (2), gradients (2), optimizer state (4 + 4) 
82 num_parameters = (h100_bytes * 8) / bytes_per_parameter 
83
84
85
86
87
88
89def tensors_basics():
90
91
92
93
94
95
96
97
[[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.
98
[[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
99
100
101 x = torch.zeros(4) # rank 1 tensor (vector) 
102 x = torch.zeros(4, 8) # rank 2 tensor (matrix) 
103 x = torch.zeros(4, 8, 2) # rank 3 tensor 
104
105
106 B = 32 # Batch size
107 S = 16 # Sequence length
108 H = 16 # Number of heads
109 D = 64 # Hidden dimension per head
110 x = torch.zeros(B, S, H, D)
111
112
113def tensors_memory():
114
115
116
117
[[Wikipedia]](https://en.wikipedia.org/wiki/Single-precision_floating-point_format)
Wikipedia
118
![](https://cs336.stanford.edu/lectures/images/fp32.png)
119
120
121
122
123
124
125 x = torch.zeros(4, 8) 
126 assert x.dtype == torch.float32 # Default type
127 assert x.numel() == 4 * 8
128 assert x.element_size() == 4 # Float is 4 bytes
129 assert get_memory_usage(x) == 4 * 8 * 4 # 128 bytes
130
131
132 assert get_memory_usage(torch.empty(12288 * 4, 12288)) == 2304 * 1024 * 1024 # 2.3 GB 
133
134
135
[[Wikipedia]](https://en.wikipedia.org/wiki/Half-precision_floating-point_format)
Wikipedia
136
![](https://cs336.stanford.edu/lectures/images/fp16.png)
137
138 x = torch.zeros(4, 8, dtype=torch.float16) 
139 assert x.element_size() == 2
140
141 x = torch.tensor([1e-8], dtype=torch.float16) 
142 assert x == 0 # Underflow!
143
144
145
146
[[Wikipedia]](https://en.wikipedia.org/wiki/Bfloat16_floating-point_format)
Wikipedia
147
![](https://cs336.stanford.edu/lectures/images/bf16.png)
148
149
150
151 x = torch.tensor([1e-8], dtype=torch.bfloat16) 
152 assert x != 0 # No underflow!
153
154
155
156
157
158
159
[[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.
160
161
162
163
[[docs]](https://pytorch.org/docs/stable/amp.html)
docs
164
165 with torch.amp.autocast("cuda", dtype=torch.bfloat16):
166 x = torch.zeros(4, 8) 
167
168
169
170
![](https://cs336.stanford.edu/lectures/var/files/image-df6d7649a3bdb77cfdc38092d8387a99-https_docs_nvidia_com_deeplearning_transformer-engine_user-guide__images_fp8_formats_png)
171
172
[[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.
173
174
175
176
177
178
179
[[Nemotron 3 Super: Open, Efficient Mixture-of-Experts Hybrid Mamba-Transformer Model for Agentic Reasoning]](https://research.nvidia.com/labs/nemotron/files/NVIDIA-Nemotron-3-Super-Technical-Report.pdf)
Nemotron 3 Super: Open, Efficient Mixture-of-Experts Hybrid Mamba-Transformer Model for Agentic Reasoning
2026-03-11
180
181
182
183
184def tensors_on_gpus():
185
186 x = torch.zeros(32, 32)
187 assert x.device == torch.device("cpu")
188
189
190
![](https://cs336.stanford.edu/lectures/images/cpu-gpu.png)
191 device = cuda_if_available() 
192
193
194 x = x.to(device)
195
196
197 with torch.device(device):
198 x = torch.zeros(32, 32)
199 assert x.device == device
200
201
202def tensor_einops():
203 einops_motivation()
204
205
206
207
[[Einops tutorial]](https://einops.rocks/1-einops-basics/)
Einops tutorial
208
209 einops_einsum()
210 einops_reduce()
211 einops_rearrange()
212
213
214def einops_motivation():
215
216 x = torch.ones(2, 2, 3) # batch seq hidden 
217 y = torch.ones(2, 2, 3) # batch seq hidden 
218 z = x @ y.transpose(-2, -1) # batch seq seq 
219
220
221
222def einops_einsum():
223
224
225 x = torch.ones(3, 4) # seq1 hidden 
226 y = torch.ones(4, 3) # hidden seq2 
227
228 # Old way
229 z = x @ y # seq1 seq2 
230
231 # New (einops) way
232 z = einsum(x, y, "seq1 hidden, hidden seq2 -> seq1 seq2") 
233
234
235
236 x = torch.ones(2, 3, 4) # batch seq1 hidden 
237 y = torch.ones(2, 3, 4) # batch seq2 hidden 
238
239 # Old way
240 z = x @ y.transpose(-2, -1) # batch seq1 seq2 
241
242 # New (einops) way
243 z = einsum(x, y, "batch seq1 hidden, batch seq2 hidden -> batch seq1 seq2") 
244
245
246 # Or can use `...` to represent broadcasting over any number of dimensions
247 z = einsum(x, y, "... seq1 hidden, ... seq2 hidden -> ... seq1 seq2") 
248
249
250def einops_reduce():
251
252 x = torch.ones(2, 3, 4) # batch seq hidden 
253
254 # Old way
255 y = x.sum(dim=-1) 
256
257 # New (einops) way
258 y = reduce(x, "... hidden -> ...", "sum") 
259
260
261def einops_rearrange():
262
263
264
265 x = torch.ones(3, 8) # seq total_hidden 
266
267 w = torch.ones(4, 4) # hidden1 hidden2 
268
269 # Break up `total_hidden` into two dimensions (`heads` and `hidden1`
270 x = rearrange(x, "... (heads hidden1) -> ... heads hidden1", heads=2) 
271
272 # Perform the transformation by `w`
273 x = einsum(x, w, "... hidden1, hidden1 hidden2 -> ... hidden2") 
274
275 # Combine `heads` and `hidden2` back together
276 x = rearrange(x, "... heads hidden2 -> ... (heads hidden2)") 
277
278
279def tensor_operations_flops():
280
281
282
283
284
285
286
287
288
289
[[article]](https://lambdalabs.com/blog/demystifying-gpt-3)
article
290
[[article]](https://patmcguinness.substack.com/p/gpt-4-details-revealed)
article
291
292
[[spec]](https://resources.nvidia.com/en-us-tensor-core/nvidia-tensor-core-gpu-datasheet)
spec
293 h100_flop_per_sec = 1979e12 / 2
294
295
296 total_flops = 8 * 2 * (60 * 60 * 24 * 7) * h100_flop_per_sec 
297
298
299 if torch.cuda.is_available():
300 B = 16384 # Number of points
301 D = 32768 # Dimension of each point
302 K = 8192 # Number of outputs
303 else:
304 B = 1024
305 D = 256
306 K = 64
307
308 x = torch.ones(B, D, device=cuda_if_available())
309 w = torch.randn(D, K, device=cuda_if_available())
310 y = x @ w
311
312
313
314 actual_num_flops = 2 * B * D * K 
315
316
317 actual_time = benchmark(lambda: x @ w) 
318
319
320 actual_flop_per_sec = actual_num_flops / actual_time 
321
322
323
[[H100 spec]](https://resources.nvidia.com/en-us-gpu-resources/h100-datasheet-24306)
H100 spec
324
325 promised_flop_per_sec = get_promised_flop_per_sec(x.dtype) 
326
327
328
329
330 mfu = actual_flop_per_sec / promised_flop_per_sec if promised_flop_per_sec else None
331
332
333
334
335
336
337
338def arithmetic_intensity():
339
![](https://cs336.stanford.edu/lectures/images/compute-memory.png)
340
341
342
343
344
345
346
347
348
349
350 assert h100_flop_per_sec == 1979e12 / 2 # Half without sparsity
351 assert h100_bytes_per_sec == 3.35e12
352
353 arithmetic_intensity_relu()
354 arithmetic_intensity_gelu()
355 arithmetic_intensity_dot_product()
356 arithmetic_intensity_matrix_vector_product()
357 arithmetic_intensity_matmul()
358
359 # Let's visualize it
360 roofline_plots()
361
362
363def arithmetic_intensity_relu():
364 n = 1024 * 1024
365 x = torch.ones(n, dtype=torch.bfloat16, device=cuda_if_available())
366 y = torch.relu(x)
367
368 bytes = (2 * n) + (2 * n) # Read x, write y (bf16 is 2 bytes/float)
369 flops = n # n comparisons
370
371 communication_time = bytes / h100_bytes_per_sec 
372 computation_time = flops / h100_flop_per_sec 
373
374
375 total_time = max(communication_time, computation_time) 
376
377
378
379
380
381
382
383
384
385 h100_accelerator_intensity = h100_flop_per_sec / h100_bytes_per_sec 
386
387
388 arithmetic_intensity = flops / bytes # ~1/4 
389
390
391
392
393
394 assert arithmetic_intensity < h100_accelerator_intensity
395
396
397
398
399
400def arithmetic_intensity_gelu():
401 n = 1024 * 1024
402 x = torch.ones(n, dtype=torch.bfloat16, device=cuda_if_available())
403 y = F.gelu(x) # GELU(x) = 0.5 x (1 + tanh(sqrt(2/pi) (x + 0.044715 x^3)))
404
405 bytes = (2 * n) + (2 * n) # Read x, write y (bf16 is 2 bytes/float)
406 flops = 20 * n # tanh can be approximated in various ways (e.g., polynomials)
407
408 arithmetic_intensity = flops / bytes
409
410 h100_accelerator_intensity = h100_flop_per_sec / h100_bytes_per_sec 
411 assert arithmetic_intensity < h100_accelerator_intensity
412
413
414
415
416
417
418def arithmetic_intensity_dot_product():
419 n = 1024 * 1024
420 x = torch.ones(n, dtype=torch.bfloat16, device=cuda_if_available())
421 w = torch.ones(n, dtype=torch.bfloat16, device=cuda_if_available())
422 y = x @ w
423
424 bytes = (2 * n) + (2 * n) + 2 # Read x, read w, write y
425 flops = 2 * n - 1 # n multiplications, n-1 additions
426
427 arithmetic_intensity = flops / bytes # ~1/2 
428
429 h100_accelerator_intensity = h100_flop_per_sec / h100_bytes_per_sec 
430 assert arithmetic_intensity < h100_accelerator_intensity
431
432
433
434def arithmetic_intensity_matrix_vector_product():
435 n = 1024
436 x = torch.ones(n, dtype=torch.bfloat16, device=cuda_if_available())
437 w = torch.ones(n, n, dtype=torch.bfloat16, device=cuda_if_available())
438 y = x @ w
439
440 bytes = (2 * n) + (2 * n * n) + (2 * n) # Read x, read w, write y
441 flops = n * (2 * n - 1) # n dot-products
442
443 arithmetic_intensity = flops / bytes # ~1 
444
445 h100_accelerator_intensity = h100_flop_per_sec / h100_bytes_per_sec 
446 assert arithmetic_intensity < h100_accelerator_intensity
447
448
449def arithmetic_intensity_matmul():
450 n = 1024
451 x = torch.ones(n, n, dtype=torch.bfloat16, device=cuda_if_available())
452 w = torch.ones(n, n, dtype=torch.bfloat16, device=cuda_if_available())
453 y = x @ w
454
455 bytes = (2 * n * n) + (2 * n * n) + (2 * n * n) # Read x, read w, write y
456 flops = n * n * (2 * n - 1) # n^2 dot products
457
458 arithmetic_intensity = flops / bytes # ~n/3 
459
460 h100_accelerator_intensity = h100_flop_per_sec / h100_bytes_per_sec 
461 assert arithmetic_intensity > h100_accelerator_intensity
462
463
464
465
466
467
468
469
470
471def roofline_plots():
472
473
![](https://cs336.stanford.edu/lectures/var/files/image-42d32b9c87939fe9a4b0a268d6d02ea7-https_jax-ml_github_io_scaling-book_assets_img_roofline-improved-1400_webp)
474
475
476
477
478
479
480
481
[[reference]](https://jax-ml.github.io/scaling-book/roofline/)
reference
482
483
484def gradients_basics():
485
486
487
488
489
490
491
492 x = torch.tensor([1., 2, 3])
493 w = torch.tensor([1., 1, 1], requires_grad=True) # Want gradient
494 pred_y = x @ w
495 loss = 0.5 * (pred_y - 5).pow(2)
496
497
498 loss.backward()
499 assert torch.equal(w.grad, torch.tensor([1, 2, 3])) 
500
501
502def gradients_flops():
503
504
505
![](https://cs336.stanford.edu/lectures/images/deep-network.png)
506
507 B = 1024 # Number of points
508 D = 256 # Dimension
509
510
511 x = torch.ones(B, D, device=cuda_if_available())
512 w1 = torch.randn(D, D, device=cuda_if_available(), requires_grad=True)
513 w2 = torch.randn(D, D, device=cuda_if_available(), requires_grad=True)
514
515 # Forward pass
516 h1 = einsum(x, w1, "batch in, in out -> batch out") # x 
517 h2 = einsum(h1, w2, "batch in, in out -> batch out") # h1 
518 loss = (h2.mean() - 0)**2 # Regress everything to 0 (arbitrary)
519
520 # Backward pass
521 h1.retain_grad() # For debugging
522 h2.retain_grad() # For debugging
523 loss.backward()
524
525
526
527
528
529 num_forward_flops = 2 * B * D * D 
530
531
532
533
534
535
536
537 h1_grad = einsum(h2.grad, w2, "batch out, in out -> batch in")
538 assert torch.allclose(h1.grad, h1_grad)
539
540 w2_grad = einsum(h2.grad, h1, "batch out, batch in -> in out")
541 assert torch.allclose(w2.grad, w2_grad)
542
543 num_backward_flops = (2 * B * D * D) + (2 * B * D * D) 
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559def deep_network():
560
![](https://cs336.stanford.edu/lectures/images/deep-network.png)
561
562
563 # Define the network
564 D = 8 # Dimensionality of input, activations, and output
565 L = 3 # Number of layers
566 model = DeepNetwork(dim=D, num_layers=L).to(cuda_if_available())
567
568 num_parameters = get_num_parameters(model) 
569 assert num_parameters == (D * D) * L
570
571 # Run the model on a batch of data
572 B = 4 # Batch size
573 x = torch.randn(B, D, device=cuda_if_available()) 
574 y = model(x) 
575
576
577class Block(nn.Module):
578 """Simple block that applies a linear transformation followed by a ReLU nonlinearity."""
579 def __init__(self, dim: int):
580 super().__init__()
581 self.weight = nn.Parameter(torch.randn(dim, dim) / math.sqrt(dim))
582
583 def forward(self, x: torch.Tensor) -> torch.Tensor:
584 x = x @ self.weight # Linear
585 x = F.relu(x) # Activation
586 return x
587
588
589class DeepNetwork(nn.Module):
590 """Map `dim`-vector to a `dim`-vector."""
591 def __init__(self, dim: int, num_layers: int):
592 super().__init__()
593 self.layers = nn.ModuleList([Block(dim) for i in range(num_layers)])
594
595 def forward(self, x: torch.Tensor) -> torch.Tensor:
596 # Apply all the layers sequentially
597 for layer in self.layers:
598 x = layer(x) 
599 return x
600
601
602def optimizer():
603
604 B = 2 # Batch size
605 D = 4 # Dimensionality of input, activations, and output
606 L = 3 # Number of layers
607 model = DeepNetwork(dim=D, num_layers=L).to(cuda_if_available()) 
608
609
610
611
612
613
614
615
[[Duchi+ 2011]](https://www.jmlr.org/papers/volume12/duchi11a/duchi11a.pdf)
616 optimizer = AdaGrad(model.parameters(), lr=0.01) 
617 state = model.state_dict() 
618
619 # Compute gradients
620 x = torch.randn(B, D, device=cuda_if_available())
621 y = torch.tensor([4., 5.], device=cuda_if_available())
622 pred_y = model(x).mean() 
623 loss = F.mse_loss(input=pred_y, target=y)
624 loss.backward()
625
626 # Take a step
627 optimizer.step()
628 optimizer_state = {i: dict(p_state) for i, (p, p_state) in enumerate(optimizer.state.items())} 
629
630 # Free up the memory
631 optimizer.zero_grad(set_to_none=True)
632
633
634
635 num_parameters = D * D * L
636 parameter_memory = 2 * num_parameters # (2 bytes for bf16) 
637 gradient_memory = 2 * num_parameters # (2 bytes for bf16) 
638 optimizer_state_memory = 4 * num_parameters # (4 bytes for fp32) 
639 activation_memory = 2 * (B * D * L) # (2 bytes for bf16) 
640
641
642
643
644
645 # Putting it all together
646 total_memory = parameter_memory + activation_memory + gradient_memory + optimizer_state_memory 
647
648
649 num_parameters = D * D * L
650 flops = 6 * B * num_parameters 
651
652
653
654
655
656
[[article]](https://erees.dev/transformer-memory/)
article
657
[[article]](https://www.adamcasson.com/posts/transformer-flops)
article
658
659
660class AdaGrad(torch.optim.Optimizer):
661 def __init__(self, params: Iterable[nn.Parameter], lr: float = 0.01):
662 super(AdaGrad, self).__init__(params, dict(lr=lr))
663
664 def step(self):
665 for group in self.param_groups:
666 lr = group["lr"]
667 for p in group["params"]:
668 # Optimizer state
669 state = self.state[p]
670 grad = p.grad.data
671
672 # Get squared gradients g2 = sum_{i<t} g_i^2
673 g2 = state.get("g2", torch.zeros_like(grad))
674
675 # Update optimizer state
676 g2 += torch.square(grad)
677 state["g2"] = g2
678
679 # Update parameters
680 p.data -= lr * grad / torch.sqrt(g2 + 1e-5)
681
682
683def train_loop():
684 # True linear function with weights (0, 1, 2, ..., D-1)
685 D = 16 # Dimensionality
686 true_w = torch.arange(D, dtype=torch.float32, device=cuda_if_available())
687
688 # Data loader that generates (x, y) pairs
689 B = 4 # Batch size
690 def get_batch() -> tuple[torch.Tensor, torch.Tensor]:
691 x = torch.randn(B, D).to(cuda_if_available())
692 true_y = x @ true_w
693 return (x, true_y)
694
695 # Define the model and optimizer
696 L = 2 # Number of layers
697 model = DeepNetwork(dim=D, num_layers=L).to(cuda_if_available()) 
698 optimizer = AdaGrad(model.parameters(), lr=0.01) 
699
700 # Train!
701 num_train_steps = 10
702 for t in range(num_train_steps):
703 # Get data
704 x, y = get_batch()
705
706 # Forward (compute loss)
707 pred_y = model(x).mean() 
708 loss = F.mse_loss(pred_y, y)
709
710 # Backward (compute gradients)
711 loss.backward()
712
713 # Update parameters
714 optimizer.step()
715 optimizer.zero_grad(set_to_none=True)
716
717
718def gradient_accumulation():
719
720
721 B = 64 # Batch size
722 D = 1024 # Dimensionality
723 L = 16 # Number of layers
724 activation_memory = 2 * B * D * L # (2 bytes for bf16) 
725
726
727
728
729 micro_batch_size = 256
730 activation_memory = 2 * micro_batch_size * D * L # (2 bytes for bf16) 
731
732
733def activation_checkpointing():
734
735
736
737
![](https://cs336.stanford.edu/lectures/images/deep-network.png)
738
739 B = 64 # Batch size
740 D = 1024 # Dimensionality
741 L = 16 # Number of layers
742
743 x = torch.randn(B, D, device=cuda_if_available(), requires_grad=True)
744 activation_memory = 2 * B * D * L 
745
746 model = DeepNetwork(dim=D, num_layers=L).to(cuda_if_available()) 
747 memory = get_max_memory_usage(lambda: model(x).sum().backward()) 
748
749
750
751
752
753
754
755
756
757 # Store all activations: x g1 h1 g2 h2 g3 h3 g4 h4
758 # Activation checkpointing: x h1 h2 h3 h4
759
760 # Define the model with checkpointing
761 model = DeepNetworkCheckpointed(dim=D, num_layers=L).to(cuda_if_available()) 
762 checkpointed_memory = get_max_memory_usage(lambda: model(x).sum().backward()) 
763
764
765
766 # Store all layers: | h1 h2 h3 h4 h5 h6 h7 h8 h9 |
767 # Store no layers: | |
768 # Store some layers: | h3 h6 h9 |
769
770
771
772
773
774
775
776class DeepNetworkCheckpointed(nn.Module):
777 """Same as DeepNetwork, but with activation checkpointing."""
778 def __init__(self, dim: int, num_layers: int):
779 super().__init__()
780 self.layers = nn.ModuleList([Block(dim) for i in range(num_layers)])
781
782 def forward(self, x: torch.Tensor) -> torch.Tensor:
783 # Apply all the layers sequentially
784 for layer in self.layers:
785 # KEY: only store activations at checkpoints, recompute the rest
786 x = torch.utils.checkpoint.checkpoint(layer, x) 
787 return x
788
789############################################################
790
791def get_memory_usage(x: torch.Tensor):
792 return x.numel() * x.element_size()
793
794
795def get_promised_flop_per_sec(dtype: torch.dtype) -> float:
796 """Return the peak FLOP/s for `device` operating on `dtype`."""
797 if not torch.cuda.is_available():
798 # No CUDA device available, so can't get FLOP/s
799 return 1
800 properties = torch.cuda.get_device_properties(cuda_if_available()) 
801
802 if "A100" in properties.name:
803 # https://www.nvidia.com/content/dam/en-zz/Solutions/Data-Center/a100/pdf/nvidia-a100-datasheet-us-nvidia-1758950-r4-web.pdf
804 if dtype == torch.float32:
805 return 19.5e12
806 if dtype in (torch.bfloat16, torch.float16):
807 return 312e12
808 raise ValueError(f"Unknown dtype: {dtype}")
809
810 if "H100" in properties.name:
811 # https://www.nvidia.com/en-us/data-center/h100/
812 if dtype == torch.float32:
813 return 67.5e12
814 if dtype in (torch.bfloat16, torch.float16):
815 return 1979e12 / 2 # 1979 is for sparse, dense is half of that
816 raise ValueError(f"Unknown dtype: {dtype}")
817
818 if "B200" in properties.name:
819 # https://www.primeline-solutions.com/media/categories/server/nach-gpu/nvidia-hgx-h200/nvidia-blackwell-b200-datasheet.pdf
820 if dtype == torch.float32:
821 return 75e12
822 if dtype in (torch.bfloat16, torch.float16):
823 return 4.5e15 / 2 # 4.5e15 is for sparse, dense is half of that
824 raise ValueError(f"Unknown dtype: {dtype}")
825
826 # Unknown GPU: return None so caller can handle gracefully
827 return None
828
829
830def benchmark(func, num_trials: int = 5) -> float:
831 """Return the number of seconds required to perform `func`."""
832
833 # Wait until previous CUDA threads are done
834 if torch.cuda.is_available():
835 torch.cuda.synchronize()
836
837 def run():
838 # Perform the operation
839 func()
840
841 # Wait until CUDA threads are done
842 if torch.cuda.is_available():
843 torch.cuda.synchronize()
844
845 # Time the operation `num_trials` times
846 total_time = timeit.timeit(run, number=num_trials)
847
848 return total_time / num_trials
849
850
851def get_num_parameters(model: nn.Module) -> int:
852 return sum(param.numel() for param in model.parameters())
853
854
855if __name__ == "__main__":
856 main()
