CS336 // FIELD MAP
← field map
ASSIGNMENT 2 · SYSTEMS107 pts · v1.0.4 · Spring 2025

Systems and Parallelism

Stanford CS336 · Spring 2025 · assignment 2 of 5 · walkthrough of the official handout, everything linked

Lectures behind it: L05 GPUs, L06 Kernels, Triton, L07 Parallelism 1, L08 Parallelism 2

TL;DR — You take the Transformer from assignment 1 and make it fast: first you measure it (a timing harness, an Nsight Systems trace, a memory snapshot), then you rewrite its attention as a fused FlashAttention-2 Triton kernel, then you spread training across GPUs with hand-rolled distributed data parallel and a sharded AdamW. Twenty-one graded problems, 107 points, and about half of them are write-up rather than code — the deliverable is very often a table of timings plus two sentences of interpretation. The hard part is flash_forward (15 pts), where you write a real Triton kernel with block pointers and online softmax, and optimizer_state_sharding (15 pts), where the tricky bit is that torch.optim.Optimizer's own constructor calls your add_param_group before your subclass has finished initialising. The handout states no wall-clock estimate; the only budget it gives is hardware — up to 6 GPUs for the communication benchmark, with each run under five minutes.

Assignment 1 asked whether you could build a language model. Assignment 2 asks whether you can afford to run one. Every problem here is downstream of one fact: modern accelerators have enormous arithmetic throughput sitting behind a comparatively narrow memory pipe, so the thing that determines your training speed is almost never FLOPs, it is bytes moved and calls issued. So the assignment forbids the shortcuts. You may not call torch.nn.functional.scaled_dot_product_attention and declare attention solved — you write the kernel. You may not wrap your model in torch.nn.parallel.DistributedDataParallel — you write the container, the backward hooks, the buckets. You may not use ZeroRedundancyOptimizer — you write the shard assignment and the post-step broadcast. The point is that after this assignment, PyTorch's production versions of all three stop being magic and start being code you could have written, with performance characteristics you can predict from first principles.

Map of the assignment

Twenty-one graded problems, 107 points total, in handout order. Eight of them are graded by a pytest adapter; the other thirteen are graded by what you write in writeup.pdf. That ratio is the assignment's real shape — it is a measurement course, and the tests only cover the four things that have a single right answer.

§ProblemPtsDeliverableGraded byLecture
1.1.3benchmarking_script4code + write-upwrite-upL06
1.1.4nsys_profile5write-upwrite-upL06
1.1.5mixed_precision_accumulation1write-upwrite-upL05
1.1.5benchmarking_mixed_precision2code + write-upwrite-upL05
1.1.6memory_profiling4code + write-upwrite-upL06
1.2.1pytorch_attention2code + write-upwrite-upL05
1.3torch_compile2code + write-upwrite-upL06
1.3.2flash_forward15codetest_attention.pyL06
1.3.2flash_backward5codetest_attention.pyL06
1.3.2flash_benchmarking5write-upwrite-upL06
2.1.1distributed_communication_single_node5code + write-upwrite-upL07
2.2naive_ddp5codewrite-upL08
2.2naive_ddp_benchmarking3write-upwrite-upL08
2.3.1minimal_ddp_flat_benchmarking2code + write-upwrite-upL07
2.3.2ddp_overlap_individual_parameters5codetest_ddp_individual_parameters.pyL07
2.3.2ddp_overlap_individual_parameters_benchmarking1write-upwrite-upL07
2.3.3ddp_overlap_bucketed8codetest_ddp.pyL07
2.3.3ddp_bucketed_benchmarking3write-upwrite-upL07
2.4communication_accounting10write-upwrite-upL07
3optimizer_state_sharding15codetest_sharded_optimizer.pyL08
3optimizer_state_sharding_accounting5write-upwrite-upL08

Setup: environment, data, tests

No data downloads. This is the one assignment in the course with no corpus to fetch. Every measurement here runs on randomly initialised weights and randomly generated batches, because you are timing arithmetic and memory traffic, not learning anything. The only fixture files in the repo are tests/fixtures/ddp_test_data.pt and ddp_test_labels.pt, a tiny synthetic batch used by the DDP tests. If you were expecting another TinyStories download, relax.

Two packages, one venv. The repo ships the staff solution to assignment 1 as a sibling package. cs336-basics/ contains a working BasicsTransformerLM, AdamW, and a module-level scaled_dot_product_attention function; cs336_systems/ is a deliberately empty module where everything you write goes. The outer pyproject.toml wires them together with a path dependency, so uv run python -c "import cs336_basics" works out of the box, and swapping in your own assignment-1 implementation is a one-line edit to [tool.uv.sources]. Python is pinned to >=3.11,<3.13 and torch to ~=2.6.0.

The fact that scaled_dot_product_attention is a module-level function matters more than it looks. It means you can monkey-patch it — cs336_basics.model.scaled_dot_product_attention = my_annotated_version — and every attention call in the model picks up your NVTX-annotated or FlashAttention-backed replacement without touching the model code. That single line is how you do the profiling problems and how you eventually benchmark your Triton kernel end to end.

The adapter pattern. As in assignment 1, tests never import your code directly. They import tests/adapters.py, which ships as eight functions that all raise NotImplementedError; you replace each body with a one-line return pointing at your class. Two wrinkles specific to A2. First, the FlashAttention adapters return the class object itself, not an instance and not .apply — the tests call .apply themselves. Second, three of the eight adapters are explicitly optional: whether you need ddp_individual_parameters_on_after_backward, ddp_bucketed_on_after_backward and ddp_bucketed_on_train_batch_start depends on how you structured your DDP container. If your forward resets bucket state itself and your finish_gradient_synchronization is called by the training loop, some of these hooks have nothing to do.

Running the tests. uv run pytest tests/ runs everything; uv run pytest -k test_flash_forward_pass_pytorch runs one. The three distributed test files spawn a two-process gloo group over localhost:12390 via mp.spawn, so they run on a CPU-only laptop — this is deliberate, and it is the fastest development loop in the whole assignment. Only the Triton tests are GPU-gated (skipif not torch.cuda.is_available()), which means a clean pytest run on a CPU box silently skips the two hardest problems. Check for s markers, not just the absence of F. The handout also recommends running each distributed test about five times, because a DDP bug that only shows up under a particular interleaving will pass once and fail on the fourth try.

Model configurations. Every benchmark in part 1 sweeps these five sizes, all with vocabulary size 10,000 and batch size 4, at context lengths that vary per problem.

Sized_modeld_ffnum_layersnum_heads
small76830721212
medium102440962416
large128051203620
xl160064004825
2.7B2560102403232

Two things to notice before you start. d_ff is exactly 4 × d_model in every row, so activation memory scales the same way parameters do. And xl has 25 heads into 1600 dimensions — an odd number of heads, giving a head dimension of 64 like everything else. The xl config is the one every distributed benchmark in part 2 uses.

1.1 · Profiling and benchmarking

The mental model for this whole section: you are building a measuring instrument before you build anything worth measuring. Three instruments, in increasing resolution. A Python timer tells you how long. Nsight Systems tells you where, kernel by kernel, on both the CPU and the GPU timeline. The PyTorch memory profiler tells you what is resident and which line of your code allocated it. Every optimization later in the assignment is justified by a number one of these three produced.

The trap that makes all of this subtle is that CUDA calls are asynchronous. torch.matmul returns to Python as soon as the kernel is enqueued, not when it finishes. If you wrap a forward pass in timeit.default_timer() without synchronising, you are timing how fast Python can submit work to a queue — which is fast, constant, and completely uninformative. torch.cuda.synchronize() after each step is the fix, and it appears in nearly every problem in this section.

Problem (benchmarking_script): 4 points

Deliverable: a command-line benchmarking script, plus timings for all five model sizes and a short discussion of what happens without warm-up.

Part (a) is the script: build a model from hyperparameters, make a random batch, run w warm-up steps, then time n measured steps of either forward-only or forward-plus-backward, synchronising after each. Part (b) runs it at all five sizes with 5 warm-up and 10 measured steps and asks for mean and standard deviation. Part (c) is the interesting one: rerun with zero warm-up steps, then with one or two, and explain the difference. The first iteration pays for CUDA context creation, cuBLAS handle initialisation, kernel autotuning and the caching allocator's first cudaMalloc calls — costs measured in hundreds of milliseconds against a steady state of single-digit milliseconds. One or two warm-up steps knock out the worst of it but the allocator is still growing its pool and cuBLAS may still be selecting algorithms, so the variance stays elevated. The handout's advice to drive everything from command-line arguments is not politeness: you will rerun this script under nsys, under the memory profiler, with autocast on, with torch.compile on, and under two flavours of DDP. Build it once, with flags, or build it six times.

The common stumble is timing with time.time() instead of timeit.default_timer() — on some platforms the former's resolution is coarse enough to quantise your small-model measurements — and forgetting that a backward pass needs a fresh graph, so you must rerun the forward inside the timed region rather than calling .backward() twice on the same graph.

Problem (nsys_profile): 5 points

Deliverable: five short written answers derived from Nsight Systems traces of forward, backward and optimizer step, across all five model sizes at context lengths 128, 256, 512 and 1024.

You prepend nsys profile -o result to the script you just wrote and open the .nsys-rep in the Nsight Systems desktop app. The five questions walk you from gross to fine: does the profiler's total forward time agree with your Python timer (a); which CUDA kernel dominates cumulative GPU time and how many times it is invoked, and whether the answer changes when you add the backward pass (b); what non-matmul kernels are eating non-trivial time (c); how the matmul fraction shifts when you add loss and optimizer step (d); and how softmax runtime compares to matmul runtime inside the attention layer, relative to how their FLOP counts compare (e).

Question (e) is the whole point of part 1.2, arriving early. Softmax is a rounding error in FLOPs and a substantial slice of wall clock, because it is memory-bound: it reads and writes an entire seq_len × seq_len matrix to produce something the same size, with about one arithmetic operation per element. That gap between FLOP share and time share is the opening FlashAttention exists to close.

The mechanical part people underestimate is NVTX annotation. Without ranges you get an undifferentiated wall of kernels and cannot answer (b) or (e) at all. Annotate at three levels: a range around the measured region (so you can filter out warm-up in the timeline), ranges for forward / backward / optimizer, and — by monkey-patching an annotated scaled_dot_product_attention over the one in cs336_basics.model — sub-ranges for the QK product, the softmax and the output matmul. nsys profile --pytorch adds automatic annotation of PyTorch C++ API calls, and --python-backtrace=cuda attaches Python stacks to CUDA calls at some overhead. Expect to run out of memory at the larger size/context combinations; the handout explicitly says to note it in the report rather than work around it.

Problem (mixed_precision_accumulation): 1 point

Deliverable: two or three sentences on the accuracy of four accumulation loops.

Cheapest point in the assignment and the most useful intuition per point. You run four variants of "add 0.01 to a running sum a thousand times", differing in the dtype of the accumulator and the dtype of the addend, and explain the results. The FP32 accumulator lands near 100 with a small drift. The pure FP16 loop is badly wrong, and the reason is worth working out precisely rather than hand-waving about "less precision": FP16 has an 11-bit significand, so once the running sum reaches the binade [64, 128) its unit in the last place is 2−4 = 0.0625. Half of that is 0.03125, which is larger than 0.01 — so from that point on, round-to-nearest turns every single addition into a no-op and the sum stops moving entirely. The two mixed variants (FP32 accumulator taking FP16 addends, with and without an explicit .type(torch.float32) cast) both recover, because the accumulator's ULP stays small even though each addend was rounded on the way in. That is the entire justification for why mixed-precision training keeps reductions and accumulations in FP32, and it is why tl.dot(..., acc=acc) with an FP32 accumulator shows up later in your Triton kernel.

Problem (benchmarking_mixed_precision): 2 points

Deliverable: the dtype of six named tensors under autocast, a discussion of layer norm's sensitivity, and BF16-vs-FP32 timings for all five model sizes.

Part (a) is a prediction exercise on a four-line toy model: under torch.autocast with FP16, what dtype are the parameters, the output of the first linear layer, the output of layer norm, the logits, the loss, and the gradients? The answer people get wrong is the parameters — autocast does not change them, they stay FP32 in the module; what gets cast is the input to each op on autocast's allow-list. Part (b) asks why layer norm is treated differently and whether BF16 changes that. Layer norm computes a mean and a variance — reductions over the feature dimension, plus a division by a square root — and reductions in FP16 both accumulate error and can overflow FP16's narrow exponent range on the sum of squares. BF16 has FP32's exponent range and only 8 significand bits, so the overflow concern disappears while the precision concern gets slightly worse; the honest answer discusses both, not just "BF16 is fine".

Part (c) puts a BF16 autocast flag on your benchmarking script and asks for the trend across sizes. The useful framing: mixed precision buys you Tensor Core throughput on the matmuls and nothing on the memory-bound ops, so the speedup grows with the fraction of time you spend in large matmuls — which grows with model size. Small models can show almost no gain, or a regression, because the cast operations themselves cost bandwidth. contextlib.nullcontext is the clean way to make the autocast context optional without branching your timing loop.

NVIDIA A100 spec, as quoted in the handoutPeak throughput
FP3219.5 TFLOP/s
FP16 / BF16312 TFLOP/s

Problem (memory_profiling): 4 points

Deliverable: two memory-timeline screenshots, a table of peak memory by context length, a mixed-precision comparison, a hand derivation of one activation's size, and a note on where the largest allocations come from.

You wrap the measured region of your script in torch.cuda.memory._record_memory_history(max_entries=1000000) / _dump_snapshot("memory_snapshot.pickle") / _record_memory_history(enabled=None) and drag the pickle onto pytorch.org/memory_viz. The whole problem uses the 2.7B config at context lengths 128, 256 and 512.

Part (a) asks you to identify the training phases from the shape of the curve alone — and they are unmistakable once you have seen them: a forward pass is a monotone staircase up as activations are stashed for backward, the backward pass is a sawtooth down as each stashed activation is consumed and freed, and the optimizer step is a step change that never comes back down because AdamW just allocated two persistent moment buffers per parameter. Part (b) tabulates peaks by context length for inference versus a full step. Part (c) asks whether mixed precision helps memory — it helps less than students expect, because autocast does not halve your parameters or your optimizer state, only the activations of the ops it casts. Part (d) is a pen-and-paper derivation of the size of a single residual-stream activation for the 2.7B model: the shape is batch × context × d_model, so with the reference batch size of 4 and d_model 2560 in single precision it is 4 × 128 × 2560 × 4 bytes at context 128, and the handout wants it in MB with a 10242 divisor — which for that case comes out to exactly 5 MB. Part (e) uses the memory_viz "Detail" slider to hide small allocations and asks you to trace the survivors back through their stack traces, where you will find the seq_len × seq_len attention score matrices staring back at you. That is the hand-off into section 1.2.

1.2–1.3 · Attention, torch.compile, and FlashAttention-2 in Triton

This is the centre of the assignment: 29 of the 107 points, and the only place you write GPU code. The argument runs in four beats. First, measure vanilla PyTorch attention and watch it die of memory at long sequence lengths. Second, try torch.compile and observe that automatic fusion helps but does not fix the asymptotics. Third, learn Triton on a toy kernel. Fourth, implement FlashAttention-2 properly, forward pass in Triton with online softmax, backward pass via recomputation.

One structural oddity worth knowing so you don't think you have lost a page: the handout numbers torch_compile under §1.3, a sibling of §1.2 rather than a child, and then nests the Triton material at §1.3.1–1.3.4 beneath it. The reading order is still linear.

Problem (pytorch_attention): 2 points

Deliverable: a timing table across the sweep, an explicit memory accounting for one OOM configuration, and a paragraph or two of analysis.

You benchmark plain attention in isolation — batch size 8, no head dimension at all, over the cartesian product below — timing 100 forward passes, measuring memory in use immediately before backward, then timing 100 backward passes, with warm-up and synchronisation throughout.

Swept dimensionValues
batch size8 (fixed, single-head)
head embedding dim d_model16, 32, 64, 128
sequence length256, 1024, 4096, 8192, 16384
timed iterations100 forward, 100 backward

The written half is the graded half. You find where it OOMs, then account for the memory by hand using the assignment-1 formulas, and the number that should jump out is that the attention score matrix is batch × seq × seq, independent of d_model. At batch 8 and sequence 16384 that is 8 × 163842 elements — over two billion floats, more than 8 GB in FP32 for a single tensor, and the graph saves several such intermediates for backward. So the memory grows quadratically in sequence length and does not depend on the embedding dimension at all, which is exactly the wrong scaling. The last question — "what would you do to eliminate this memory cost?" — is asking you to invent recomputation before being taught it, and a good answer names both tiling and recomputing P in the backward pass rather than storing it.

Problem (torch_compile): 2 points

Deliverable: two comparison tables — compiled versus uncompiled attention on the sweep above, and compiled versus uncompiled full Transformer end to end.

A one-line intervention, torch.compile(module), with two lessons. On the attention microbenchmark, TorchInductor will fuse the softmax with its neighbours and generate Triton kernels for you, which cuts a real fraction of the memory traffic — but it will not restructure the algorithm, so the quadratic materialisation of the score matrix survives and so does the OOM cliff. On the full model, the picture is muddier: the forward pass usually improves modestly, forward-plus-backward-plus-optimizer often improves more because there is more elementwise work to fuse, and both measurements are easy to ruin by including compilation time. Compilation is lazy and happens on the first call with a given input shape, so your warm-up steps must run with the exact shapes you will measure, or you are timing the compiler. If your compiled numbers look catastrophic, that is almost always why — or it is recompilation triggered by changing shapes across your sweep.

§1.3.1 · The weighted-sum example (ungraded)

Before the graded Triton work, the handout walks through a complete forward-and-backward Triton implementation of (weight * x).sum(axis=-1). It is not graded and it is not optional in any practical sense — it introduces every mechanism you need: tl.program_id to identify which thread block you are, tl.make_block_ptr to describe an N-dimensional window into a tensor from a base pointer plus shape, strides, offsets, block shape and memory order, .advance() to slide that window along an axis, tl.load / tl.store with boundary_check and padding_option, and the pattern of wrapping the kernel in a torch.autograd.Function so it participates in autograd.

The one genuinely instructive idea in it is how the backward handles a reduction across thread blocks. Each program instance owns a tile of rows and can compute its own contribution to the weight gradient, but the true gradient sums over all rows — across blocks that have no ordering guarantees relative to each other. Rather than use atomics, the kernel writes a partial buffer of shape n_row_tiles × D and reduces it with a plain torch.sum outside the kernel. Remember that trick: FlashAttention-2's backward pass faces the identical problem with dQ, and the handout's optional §1.3.4 solves it the same way, with two passes over the input instead of one.

Problem (flash_forward): 15 points

Deliverable: two torch.autograd.Function subclasses implementing the FlashAttention-2 forward pass — one pure PyTorch, one Triton — plus a causal-masking flag on the Triton one.

The single largest problem in the assignment, and it is structured as a ladder so that you are never debugging two things at once. Part (a) asks for the tiled algorithm in pure PyTorch. It will be slower than plain attention — that is expected and fine — because its purpose is to be a reference implementation you can compare against tile by tile when the Triton version produces garbage. Part (b) is the same algorithm as a real kernel. Part (c) adds causal masking.

The algorithm is FlashAttention-2's forward pass, and the thing that makes it work is online softmax. You cannot normalise a tile of scores without seeing the whole row, so instead you carry two running statistics per query row — a running maximum m for numerical stability and a running denominator l — and each time you process a new key tile you rescale the accumulated output by exp(m_prev − m_new) before adding the new tile's contribution. After the last key tile you divide the accumulator by the final l once, and you write out L = m + log(l), the log-sum-exp of the row. That L is not a debugging aid; it is the checkpoint that makes the backward pass cheap, because with it you can recompute the probabilities from Q, K and L without ever having stored them.

The interface is then def forward(ctx, Q, K, V, is_causal=False). Determine your own tile sizes, but make sure they are at least of size 16 × 16. We will always test your code with dimensions that are clean powers of 2 and at least 16, so you don't need to worry about out-of-bounds accesses.
— handout, problem flash_forward (a)

That last sentence is a real gift. Boundary handling is the most tedious part of writing tile-based kernels, and the tests are guaranteed to hand you clean powers of two. You can still pass boundary_check to your loads if you want the safety, but you are not being graded on ragged shapes.

Your launch grid should be set as (Tq, batch_size), meaning each Triton program instance will load only elements from a single batch index, and only read/write to a single query tile of Q, O, and L. The kernel should only have a single loop, which will iterate key tiles 1 ≤ j ≤ Tk.
— handout, problem flash_forward (b)

That grid choice is the parallelisation decision that distinguishes FlashAttention-2 from FlashAttention-1: parallelise over query tiles, loop over key tiles. Every program instance owns its output tile outright, so nothing needs to be atomic and no block ever waits on another. It also explains why the batch dimension is flattened before launch — the kernel signature takes a single batch stride, so a four-dimensional (batch, heads, seq, dim) tensor should be reshaped to three dimensions with batch-times-heads in front.

Two precision rules from the handout that are easy to skip and expensive to debug. The on-chip buffers for O, l and m must be tl.float32 even when your inputs are BF16 — use acc = tl.dot(..., acc=acc) so the accumulation happens in the register file at full precision. And you must explicitly cast the probability tile down to V's dtype before multiplying, and cast O before storing, using tensor.to(...) with block_ptr.type.element_ty as the target. Mixing dtypes into tl.dot is a compile error at best and a silent precision loss at worst.

Causal masking in part (c) is a comparison of two index vectors — query positions against key positions — forming a Bq × Bk boolean tile.

For elements that are masked out, add the constant value of −1e6 to the corresponding elements of the attention score matrix.
— handout, problem flash_forward (c)

Add, not assign, and −1e6 rather than −inf — because −inf propagates NaN through the exp(S − m) when an entire tile is masked, and because the reference implementation in the test file does exactly the same thing with the same constant. The flag must be annotated is_causal: tl.constexpr (Triton needs it at compile time to specialise the kernel), must default to False so the non-causal tests still pass, and must be stashed as ctx.is_causal for the backward pass.

The test's hidden requirement. _test_flash_forward_pass reaches into o.grad_fn.saved_tensors and looks for tensors whose shape is exactly (batch, n_queries). It asserts that there is exactly one. So you must save L via ctx.save_for_backward, and you must not save any other tensor of that shape — a stray m or row-sum buffer with the same shape will fail the test with a confusing message about tensor counts. The tests use batch 4, 128 queries, 128 keys and D = 64, and compare against a reference attention plus torch.logsumexp at rtol/atol of 1e-2.

Handout/code drift, worth knowing before you go hunting. The v1.0.4 handout names the adapter for part (b) adapters.get_flash_autograd_function_triton. No such function exists — the file at the spring2025 ref defines get_flashattention_autograd_function_triton. The changelog records this under 1.0.5 as "typos regarding FlashAttention2 in the adapters". Trust adapters.py, not the PDF.

Problem (flash_backward): 5 points

Deliverable: a working backward pass for your FlashAttention-2 autograd function, in PyTorch with torch.compile — Triton is explicitly not required.

This is the payoff for having stored L. The recomputation identity is that you can rebuild the probability matrix from the saved values with a single elementwise exponential, P = exp(S − L), where S is recomputed from Q and K. No softmax, no maximum-tracking, no online anything — the backward pass becomes five straightforward matrix expressions. The one piece of extra bookkeeping is the vector D = rowsum(dO ∘ O), precomputed once, which appears in the dS expression: dS = P ∘ (dP − D). The handout derives why rowsum(O ∘ dO) equals rowsum(P ∘ dP), and that identity is the reason you never need P's row sums from the forward pass.

Because there is no online trick left, the handout explicitly permits — and recommends — writing this as an ordinary PyTorch function and wrapping it in torch.compile to get fusion for free. Doing it in Triton is deferred to the optional §1.3.4 and is mainly a leaderboard play. Take the easy path first; a correct torch.compile backward is worth the same five points as a hand-written kernel.

Two practical notes. The scale factor 1/√d appears in dQ and dK and it is easy to apply it twice (once when recomputing S, once in the gradient) or not at all — check against the test rather than against your algebra. And the pure-PyTorch autograd function needs a real backward too: part (a) of flash_forward lets you stub it with NotImplementedError, but test_flash_backward_pytorch runs on CPU and will call it.

Problem (flash_benchmarking): 5 points

Deliverable: a table of forward, backward and end-to-end latencies for your FlashAttention-2 against vanilla PyTorch attention, over the sweep below, measured with triton.testing.do_bench on a single H100.

SettingValues
batch size1
maskingcausal, always
sequence lengthpowers of 2 from 128 to 65536
embedding dimensionpowers of 2 from 16 to 128
precisiontorch.bfloat16 and torch.float32
hardwarea single H100
measuredforward, backward, forward+backward

triton.testing.do_bench handles warm-up, repetition and cache clearing between runs, which is why the handout switches to it here rather than reusing your timeit harness. The sweep is large — seven sequence lengths × four embedding dimensions × two precisions × two implementations — and the baseline will OOM at the top end, which is itself a result worth tabulating rather than an error to suppress. The handout warns that you will likely need to vary tile sizes across the sweep; a tile configuration that is optimal at sequence 512 and d = 16 can exhaust shared memory at d = 128, so make Bq and Bk parameters and pick them as a function of the input shape.

§1.3.3 · The FlashAttention-2 leaderboard (extra credit)

Optional, and the only competitive element in the assignment. The Spring 2025 leaderboard times a fused forward-plus-backward on a fixed configuration.

Leaderboard configuration (Spring 2025)Value
hardwaresingle H100
batch size1
sequence length16384 (Q, K and V)
d_model / heads / d_head1024 / 16 / 64
precision, maskingBF16, causal
timing harnesstriton.testing.do_bench(fn, rep=10000, warmup=1000)
naive baseline, verified80 ms
best verified 2025 submission5.364 ms
The restrictions are that you cannot change the input/outputs of the function, and you must use Triton (no CUDA, unfortunately). Your inputs will be tested at BF16 with causal masking, and it must pass the same tests as your regular implementation. The implementation must also be your own, and you cannot use pre-existing implementations.
— handout, §1.3.3

The handout's own list of ideas is a decent optimisation curriculum: autotune the tile sizes, tune the other Triton config knobs, move the backward pass into Triton, split the backward into two passes (one for dQ, one for dK and dV) to avoid atomics, terminate program instances early on causal tiles that are entirely masked, separate fully-unmasked tiles from the diagonal tiles so only the diagonal pays for index comparisons, and use the H100's Tensor Memory Accelerator following the persistent matmul tutorial. The early-exit trick alone is worth close to a factor of two under causal masking, since half the score matrix is discarded.

The leaderboard repo has moved on. Its main branch now hosts the Spring 2026 task, which is a completely different benchmark — a full training step of an 8B model on two B200s, with a 10-second naive baseline. The Spring 2025 FlashAttention-2 board, with the numbers in the table above, is preserved at commit a98a9c78. Link that one, not the repo root, if you want the rules the v1.0.4 handout describes.

2 · Distributed data parallel training

Part 2 is a single idea developed in four stages, and the pleasure of it is that each stage is a small, obvious improvement over the previous one with a measurable payoff you produce yourself. Stage zero: measure the collectives in isolation so you know what communication costs. Stage one: naive DDP — one all-reduce per parameter tensor after the backward pass finishes. Stage two: flatten all the gradients into one buffer and issue a single all-reduce, trading many small calls for one big one. Stage three: fire the all-reduces from backward hooks as each gradient becomes ready, overlapping communication with the rest of the backward pass. Stage four: bucket the parameters so you get both — few calls and overlap. That final design is, essentially, what PyTorch's real DistributedDataParallel does.

All of it runs on gloo/CPU for correctness and NCCL/GPU for the benchmarks. The handout's benchmarking discipline for this section is worth internalising: same machine for every comparison, five warm-up iterations before timing (NCCL especially needs them), torch.cuda.synchronize() even around calls made with async_op=False — because that flag only means "queued", not "finished" — and aggregation of timings across ranks via dist.all_gather_object because ranks drift.

Problem (distributed_communication_single_node): 5 points

Deliverable: plots and/or tables over the grid below, plus two or three sentences on how the factors interact.

Swept factorValues
backend + deviceGloo + CPU, NCCL + GPU
all-reduce payload (float32)1 MB, 10 MB, 100 MB, 1 GB
number of processes2, 4, 6
resource budgetup to 6 GPUs; each run under 5 minutes

You are measuring the collective by itself, with no model attached, so the shape of the curve tells you which regime you are in. At 1 MB you are measuring fixed per-call overhead — latency, kernel launch, the ring's synchronisation — and the time barely moves with payload. At 1 GB you are measuring bandwidth and the time is close to linear in bytes. Somewhere between those two the crossover happens, and that crossover point is precisely the number that determines the right bucket size in §2.3.3. Process count matters differently for the two backends: a ring all-reduce moves 2(n−1)/n × payload bytes per rank, so the per-rank volume is nearly independent of n, but the number of hops and therefore the latency component grows.

The measurement error everyone makes here is timing rank 0 only and reporting it as "the" time. Ranks finish at different moments, the barrier semantics mean rank 0 can appear artificially fast or slow, and the honest number is an aggregate across ranks — hence the handout's pointer to dist.all_gather_object.

Problem (naive_ddp): 5 points

Deliverable: a script that does DDP by all-reducing each parameter gradient after the backward pass, plus a correctness check against single-process training.

The whole algorithm in four steps: broadcast rank 0's parameters to everyone so all ranks start identical; give each rank a disjoint 1/d slice of the batch; run forward and backward locally; all-reduce and average every parameter's .grad; step the local optimizer. Because every rank started from the same weights and applies the same averaged gradients, they stay bit-identical forever, and no parameter communication is needed after the initial broadcast.

The verification the handout asks for is the part that teaches: train a toy model both ways and assert the weights match. Three things routinely break it. Averaging — dist.all_reduce sums, so you must divide by world size, and dividing in the wrong place gives you a learning rate that is silently d times too large. Data sharding — the ranks must see disjoint shards of the same batch, which usually means seeding identically and slicing by rank, not seeding differently. And initial synchronisation — if you construct the model independently on each rank without broadcasting, the ranks diverge from step zero and no amount of gradient averaging brings them back. The handout points you at test_ddp_individual_parameters.py as a worked example of how to structure this comparison; its validate_ddp_net_equivalence helper in tests/common.py is the pattern — all-gather every entry of the state dict and assert they are all close.

Problem (naive_ddp_benchmarking): 3 points

Deliverable: a description of your setup, the measured time per training iteration, and the fraction of that time spent communicating gradients — for the xl model on 1 node × 2 GPUs.

This establishes the baseline that the next three problems try to beat. The xl config has 48 layers at d_model 1600, so you are issuing hundreds of separate all-reduce calls per step, most of them for small tensors — biases, norm weights — where per-call overhead dominates the payload entirely. Isolating the communication time cleanly requires synchronising before you start the timer and after you stop it, otherwise you attribute queueing time to the wrong bucket. The 1 node × 2 GPU, xl-model configuration is the fixed setting for every remaining benchmark in part 2, so build the harness to be reused.

Problem (minimal_ddp_flat_benchmarking): 2 points

Deliverable: time per iteration and communication time with a single flattened all-reduce, and one or two sentences comparing against per-parameter communication.

Concatenate every gradient into one contiguous buffer, all-reduce once, scatter the result back. torch._utils._flatten_dense_tensors and _unflatten_dense_tensors do the packing and unpacking; they are private API and the handout says to use them anyway. The expected result is a solid improvement, because you have replaced hundreds of latency-bound calls with one bandwidth-bound call, and from the previous problem you already know roughly where that crossover sits. What you have not fixed is the serialisation: the flattened all-reduce cannot begin until the last gradient exists, so the entire communication cost is still tacked onto the end of the step. That is the observation §2.3.2 attacks.

Problem (ddp_overlap_individual_parameters): 5 points

Deliverable: a DDP container class that overlaps gradient communication with backward computation, with three methods — __init__(self, module), forward(self, *inputs, **kwargs), and finish_gradient_synchronization(self).

The insight is that backward propagates from the loss toward the input, so the last layer's gradients are ready long before the first layer's. Instead of waiting, you register a hook on every parameter with Tensor.register_post_accumulate_grad_hook that fires an async_op=True all-reduce the moment that parameter's gradient lands, stash the returned handle, and drain all the handles in finish_gradient_synchronization() just before the optimizer step. Communication for the deep layers then happens while the shallow layers are still computing, and only the tail is exposed.

The container must also broadcast parameters from rank 0 at construction time, so that the class is a complete replacement for the setup half of naive_ddp and not just the communication half. Note the deliberate choice of register_post_accumulate_grad_hook over the older register_hook: the post-accumulate variant fires after .grad has been written, which is what you want, and it fires once per accumulation rather than once per gradient contribution.

The test that catches sloppy implementations is the tied-weights one. ToyModelWithTiedWeights shares one weight tensor between two layers, so that parameter's gradient is accumulated twice during a single backward. If your hook fires on the first accumulation and you all-reduce a half-finished gradient, you get wrong answers that a non-tied model would never surface. ToyModel has a matching trap in the other direction: it contains a bias with requires_grad=False and a buffer-like parameter named no_grad_fixed_param, and hooking or communicating parameters that never receive gradients will either crash or hang. Filter on requires_grad.

Problem (ddp_overlap_individual_parameters_benchmarking): 1 point

Deliverable: time per iteration for the overlapped implementation compared with the two earlier ones, plus two Nsight screenshots showing overlap happening and not happening.

Part (a) is one more row in the table you have been building. Part (b) is the more interesting half: profile both DDP implementations under Nsight and show, visually, that the NCCL kernels interleave with the backward kernels in one trace and sit in a solid block at the end in the other. This is the assignment asking you to make an argument with a picture rather than a number, and it is the clearest single image in the whole course — the naive trace has an obvious communication tail, the overlapped trace has communication smeared underneath the compute. Changelog note: v1.0.4 switched this problem from the PyTorch profiler to Nsight, so older solutions you might find online use a different tool.

Problem (ddp_overlap_bucketed): 8 points

Deliverable: a DDP container that buckets gradients, with the same interface as the previous one plus a bucket_size_mb: float constructor argument.

Both optimisations at once. Group parameters into buckets of at most bucket_size_mb megabytes; count gradients as they arrive; when a bucket's last gradient lands, flatten that bucket and issue one asynchronous all-reduce for it. You get few calls, and you get overlap, and the tuning knob between the two extremes is the bucket size — a single unbounded bucket degenerates to §2.3.1, and one parameter per bucket degenerates to §2.3.2.

We suggest allocating parameters to buckets using the reverse order of model.parameters(), since the gradients will become ready in approximately that order during the backward pass.
— handout, problem ddp_overlap_bucketed

That hint is doing real work. Bucketing in forward order would put the first layer's parameters — whose gradients arrive last — into the same bucket as parameters whose gradients arrive early, and that bucket could not be flushed until the very end of backward, destroying the overlap you built the bucketing to preserve. Reverse order approximately matches gradient-readiness order, so buckets fill and fire in sequence.

The state machine is where implementations go wrong. Each bucket needs a counter of how many of its parameters still owe a gradient, and that counter has to be reset at the start of every training step — which is what the optional ddp_bucketed_on_train_batch_start adapter exists for, if you did not reset it inside forward. Tied weights make the counting subtle again: a shared parameter appears once in model.parameters() but accumulates twice, so a naive counter can hit zero early. And the flattened buffer must be unflattened back into the individual .grad tensors before the optimizer runs, or you will step on stale gradients.

The test parametrizes bucket sizes deliberately, to force all three regimes on a three-parameter toy model:

bucket_size_mbWhat it exercises
0.011 bucket holding 3 parameter tensors
0.00162 buckets, split 2 and 2
0.00013 buckets, one parameter tensor each

Those numbers are tiny on purpose. If your bucket-assignment logic special-cases "parameter larger than the bucket" incorrectly — dropping it, or refusing to place it — the 0.0001 case will fail.

Problem (ddp_bucketed_benchmarking): 3 points

Deliverable: time per iteration at bucket sizes 1, 10, 100 and 1000 MB, three or four sentences on whether the results match your expectations, and an analytical model of DDP overhead with a derived optimal bucket size.

Part (a) is set up to disappoint you, and the handout says so with unusual candour — it asks what you would change about the experimental setup to make the results match theory. On two GPUs in one node with NVLink between them, bandwidth is enormous and communication is barely a bottleneck, so the bucket-size curve is often flat or non-monotonic and swamped by allocator noise. Ordering effects matter too: NCCL executes calls in issue order on a stream, so a bucket that becomes ready early but was issued late still waits. A good answer names the conditions that would sharpen the effect — more ranks, cross-node links instead of NVLink, a slower interconnect, a model with more small tensors.

Part (b) is the analysis and it is straightforward once you accept the stated assumption that gradient computation for a bucket takes exactly as long as communicating it. With total parameter bytes s, algorithmic bandwidth w, per-call overhead o and nb buckets: every bucket except the last hides behind subsequent computation, so the exposed overhead is the fixed cost of all the calls plus the transfer of the final bucket — roughly nb·o + s/(nb·w). More buckets means more overhead calls but a smaller exposed tail. Differentiate with respect to nb, set to zero, and you get the classic square-root balance; converting nb back into a bucket size gives the answer the problem wants. It is the same shape of result as the latency/bandwidth crossover you measured empirically in distributed_communication_single_node, which is a satisfying place for the section to land.

2.4 · 4D parallelism and communication accounting

A theory section with one large problem attached. The handout lays out five axes of parallelism — data, fully-sharded data, tensor, pipeline and expert — then argues them down to four by pointing out that FSDP and TP are almost always combined along the same mesh dimension, and drops expert parallelism because the course's models are dense. The picture to hold is a device mesh: 16 GPUs arranged 4 × 4, one axis data parallel, the other axis FSDP-plus-TP, with each axis carrying its own collectives and its own communication cost.

Problem (communication_accounting): 10 points

Deliverable: four written parts with the arithmetic shown — single-device memory, sharded memory and the required shard count, the compute-bound batch size, and a paragraph on reducing batch size without becoming communication bound.

The largest write-up problem in the assignment, and worth more than anything except the two 15-point implementations. It uses a hypothetical XXL config with simplifying assumptions: no attention, no embeddings, no output projection, each block just two linear layers, no activation checkpointing, activations and gradient communication in BF16, master weights and optimizer state in FP32.

XXL configurationValue
d_model16384
d_ff53248
num_blocks126
per-block parameters2 × d_model × d_ff
FP32 state per parameter16 bytes (master weight + gradient + 2 AdamW moments)
H100 memory, part (a)80 GB per device
TPU v5p memory, part (b)95 GB per device
TPU v5p bandwidth Wici, part (c)2 · 9 · 1010
TPU v5p compute C, part (c)4.6 · 1014 FLOP/s
mesh, part (c)MX = 2, MY = 1; X = 16 (FSDP), Y = 4 (TP)

Part (a) is bookkeeping that produces a shocking number: those dimensions give roughly 220 billion parameters, and at 16 bytes each the FP32 training state alone is a few thousand gigabytes — several dozen H100s' worth of memory before a single activation is stored. Part (b) shards master weights, optimizer state, gradients and half the activations across NFSDP devices and asks how large NFSDP must be to fit in 95 GB. Part (c) is the one that requires actually reading Part 5 of the TPU Scaling Book, because it uses that book's notation for mesh dimensions and its arithmetic-intensity argument to find the per-device batch size at which compute time exceeds communication time. Part (d) is open-ended: what else lets you shrink the global batch size without going communication bound? Sequence parallelism, activation checkpointing traded against recompute, communication-computation overlap, gradient accumulation, better collectives, lower-precision communication — the handout asks you to back the claims with references or equations, so cite.

One broken link to know about. The handout's citation for the Ultra-Scale Playbook's pipeline-parallel appendix points at a static.hf.space deep link that now 404s. The live document is the Hugging Face space, and the same ?section= anchor works there.

3 · Optimizer state sharding

The last section, and the one that most directly previews production practice. DDP has every rank hold a full copy of the parameters, the gradients and the optimizer state — and for AdamW the optimizer state is two floats per parameter, so it is twice the size of the model. Across d ranks that is d redundant copies of the largest single item in the memory budget. The fix in this assignment is the simplest member of the ZeRO family: partition the optimizer state so each rank only owns and updates roughly 1/world_size of the parameters, then broadcast the updated slices so everyone's weights stay in sync.

Problem (optimizer_state_sharding): 15 points

Deliverable: a torch.optim.Optimizer subclass wrapping an arbitrary optimizer class, with __init__(self, params, optimizer_cls, **kwargs), step(self, closure, **kwargs) and add_param_group(self, param_group).

Tied for the largest problem in the assignment, and the difficulty is almost entirely structural rather than algorithmic. The algorithm is short: assign each parameter to a rank, construct the inner optimizer over only this rank's parameters, and after every step() broadcast each parameter from the rank that owns it so all ranks agree again.

The structural trap is initialisation order, and it catches nearly everyone. torch.optim.Optimizer.__init__ calls self.add_param_group(...) for each group it is given — so your override runs during the superclass constructor, before your subclass has had a chance to set up whatever state it wants to use. Any attribute your add_param_group touches must therefore exist before you call super().__init__(), or be created lazily with a getattr guard. The handout is explicit that you must call the superclass constructor, so you cannot dodge it.

add_param_group also has to work when called later during training — the handout gives gradual layer unfreezing as the motivating case — so shard assignment cannot be a one-shot pass over all parameters in the constructor. It has to be an incremental policy: as each new group arrives, hand its parameters out to ranks and register the ones this rank owns with the inner optimizer.

Design decisions the tests will probe. Assignment must be deterministic and identical on every rank, since rank i needs to know which parameters rank j owns in order to receive its broadcast. Every rank still needs full parameters and full gradients — only the optimizer state is sharded, which is exactly what makes this ZeRO stage 1 rather than 2 or 3. Tied weights appear once in parameters() but are reachable from two modules, so they must be owned by exactly one rank and broadcast once. And the closure argument of step has to be forwarded to the inner optimizer, not swallowed.

Problem (optimizer_state_sharding_accounting): 5 points

Deliverable: peak memory at three points in the step with and without sharding plus a component breakdown, a runtime comparison, and a written contrast with ZeRO stage 1.

Part (a) profiles peak memory at three moments — after model initialisation, immediately before the optimizer step, and immediately after — on the standard 1 node × 2 GPU xl setup, and asks for a breakdown into parameters, gradients, optimizer state and activations. The expected shape: initialisation is dominated by parameters, the pre-step peak adds gradients and saved activations, and the post-step jump is AdamW's moment buffers materialising for the first time. With sharding, that last jump is roughly halved on two ranks while everything else is unchanged, which makes the accounting legible.

Part (b) measures the speed cost, and there is one: you have added a broadcast per parameter after every step, which pure DDP did not need. On two ranks with a fast interconnect it is small; the point is that the memory saving is not free.

Part (c) is the conceptual close. The honest comparison with ZeRO stage 1 as described in Rajbhandari et al. is that this implementation gets the same memory saving but pays more communication. Real ZeRO-1 replaces the all-reduce of gradients with a reduce-scatter — each rank ends up with only the reduced gradients for its own shard — and then all-gathers the updated parameters, and the reduce-scatter plus all-gather pair moves the same total bytes as a single all-reduce did. This assignment's version keeps the full gradient all-reduce and then adds a broadcast per parameter on top, so it costs strictly more bandwidth than ZeRO-1 for the same memory win. It also still stores full-size gradients on every rank, where ZeRO stage 2 would shard those too. Saying this clearly is the answer; lecture 8 derives the reduce-scatter/all-gather identity that makes it precise.

What you hand in

Two artefacts to Gradescope. writeup.pdf with typeset answers to every written question — which, given that thirteen of the twenty-one problems are write-up only and most of the rest have a written part, is the bulk of the grade. And code.zip containing everything you wrote, produced by running test_and_make_submission.sh, which runs the full test suite with --junitxml=test_results.xml and then zips the directory minus caches, checkpoints and data files. Note two harmless oddities in that script: it names its output cs336-spring2024-assignment-2-submission.zip, a leftover from the previous year, and it runs pytest with || true so a failing test does not stop the packaging. Failing tests are recorded in the XML, not hidden — but do not read a successful script run as evidence that your tests passed.

Because so much of the grade is tables, take the handout's advice from §1.1.2 seriously and generate them from code. pandas.DataFrame.to_latex() and .to_markdown() exist precisely so that you are not hand-transcribing sixty timing numbers at 2am, and you will re-run several of these sweeps after fixing a bug.

The leaderboard is a separate, optional submission: a pull request against the leaderboard repo adding your name and time to the table, with a description of what you did.

A note on hardware honesty. Several problems specify hardware you may not have — a single H100 for the FlashAttention benchmarks, up to 6 GPUs for the collectives sweep, two GPUs for everything in part 2. The handout's own stance, stated for the OOM cases in nsys_profile, is to report what happened rather than fabricate a clean grid. Missing cells with an explanation read better than invented numbers.

Materials, verified

Next: A3 Scaling · Back to the assignments.