GPUs & Kernels
Almost every performance question in modern ML ("why is decode slow?", "why did FlashAttention help?", "why does FP8 not give 2× end to end?") comes down to one picture: a chip that can do arithmetic hundreds of times faster than it can move bytes from memory. This page builds the mental model an AI engineer needs: how a GPU executes work, where the bytes live, how to tell whether an op is compute-, memory- or overhead-bound, and how the important kernels (GEMM, FlashAttention, decode attention) are engineered around those limits. Interviewers ask about this because it separates people who can use models from people who can make them cheap and fast.
TL;DR: the things to be able to say out loud
- A GPU is ~100+ SMs, each running many warps (32 threads in lockstep, SIMT). Latency is hidden by switching between warps, not by caches or branch prediction.
- Programming model: grid → blocks (CTAs) → warps → threads. A block lives on one SM and shares its shared memory; Hopper adds clusters of blocks.
- Tensor cores do almost all the FLOPs. On H100, dense BF16 is ~989 TFLOPS vs ~67 TFLOPS for ordinary FP32 lanes, a ~15× gap. Anything that isn't a matmul is "slow math".
- Memory hierarchy: registers → shared memory/L1 (~228 KB/SM) → L2 (50 MB) → HBM (80 GB at ~3.35 TB/s). Each level down is bigger, slower and more contended.
- Roofline: attainable FLOP/s = min(peak compute, bandwidth × arithmetic intensity). H100's BF16 ridge point is ~295 FLOP/byte. Elementwise ops sit at <1, decode GEMV at ~batch size, big GEMMs at 1000+.
- Three regimes: compute-bound, memory-bound, overhead-bound. The fix differs: better tensor-core use; fusion and fewer bytes; CUDA graphs, bigger batches, less Python.
- Fusion wins because it removes HBM round trips of intermediates. torch.compile/Inductor generates fused Triton kernels automatically.
- GEMM = hierarchical tiling (HBM → SMEM block tiles → register tiles → tensor-core fragments) plus a pipeline that keeps loads (cp.async/TMA) in flight while tensor cores compute. Hopper adds TMA, WGMMA and warp specialization.
- FlashAttention = tile Q/K/V into SRAM, use the online-softmax recurrence (running max \(m\), running sum \(\ell\)) so the \(N\times N\) score matrix never touches HBM, and recompute it in the backward pass. v2 is better parallelism, v3 is Hopper async + FP8, v4 targets Blackwell.
- Decode is memory-bound: each token streams all weights and the KV cache. FlashDecoding splits the KV sequence across SMs (split-K) and merges partials with log-sum-exp.
- Kernel languages: CUDA C++/CUTLASS (max control), Triton (block-level, compiler handles intra-block details), and newer tile DSLs (CuTe DSL, ThunderKittens, TileLang, cuTile, Mojo).
- Numerics: BF16 for range, FP32 accumulation, FP8 (E4M3 forward/E5M2 gradients) with scaling, block-scaled FP4 (MXFP4/NVFP4) on Blackwell.
1. GPU architecture for ML engineers
A CPU spends its transistors making one thread fast: big caches, branch prediction, out-of-order execution. A GPU spends them on throughput: thousands of simple lanes, and enough concurrent threads that when some are waiting on memory, others have work ready. That one design decision explains most GPU programming advice.
Streaming multiprocessors (SMs)
The GPU is a grid of identical SMs. An H100 SXM has 132 of them NVIDIA 2022. Each SM is split into 4 sub-partitions, and each sub-partition has a warp scheduler, a register-file slice, FP32/INT32 lanes, special-function units (exp, sin, rsqrt) and one tensor core. Per SM you also get a 256 KB block of on-chip SRAM, split between L1 cache and programmer-managed shared memory (up to 228 KB of shared memory), plus a 256 KB register file NVIDIA 2022. One way to picture it: an SM is a small vector processor with a scratchpad, and the GPU is 132 of those sharing an L2 and HBM.
Threads, warps, blocks, grids (and clusters)
You launch a kernel over a grid of thread blocks (also called CTAs, cooperative thread arrays). The hardware assigns each block to one SM, where it stays until it finishes. Threads within a block can synchronize (__syncthreads()) and share data through shared memory. Threads in different blocks cannot, except through global memory and atomics, or on Hopper and later, within a thread block cluster: a group of blocks scheduled together on neighbouring SMs that can read each other's shared memory ("distributed shared memory") NVIDIA 2022.
The warp is the real unit of execution: 32 threads that share an instruction stream. NVIDIA calls this SIMT (single instruction, multiple threads). It's SIMD with per-lane addressing and the illusion of independent control flow. Each cycle, each scheduler picks a warp whose operands are ready and issues its next instruction.
Latency hiding and occupancy
A global-memory load takes hundreds of cycles (roughly 400–800). The GPU doesn't stall. It switches to another resident warp at zero cost, because every resident warp's registers stay live in the register file. Occupancy is the ratio of resident warps to the maximum the SM supports (64 warps = 2048 threads per SM on H100 NVIDIA docs). Three resources limit it:
- Registers. 65,536 32-bit registers per SM. A kernel using 128 registers/thread fits \(65536/128 = 512\) threads = 16 warps → 25% occupancy. Using 255 registers (the per-thread max) gives ~8 warps.
- Shared memory. If each block needs 100 KB, only 2 blocks fit in 228 KB.
- Block/thread limits. Max blocks per SM and max threads per block.
The key point: high occupancy is a means, not a goal. Good GEMM and attention kernels often run at low occupancy (1–2 big blocks per SM, ~8–16 warps). They get latency tolerance from instruction-level parallelism and explicit asynchronous copy pipelines instead of from many warps. Low-occupancy, high-ILP kernels are the norm at the top end. Memory-bound elementwise kernels, on the other hand, usually do need enough warps (or enough bytes in flight per warp) to saturate bandwidth. By Little's law, bytes in flight = bandwidth × latency: \(3.35\,\text{TB/s} \times {\sim}600\,\text{ns} \approx 2\,\text{MB}\) in flight across the chip, or roughly 15 KB per SM.
Warp divergence
If threads in one warp take different branches of an if, the warp runs both paths one after the other, with lanes masked off. A 50/50 split halves throughput for that region. Since Volta, threads have independent program counters (so intra-warp locking and the like is legal), but divergent paths still serialize. Divergence between warps costs nothing. In ML kernels divergence mostly appears at boundaries (masking the last partial tile, causal masks). The standard trick is to handle full tiles on a branch-free fast path and only mask the edge tiles, which is exactly what FlashAttention does for causal attention.
Tensor cores: where the FLOPs are
Tensor cores are matrix-multiply-accumulate units: each instruction computes \(D = A B + C\) on a small tile (for example 16×8×16 per warp on Ampere-style mma.sync). On Hopper the instruction is WGMMA, issued asynchronously by a whole warpgroup (128 threads) with operands that can come straight from shared memory, on tiles up to 64×256×16. On Blackwell, the tcgen05 MMA instructions write accumulators into a new per-SM tensor memory (TMEM) instead of registers, and a "2-CTA" mode lets a pair of SMs cooperate on one larger MMA Zadouri+ 2026.
| Spec (dense, no sparsity) | A100 SXM 80GB | H100 SXM | B200 (per GPU, DGX) |
|---|---|---|---|
| SMs | 108 | 132 | 148 (2 dies) |
| BF16/FP16 tensor TFLOPS | 312 | ~989 | ~2,250 |
| FP8 tensor TFLOPS | n/a | ~1,979 | ~4,500 |
| FP4 tensor TFLOPS | n/a | n/a | ~9,000 (72 PF dense / 8 GPUs) NVIDIA DGX B200 |
| FP32 non-tensor TFLOPS | 19.5 | ~67 | ~75–80 |
| HBM capacity / bandwidth | 80 GB / ~2.0 TB/s | 80 GB / ~3.35 TB/s | 180 GB / ~8 TB/s (1,440 GB, 64 TB/s per 8) NVIDIA DGX B200 |
| L2 cache | 40 MB | 50 MB | ~126 MB (GB200 per tuning guide) NVIDIA docs |
| BF16 ridge point (FLOP/byte) | ~153 | ~295 | ~280 |
| NVLink per GPU | 600 GB/s | 900 GB/s | 1.8 TB/s |
Sources for H100: NVIDIA H100 NVIDIA 2022. A100 numbers as quoted in He 2022. Vendor peaks are marketing numbers measured at boost clocks. Real GEMMs typically reach 60–80% of peak, and sustained clocks drop under power limits. One measurement found an H100 matmul runs measurably faster on all-zeros input than on random data, because switching power triggers throttling He 2024.
Tensor-core FLOPs grew roughly 3× from A100 to H100 and again ~2× to B200. HBM bandwidth grew about 1.7× and then ~2.4×. Exp/special-function throughput and shared-memory bandwidth grew more slowly still. Each generation, the "everything that isn't a matmul" part of a kernel becomes relatively more expensive. That's why FlashAttention-3 and -4 are mostly about hiding softmax and data movement behind matmuls.
"Explain how a GPU hides memory latency" is a classic. A strong answer: many resident warps with registers kept on chip, zero-cost warp switching, occupancy limited by registers and shared memory, and the caveat that top kernels instead use asynchronous copies and ILP at low occupancy. Bonus points: work out the occupancy for a given register count, and say that warp divergence only matters within a warp.
- NVIDIA Hopper Architecture In-Depth: authoritative per-SM numbers, TMA, clusters, FP8.
- Modal GPU Glossary: concise definitions of every term on this page, with diagrams.
- CUDA C++ Programming Guide: the programming model straight from the source.
2. The memory hierarchy
Performance engineering on GPUs is mostly data-movement engineering. Each level of the hierarchy is roughly an order of magnitude bigger and slower than the one above it, and the art is to load each byte from HBM once and reuse it many times from faster levels.
| Level | Size (H100 SXM) | Bandwidth | Latency | Managed by |
|---|---|---|---|---|
| Registers | 64K × 32-bit = 256 KB per SM | hundreds of TB/s aggregate | ~1 cycle | Compiler (spills go to "local memory", which is actually HBM-backed) |
| Shared memory / L1 | 256 KB per SM, up to 228 KB as SMEM | ~128 B/clk/SM ≈ 30 TB/s aggregate | ~20–30 cycles | Programmer (SMEM) / hardware (L1) |
| L2 | 50 MB | ~2–4× HBM | ~200–300 cycles | Hardware (with persistence hints) |
| HBM3 | 80 GB | ~3.35 TB/s | ~400–800 cycles | Programmer (allocations) |
| NVLink (peer GPU) | n/a | 900 GB/s total per GPU | µs-scale | NCCL / NVSHMEM |
| PCIe Gen5 x16 (host) | n/a | ~64 GB/s per direction | µs-scale | cudaMemcpy, pinned memory |
Treating "L1/SMEM is fast" as unlimited. Shared-memory bandwidth per SM is finite (about one 128-byte wavefront per clock), and for tensor-core kernels it's often the actual bottleneck once HBM traffic is solved. That's why Hopper lets WGMMA read operands straight from SMEM in a swizzled layout, and why Blackwell adds TMEM and 2-CTA MMA "to decrease shared memory traffic" Zadouri+ 2026.
Coalescing: how global memory wants to be read
HBM is accessed in 32-byte sectors (grouped into 128-byte cache lines). When the 32 threads of a warp load 32 consecutive 4-byte words, the hardware merges them into 4 sector transactions: a coalesced access. If they stride by 1 KB, every thread touches a different sector and you waste up to 8× bandwidth, because 32 bytes are fetched for every 4 used Harris 2013. Rules of thumb: map the fastest-varying thread index (threadIdx.x) to the contiguous dimension, use vectorized 128-bit loads (float4, 8 BF16 values) per thread, and keep row starts aligned (padding dimensions to multiples of 8/16/64 helps both alignment and tensor-core tile shapes).
Shared-memory bank conflicts
Shared memory is split into 32 banks, each 4 bytes wide. Address \(a\) maps to bank \((a/4) \bmod 32\). If the threads of a warp hit distinct banks (or the same word, which broadcasts), the access takes one cycle. If \(k\) threads hit different words in the same bank, it replays \(k\) times. The textbook case is reading a column of a 32×32 float tile: every element of the column sits in the same bank, giving a 32-way conflict. The fixes are padding (declare tile[32][33] so columns shift banks) or swizzling, which XORs the column index with bits of the row index so that both row-wise and column-wise accesses spread across banks. CUTLASS/CuTe and TMA use swizzled layouts (32/64/128-byte swizzle modes) because padding wastes space and breaks the layouts tensor cores expect Harris 2013.
"What is a bank conflict and how do you fix it?" plus "what does coalescing mean?" are standard filters for kernel roles. Say: 32 banks × 4 B, conflicts serialize, padding versus XOR swizzling. Coalescing means consecutive lanes touch consecutive addresses, so a warp's load becomes a few 32 B sectors. Also mention that tensor-core kernels avoid shared-memory round trips where they can (WGMMA from SMEM descriptors, TMEM on Blackwell).
- CUDA C++ Best Practices Guide: coalescing, shared memory, occupancy chapters.
- Using Shared Memory in CUDA C/C++: bank conflicts with worked examples.
- Hopper Tuning Guide: per-architecture limits (warps/SM, SMEM configs).
3. The roofline model and arithmetic intensity
The roofline model Williams+ 2009 is the single most useful tool for reasoning about kernel performance. Define the arithmetic intensity of a kernel as the FLOPs it performs per byte it moves to or from (some level of) memory:
$$ I = \frac{\text{FLOPs}}{\text{bytes moved}}, \qquad \text{attainable FLOP/s} = \min\big(P_{\text{peak}},\; B_{\text{mem}} \cdot I\big). $$The ridge point \(I^* = P_{\text{peak}}/B_{\text{mem}}\) is the intensity at which a kernel switches from memory-bound to compute-bound. For H100 SXM in dense BF16: \(I^* \approx 989\times10^{12} / 3.35\times10^{12} \approx 295\) FLOP/byte. In FP8 it's about 590. A kernel has to do roughly 300 multiply-adds' worth of work for every byte it reads from HBM before the tensor cores become the limit. Equivalently, as a time bound:
$$ t \;\gtrsim\; \max\!\left(\frac{\text{FLOPs}}{P_{\text{peak}}},\; \frac{\text{bytes}}{B_{\text{mem}}}\right) \;(+\;\text{overhead}). $$This "max of compute time and memory time" form is what you use for back-of-envelope estimates Austin+ 2025 (Scaling Book).
Worked examples
Assume BF16 (2 bytes per element) and count only compulsory HBM traffic (each input read once, each output written once).
(a) Elementwise add, \(z = x + y\), \(n\) elements
FLOPs \(= n\). Bytes \(= 3 \times 2n = 6n\). So \(I = 1/6 \approx 0.17\). At 3.35 TB/s that's ~0.56 TFLOP/s, about 0.06% of tensor peak. Even GELU (~10–20 FLOPs/element, on CUDA cores rather than tensor cores) stays far below the ridge. For a 64M-element GELU: read 128 MB + write 128 MB = 256 MB → ~76 µs. The math takes ~1–10 µs. Every elementwise and normalization op is memory-bound. Its cost is its bytes.
(b) GEMM, \(C_{M\times N} = A_{M\times K} B_{K\times N}\)
$$ \text{FLOPs} = 2MNK, \qquad \text{bytes} = 2(MK + KN + MN), \qquad I = \frac{MNK}{MK+KN+MN}. $$For a square \(n^3\) GEMM this is \(n/3\). At \(n=4096\), \(I \approx 1365\), well past the ridge, so compute-bound. Time ≈ \(2 \cdot 4096^3 / 989\text{T} \approx 139\,\mu s\) at peak, realistically ~170–200 µs. Now consider a skinny GEMM, where \(M\) is the number of tokens: with \(M \ll K, N\), \(I \approx \frac{MNK}{KN} = M\). Intensity equals the number of tokens in the batch. That's the whole story of LLM decode.
(c) Decode GEMV: one token through a weight matrix
For batch \(B\) tokens through a \(d\times d\) BF16 weight: FLOPs \(= 2Bd^2\), bytes ≈ \(2d^2\) (weights dominate), so \(I \approx B\). At \(B=1\), \(I \approx 1\), about 300× below the ridge. A model with \(P\) parameters in BF16 must stream \(2P\) bytes per decode step:
$$ t_{\text{token}} \gtrsim \frac{2P}{B_{\text{mem}}} \quad\Rightarrow\quad \text{8B model on one H100: } \frac{16\,\text{GB}}{3.35\,\text{TB/s}} \approx 4.8\,\text{ms} \;\Rightarrow\; \lesssim 210\ \text{tokens/s at batch 1}. $$Batching is nearly free until \(B\) approaches the ridge (~150–300 tokens): you stream the same weights once and do \(B\times\) the work. That's why serving systems batch aggressively, and why quantizing weights to 8 or 4 bits speeds up decode almost proportionally (fewer bytes), even when the math still runs in BF16. Speculative decoding also exploits this: verifying \(k\) draft tokens costs about the same as one memory-bound step. (See the inference page.)
(d) Attention
Per head, with sequence length \(N\) and head dim \(d\): FLOPs ≈ \(4N^2 d\) (\(QK^\top\) and \(PV\)).
- Naive attention writes the \(N\times N\) score matrix \(S\) to HBM, reads it for softmax, writes \(P\), reads \(P\) again: roughly \(4 \times 2N^2\) extra bytes. Intensity ≈ \(4N^2 d / 8N^2 = d/2 \approx 64\) for \(d=128\). Memory-bound, and \(O(N^2)\) memory besides.
- FlashAttention keeps \(S, P\) on chip, so bytes ≈ \(4Nd \times 2\) (read Q, K, V, write O), ignoring re-reads of K/V per Q-tile. Intensity ≈ \(4N^2 d / 8Nd = N/2\). At \(N = 4096\), ~2000: compute-bound. (Re-reading K/V once per Q tile multiplies traffic by \(N/B_r\) but still leaves prefill compute-bound.)
- Decode attention (1 query token vs \(N\) cached keys): FLOPs ≈ \(4Nd\), bytes ≈ \(2Nd \times 2\) (K and V cache), so \(I \approx 1\). With grouped-query attention, \(g\) query heads share each KV head, so \(I \approx g\) (e.g. 8). Still memory-bound. This is why KV-cache size, GQA/MLA and KV quantization matter so much.
| Op (BF16) | Arithmetic intensity | Regime on H100 (ridge ~295) | Lever |
|---|---|---|---|
| Elementwise add/mul, activation | 0.1–2 | Memory | Fuse into neighbours |
| LayerNorm/RMSNorm, softmax (standalone) | ~1–5 | Memory | Fuse; one pass over data |
| Decode GEMV (batch \(B\)) | ≈ \(B\) | Memory until \(B\sim 150\)–\(300\) | Batch; quantize weights |
| Decode attention | ≈ GQA group size | Memory | GQA/MLA, KV quant, FlashDecoding |
| Naive attention prefill | ≈ \(d/2\) | Memory (and \(O(N^2)\) memory) | FlashAttention |
| FlashAttention prefill | ≈ \(N/2\)-ish | Compute (for long \(N\)) | Overlap softmax with MMA (FA3/FA4) |
| Large GEMM \(n^3\) | ≈ \(n/3\) | Compute | Tile well, FP8, avoid wave quantization |
In a transformer, more than 99% of FLOPs are matmuls, yet in an unfused eager-mode model a large share of runtime can go to the "trivial" elementwise and norm ops, because they're memory-bound. Fusion closes that gap. In decode, even the matmuls are memory-bound, so the whole model runs at HBM speed.
Using the wrong memory level. The roofline has a roof per level (HBM, L2, SMEM). A GEMM whose 128×128 block tile gives only ~64 FLOP/byte against global memory can still be compute-bound, because most of those "global" loads hit in L2 (neighbouring blocks share A rows and B columns). Say which level you're computing intensity against.
Expect a live back-of-envelope: "How long does one decode step of a 70B model take on 8×H100?" Answer: 140 GB of BF16 weights / (8 × 3.35 TB/s) ≈ 5.2 ms lower bound, plus KV-cache reads, plus all-reduce latency for tensor parallelism. So ~6–10 ms/token in practice at small batch. Then: "what changes at batch 256?" Intensity rises to ~256, near the ridge, so you approach compute-bound. KV-cache traffic grows linearly with batch × context and becomes the dominant byte cost.
- How to Scale Your Model: All About Rooflines: the best modern treatment with LLM-specific worked problems.
- Williams, Waterman, Patterson (2009), Roofline: the original paper.
- Making Deep Learning Go Brrrr From First Principles: compute vs memory vs overhead, with intuition.
4. Compute-bound, memory-bound, overhead-bound
Horace He's "Making Deep Learning Go Brrrr From First Principles" splits runtime into three buckets: compute (time on FLOPs), memory (time moving tensors) and overhead (everything else: Python interpreter, framework dispatch, kernel launches) He 2022. The practical value is diagnosis: each regime has a different fix, and optimizing the wrong one does nothing.
| Regime | Symptom | How to confirm | Fixes |
|---|---|---|---|
| Compute-bound | Tensor-core utilization high; time scales with FLOPs | Nsight Compute "SM / tensor pipe" throughput near peak; achieved TFLOPS close to spec | Lower precision (BF16 → FP8 → FP4), better tile shapes, avoid wave/tile quantization, reduce FLOPs (algorithmic: GQA, MoE, sparsity) |
| Memory-bound | DRAM throughput near 3 TB/s; time scales with bytes | ncu "Memory throughput" % of peak high, compute low; roofline places kernel on sloped part | Fuse ops, avoid materializing intermediates, quantize weights/KV, increase batch (more reuse), better layouts |
| Overhead-bound | GPU idle between kernels; doubling batch barely changes step time | Nsight Systems timeline shows gaps; many µs-sized kernels; CPU thread pegged at 100% | CUDA graphs, torch.compile (fewer, fused kernels), bigger batches, move logic off the Python hot path, async data loading |
The "doubling test" from the same essay is a great diagnostic to quote: if doubling the batch size increases runtime by much less than 2×, you're overhead-bound. Once runtime grows proportionally, you've left the overhead regime. Similarly, if switching an op from FP32 to BF16 halves its time, it was memory-bound (half the bytes). If it speeds up 8–15×, it moved onto tensor cores and was compute-bound.
Overhead magnitudes: a PyTorch eager op costs on the order of several µs of CPU-side dispatch, and a raw CUDA kernel launch costs a few µs. Meanwhile an H100 can do ~1 GFLOP per µs of tensor math, so a small decode-step kernel can finish faster than the CPU can enqueue the next one. Overhead dominates small-batch inference, small models, RL rollouts with tiny batches, and anything with lots of tiny ops (optimizers over many parameter tensors, which is why "foreach"/fused optimizers exist).
"Our training step is slower than expected. How do you investigate?" Strong answer: profile first (Nsight Systems / torch profiler). Classify the time into compute, memory and overhead. Check MFU (model FLOPs utilization) against a realistic target of ~35–55% for large dense training. Look for gaps (overhead, data loader, synchronizations like .item()), exposed communication, and top memory-bound kernels that could be fused. Name the fix per regime, not a generic "use a bigger GPU."
5. Kernel fusion, torch.compile and CUDA graphs
Why fusion matters
Consider y = gelu(x @ W + b) followed by dropout and a residual add, in eager mode. The matmul writes its output to HBM. The bias add reads it and writes it back. GELU reads and writes again. Dropout reads, writes, and also writes a mask. The residual add reads two tensors and writes one. Each memory-bound step costs a full HBM round trip of the activation. Fusing all the elementwise work into the GEMM's epilogue (applied to the accumulator tile while it's still in registers) leaves exactly one write. For a chain of \(k\) elementwise ops on a tensor of \(S\) bytes, unfused traffic is ~\(2kS\) and fused traffic is ~\(2S\): a \(k\times\) speedup for the memory-bound part. He's essay shows x.cos().cos() fused running about as fast as a single cos He 2022.
Vertical fusion
Producer → consumer chains (matmul → bias → GELU; the ops inside RMSNorm). Intermediate values stay in registers/SMEM. This is the common, high-value kind.
Horizontal fusion
Independent ops of the same shape launched together (e.g. Q, K and V projections as one GEMM with concatenated weights; multi-tensor optimizer updates). This cuts launches and improves GPU fill.
Fusion has limits. Fusing a reduction (softmax, layernorm) with something that needs the finished reduction requires either the whole row to fit on chip or a clever recurrence (online softmax, which section 7 covers). Fusing two GEMMs back to back is hard because the intermediate is large and the second GEMM needs all of it. FlashAttention is precisely "fuse GEMM → softmax → GEMM", which is why it was a research result rather than a compiler pass.
torch.compile and TorchInductor
PyTorch 2's torch.compile has two stages Ansel+ 2024:
- TorchDynamo captures the Python program into FX graphs by intercepting bytecode at runtime. It installs guards (shapes, dtypes, Python values) that trigger recompiles when violated, and inserts graph breaks around untraceable code (data-dependent control flow, some Python side effects,
.item()). - TorchInductor lowers the graph, decides fusion groups, and generates Triton kernels for GPUs (C++/OpenMP for CPUs). Matmuls usually go to cuBLAS/CUTLASS, or optionally autotuned Triton templates with fused epilogues (
mode="max-autotune").
Typical gains are 1.3–2× for training on memory-bound-heavy models. Practical gotchas: compile time (minutes for large models), recompilation storms with dynamic shapes (mark dims dynamic, pad or bucket sequence lengths), and graph breaks that silently split the graph (inspect with TORCH_LOGS="graph_breaks" or torch._dynamo.explain). Treat exact flag names as something to check in the current docs PyTorch docs.
CUDA graphs: removing launch overhead
A CUDA graph records a sequence of kernel launches (with their arguments and memory addresses) once, then replays the whole DAG with a single launch NVIDIA 2019. The per-kernel CPU cost disappears, and the driver can launch dependent kernels back to back. Constraints: static shapes and static memory addresses (inputs are copied into fixed buffers), no CPU synchronization or data-dependent control flow inside the graph. LLM inference servers capture one graph per batch-size bucket for the decode step, which is a big part of why decode with small batches isn't launch-bound. In PyTorch, torch.cuda.graphs / torch.cuda.CUDAGraph provide manual capture, and torch.compile(mode="reduce-overhead") applies CUDA graphs automatically PyTorch 2021.
Benchmarking GPU code without synchronization. CUDA launches are asynchronous, so time.time() around an op measures the enqueue. Use CUDA events or torch.cuda.synchronize(), warm up first (autotuning, compilation, caches), and take a median over many iterations (triton.testing.do_bench handles this). Also clear L2 between iterations if you want a cold-cache number. Otherwise small inputs look impossibly fast.
- PyTorch 2 paper (ASPLOS 2024): Dynamo + Inductor design and measured speedups.
- Getting Started with CUDA Graphs: what is captured and why it removes launch overhead.
- Accelerating PyTorch with CUDA Graphs: the PyTorch API and pitfalls.
6. GEMM: tiling all the way down
Matrix multiplication is the kernel that matters most and the canonical example of hierarchical tiling. The naive kernel (one thread per output element, loop over \(K\) reading A and B from global memory) has an intensity of about 0.25 FLOP/byte and reaches a few percent of peak. Simon Boehm's step-by-step walkthrough goes from ~1% to ~90%+ of cuBLAS through a sequence of changes, and it's worth knowing that sequence Boehm 2022.
Step 1: block tiling in shared memory
Each block loads a \(BM\times BK\) slab of A and a \(BK\times BN\) slab of B into shared memory, synchronizes, then every thread computes from SMEM. Each element loaded from global memory is now reused \(BN\) times (for A) or \(BM\) times (for B). Intensity against global memory per K-step:
$$ I_{\text{tile}} = \frac{2\,BM\,BN\,BK}{2\,(BM + BN)\,BK} = \frac{BM \cdot BN}{BM + BN}\ \text{FLOP/byte (BF16)}. $$A 128×128 tile gives 64 FLOP/byte and 128×256 gives ~85. That's below the HBM ridge of ~295, which is why the order in which blocks are scheduled matters. Blocks that run at the same time and share A rows or B columns get L2 hits, so actual HBM traffic is several times lower. The "grouped" or "swizzled" block ordering in the Triton matmul tutorial and CUTLASS's rasterization options exist for exactly this L2 reuse. Hopper's TMA multicast within a cluster goes further: one load from L2 is delivered to the shared memory of several SMs.
Step 2: register tiling
Shared memory is fast but finite (see the bank discussion). So each thread (or with tensor cores, each warp) computes a \(TM\times TN\) micro-tile of outputs held in registers. Per \(k\) step it loads \(TM + TN\) values from SMEM and does \(TM\cdot TN\) FMAs, giving a SMEM-level intensity of \(TM\cdot TN/(TM+TN)\). With tensor cores, the "register tile" is the MMA fragment, and the hardware dictates its shape.
Step 3: coalesced, vectorized, conflict-free loads
Global → shared copies use 128-bit vector loads with consecutive threads on consecutive addresses. The SMEM layout is swizzled so that the tensor-core fragment loads (ldmatrix on Ampere) are bank-conflict-free. Transposing on the fly (e.g. when B is stored row-major but the MMA wants column-major) is done in the layout, not with scalar shuffles.
Step 4: software pipelining (double/multi-buffering)
Without pipelining, a block alternates between "load tile k" (tensor cores idle) and "compute tile k" (memory idle). With a ring of \(S\) SMEM stages, the block issues loads for tiles \(k+1 \ldots k+S-1\) while computing tile \(k\). Ampere introduced cp.async: an asynchronous global → shared copy that bypasses registers, completion-tracked by commit groups. Hopper introduced the Tensor Memory Accelerator (TMA): a single thread issues a bulk copy of a whole multi-dimensional tile, described by a tensor map (base, shape, strides, swizzle), and the hardware computes addresses and signals completion through an mbarrier in shared memory NVIDIA 2022. This frees the other threads from address arithmetic and enables multicast.
Step 5: warp specialization and WGMMA (Hopper)
With TMA and async WGMMA, the best Hopper kernels split warps into roles. Producer warps (often one warp, with registers deallocated via setmaxnreg) issue TMA loads into the stage ring. Consumer warpgroups wait on the "full" barrier for a stage, issue WGMMA, and release the stage through the "empty" barrier. This is a classic bounded producer–consumer queue implemented in hardware barriers.
Two further patterns you'll hear about: persistent kernels (launch exactly one CTA per SM, each looping over many output tiles, so tile setup and epilogue overlap with the next tile's mainloop; the scheduler is in software) and Stream-K / split-K (split the K dimension across CTAs when \(M\times N\) is too small to fill the GPU, then reduce the partials).
Tile and wave quantization
Two shape effects explain "why is my GEMM slower with N=4097 than N=4096?" Tile quantization: if a dimension isn't a multiple of the tile size, the edge tiles do wasted work. Wave quantization: with 132 SMs and one tile per SM per wave, 133 tiles take two waves, the second nearly empty, so efficiency drops by about half. Keep dims multiples of 64/128, and pad vocabulary sizes (a known trick is padding vocab to a multiple of 64 or 128).
CUTLASS and CuTe
CUTLASS is NVIDIA's open-source C++ template library for GEMMs and convolutions at every level of this hierarchy NVIDIA CUTLASS. CUTLASS 3.x is built on CuTe, an algebra of layouts: a layout is a (shape, stride) pair mapping logical coordinates to offsets, and layouts compose, tile (local_tile, local_partition) and swizzle. This lets one express "partition this tile across these threads in this tensor-core fragment layout" uniformly NVIDIA blog. In 2025 NVIDIA shipped the CuTe DSL, a Python-embedded way to write CuTe kernels with much faster compile times than C++ templates NVIDIA docs. FlashAttention-4 is written in it Zadouri+ 2026.
| Generation | Load path | MMA | Accumulator | Key new tricks |
|---|---|---|---|---|
| Volta/Turing | LDG → regs → STS | wmma/mma.sync (warp) | Registers | First tensor cores |
| Ampere (A100) | cp.async global → SMEM | mma.sync (warp) | Registers | Multi-stage async pipelines, ldmatrix, BF16/TF32 |
| Hopper (H100) | TMA (bulk, multicast) | WGMMA (warpgroup, async, operands from SMEM) | Registers | Warp specialization, clusters/DSMEM, FP8 |
| Blackwell (B200) | TMA | tcgen05.mma (single thread issues, 2-CTA mode) | Tensor memory (TMEM) | FP4/FP6 + block-scaled formats, frees registers for epilogue/softmax |
"Walk me through optimizing a matmul kernel" is the canonical kernel-engineer question. Hit these steps in order: naive → coalescing → SMEM block tiling (state the reuse formula) → register tiling → vectorized loads and bank-conflict-free/swizzled layouts → double-buffering with async copies → tensor cores → (Hopper) TMA + WGMMA + warp specialization + persistent scheduling → epilogue fusion → autotuning tile sizes per shape. Then mention wave quantization and L2-aware rasterization. You don't have to write all of it. Showing you know why each step helps is the signal.
- How to Optimize a CUDA Matmul Kernel for cuBLAS-like Performance: the classic step-by-step worklog.
- NVIDIA CUTLASS: examples directory includes Hopper warp-specialized and Blackwell kernels.
- DeepGEMM: a compact, readable FP8 GEMM library for Hopper/Blackwell with fine-grained scaling.
7. FlashAttention and FlashDecoding
Standard attention for one head computes \(S = QK^\top/\sqrt{d}\) (\(N\times N\)), \(P = \text{softmax}(S)\), \(O = PV\). Implemented as three kernels, it writes and reads \(S\) and \(P\) through HBM: \(O(N^2)\) memory and traffic. FlashAttention made attention IO-aware. It computes exact attention (no approximation) in tiles that live in SRAM and never materializes \(S\) or \(P\) in HBM Dao+ 2022.
The online softmax recurrence
The obstacle is that softmax normalizes over the whole row, but we want to process K/V in blocks. The fix comes from the online normalizer trick Milakov & Gimelshein 2018, also used for memory-efficient attention Rabe & Staats 2021. Keep, for each query row, a running max \(m\), a running denominator \(\ell\), and an unnormalized output accumulator \(\tilde O\). When a new block of scores \(S^{(j)} = Q_i K_j^\top\) arrives:
$$ \begin{aligned} m^{\text{new}} &= \max\!\big(m,\ \operatorname{rowmax}(S^{(j)})\big) \\ \tilde P^{(j)} &= \exp\!\big(S^{(j)} - m^{\text{new}}\big) \\ \ell^{\text{new}} &= e^{\,m - m^{\text{new}}}\,\ell \;+\; \operatorname{rowsum}(\tilde P^{(j)}) \\ \tilde O^{\text{new}} &= e^{\,m - m^{\text{new}}}\,\tilde O \;+\; \tilde P^{(j)} V_j \end{aligned} \qquad\text{and finally}\qquad O = \tilde O / \ell . $$Why it's exact: every term ever added to \(\ell\) and \(\tilde O\) carries a factor \(e^{-m_{\text{old}}}\). Multiplying by \(e^{m_{\text{old}} - m_{\text{new}}}\) re-bases everything to the new max, so at the end both numerator and denominator are referenced to the true row max and the ratio equals the standard softmax. Subtracting the max keeps every \(\exp\) argument ≤ 0, so there's no overflow. For the backward pass, store just the per-row log-sum-exp \(L = m + \log \ell\) (\(N\) floats, not \(N^2\)).
FlashAttention (v1, 2022): tiling + recomputation + IO analysis
- Tiling. The forward pass uses the recurrence above. Block sizes are chosen so that \(Q_i\), \(K_j\), \(V_j\) and \(S_{ij}\) fit in on-chip SRAM of size \(M\).
- Recomputation in backward. Instead of storing \(P\) (\(N^2\)), the backward pass recomputes \(S_{ij}\) and \(P_{ij} = \exp(S_{ij} - L)\) tile by tile from Q, K and the saved LSE. This adds FLOPs but saves far more HBM traffic, a trade that's clearly right when you're memory-bound. It's essentially selective activation checkpointing inside the kernel.
- IO complexity. Standard attention needs \(\Theta(Nd + N^2)\) HBM accesses, and FlashAttention needs \(\Theta(N^2 d^2 / M)\). For typical \(d\) (64–128) and \(M\) (~100–200 KB), \(d^2/M\) is well below 1, so traffic drops by many times. The paper also proves this is asymptotically optimal over a range of SRAM sizes Dao+ 2022.
- Memory is linear in \(N\), which is what made long context practical to train.
The reported results were roughly 3× faster GPT-2 training and 15% faster BERT-large than the baselines of the time Dao+ 2022.
FlashAttention-2 (2023): parallelism and work partitioning
FA1 reached only ~25–40% of A100 peak. FA2 made three changes Dao 2023:
- Fewer non-matmul FLOPs. Don't rescale \(\tilde O\) by \(1/\ell\) on every iteration. Keep it unnormalized and divide once at the end (as in the equations above). Non-matmul FLOPs cost up to ~16× more per FLOP than tensor-core FLOPs, so this matters.
- Parallelize over sequence length. FA1 parallelized over batch × heads only, which underfills 108–132 SMs at long context and small batch. FA2 also launches one CTA per Q block, with the Q-block loop as the outer, parallel loop.
- Better warp partitioning within a block. FA1 split K/V across warps ("split-K"), so warps had to exchange partial results through shared memory. FA2 splits Q across warps, all sharing K/V. Each warp owns complete rows and no inter-warp reduction is needed.
The result was about 2× over FA1 and 50–73% of A100 theoretical peak in the forward pass Dao 2023.
FlashAttention-3 (2024): Hopper asynchrony and FP8
On H100, FA2 reached only ~35% of peak because it didn't use TMA/WGMMA, and because softmax became a big fraction of time. H100 has ~989 TFLOPS of matmul but only ~3.9 TFLOPS of special-function (exp) throughput, a 256× gap Dao+ 2024 (blog). FA3's ideas Shah+ 2024:
- Warp specialization: producer warps issue TMA loads of K/V; consumer warpgroups do WGMMA + softmax (the GEMM pipeline pattern from section 6).
- Ping-pong scheduling between warpgroups: while warpgroup 1 does softmax (exp on the SFUs), warpgroup 2 runs its GEMMs on the tensor cores, and vice versa, enforced with named barriers.
- Intra-warpgroup pipelining: overlap the softmax of block \(j\) with the \(QK^\top\) GEMM of block \(j+1\), at the cost of extra registers.
- FP8: block quantization plus "incoherent processing" (multiply Q and K by a random orthogonal/Hadamard matrix to spread outliers before quantizing), giving lower error than naive per-tensor FP8.
Reported results: up to 740 TFLOPS in FP16 (~75% utilization) and close to 1.2 PFLOPS in FP8 on H100 Shah+ 2024.
FlashAttention-4 (2025–26): Blackwell
FA4 exists. It was first released in the flash-attention repo in 2025 and described in a March 2026 paper, "FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling" Zadouri+ 2026. The premise is that on Blackwell, tensor-core throughput doubled while exp throughput and shared-memory bandwidth scaled much less, so the non-matmul work must shrink further. Its main ideas, per the abstract:
- Pipelines redesigned around fully asynchronous
tcgen05MMAs with accumulators in tensor memory, and larger tiles. - Software-emulated exponential (a polynomial approximation on the FMA units) to offload part of the exp work from the SFUs.
- Conditional softmax rescaling: skip rescaling \(\tilde O\) unless the running max has grown enough to threaten numerical stability. The math stays exact because the final \(\tilde O/\ell\) uses a consistent reference max.
- TMEM and 2-CTA MMA in the backward pass to reduce shared-memory traffic and atomics.
- Written entirely in CuTe DSL (Python), with 20–30× faster compilation than C++ templates.
Reported: up to ~1613 TFLOPS BF16 forward on B200 (~71% utilization), up to 1.3× over cuDNN 9.13 and 2.7× over Triton Zadouri+ 2026.
FA4 and Blackwell attention kernels (cuDNN, FlashInfer, vendor FP4 attention variants) are moving quickly in 2026, and follow-ups such as FP4 attention papers are appearing. Check the flash-attention repo for current support (forward/backward, head dims, FP8/FP4, Hopper vs Blackwell) before quoting specifics.
| Version | Target | Core idea | Headline result (as reported) |
|---|---|---|---|
| FA1 (2022) | A100 | Tiling + online softmax + backward recomputation; IO-optimal | Linear memory; ~2–4× faster attention |
| FA2 (2023) | A100 | Fewer non-matmul FLOPs; parallelize over seq; split Q not K across warps | ~2× FA1; 50–73% of peak |
| FA3 (2024) | H100 | TMA/WGMMA, warp specialization, ping-pong softmax/GEMM overlap, FP8 | ~740 TFLOPS FP16, ~1.2 PF FP8 |
| FA4 (2025–26) | B200 | Async tcgen05 + TMEM, emulated exp, conditional rescale, 2-CTA backward, CuTe DSL | ~1613 TFLOPS BF16 (~71%) |
FlashDecoding and split-K for decode
At decode time the query has length 1 per sequence. FA2's grid is batch × heads × (Q blocks = 1). With batch 1 and 32 heads that's only 32 CTAs on 132 SMs, each streaming a long KV cache serially, so most of the GPU idles. Flash-Decoding also splits the KV sequence into chunks processed by different CTAs, each producing a partial output \(O_s\) and its log-sum-exp \(\text{lse}_s\). A small second kernel then merges them Dao+ 2023 (CRFM blog):
$$ \text{lse} = \log \sum_s e^{\text{lse}_s}, \qquad O = \sum_s e^{\,\text{lse}_s - \text{lse}}\; O_s . $$This is the same associativity that makes online softmax work, applied across CTAs instead of across loop iterations. The authors reported up to 8× faster end-to-end generation for very long sequences, and attention itself up to ~50× faster than FlashAttention at 64k context and batch 1 CRFM 2023. Choosing the number of splits is a heuristic: enough CTAs to fill the SMs, but not so many that the reduction dominates. Follow-ups such as FlashDecoding++ Hong+ 2023 attack the remaining synchronization. Production engines (FlashInfer, vLLM, SGLang, TensorRT-LLM) combine split-KV with paged KV layouts and GQA packing, where all query heads in a group share one pass over the KV block so intensity rises to ~\(g\). DeepSeek's FlashMLA does the analogous thing for multi-head latent attention DeepSeek FlashMLA.
A related tool is FlexAttention in PyTorch, which lets you write a score_mod/mask_mod in Python (causal, sliding window, ALiBi, document masking) and compiles it into a fused FlashAttention-style Triton kernel, skipping fully masked blocks via a block mask Dong+ 2024.
The most common deep question here is "derive or explain online softmax and why FlashAttention is exact." Write the four update equations, explain the re-basing factor \(e^{m-m'}\), say what's stored for backward (LSE), and state the IO complexity. Follow-ups: "Is FlashAttention approximate?" (No.) "Does it reduce FLOPs?" (No. It increases them slightly in backward. It reduces bytes.) "Why doesn't FA2 help decode?" (Not enough parallelism, so split the KV, i.e. FlashDecoding.) "What changed on Hopper/Blackwell?" (Async TMA/WGMMA/TMEM and hiding softmax, because exp throughput didn't scale with tensor cores.)
- FlashAttention (Dao+ 2022): section 3 has the algorithm and IO-complexity proof.
- FlashAttention-3 blog post: clear diagrams of ping-pong and intra-warpgroup overlap.
- FlashAttention-4 (Zadouri+ 2026): Blackwell-specific co-design.
- Flash-Decoding for long-context inference: split-KV with animations.
8. Writing kernels: CUDA C++, Triton and the new DSLs
CUDA C++: the thread-level model
In CUDA you write the program for one thread. You compute your global index from blockIdx, blockDim and threadIdx, and manage shared memory, synchronization, vectorization, tensor-core fragments and async copies yourself (or through CUTLASS/CuTe). You get maximum control, including inline PTX and every new hardware feature on day one. The cost is a lot of code and room for subtle bugs (races, bank conflicts, misaligned TMA descriptors).
// CUDA: one thread per element, grid-stride not shown
__global__ void add(const float* x, const float* y, float* z, int n) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i < n) z[i] = x[i] + y[i]; // coalesced: lane i touches address i
}
// launch: add<<<(n + 255) / 256, 256>>>(x, y, z, n);
Triton: the block-level model
Triton Tillet+ 2019 moves the unit of programming from the thread to the block (program). You write code that operates on whole tiles (tl.arange vectors and 2D blocks). The compiler decides how to map them onto threads, coalesce loads, use shared memory, insert synchronization, pipeline loads (num_stages) and emit tensor-core instructions for tl.dot. You still choose the tiling, the grid and the fusion strategy, which is where most of the performance comes from. It's the language TorchInductor emits, and it runs on NVIDIA and AMD GPUs.
import torch, triton
import triton.language as tl
@triton.jit
def add_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr):
pid = tl.program_id(axis=0) # which block am I?
offs = pid * BLOCK + tl.arange(0, BLOCK) # a vector of BLOCK indices
mask = offs < n # guard the ragged last block
x = tl.load(x_ptr + offs, mask=mask)
y = tl.load(y_ptr + offs, mask=mask)
tl.store(out_ptr + offs, x + y, mask=mask)
def add(x, y):
out = torch.empty_like(x)
n = out.numel()
grid = lambda meta: (triton.cdiv(n, meta["BLOCK"]),)
add_kernel[grid](x, y, out, n, BLOCK=1024)
return out
A fused row-wise softmax, one program per row, with the whole row held on chip (the pattern from the official tutorial Triton docs):
@triton.jit
def softmax_kernel(out_ptr, in_ptr, in_stride, out_stride, n_cols,
BLOCK: tl.constexpr):
row = tl.program_id(0)
cols = tl.arange(0, BLOCK) # BLOCK = next_power_of_2(n_cols)
mask = cols < n_cols
x = tl.load(in_ptr + row * in_stride + cols, mask=mask, other=-float("inf"))
x = x.to(tl.float32) # do the math in fp32
x = x - tl.max(x, axis=0) # numerical stability
num = tl.exp(x)
y = num / tl.sum(num, axis=0)
tl.store(out_ptr + row * out_stride + cols, y, mask=mask)
# grid = (n_rows,); BLOCK = triton.next_power_of_2(n_cols)
# Traffic: 1 read + 1 write per element (vs ~5 passes for eager max/sub/exp/sum/div).
When a row doesn't fit on chip (e.g. a 256k vocabulary), you loop over column chunks using the online-softmax recurrence: two passes (max and sum fused online, then normalize) instead of three.
The newer tile DSLs
Between "Triton hides too much for Hopper/Blackwell peak" and "CUDA C++ templates are painful", a crop of tile-level languages appeared in 2024–2026:
| Tool | Level | What it's for | Status (as of late 2026, verify) |
|---|---|---|---|
| CUDA C++ + CUTLASS/CuTe | Thread/warp + layout algebra | Peak GEMM/attention, new hardware first | Mature; the reference for NVIDIA |
| Triton | Block | Fused elementwise/reduction/matmul/attention; Inductor backend | Mature, NVIDIA + AMD; lower-level "Gluon" dialect exposes more control Triton repo |
| CuTe DSL (CUTLASS Python) | CuTe layouts in Python | Hopper/Blackwell peak kernels with fast compile | Released 2025; used by FA4 NVIDIA docs |
| cuTile / CUDA Tile | Tile | NVIDIA's own tile model aimed at tensor-core portability | Python available, C++ planned NVIDIA; newer, adoption still forming |
| ThunderKittens | Warp/tile primitives in C++ | Small library of 16×16 tile abstractions; simple, fast attention/GEMM kernels | Research (Stanford Hazy Research) Spector+ 2024 |
| TileLang | Tile, TVM-based | Separates dataflow from scheduling annotations; used for e.g. MLA/sparse-attention kernels | Open source, active Wang+ 2025 |
| Mojo (Modular) | Systems language, Python-like | Portable GPU kernels (NVIDIA/AMD) inside the MAX stack | Commercial/open components; evolving Modular |
| Pallas (JAX) | Block | Custom kernels for TPU (and GPU via Triton/Mosaic) | Part of JAX JAX docs |
| NKI | Tile | Kernels for AWS Trainium/Inferentia | AWS Neuron SDK AWS docs |
The kernel-DSL landscape changes month to month (cuTile, Gluon, CuTe DSL, TileLang and Mojo all had major releases in 2025–26). There's also active work on LLM-generated kernels and kernel-generation benchmarks. Treat the status column as a snapshot and check each project's repo before claiming feature support.
"When would you write Triton vs CUDA?" Triton for fused memory-bound ops, custom attention variants and quick iteration with autotuning (and anything that should run on AMD too). CUDA/CUTLASS/CuTe DSL when you need the last 10–30% on the latest hardware (warp specialization, TMEM, cluster multicast) or features the compiler doesn't expose yet. First, always check whether cuBLAS/cuDNN/FlashAttention/FlashInfer or torch.compile already does it.
- Triton tutorials: vector add, fused softmax, matmul (with grouped ordering), layernorm, attention.
- GPUs Go Brrr (ThunderKittens blog): an opinionated tour of what matters on H100.
- GPU MODE lectures: community lecture series on CUDA, Triton, CUTLASS, profiling.
9. Profiling: finding where the time goes
Nsight Systems (nsys)
System-wide timeline. CPU threads, CUDA API calls, kernel launches, memcpys, NCCL, NVTX ranges, all on one time axis. Use it to find gaps (overhead, sync points, data loading), stream concurrency, and whether communication overlaps compute. Low overhead. Profile whole training steps NVIDIA docs.
Nsight Compute (ncu)
Per-kernel deep dive. Replays a kernel many times to collect hardware counters: "Speed of Light" compute vs memory throughput as % of peak, roofline chart, achieved occupancy, warp stall reasons, L1/L2/DRAM traffic, bank conflicts, uncoalesced accesses. High overhead, so filter to specific kernels NVIDIA docs.
PyTorch profiler (torch.profiler.profile with ProfilerActivity.CPU/CUDA, record_function labels, export_chrome_trace) gives an op-level view mapped back to your Python code, plus memory profiling. It's the right first tool for model code PyTorch docs.
| What you see | Likely meaning | Next step |
|---|---|---|
| GPU idle gaps between short kernels in nsys | Overhead-bound (Python/dispatch/launch) | CUDA graphs, torch.compile, larger batch |
Long cudaStreamSynchronize / .item() on the CPU | Hidden host–device syncs serialize the pipeline | Remove syncs from the hot loop; keep values on device |
| ncu: DRAM throughput ~90%, compute ~10% | Memory-bound and already efficient | Only fusion or fewer bytes (quantization) will help |
| ncu: both throughputs low | Latency-bound: low occupancy, poor ILP, too few CTAs, tail effect | More parallelism (split-K), more bytes in flight, check stall reasons |
| Stall reason "long scoreboard" | Waiting on global memory loads | Prefetch/pipeline, async copies |
| Stall reason "MIO throttle" / "short scoreboard"; bank-conflict counters high | Shared-memory pressure | Swizzle/pad layouts, fewer SMEM round trips |
| NCCL kernels not overlapping GEMMs | Exposed communication | Bucketing, separate streams, overlap scheduling (FSDP prefetch) |
| Kernel uses 255 registers, local-memory traffic | Register spills | Smaller tiles, fewer live values, launch_bounds/maxnreg |
Be clear on which tool answers which question: nsys answers "where is my step's time going, and is the GPU busy?", ncu answers "why is this one kernel slow?", and the torch profiler answers "which line of my model launched it?" Mention MFU as the top-line metric and achieved bandwidth (GB/s ÷ peak) for memory-bound kernels.
10. Numerics: formats, accumulation and rounding
| Format | Bits (sign/exp/mantissa) | Max normal (approx) | Typical use |
|---|---|---|---|
| FP32 | 1/8/23 | 3.4e38 | Master weights, optimizer states, accumulators, reductions |
| TF32 | 1/8/10 (19 bits used) | 3.4e38 | Tensor-core input mode for "FP32" matmuls (Ampere+); torch.backends.cuda.matmul.allow_tf32 |
| BF16 | 1/8/7 | 3.4e38 | Default training/inference dtype; same range as FP32, so usually no loss scaling |
| FP16 | 1/5/10 | 65,504 | More precision, little range; needs loss scaling in training Micikevicius+ 2017 |
| FP8 E4M3 | 1/4/3 | 448 | Weights/activations (forward) Micikevicius+ 2022 |
| FP8 E5M2 | 1/5/2 | 57,344 | Gradients (need range more than precision) |
| FP4 E2M1 | 1/2/1 | 6 | Only usable with block scaling: MXFP4 (32-elt blocks, E8M0 scale) Rouhani+ 2023; NVFP4 (16-elt blocks, E4M3 scale + FP32 tensor scale) NVIDIA 2025 |
Key ideas
- Accumulate in higher precision. Tensor cores multiply low-precision inputs and accumulate in FP32 (or FP16 optionally). A dot product of length \(K = 8192\) accumulated in BF16 loses precision as the running sum grows, because each add rounds to 8 mantissa bits. Hopper's FP8 tensor-core accumulation turned out to keep only ~14 bits of precision. DeepSeek-V3 worked around this by periodically promoting partial sums to FP32 CUDA-core registers (every 128 elements along K) DeepSeek-AI 2024.
- Scaling is what makes FP8/FP4 work. A format with 3 or 1 mantissa bits and narrow range needs a scale so values fill the representable range: \(x \approx s \cdot q\). The granularity is a trade-off. Per-tensor (cheap, but one outlier wrecks it; "delayed scaling" uses the amax history), per-channel/per-token, or fine-grained blocks: DeepSeek-V3 used 1×128 activation tiles and 128×128 weight blocks, and MX/NVFP4 bake 32-/16-element block scales into the hardware format so the tensor core applies them.
- Reductions and softmax in FP32. Norm statistics, softmax max/sum, loss and logits are commonly kept in FP32 even in "BF16" models. The Triton softmax above upcasts for this reason.
- Stochastic rounding rounds up with probability equal to the fractional distance, so the rounding is unbiased in expectation. It matters when many tiny updates are added to a large value (e.g. BF16 weights + small gradient updates, where round-to-nearest silently drops them) Gupta+ 2015. Some accelerators support it in hardware. On GPUs it's usually done in software inside optimizer kernels, or avoided by keeping FP32 master weights.
- Non-determinism. Floating-point addition isn't associative, so split-K, atomics and different tile configs produce bitwise-different results run to run. This shows up in RL and evals as train/inference mismatch. "Batch-invariant" or deterministic kernels trade some speed for reproducibility.
Why doesn't FP8 give 2× end to end? Only the GEMMs get 2× peak. Memory-bound ops get at most 2× from halved bytes, and only if their inputs and outputs are actually FP8. Quantize/dequantize and scale computation add work, attention may stay in BF16, and communication and overhead don't change. Realistic training gains have been reported in the ~1.3–1.5× range rather than 2×.
Low-precision training recipes (FP8 with block scaling, MXFP8, NVFP4 training) are an active 2025–26 area, and which recipes are "standard" is still changing. Check current Transformer Engine / framework docs and recent model reports.
- FP8 Formats for Deep Learning: E4M3/E5M2 definitions and training results.
- Microscaling Data Formats: MXFP8/6/4 and shared block exponents.
- DeepSeek-V3 Technical Report: section on FP8 training infrastructure (fine-grained scaling, accumulation promotion).
11. Communication kernels and compute–comm overlap
Multi-GPU training and inference add a fourth cost: moving bytes between GPUs. Collectives (all-reduce, all-gather, reduce-scatter, all-to-all) are implemented by libraries such as NCCL NVIDIA NCCL as GPU kernels that read and write HBM and drive NVLink/InfiniBand.
- Bandwidth cost. A ring all-reduce of \(S\) bytes over \(n\) GPUs sends and receives \(2\frac{n-1}{n}S\) per GPU (a reduce-scatter, then an all-gather). With 1 GB of BF16 gradients over 8 H100s at an effective ~350–450 GB/s per direction: about 1.75 GB / 400 GB/s ≈ 4–5 ms.
- Latency cost. Small messages (TP all-reduces in decode) are latency-bound (µs per hop), so tree algorithms, fused "one-shot" all-reduce kernels and NVLink-SHARP-style in-switch reduction matter more than bandwidth there.
- Overlap. Run comm on a separate stream concurrent with independent compute: DDP buckets gradients and all-reduces them while backward continues, and FSDP prefetches the next layer's all-gather. The catch is that NCCL kernels occupy SMs and HBM bandwidth, so "overlapped" compute slows down. That's why some systems restrict comm to a few SMs or use copy engines. DeepSeek's DeepEP is a good example of custom all-to-all kernels for MoE expert parallelism with tunable SM usage DeepSeek DeepEP.
- Fused compute–comm kernels (e.g. GEMM + reduce-scatter tiled so each tile is sent as soon as it's ready, as in async tensor parallelism) push overlap into the kernel itself.
Know the \(2(n-1)/n\) factor, that reduce-scatter + all-gather = all-reduce (why ZeRO/FSDP costs ~1.5× DDP's traffic rather than 3×), and that overlap isn't free because comm kernels steal SMs and bandwidth. The inference page and the distributed-training material cover parallelism strategies in depth.
12. Beyond NVIDIA: AMD, TPU, Trainium
| AMD Instinct (MI300X / MI3xx) | Google TPU | AWS Trainium | |
|---|---|---|---|
| Compute unit | CUs with matrix cores (MFMA instructions); wavefront of 64 threads on CDNA (vs warp of 32) | Few big cores; MXU systolic arrays (128×128 on v5e/v5p, 256×256 on v6e) plus a vector unit Scaling Book | NeuronCores with tensor, vector, scalar and general-purpose engines |
| On-chip memory | LDS (shared memory) per CU, large L2 + Infinity Cache | VMEM scratchpad, compiler-managed | SBUF/PSUM scratchpads, software-managed |
| HBM (example) | MI300X: 192 GB, ~5.3 TB/s AMD | v5e ~820 GB/s, v5p ~2.8 TB/s per chip Scaling Book | Varies by generation |
| Programming | ROCm; HIP (CUDA-like, hipify), Composable Kernel, Triton backend, rocBLAS/hipBLASLt AMD ROCm | XLA compiler (whole-graph fusion and layout), JAX/PyTorch-XLA; Pallas for custom kernels OpenXLA | Neuron compiler; NKI for custom kernels |
| Interconnect | Infinity Fabric | ICI torus (2D/3D), optical reconfigurable on v4+ Jouppi+ 2023 | NeuronLink |
Systolic arrays, briefly. A TPU MXU is a 2D grid of multiply-accumulate cells. In the weight-stationary scheme, weights are preloaded into the cells, activations flow in from one side, and partial sums flow down. Each value is reused across a whole row or column of cells without touching memory, so a 128×128 array does 16K MACs per cycle with very little data movement Jouppi+ 2017. The trade-off: it's superb for large dense matmuls with friendly shapes (dims multiples of 128/256) and less flexible for irregular work. More of the optimization burden sits in the compiler (XLA), which explains the JAX-plus-compiler culture around TPUs. Roofline reasoning carries over unchanged: same equations, different peak and bandwidth numbers.
AMD (MI350/MI355-class and later), Google (TPU v7 "Ironwood"-era parts) and AWS (Trainium2/3) all shipped or announced new generations in 2025–26. Check vendor pages for current specs before quoting numbers.
13. Common kernel interview tasks (with sketched solutions)
Task 1: "Estimate the runtime of this kernel"
Recipe: count compulsory bytes and FLOPs, compute both times, take the max, then sanity-check overhead. Example: RMSNorm over a \(16384 \times 8192\) BF16 activation. Bytes = read 256 MB + write 256 MB (+ tiny weight) ≈ 512 MB, so 512 MB / 3.35 TB/s ≈ 153 µs at peak. Well-written kernels reach ~80–90% of peak bandwidth, so ~170–190 µs. FLOPs ≈ 4 per element ≈ 0.5 GFLOP, negligible. If the measured time is 600 µs, the kernel is making multiple passes or reads uncoalesced. Quote achieved bandwidth = bytes / time as your efficiency metric.
Task 2: Fused softmax in Triton
See section 8. Talking points: one program per row, BLOCK = next_power_of_2(n_cols), other=-inf for masked lanes (so they don't affect the max and exp to 0), subtract the max, FP32 math, a single read and write. For very wide rows, loop over chunks with online max/sum and do a second pass to write normalized outputs. Expected performance: near peak bandwidth. The eager version makes ~4–5 passes.
Task 3: LayerNorm forward in Triton
@triton.jit
def layernorm_fwd(X, Y, W, B, Mean, Rstd, stride, N, eps, BLOCK: tl.constexpr):
row = tl.program_id(0)
cols = tl.arange(0, BLOCK)
mask = cols < N
x = tl.load(X + row * stride + cols, mask=mask, other=0.).to(tl.float32)
mean = tl.sum(x, axis=0) / N
xc = tl.where(mask, x - mean, 0.)
var = tl.sum(xc * xc, axis=0) / N
rstd = 1 / tl.sqrt(var + eps)
w = tl.load(W + cols, mask=mask); b = tl.load(B + cols, mask=mask)
y = xc * rstd * w + b
tl.store(Y + row * stride + cols, y, mask=mask)
tl.store(Mean + row, mean); tl.store(Rstd + row, rstd) # saved for backward
Talking points: masked lanes must not pollute the variance (tl.where), save mean and rstd for backward, and the backward pass needs a reduction across rows for \(dW\) and \(dB\) (partial sums per program, then a second reduction kernel, or atomics; the Triton tutorial uses a lock-based accumulation). Welford's algorithm is the numerically stable alternative when you can't hold the row.
Task 4: Tiled matmul in Triton
@triton.jit
def matmul(A, B, C, M, N, K, sam, sak, sbk, sbn, scm, scn,
BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr):
pid_m = tl.program_id(0); pid_n = tl.program_id(1)
rm = pid_m * BM + tl.arange(0, BM)
rn = pid_n * BN + tl.arange(0, BN)
rk = tl.arange(0, BK)
a_ptrs = A + rm[:, None] * sam + rk[None, :] * sak
b_ptrs = B + rk[:, None] * sbk + rn[None, :] * sbn
acc = tl.zeros((BM, BN), dtype=tl.float32) # FP32 accumulator
for k in range(0, tl.cdiv(K, BK)):
a = tl.load(a_ptrs, mask=(rm[:, None] < M) & (rk[None, :] + k * BK < K), other=0.)
b = tl.load(b_ptrs, mask=(rk[:, None] + k * BK < K) & (rn[None, :] < N), other=0.)
acc += tl.dot(a, b) # tensor cores
a_ptrs += BK * sak; b_ptrs += BK * sbk
# epilogue: fuse bias / activation here before the cast
c = acc.to(C.dtype.element_ty)
tl.store(C + rm[:, None] * scm + rn[None, :] * scn, c,
mask=(rm[:, None] < M) & (rn[None, :] < N))
Talking points: FP32 accumulator, masks for ragged edges, @triton.autotune over (BM, BN, BK, num_warps, num_stages), grouped program ordering (a 1D pid remapped so consecutive programs share A rows) for L2 reuse, and epilogue fusion. Compare to cuBLAS at large sizes, where you should reach a large fraction of it, and at skinny shapes, where split-K helps.
Task 5: "Why is this kernel slow?" (diagnosis)
Common planted bugs: column-major access by consecutive threads (uncoalesced, fix by swapping index roles or transposing through SMEM), a reduction done with global atomics per element (fix with a warp-shuffle reduction __shfl_down_sync, then a block reduction in SMEM, then one atomic per block), a 32-way bank conflict on transposed SMEM reads (pad to 33), too few blocks for the GPU (grid of 16 on 132 SMs, so split the work), and register spills from huge per-thread arrays.
Task 6: Parallel reduction / sum
The standard hierarchy: each thread sums a strided chunk in registers (grid-stride loop with vector loads, so bandwidth-bound), then a warp reduction via shuffles (5 steps for 32 lanes, no SMEM), then a block reduction through SMEM (one value per warp), then either a second kernel or one atomic per block. The target is ~peak HBM bandwidth, because a sum is \(I \approx 0.25\)–\(0.5\) FLOP/byte.
Task 7: Back-of-envelope serving questions
"Max decode tokens/s for a 70B BF16 model on 8×H100 at batch 1?" ≈ 1 / 5.2 ms ≈ 190 tokens/s upper bound, realistically less. "What's the KV cache per token for Llama-3-70B-like settings (80 layers, 8 KV heads, d=128, BF16)?" \(2 \times 80 \times 8 \times 128 \times 2\,\text{B} = 327{,}680\,\text{B} \approx 320\,\text{KB/token}\), so a 32k-token sequence needs ~10 GB. At batch 32 with 32k contexts the KV cache is ~340 GB, so every decode step reads more than twice the 140 GB of weights (and it won't even fit next to the weights on 8×80 GB). Long-context batch decode is KV-bandwidth-bound, which is why GQA/MLA, KV quantization and paged attention exist.
Interview question bank
1. What is a warp, and why does warp divergence hurt performance?
A warp is 32 threads that share one instruction stream and execute in lockstep (SIMT). The SM's schedulers issue one instruction per warp per cycle across all 32 lanes. If threads in a warp take different branches, the warp executes each path in turn with the inactive lanes masked, so a two-way split roughly halves throughput for that code. Divergence across warps is free, since each warp has its own instruction stream. Since Volta each thread has its own program counter, which makes intra-warp synchronization patterns legal, but divergent paths still serialize. In practice you structure kernels so that branches depend on warp- or block-uniform values, and you handle boundary or mask cases on separate code paths (e.g. only edge tiles are masked).
2. Explain occupancy. Is higher always better? Compute it for a kernel using 96 registers/thread and 256-thread blocks on H100.
Occupancy = resident warps / max warps per SM (64 on H100). Registers: 96 × 256 = 24,576 per block, and 65,536 / 24,576 = 2.67, so 2 blocks = 512 threads = 16 warps → 25% (ignoring register allocation granularity, which can round slightly differently). Higher occupancy helps hide latency when you rely on warp switching, as memory-bound kernels with simple per-thread work do. But it isn't always better. GEMM and attention kernels deliberately use many registers per thread (big accumulator tiles) and large SMEM tiles, run at low occupancy, and hide latency with ILP and async copy pipelines instead. The right question is whether enough work is in flight to saturate the bottleneck resource, and Nsight Compute's stall reasons tell you.
3. Sketch the H100 memory hierarchy with rough numbers.
Registers: 256 KB per SM (64K 32-bit), ~1 cycle. Shared memory/L1: 256 KB per SM, up to 228 KB as programmer-managed SMEM, tens of cycles latency, ~128 bytes/clock/SM. L2: 50 MB shared by all 132 SMs, a few hundred cycles, a few times HBM bandwidth. HBM3: 80 GB at ~3.35 TB/s, hundreds of cycles latency. Beyond that, NVLink to peers at 900 GB/s and PCIe to the host at ~64 GB/s per direction. The key ratio is that tensor cores can consume operands far faster than HBM supplies them (ridge ~295 FLOP/byte), so kernels must reuse data from registers and SMEM. Latency numbers are approximate and vary by measurement.
4. Define arithmetic intensity and the ridge point. Where do GEMM, elementwise ops and decode sit?
Arithmetic intensity is FLOPs per byte moved from a given memory level. Attainable throughput is \(\min(P_{peak}, B \cdot I)\), and the ridge \(P/B\) is where the two meet: about 295 FLOP/byte for H100 BF16 against HBM. Elementwise ops are at ~0.1–1 (hopelessly memory-bound, so fuse them). A square GEMM is \(n/3\) (≈1365 at 4096, compute-bound). A skinny GEMM with \(M\) tokens is ≈ \(M\), so decode at batch 1 sits at ~1 and needs a batch of ~150–300 to approach the ridge. Decode attention is ~1 (or ~GQA group size), regardless of batch, because each sequence has its own KV cache.
5. Estimate time per output token for a 13B BF16 model on one H100 at batch 1 and at batch 64.
Weights = 26 GB. Batch 1: 26 GB / 3.35 TB/s ≈ 7.8 ms lower bound (~128 tokens/s), plus KV-cache reads and overheads, so maybe 9–11 ms in practice. Batch 64: weight traffic is unchanged (7.8 ms), and compute is 2 × 13e9 × 64 ≈ 1.66 TFLOP → ~1.7 ms at peak, still below the memory time, so the step stays ~8 ms plus extra KV traffic (64 sequences' caches). Per-sequence latency barely moves while aggregate throughput rises ~60×. That's the core economics of batching in LLM serving. With long contexts, the KV-cache bytes at batch 64 can exceed the weight bytes and dominate.
6. What are the three regimes of deep learning performance, and how do you tell which you're in?
Compute-bound (tensor cores busy), memory-bandwidth-bound (HBM saturated), and overhead-bound (GPU idle while the CPU, Python, dispatcher or launch path catches up), following Horace He's framing. Diagnose with a profiler. Gaps on the nsys timeline mean overhead. ncu "speed of light" shows whether DRAM or compute throughput is near peak. Quick experiments: if doubling the batch barely changes step time, you're overhead-bound. If halving bytes (BF16 vs FP32) halves time, you're memory-bound. Fixes differ: CUDA graphs/compile/bigger batches for overhead, fusion and quantization for memory, better tiling and lower precision for compute.
7. Why does kernel fusion speed things up? Give a quantitative example.
Each unfused memory-bound op reads its inputs from and writes its output to HBM. A chain of \(k\) elementwise ops on an \(S\)-byte tensor costs about \(2kS\) bytes unfused and about \(2S\) fused, because intermediates stay in registers. For bias + GELU + dropout + residual on a 256 MB activation, that's ~2.3 GB versus ~0.77 GB (the residual adds one extra input; dropout's mask adds more if saved), roughly 700 µs versus 230 µs on H100. Fusing into a GEMM epilogue removes even the separate write and read of the GEMM output. Fusion also cuts kernel launches (overhead). Limits: reductions need the full row or an online recurrence, and GEMM→GEMM fusion requires something like the FlashAttention structure.
8. What do torch.compile and CUDA graphs each fix, and what are their pitfalls?
torch.compile (Dynamo capture + Inductor codegen) fixes memory-bound and overhead problems by fusing ops into generated Triton kernels and reducing Python dispatch. Pitfalls: compile time, graph breaks from untraceable code, recompilation on shape changes (use dynamic shapes or bucketing), and harder debugging. CUDA graphs fix launch overhead by recording a sequence of kernels once and replaying it with a single launch. They require static shapes and addresses and no host syncs inside, and they use extra memory for static buffers. Inference servers capture a graph per batch-size bucket for decode. mode="reduce-overhead" combines the two. Neither helps a kernel that's already compute-bound.
9. Walk through optimizing a GEMM kernel from naive to near-cuBLAS.
Naive: one thread per output, reading A and B from global memory each k-step (intensity ~0.25). Then: (1) coalesce loads. (2) Block tiling, where each CTA stages BM×BK and BK×BN tiles in SMEM for reuse factors of BN and BM (intensity \(BM\cdot BN/(BM+BN)\)). (3) Register tiling, where each thread or warp computes a micro-tile from SMEM fragments. (4) Vectorized 128-bit loads and swizzled, conflict-free SMEM layouts. (5) Multi-stage pipelining with cp.async (Ampere) or TMA (Hopper) so loads overlap MMAs. (6) Tensor cores (mma.sync → WGMMA → tcgen05). (7) On Hopper and later, warp specialization (producer TMA warp, consumer MMA warpgroups), persistent CTAs, and cluster multicast. (8) Epilogue fusion and L2-aware tile rasterization. (9) Autotune tile shapes per problem size, watching tile and wave quantization. Skinny shapes need split-K or Stream-K.
10. What is a shared-memory bank conflict? How do padding and swizzling fix it?
SMEM has 32 banks of 4 bytes. Address \(a\) lives in bank \((a/4) \bmod 32\). If multiple lanes of a warp access different addresses in the same bank, the accesses serialize (an n-way conflict costs n cycles). Reading a column of a row-major 32×32 float tile puts all 32 lanes in one bank. Padding each row to 33 floats offsets each row by one bank, so a column spans all banks. Swizzling permutes the column index (XOR with row bits) so both row and column patterns are conflict-free without wasting memory. Hardware paths like TMA and WGMMA expect specific swizzle modes, which is why CUTLASS/CuTe expresses layouts algebraically.
11. Derive the online softmax update used in FlashAttention and explain why the result is exact.
Keep per row a running max \(m\), denominator \(\ell\) and unnormalized output \(\tilde O\). For a new score block \(S\): \(m' = \max(m, \text{rowmax}(S))\), \(\tilde P = e^{S - m'}\), \(\ell' = e^{m-m'}\ell + \text{rowsum}(\tilde P)\), \(\tilde O' = e^{m-m'}\tilde O + \tilde P V\). At the end, \(O = \tilde O/\ell\). Exactness: all previous contributions were computed relative to \(e^{-m}\), and multiplying by \(e^{m-m'}\) re-bases them to \(e^{-m'}\). By the end, numerator and denominator share the true global max as reference, so the ratio equals softmax(S)V exactly (up to float rounding). Subtracting the max keeps exponentials ≤ 1, avoiding overflow. The backward pass stores only the LSE \(m + \log \ell\) per row.
12. FlashAttention doesn't reduce FLOPs, so why is it faster? What's its IO complexity?
Standard attention is memory-bound because it writes and reads the \(N\times N\) score and probability matrices through HBM: \(\Theta(Nd + N^2)\) accesses. FlashAttention fuses QKᵀ, softmax and PV into one kernel that keeps tiles in SRAM, giving \(\Theta(N^2 d^2/M)\) accesses for SRAM size \(M\), which is many times fewer for typical \(d\) and \(M\). Memory use becomes linear in \(N\). The backward pass recomputes \(S\) and \(P\) from Q, K and the saved LSE, which adds FLOPs but avoids storing \(N^2\) values. Since attention was bandwidth-bound, trading compute for bytes wins. It's exact, not an approximation.
13. What did FlashAttention-2 and -3 change, and why were those changes needed?
FA2: (a) defers the \(1/\ell\) normalization to the end to cut non-matmul FLOPs, which are ~16× more expensive than tensor-core FLOPs; (b) parallelizes over the sequence (Q blocks) as well as batch × heads, to fill the SMs at long context and small batch; (c) splits Q rather than K/V across warps, so no inter-warp reduction through SMEM. That gave ~2× over FA1 and 50–73% of A100 peak. FA3 targets Hopper, where FA2 got ~35%: it uses TMA and async WGMMA with warp specialization, ping-pong scheduling so one warpgroup's softmax overlaps another's GEMM (exp throughput is ~256× below matmul throughput), intra-warpgroup softmax/GEMM pipelining, and FP8 with block quantization and Hadamard-style incoherent processing. Reported: up to ~740 TFLOPS FP16 and ~1.2 PFLOPS FP8.
14. Does FlashAttention-4 exist? What's different on Blackwell?
Yes. It was released in the flash-attention repo in 2025 and described in an arXiv paper (2603.05451, March 2026) by Zadouri, Hoehnerbach, Shah, Liu, Thakkar and Dao. The motivation is asymmetric scaling: Blackwell roughly doubles tensor-core throughput while exp units and shared-memory bandwidth scale less, so non-matmul work and data movement dominate further. FA4 redesigns the pipeline around asynchronous tcgen05 MMAs with accumulators in tensor memory, emulates part of the exponential with polynomials on FMA units, skips most online-softmax rescales unless the max shift is large, and uses TMEM and 2-CTA MMA in the backward pass to cut SMEM traffic and atomics. It's written in CuTe DSL. Reported ~1613 TFLOPS BF16 on B200 (~71%). Details may have changed since, so check the repo.
15. Why is FlashAttention-2 inefficient for decode, and how does FlashDecoding fix it?
In decode each sequence has one query token, so FA2's parallelism (batch × heads × Q-blocks) collapses to batch × heads, e.g. 32 CTAs on 132 SMs at batch 1. Each CTA streams a long KV cache serially, leaving most SMs idle and HBM underused. FlashDecoding adds a split over the KV sequence: each CTA handles a chunk and emits a partial output plus its log-sum-exp, and a small reduction kernel combines them with weights \(e^{lse_s - lse}\). This is the same associative softmax merge, done across CTAs. It's split-K applied to attention, with reported attention speedups up to ~50× at 64k context and batch 1. Production kernels also pack GQA query heads together so each KV byte is used \(g\) times.
16. Compare CUDA C++, Triton and CuTe DSL / other tile DSLs. When would you choose each?
CUDA C++ is the thread-level model with full control (warp specialization, TMA, inline PTX), at the price of verbosity and bug surface. CUTLASS/CuTe give reusable building blocks for peak GEMM and attention. Triton is block-level: you write tile operations and the compiler handles thread mapping, SMEM, pipelining and tensor-core lowering. It's great for fused memory-bound kernels, custom attention variants and autotuning, it's what torch.compile emits, and it supports AMD. It sometimes trails hand-tuned kernels on the newest features. CuTe DSL brings CuTe's layout algebra to Python with fast compiles (FA4 uses it). ThunderKittens, TileLang, cuTile and Mojo offer different tile abstractions. Choose by need: library first, Triton for most custom work, CUTLASS/CuTe for the last 20% on new hardware.
17. You see a kernel in Nsight Compute at 15% of DRAM bandwidth and 10% of compute. What's going on, and what do you try?
Neither roof is reached, so it's latency- or parallelism-bound. Common causes: too few CTAs for the GPU (small grid, tail effects), low occupancy with no ILP or async copies, serialized dependent loads, heavy synchronization, or uncoalesced access wasting sectors (check sectors-per-request). Look at warp stall reasons ("long scoreboard" means waiting on memory, "barrier" means sync), achieved occupancy, and the launch configuration. Fixes: increase parallelism (split-K, smaller tiles, more blocks), raise bytes in flight (vector loads, multiple loads per thread, cp.async/TMA pipelines), fix coalescing, reduce register usage if occupancy is limited by registers. Very small kernels may simply be launch-overhead dominated and belong in a fused kernel or CUDA graph.
18. Why do we accumulate in FP32 even when inputs are BF16/FP8? What did DeepSeek-V3 discover about Hopper FP8?
Summing thousands of products in a low-precision accumulator loses small contributions once the running sum is large, because each addition rounds to few mantissa bits, so error grows with K. Tensor cores therefore multiply in low precision and accumulate in FP32. DeepSeek-V3 reported that H800 FP8 GEMMs accumulate with limited internal precision (around 14 bits), causing noticeable error for large K. They promoted partial results to FP32 on CUDA cores at fixed intervals along K (128 elements), combined with fine-grained scaling (1×128 activation tiles, 128×128 weight blocks). It's a good example of numerics shaping kernel design.
19. Explain E4M3 vs E5M2 and block-scaled FP4 (MXFP4 vs NVFP4).
E4M3 has 4 exponent and 3 mantissa bits: more precision, max ~448. It's used for weights and activations. E5M2 has more range (max ~57,344) but less precision, which suits gradients with their wide dynamic range. Both need scale factors. FP4 E2M1 can only represent a handful of values (±0, 0.5, 1, 1.5, 2, 3, 4, 6), so it's usable only with fine block scales. MXFP4 (OCP microscaling) shares one power-of-two E8M0 scale per 32 elements. NVFP4 uses 16-element blocks with an E4M3 (non-power-of-two) scale plus a per-tensor FP32 scale, which lowers quantization error. Blackwell tensor cores apply these block scales in hardware.
20. What is stochastic rounding and when does it matter?
Instead of rounding to the nearest representable value, round up with probability equal to the fractional distance to the upper neighbour. The expected rounded value then equals the true value, so the rounding is unbiased. It matters when you repeatedly add small increments to a large accumulator in low precision. With round-to-nearest, updates smaller than half an ULP vanish every time (e.g. BF16 weights with small learning-rate updates), and training stalls. With stochastic rounding they survive on average. The alternatives are FP32 master weights (the usual mixed-precision recipe) or Kahan-style compensated summation. Some accelerators implement it in hardware. On GPUs it's typically done in software inside optimizer or quantization kernels.
21. How do you overlap communication with computation in data-parallel training, and why isn't the overlap free?
DDP groups gradients into buckets and launches an all-reduce for each bucket as soon as its gradients are ready, on a separate stream, while backward continues on earlier layers. FSDP/ZeRO prefetches the next layer's parameter all-gather during the current layer's compute and reduce-scatters gradients as they're produced. It isn't free because NCCL collectives are GPU kernels that occupy SMs and consume HBM bandwidth, so concurrent GEMMs slow down. Overlap also requires enough independent compute to hide the transfer. Ring all-reduce moves \(2(n-1)/n\) times the buffer per GPU, so estimate comm time against backward compute time to see if it can be hidden. Remedies: tuning bucket sizes, limiting comm SMs, copy-engine transfers, or fused GEMM+collective kernels.
22. How does a TPU differ from a GPU for matmul, and what carries over?
A TPU core centres on large systolic arrays (MXUs: 128×128 on v5e/v5p, 256×256 on v6e), where operands flow through a grid of MAC cells and get reused in place, plus a vector unit and a compiler-managed VMEM scratchpad. There are no thousands of independent threads. XLA schedules everything statically, including fusion and layouts, and Pallas is the escape hatch for custom kernels. GPUs are many SMs running SIMT warps with tensor cores, programmer-managed shared memory and dynamic scheduling. What carries over is roofline thinking: arithmetic intensity, fusion to avoid HBM round trips, padding shapes to hardware tile multiples, and overlapping communication over the interconnect (ICI torus on TPU, NVLink/IB on GPU).
23. Design question: you need a custom fused attention variant (e.g. with a learned bias and sliding window) for training. How do you approach it?
First check whether existing tools already express it: FlexAttention (score_mod/mask_mod compiled into a fused Triton kernel with block sparsity) or FlashAttention options (sliding window, ALiBi, softcap). If not, write a Triton kernel based on the FlashAttention forward. Tile Q per program, loop over K/V blocks, apply the bias in-register before the online softmax, and skip fully masked blocks. Save the LSE. The backward pass recomputes scores and needs atomics or a separate pass for dK/dV. Validate against a naive FP32 reference across shapes, including ragged edges, with tolerances appropriate to BF16. Benchmark achieved TFLOPS against FA2/FA3 on the same shapes. Only move to CUTLASS/CuTe DSL if profiling shows you're far from the roofline on Hopper/Blackwell in ways Triton can't fix.
24. Your decode service at batch 8 shows the GPU at 40% "utilization" but throughput is poor. What's happening and what do you do?
"GPU utilization" from nvidia-smi only means a kernel was running. It says nothing about SM or bandwidth saturation. At batch 8, decode is memory-bound (intensity ~8) and likely overhead-bound too: many small kernels per layer with Python dispatch between them. Steps: take an nsys trace to check for gaps (if present, capture the decode step in CUDA graphs and use fused kernels for RMSNorm, rotary, attention and sampling), check that attention uses a split-KV decode kernel with paged KV, then raise batch size with continuous batching (throughput scales nearly linearly until KV bandwidth or memory runs out), quantize weights and KV to cut bytes, and consider speculative decoding to use idle compute. Report bandwidth utilization (bytes per step / step time ÷ 3.35 TB/s) as the real efficiency metric.
25. Why might a GEMM with N=4097 run much slower than N=4096?
Two shape effects. Tile quantization: with 128-wide tiles, 4097 needs 33 tiles in that dimension, and the last tile does 1 useful column out of 128, so roughly 3% wasted work, more for smaller shapes. Wave quantization: the total tile count may now slightly exceed a multiple of the SM count (132), adding an extra nearly empty wave, which can cost up to ~2× for shapes near the boundary. Odd sizes also break alignment requirements for vectorized/TMA loads (16-byte alignment), which can force slower code paths or kernels. Fixes: pad dimensions to multiples of 64/128 (common for vocab size), choose tile shapes per problem (autotuning), or use Stream-K-style scheduling that balances partial tiles across SMs.