Chapter 25 β CUDA & Kernel Development
βIf you want to understand why Flash Attention is fast, or why fused kernels matter, you need to understand how GPU programs work at the thread level.β
The CUDA Programming Model
CUDA organizes computation into a hierarchy:
Each thread knows its position via built-in variables:
threadIdx.xβ thread ID within the blockblockIdx.xβ block ID within the gridblockDim.xβ number of threads per block
A First Kernel: Vector Addition
1
2
3
4
5
6
7
8
9
10
// vector_add.cu β simplest possible CUDA kernel
__global__ void vector_add(float* a, float* b, float* c, int n) {
int idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx < n) {
c[idx] = a[idx] + b[idx];
}
}
// Launch: ceil(n/256) blocks, 256 threads each
// vector_add<<<(n+255)/256, 256>>>(a, b, c, n);
This kernel is memory-bound: each thread does 1 add but reads 2 floats and writes 1 float (12 bytes for 1 FLOP).
Why Fused Kernels Matter
Every separate kernel launch reads from and writes to HBM. Fusing operations into a single kernel keeps intermediate values in fast SRAM:
- Kernel 1: matmul β write to HBM
- Kernel 2: read from HBM β add bias β write to HBM
- Kernel 3: read from HBM β GELU β write to HBM
- 6 HBM reads/writes total
- Kernel: matmul β bias (in SRAM) β GELU (in SRAM) β write to HBM
- Only the final result hits HBM
- 2 reads + 1 write total
- 3Γ less memory traffic
This is why torch.compile and frameworks like Triton exist β they automatically fuse element-wise operations.
Tiled Matrix Multiplication
The naive matrix multiply reads each element from HBM many times. Tiling loads blocks into shared memory, reusing each element multiple times:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
// Simplified tiled matmul (C = A Γ B)
__global__ void matmul_tiled(float* A, float* B, float* C, int M, int K, int N) {
__shared__ float tileA[TILE_SIZE][TILE_SIZE];
__shared__ float tileB[TILE_SIZE][TILE_SIZE];
int row = blockIdx.y * TILE_SIZE + threadIdx.y;
int col = blockIdx.x * TILE_SIZE + threadIdx.x;
float sum = 0.0f;
for (int t = 0; t < K / TILE_SIZE; t++) {
// Load tiles from HBM to shared memory (collaborative)
tileA[threadIdx.y][threadIdx.x] = A[row * K + t * TILE_SIZE + threadIdx.x];
tileB[threadIdx.y][threadIdx.x] = B[(t * TILE_SIZE + threadIdx.y) * N + col];
__syncthreads();
// Compute partial dot product using shared memory
for (int k = 0; k < TILE_SIZE; k++) {
sum += tileA[threadIdx.y][k] * tileB[k][threadIdx.x];
}
__syncthreads();
}
C[row * N + col] = sum;
}
Each element of A and B is loaded from HBM once per tile but reused TILE_SIZE times in shared memory. This reduces HBM traffic by TILE_SIZEΓ.
Triton: GPU Kernels in Python
Triton (OpenAI) lets you write GPU kernels in Python, with automatic tiling and shared memory management:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
import triton
import triton.language as tl
import torch
@triton.jit
def softmax_kernel(input_ptr, output_ptr, n_cols, BLOCK_SIZE: tl.constexpr):
"""Fused online softmax in Triton."""
row_idx = tl.program_id(0)
col_offsets = tl.arange(0, BLOCK_SIZE)
mask = col_offsets < n_cols
# Load row from HBM
row = tl.load(input_ptr + row_idx * n_cols + col_offsets, mask=mask, other=-float('inf'))
# Compute softmax in SRAM (numerically stable)
row_max = tl.max(row, axis=0)
numerator = tl.exp(row - row_max)
denominator = tl.sum(numerator, axis=0)
softmax_output = numerator / denominator
# Write back to HBM
tl.store(output_ptr + row_idx * n_cols + col_offsets, softmax_output, mask=mask)
def triton_softmax(x):
n_rows, n_cols = x.shape
BLOCK_SIZE = triton.next_power_of_2(n_cols)
output = torch.empty_like(x)
softmax_kernel[(n_rows,)](x, output, n_cols, BLOCK_SIZE=BLOCK_SIZE)
return output
Flash Attention Kernel Design
Flash Attention (Chapter 4) is the most impactful fused kernel in LLM inference. Its key innovations:
Flash Attention reduces memory from O(TΒ²) to O(T) and is 2β4Γ faster than unfused attention on long sequences.
torch.compile
PyTorchβs compiler automatically fuses operations and generates optimized kernels:
1
2
3
4
5
6
7
8
import torch
model = MyTransformerBlock(d_model=4096, n_heads=32)
model = torch.compile(model, mode="reduce-overhead")
# First call triggers compilation (slow)
# Subsequent calls use compiled, fused kernels
output = model(input_tensor)
torch.compile with Inductor backend:
- Fuses element-wise operations (LayerNorm components, activations, residuals)
- Generates Triton kernels for fused operations
- Doesnβt replace Flash Attention (thatβs already hand-optimized)
- 10β30% speedup on typical transformer workloads
Whatβs Next
Single-GPU performance has limits. The next chapter covers distributed training β how to split work across hundreds or thousands of GPUs using data, tensor, and pipeline parallelism.
β Previous: Chapter 24 β GPU Architecture for ML Β· Next: Chapter 26 β Distributed Training β
Last updated: April 2026