# Trace - lecture_07

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

lecture_07.py☀️⚪️🅴⬛⬅️➡️↖️↗️⤴️
1import torch
2import time
3import math
4import sys
5import os
6from inspect import isfunction
7from typing import Callable
8from torch import nn, tensor
9import torch.nn.functional as F
10import torch.distributed as dist
11import torch.multiprocessing as mp
12from edtrace import text, image, link
13from gpu_util import cuda_if_available
14from lecture_util import article_link
15
16if not torch.cuda.is_available():
17 torch.cuda.synchronize = lambda: None # No-op if CUDA is not available
18
19def main():
20
21
22
23
![](https://cs336.stanford.edu/lectures/images/gpu-node-overview.png)
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41 # When you execute this lecture directly (python lecture_07.py), it uses multiprocessing, which produces output from each process (below).
42 # However, when you trace this lecture (python -m edtrace.execute -m lecture_07), we turn off multiprocessing.
43
[[stdout for this lecture]](https://cs336.stanford.edu/lectures/var/traces/lecture_07_stdout.txt)
stdout for this lecture
44
45
46 collective_operations() # Programming model
47 hardware() # Hardware: how GPUs are connected
48 torch_distributed() # How this is implemented in NCCL/PyTorch
49 benchmarking() # Measure actual NCCL bandwidth
50
51
52
53
54
55 data_parallelism() # Cut up along the batch dimension
56 tensor_parallelism() # Cut up along the width dimension
57 pipeline_parallelism() # Cut up along the depth dimension
58
59
60
61
62
63
[[levanter]](https://crfm.stanford.edu/2023/06/16/levanter-1_0-release.html)
levanter
64
65
66
67
68
69
70
71
72
73
74
75def collective_operations():
76
[[article]](https://en.wikipedia.org/wiki/Collective_operation)
article
77
78
79
80
81
82
![](https://cs336.stanford.edu/lectures/images/ranks.png)
83
84
85
86
87
88
89
90
91
92 # Input
93 rank0 = tensor([0., 1, 2, 3])
94
95 # Output
96 rank0 = tensor([0., 1, 2, 3])
97 rank1 = tensor([0., 1, 2, 3])
98 rank2 = tensor([0., 1, 2, 3])
99 rank3 = tensor([0., 1, 2, 3])
100
101
102
103
104 # Input
105 rank0 = tensor([0., 1, 2, 3])
106
107 # Output
108 rank0 = tensor([0.])
109 rank1 = tensor([1.])
110 rank2 = tensor([2.])
111 rank3 = tensor([3.])
112
113
114
115
116 # Input
117 rank0 = tensor([0.])
118 rank1 = tensor([1.])
119 rank2 = tensor([2.])
120 rank3 = tensor([3.])
121
122 # Output
123 rank0 = tensor([0., 1, 2, 3])
124
125
126
127
128 # Input
129 rank0 = tensor([0.])
130 rank1 = tensor([1.])
131 rank2 = tensor([2.])
132 rank3 = tensor([3.])
133
134 # Output
135 rank0 = tensor([6.]) # Sum of all ranks (0 + 1 + 2 + 3)
136
137
138
139
140 # Input
141 rank0 = tensor([0.])
142 rank1 = tensor([1.])
143 rank2 = tensor([2.])
144 rank3 = tensor([3.])
145
146 # Output
147 rank0 = tensor([0., 1, 2, 3])
148 rank1 = tensor([0., 1, 2, 3])
149 rank2 = tensor([0., 1, 2, 3])
150 rank3 = tensor([0., 1, 2, 3])
151
152
153
154
155 # Input
156 rank0 = tensor([0., 1, 2, 3])
157 rank1 = tensor([1., 2, 3, 4])
158 rank2 = tensor([2., 3, 4, 5])
159 rank3 = tensor([3., 4, 5, 6])
160
161 # Output
162 rank0 = tensor([6.]) # Sum along dim 0 (0 + 1 + 2 + 3)
163 rank1 = tensor([10.]) # Sum along dim 1 (1 + 2 + 3 + 4)
164 rank2 = tensor([14.]) # Sum along dim 2 (2 + 3 + 4 + 5)
165 rank3 = tensor([18.]) # Sum along dim 3 (3 + 4 + 5 + 6)
166
167
168
169
170 # Input
171 rank0 = tensor([0., 1, 2, 3])
172 rank1 = tensor([1., 2, 3, 4])
173 rank2 = tensor([2., 3, 4, 5])
174 rank3 = tensor([3., 4, 5, 6])
175
176 # Output
177 rank0 = tensor([6., 10, 14, 18])
178 rank1 = tensor([6., 10, 14, 18])
179 rank2 = tensor([6., 10, 14, 18])
180 rank3 = tensor([6., 10, 14, 18])
181
182
183
184
185
186 # Input
187 rank0 = tensor([0., 1, 2, 3]) # send 0 to rank 0, 1 to rank 1, 2 to rank 2, 3 to rank 3
188 rank1 = tensor([4., 5, 6, 7]) # send 4 to rank 0, 5 to rank 1, 6 to rank 2, 7 to rank 3
189 rank2 = tensor([8., 9, 10, 11]) # send 8 to rank 0, 9 to rank 1, 10 to rank 2, 11 to rank 3
190 rank3 = tensor([12., 13, 14, 15]) # send 12 to rank 0, 13 to rank 1, 14 to rank 2, 15 to rank 3
191
192 # Output
193 rank0 = tensor([0, 4, 8, 12])
194 rank1 = tensor([1, 5, 9, 13])
195 rank2 = tensor([2, 6, 10, 14])
196 rank3 = tensor([3, 7, 11, 15])
197
198
199
200
201
202
203
204
205
206
207
208
209def hardware():
210
211
![](https://cs336.stanford.edu/lectures/var/files/image-b0641f11a73711b3078acbd257b0c805-https_media_springernature_com_lw685_springer-static_image_art_3A10_1186_2Fs42774-021-00098-3_MediaObjects_42774_2021_98_Fig1_HTML_png_as_webp)
212
[[article]](https://en.wikipedia.org/wiki/PCI_Express)
article
213
214
215
216
![](https://cs336.stanford.edu/lectures/images/gpu-node-overview.png)
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
[[talk]](https://www.nvidia.com/en-us/on-demand/session/gtcspring21-s31880/)
talk
234
235
236
237
238
239def torch_distributed():
240
[[documentation]](https://pytorch.org/docs/stable/distributed.html)
documentation
241
242
243
244
245
246 spawn(collective_operations_main, world_size=4)
247
248
249def collective_operations_main(rank: int, world_size: int): 
250 """This function is running asynchronously for each process (rank = 0, ..., world_size - 1)."""
251 setup(rank, world_size)
252
253 ### All-reduce (dist = torch.distributed)
254 dist.barrier() # Waits for all processes to get to this point (in this case, for print statements)
255
256 data = tensor([0., 1, 2, 3], device=cuda_if_available(rank)) + rank # Both input and output
257
258 print(f"Rank {rank} [before all-reduce]: {data}", flush=True)
259 dist.all_reduce(tensor=data, op=dist.ReduceOp.SUM, async_op=False) # Modifies tensor in place
260 print(f"Rank {rank} [after all-reduce]: {data}", flush=True)
261
262 ### Reduce-scatter
263 dist.barrier()
264
265 input = torch.arange(world_size, dtype=torch.float32, device=cuda_if_available(rank)) + rank # Input
266 output = torch.empty(1, device=cuda_if_available(rank)) # Allocate output
267
268 print(f"Rank {rank} [before reduce-scatter]: input = {input}, output = {output}", flush=True)
269 dist.reduce_scatter_tensor(output=output, input=input, op=dist.ReduceOp.SUM, async_op=False)
270 print(f"Rank {rank} [after reduce-scatter]: input = {input}, output = {output}", flush=True)
271
272 ### All-gather
273 dist.barrier()
274
275 input = output # Input is the output of reduce-scatter
276 output = torch.empty(world_size, device=cuda_if_available(rank)) # Allocate output
277
278 print(f"Rank {rank} [before all-gather]: input = {input}, output = {output}", flush=True)
279 dist.all_gather_into_tensor(output_tensor=output, input_tensor=input, async_op=False)
280 print(f"Rank {rank} [after all-gather]: input = {input}, output = {output}", flush=True)
281
282
283
284 cleanup()
285
286
287def benchmarking():
288
289
290 # All-reduce
291 spawn(all_reduce, world_size=4, num_elements=100 * 1024**2)
292
293 # Reduce-scatter
294 spawn(reduce_scatter, world_size=4, num_elements=100 * 1024**2)
295
296
297
[[How to reason about collective operations]](https://github.com/NVIDIA/nccl-tests/blob/master/doc/PERFORMANCE.md#allreduce)
How to reason about collective operations
298
[[Sample benchmarking code]](https://github.com/stas00/ml-engineering/blob/master/network/benchmarks/all_reduce_bench.py)
Sample benchmarking code
299
300
301def all_reduce(rank: int, world_size: int, num_elements: int):
302 setup(rank, world_size) 
303
304 # Create tensor
305 data = torch.randn(num_elements, device=cuda_if_available(rank))
306
307 # Warmup
308 dist.all_reduce(tensor=data, op=dist.ReduceOp.SUM, async_op=False)
309 torch.cuda.synchronize() # Wait for CUDA kernels to finish
310 dist.barrier() # Wait for all the processes to get here
311
312 # Perform all-reduce
313 start_time = time.time()
314 dist.all_reduce(tensor=data, op=dist.ReduceOp.SUM, async_op=False)
315 torch.cuda.synchronize() # Wait for CUDA kernels to finish
316 dist.barrier() # Wait for all the processes to get here
317 end_time = time.time()
318
319 duration = end_time - start_time
320 print(f"[all_reduce] Rank {rank}: all_reduce(world_size={world_size}, num_elements={num_elements}) took {render_duration(duration)}", flush=True) 
321
322 # Measure the effective bandwidth
323 dist.barrier()
324 size_bytes = data.element_size() * data.numel()
325 sent_bytes = size_bytes * 2 * (world_size - 1) # 2x because send + receive, world_size-1 steps in all-reduce
326 total_duration = world_size * duration
327 bandwidth = sent_bytes / total_duration
328 print(f"[all_reduce] Rank {rank}: all_reduce measured bandwidth = {round(bandwidth / 1024**3)} GB/s", flush=True)
329
330 # Notes:
331 # - Effective bandwidth ~ 2 * size_bytes / total_duration
332 # - Independent of world_size
333 # - Independent of topology (ring or tree)
334
335 cleanup() 
336
337
338def reduce_scatter(rank: int, world_size: int, num_elements: int):
339 setup(rank, world_size) 
340
341 # Create input and outputs
342 input = torch.randn(world_size, num_elements, device=cuda_if_available(rank)) # Each rank has a matrix
343 output = torch.empty(num_elements, device=cuda_if_available(rank))
344
345 # Warmup
346 dist.reduce_scatter_tensor(output=output, input=input, op=dist.ReduceOp.SUM, async_op=False)
347 torch.cuda.synchronize() # Wait for CUDA kernels to finish
348 dist.barrier() # Wait for all the processes to get here
349
350 # Perform reduce-scatter
351 start_time = time.time()
352 dist.reduce_scatter_tensor(output=output, input=input, op=dist.ReduceOp.SUM, async_op=False)
353 torch.cuda.synchronize() # Wait for CUDA kernels to finish
354 dist.barrier() # Wait for all the processes to get here
355 end_time = time.time()
356
357 duration = end_time - start_time
358 print(f"[reduce_scatter] Rank {rank}: reduce_scatter(world_size={world_size}, num_elements={num_elements}) took {render_duration(duration)}", flush=True) 
359
360 # Measure the effective bandwidth
361 dist.barrier()
362 data_bytes = input.element_size() * input.numel() # How much data in the input
363 sent_bytes = data_bytes * (world_size - 1) # How much needs to be sent (no 2x here)
364 total_duration = world_size * duration # Total time for transmission
365 bandwidth = sent_bytes / total_duration
366 print(f"[reduce_scatter] Rank {rank}: reduce_scatter measured bandwidth = {round(bandwidth / 1024**3)} GB/s", flush=True)
367
368 # Notes:
369 # - all-reduce = reduce-scatter + all-gather
370 # - all-reduce moves 2x the data in 2x the time compared to reduce-scatter, so similar bandwidth
371
372 cleanup() 
373
374
375def data_parallelism():
376
![](https://cs336.stanford.edu/lectures/images/data-parallelism.png)
377
378
379 data = generate_sample_data()
380 spawn(data_parallelism_main, world_size=4, data=data, num_layers=4, num_steps=1)
381
382
383
384
385
386
387
388
389
390def generate_sample_data():
391 batch_size = 128
392 num_dim = 1024
393 data = torch.randn(batch_size, num_dim)
394 return data
395
396
397def data_parallelism_main(rank: int, world_size: int, data: tensor, num_layers: int, num_steps: int):
398 setup(rank, world_size) 
399
400 # Get the slice of data for this rank (in practice, each rank should load only its own data)
401 # --- B0 ---
402 # --- B1 ---
403 # --- B2 ---
404 # --- B3 ---
405 batch_size = data.size(0) 
406 num_dim = data.size(1) 
407 local_batch_size = int_divide(batch_size, world_size) 
408 start_index = rank * local_batch_size 
409 end_index = start_index + local_batch_size 
410 data = data[start_index:end_index].to(cuda_if_available(rank))
411
412 # Create MLP parameters params[0], ..., params[num_layers - 1] (each rank has all parameters)
413 params = [get_init_params(num_dim, num_dim, rank) for layer in range(num_layers)]
414 optimizer = torch.optim.AdamW(params, lr=1e-3) # Each rank has own optimizer state
415
416 for step in range(num_steps):
417 # Forward pass
418 x = data
419 for param in params:
420 x = x @ param
421 x = F.gelu(x)
422 loss = x.square().mean() # Loss function is average squared magnitude
423
424 # Backward pass
425 loss.backward()
426
427 # Sync gradients across workers (ONLY difference between standard training and DDP)
428 for param in params:
429 dist.all_reduce(tensor=param.grad, op=dist.ReduceOp.AVG, async_op=False)
430
431 # Update parameters
432 optimizer.step()
433
434 print(f"[data_parallelism] Rank {rank}: step = {step}, loss = {loss.item()}, params = {[summarize_tensor(params[layer]) for layer in range(num_layers)]}", flush=True) 
435
436 cleanup() 
437
438
439def tensor_parallelism():
440
![](https://cs336.stanford.edu/lectures/images/tensor-parallelism.png)
441
442
443 data = generate_sample_data()
444 spawn(tensor_parallelism_main, world_size=4, data=data, num_layers=4)
445
446
447def tensor_parallelism_main(rank: int, world_size: int, data: tensor, num_layers: int):
448 setup(rank, world_size) 
449
450 data = data.to(cuda_if_available(rank)) # All ranks get the data (batch_size x num_dim)
451 batch_size = data.size(0) 
452 num_dim = data.size(1) 
453 local_num_dim = int_divide(num_dim, world_size) # Shard `num_dim` 
454
455 # Create model (each rank gets 1/world_size of the parameters)
456 # | | | |
457 # W0 W1 W2 W3
458 # | | | |
459 params = [get_init_params(num_dim, local_num_dim, rank) for layer in range(num_layers)]
460
461 # Forward pass
462 x = data
463 for layer in range(num_layers):
464 # Compute activations (batch_size x local_num_dim)
465 x = x @ params[layer] # Note: this is only on a slice of the parameters
466 x = F.gelu(x)
467
468 # Allocate memory for activations (world_size x batch_size x local_num_dim)
469 activations = [torch.empty(batch_size, local_num_dim, device=cuda_if_available(rank)) for _ in range(world_size)]
470
471 # Send activations via all gather
472 dist.all_gather(tensor_list=activations, tensor=x, async_op=False)
473
474 # Concatenate them to get batch_size x num_dim
475 x = torch.cat(activations, dim=1)
476
477 print(f"[tensor_parallelism] Rank {rank}: forward pass produced activations {summarize_tensor(x)}", flush=True) 
478
479 # Backward pass: homework exercise
480
481 cleanup() 
482
483
484def pipeline_parallelism():
485
![](https://cs336.stanford.edu/lectures/images/pipeline-parallelism.png)
486
487
488 data = generate_sample_data()
489 spawn(pipeline_parallelism_main, world_size=2, data=data, num_layers=4, num_micro_batches=4)
490
491
492def pipeline_parallelism_main(rank: int, world_size: int, data: tensor, num_layers: int, num_micro_batches: int):
493 setup(rank, world_size) 
494
495 # Use all the data
496 data = data.to(cuda_if_available(rank))
497 batch_size = data.size(0) 
498 num_dim = data.size(1) 
499
500 # Split up layers
501 local_num_layers = int_divide(num_layers, world_size) 
502
503 # Each rank gets a subset of layers
504 local_params = [get_init_params(num_dim, num_dim, rank) for layer in range(local_num_layers)] 
505
506 # Forward pass
507
508 # Break up into micro batches to minimize the bubble
509 micro_batch_size = int_divide(batch_size, num_micro_batches) 
510 if rank == 0:
511 # The data
512 micro_batches = data.chunk(chunks=num_micro_batches, dim=0)
513 else:
514 # Allocate memory for activations
515 micro_batches = [torch.empty(micro_batch_size, num_dim, device=cuda_if_available(rank)) for _ in range(num_micro_batches)]
516
517 for x in micro_batches:
518 # Get activations from previous rank
519 if rank - 1 >= 0:
520 dist.recv(tensor=x, src=rank - 1)
521
522 # Compute layers assigned to this rank
523 for param in local_params:
524 x = x @ param
525 x = F.gelu(x)
526
527 # Send to the next rank
528 if rank + 1 < world_size:
529 print(f"[pipeline_parallelism] Rank {rank}: sending {summarize_tensor(x)} to rank {rank + 1}", flush=True) 
530 dist.send(tensor=x, dst=rank + 1)
531
532
533
534 # Backward pass: homework exercise
535
536 cleanup() 
537
538############################################################
539
540def setup(rank: int, world_size: int):
541 """Initializes the distributed environment (called at start of process)."""
542 # Specify where master lives (rank 0), used to coordinate (actual data goes through NCCL)
543 os.environ["MASTER_ADDR"] = "localhost"
544 os.environ["MASTER_PORT"] = "15623"
545
546 if torch.cuda.is_available():
547 dist.init_process_group("nccl", rank=rank, world_size=world_size)
548 else:
549 dist.init_process_group("gloo", rank=rank, world_size=world_size)
550
551
552def cleanup():
553 """Cleans up the distributed environment (called at end of process)."""
554 torch.distributed.destroy_process_group()
555
556
557class DisableDistributed:
558 """
559 Context manager that temporarily disables distributed functions (replaces with no-ops).
560 This is for when we're tracing the lecture, since we can't trace through
561 multiprocessing, so we just want to run the function directly without
562 distributed communication.
563 """
564 def __enter__(self):
565 self.old_functions = {}
566 for name in dir(dist):
567 value = getattr(dist, name, None)
568 if isfunction(value):
569 self.old_functions[name] = value
570 setattr(dist, name, lambda *args, **kwargs: None)
571
572 def __exit__(self, exc_type, exc_value, traceback):
573 for name in self.old_functions:
574 setattr(dist, name, self.old_functions[name])
575
576
577def spawn(func: Callable, world_size: int, *args, **kwargs):
578 """
579 Launches `world_size` processes that each calls `func` on world_size, args, kwargs.
580 Note: if we are being traced (inside edtrace), we just run the function directly without multiprocessing and disable distributed functions.
581 """
582 # Note: assume kwargs are in the same order as what main needs
583 if not sys.gettrace():
584 # This is the normal code path for multiprocessing
585 args = (world_size,) + args + tuple(kwargs.values())
586 mp.spawn(func, args=args, nprocs=world_size, join=True)
587 else:
588 # If we're being traced (inside edtrace), just run the function directly.
589 with DisableDistributed(): 
590 args = (0, world_size,) + args + tuple(kwargs.values())
591 func(*args)
592
593
594def get_init_params(num_inputs: int, num_outputs: int, rank: int) -> nn.Parameter:
595 """Create parameters and put them on the `rank`-th GPU."""
596 torch.random.manual_seed(0) # For reproducibility
597 return nn.Parameter(torch.randn(num_inputs, num_outputs, device=cuda_if_available(rank)) / math.sqrt(num_outputs))
598
599
600def int_divide(a: int, b: int):
601 """Return a / b and throw an error if there's a remainder."""
602 assert a % b == 0
603 return a // b
604
605
606def summarize_tensor(tensor: tensor) -> str:
607 return "x".join(map(str, tensor.shape)) + "[" + str(round(tensor.view(-1)[0].item(), 4)) + "...]"
608
609
610def render_duration(duration: float) -> str:
611 if duration < 1e-3:
612 return f"{duration * 1e6:.2f}us"
613 if duration < 1:
614 return f"{duration * 1e3:.2f}ms"
615 return f"{duration:.2f}s"
616
617
618if __name__ == "__main__":
619 main()
