Track A · Model internals

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.

Grid (whole kernel launch, e.g. 4096 blocks) └─ Cluster (Hopper+, optional: up to 8 blocks portable, 16 non-portable on some parts) └─ Thread block / CTA (≤1024 threads) ── pinned to ONE SM, owns a slice of shared memory └─ Warpgroup (Hopper: 4 warps = 128 threads, the unit WGMMA issues on) └─ Warp (32 threads, one instruction stream: SIMT) └─ Thread (own registers; lane id 0..31)

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:

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 80GBH100 SXMB200 (per GPU, DGX)
SMs108132148 (2 dies)
BF16/FP16 tensor TFLOPS312~989~2,250
FP8 tensor TFLOPSn/a~1,979~4,500
FP4 tensor TFLOPSn/an/a~9,000 (72 PF dense / 8 GPUs) NVIDIA DGX B200
FP32 non-tensor TFLOPS19.5~67~75–80
HBM capacity / bandwidth80 GB / ~2.0 TB/s80 GB / ~3.35 TB/s180 GB / ~8 TB/s (1,440 GB, 64 TB/s per 8) NVIDIA DGX B200
L2 cache40 MB50 MB~126 MB (GB200 per tuning guide) NVIDIA docs
BF16 ridge point (FLOP/byte)~153~295~280
NVLink per GPU600 GB/s900 GB/s1.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.

Intuition

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.

Interview angle

"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.

Go deeper

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.

Registers SMEM / L1 L2 cache HBM (global) Host DRAM / peers (PCIe, NVLink) 256 KB/SM (~33 MB total) · ~1 cycle private per thread; fastest operand source ≤228 KB/SM (256 KB with L1) · ~20–30 cycles per block; ~30 TB/s aggregate (approx) 50 MB, shared by all SMs · ~200+ cycles roughly 2–4× HBM bandwidth (approx) 80 GB · ~3.35 TB/s · ~400–800 cycles where weights, activations, KV cache live NVLink 900 GB/s · PCIe5 ~64 GB/s/dir faster, smaller, closer to tensor cores
H100-class memory hierarchy. Sizes are from NVIDIA; latencies and the SMEM/L2 bandwidths are rough, measured-ballpark numbers (see table).
LevelSize (H100 SXM)BandwidthLatencyManaged by
Registers64K × 32-bit = 256 KB per SMhundreds of TB/s aggregate~1 cycleCompiler (spills go to "local memory", which is actually HBM-backed)
Shared memory / L1256 KB per SM, up to 228 KB as SMEM~128 B/clk/SM ≈ 30 TB/s aggregate~20–30 cyclesProgrammer (SMEM) / hardware (L1)
L250 MB~2–4× HBM~200–300 cyclesHardware (with persistence hints)
HBM380 GB~3.35 TB/s~400–800 cyclesProgrammer (allocations)
NVLink (peer GPU)n/a900 GB/s total per GPUµs-scaleNCCL / NVSHMEM
PCIe Gen5 x16 (host)n/a~64 GB/s per directionµs-scalecudaMemcpy, pinned memory
Common mistake

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.

Interview angle

"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).

Go deeper

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).

0.1110 100100010⁴ arithmetic intensity (FLOP / byte of HBM traffic, log scale) 100010010 10.1 attainable TFLOP/s (log) FP8 roof ~1979 BF16 roof ~989 HBM roof: 3.35 TB/s × I ridge ≈ 295 add / GELU (~0.17–1) decode GEMV, B=1 (~1) LayerNorm (~2) decode, batch 64 (~64) FlashAttn prefill GEMM 4096³ (~1365)
Roofline for an H100 SXM (dense BF16 solid, FP8 dashed). Red = deeply memory-bound, orange = memory-bound but improvable by batching, green = compute-bound. Points sit on the roof for illustration; real kernels land below it.

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\)).

Op (BF16)Arithmetic intensityRegime on H100 (ridge ~295)Lever
Elementwise add/mul, activation0.1–2MemoryFuse into neighbours
LayerNorm/RMSNorm, softmax (standalone)~1–5MemoryFuse; one pass over data
Decode GEMV (batch \(B\))≈ \(B\)Memory until \(B\sim 150\)–\(300\)Batch; quantize weights
Decode attention≈ GQA group sizeMemoryGQA/MLA, KV quant, FlashDecoding
Naive attention prefill≈ \(d/2\)Memory (and \(O(N^2)\) memory)FlashAttention
FlashAttention prefill≈ \(N/2\)-ishCompute (for long \(N\))Overlap softmax with MMA (FA3/FA4)
Large GEMM \(n^3\)≈ \(n/3\)ComputeTile well, FP8, avoid wave quantization
Intuition

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.

Common mistake

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.

Interview angle

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.

Go deeper

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.

RegimeSymptomHow to confirmFixes
Compute-boundTensor-core utilization high; time scales with FLOPsNsight Compute "SM / tensor pipe" throughput near peak; achieved TFLOPS close to specLower precision (BF16 → FP8 → FP4), better tile shapes, avoid wave/tile quantization, reduce FLOPs (algorithmic: GQA, MoE, sparsity)
Memory-boundDRAM throughput near 3 TB/s; time scales with bytesncu "Memory throughput" % of peak high, compute low; roofline places kernel on sloped partFuse ops, avoid materializing intermediates, quantize weights/KV, increase batch (more reuse), better layouts
Overhead-boundGPU idle between kernels; doubling batch barely changes step timeNsight 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).

Interview angle

"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:

  1. 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()).
  2. 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.

Common mistake

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.

Go deeper

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.

HBM: A (M×K), B (K×N) C (M×N) │ each CTA owns one BM×BN output tile, marches over K in steps of BK ▼ SMEM: A_tile[BM×BK], B_tile[BK×BN] (×2–4 stages: ring buffer) e.g. BM=128, BN=256, BK=64 │ each warp(group) owns a WM×WN sub-tile ▼ REGS: fragments of A, B ──► TENSOR CORE MMA ──► accumulator tile (FP32) in registers / TMEM │ after the K loop: epilogue (scale, bias, activation, cast to BF16) fused here ▼ HBM: C tile written once (ideally via TMA store)

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.

stage ring in SMEM: [ s0 ][ s1 ][ s2 ][ s3 ] ▲ full_bar[i] ▼ empty_bar[i] Producer warp : wait empty[i] → TMA load A,B tile k → (TMA arrives on full[i]) Consumer WG 0 : wait full[i] → WGMMA (async) → wait MMA done → arrive empty[i] Consumer WG 1 : (ping-pong: runs its MMAs while WG 0 does its epilogue / softmax) Epilogue : regs → (SMEM) → TMA store of C tile, overlapped with next tile's mainloop

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.

GenerationLoad pathMMAAccumulatorKey new tricks
Volta/TuringLDG → regs → STSwmma/mma.sync (warp)RegistersFirst tensor cores
Ampere (A100)cp.async global → SMEMmma.sync (warp)RegistersMulti-stage async pipelines, ldmatrix, BF16/TF32
Hopper (H100)TMA (bulk, multicast)WGMMA (warpgroup, async, operands from SMEM)RegistersWarp specialization, clusters/DSMEM, FP8
Blackwell (B200)TMAtcgen05.mma (single thread issues, 2-CTA mode)Tensor memory (TMEM)FP4/FP6 + block-scaled formats, frees registers for epilogue/softmax
Interview angle

"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.

Go deeper

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\)).

Q (N×d) Qᵢ outer loop: one CTA per Qᵢ (v2+) Kᵀ (d×N) Kⱼᵀ Sᵢⱼ S = QKᵀ (N×N) never written to HBM inner loop j V (N×d) Vⱼ On-chip (SMEM / regs / TMEM) Qᵢ (resident whole loop) Kⱼ, Vⱼ (streamed, 2+ stages) Sᵢⱼ = Qᵢ Kⱼᵀ (MMA #1) m, ℓ ← online update P̃ = exp(Sᵢⱼ − m) Õ ← α·Õ + P̃ Vⱼ (MMA #2) end: Oᵢ = Õ/ℓ, LSE → HBM
FlashAttention forward. Each CTA owns a query block \(Q_i\) and streams K/V blocks through SRAM. Only \(O_i\) and the per-row log-sum-exp return to HBM.

FlashAttention (v1, 2022): tiling + recomputation + IO analysis

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:

  1. 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.
  2. 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.
  3. 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:

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:

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.

May be out of date

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.

VersionTargetCore ideaHeadline result (as reported)
FA1 (2022)A100Tiling + online softmax + backward recomputation; IO-optimalLinear memory; ~2–4× faster attention
FA2 (2023)A100Fewer non-matmul FLOPs; parallelize over seq; split Q not K across warps~2× FA1; 50–73% of peak
FA3 (2024)H100TMA/WGMMA, warp specialization, ping-pong softmax/GEMM overlap, FP8~740 TFLOPS FP16, ~1.2 PF FP8
FA4 (2025–26)B200Async 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.

decode, 1 query × 32k cached tokens, split into 8 chunks: CTA 0: q · K[0:4k] → (O₀, lse₀) ┐ CTA 1: q · K[4k:8k] → (O₁, lse₁) │ ... ├──► reduce: lse = logΣe^lseₛ ; O = Σ e^(lseₛ−lse)·Oₛ CTA 7: q · K[28k:32k]→ (O₇, lse₇) ┘ (× batch × kv-heads CTAs total, enough to fill 132 SMs)

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.

Interview angle

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.)

Go deeper

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:

ToolLevelWhat it's forStatus (as of late 2026, verify)
CUDA C++ + CUTLASS/CuTeThread/warp + layout algebraPeak GEMM/attention, new hardware firstMature; the reference for NVIDIA
TritonBlockFused elementwise/reduction/matmul/attention; Inductor backendMature, NVIDIA + AMD; lower-level "Gluon" dialect exposes more control Triton repo
CuTe DSL (CUTLASS Python)CuTe layouts in PythonHopper/Blackwell peak kernels with fast compileReleased 2025; used by FA4 NVIDIA docs
cuTile / CUDA TileTileNVIDIA's own tile model aimed at tensor-core portabilityPython available, C++ planned NVIDIA; newer, adoption still forming
ThunderKittensWarp/tile primitives in C++Small library of 16×16 tile abstractions; simple, fast attention/GEMM kernelsResearch (Stanford Hazy Research) Spector+ 2024
TileLangTile, TVM-basedSeparates dataflow from scheduling annotations; used for e.g. MLA/sparse-attention kernelsOpen source, active Wang+ 2025
Mojo (Modular)Systems language, Python-likePortable GPU kernels (NVIDIA/AMD) inside the MAX stackCommercial/open components; evolving Modular
Pallas (JAX)BlockCustom kernels for TPU (and GPU via Triton/Mosaic)Part of JAX JAX docs
NKITileKernels for AWS Trainium/InferentiaAWS Neuron SDK AWS docs
May be out of date

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.

Interview angle

"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.

Go deeper

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 seeLikely meaningNext step
GPU idle gaps between short kernels in nsysOverhead-bound (Python/dispatch/launch)CUDA graphs, torch.compile, larger batch
Long cudaStreamSynchronize / .item() on the CPUHidden host–device syncs serialize the pipelineRemove syncs from the hot loop; keep values on device
ncu: DRAM throughput ~90%, compute ~10%Memory-bound and already efficientOnly fusion or fewer bytes (quantization) will help
ncu: both throughputs lowLatency-bound: low occupancy, poor ILP, too few CTAs, tail effectMore parallelism (split-K), more bytes in flight, check stall reasons
Stall reason "long scoreboard"Waiting on global memory loadsPrefetch/pipeline, async copies
Stall reason "MIO throttle" / "short scoreboard"; bank-conflict counters highShared-memory pressureSwizzle/pad layouts, fewer SMEM round trips
NCCL kernels not overlapping GEMMsExposed communicationBucketing, separate streams, overlap scheduling (FSDP prefetch)
Kernel uses 255 registers, local-memory trafficRegister spillsSmaller tiles, fewer live values, launch_bounds/maxnreg
Interview angle

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

FormatBits (sign/exp/mantissa)Max normal (approx)Typical use
FP321/8/233.4e38Master weights, optimizer states, accumulators, reductions
TF321/8/10 (19 bits used)3.4e38Tensor-core input mode for "FP32" matmuls (Ampere+); torch.backends.cuda.matmul.allow_tf32
BF161/8/73.4e38Default training/inference dtype; same range as FP32, so usually no loss scaling
FP161/5/1065,504More precision, little range; needs loss scaling in training Micikevicius+ 2017
FP8 E4M31/4/3448Weights/activations (forward) Micikevicius+ 2022
FP8 E5M21/5/257,344Gradients (need range more than precision)
FP4 E2M11/2/16Only 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

Intuition

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×.

May be out of date

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.

Go deeper

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.

Interview angle

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 TPUAWS Trainium
Compute unitCUs 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 BookNeuronCores with tensor, vector, scalar and general-purpose engines
On-chip memoryLDS (shared memory) per CU, large L2 + Infinity CacheVMEM scratchpad, compiler-managedSBUF/PSUM scratchpads, software-managed
HBM (example)MI300X: 192 GB, ~5.3 TB/s AMDv5e ~820 GB/s, v5p ~2.8 TB/s per chip Scaling BookVaries by generation
ProgrammingROCm; HIP (CUDA-like, hipify), Composable Kernel, Triton backend, rocBLAS/hipBLASLt AMD ROCmXLA compiler (whole-graph fusion and layout), JAX/PyTorch-XLA; Pallas for custom kernels OpenXLANeuron compiler; NKI for custom kernels
InterconnectInfinity FabricICI torus (2D/3D), optical reconfigurable on v4+ Jouppi+ 2023NeuronLink

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.

May be out of date

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.