From 1% to 94% of cuBLAS, one kernel at a time
A visual walkthrough of Simon Boehm's CUDA matrix multiplication worklog. Each kernel fixes one specific bottleneck. This page explains what that bottleneck is, what the fix looks like on the hardware, and why the number moves. It starts from zero: threads, warps, blocks, grids, and the memory hierarchy are all explained before the first kernel.
Source article: siboehm.com/articles/22/CUDA-MMM (December 2022). Code: github.com/siboehm/SGEMM_CUDA. All throughput numbers are the article's measurements on an NVIDIA RTX A6000 multiplying two 4092 x 4092 fp32 matrices. Diagrams and text on this page are original; code sketches are simplified rewrites, not the repository's code.
The shape of the climb is the lesson. The first three kernels are about not wasting memory bandwidth: reading global memory in the pattern the hardware wants, then caching in shared memory. Kernels 4 to 6 are about arithmetic intensity: doing more multiply-adds per byte loaded, by giving each thread a tile of outputs and keeping operands in registers. Kernels 9 and 10 are about matching the hardware's real structure: tuning tile sizes per GPU and organising work around warps, the unit the scheduler actually sees.
If you already know what a warp is and how shared memory differs from global memory, skip to the problem. Otherwise the primer below defines every term used later. The final section maps every level of this hierarchy onto what Triton does for you, and shows where the flash-linear-attention (FLA) kernels sit on top of it.
Primer: the vocabulary
CUDA describes a computation in a hierarchy that exists for correctness (grid, block, thread) and a hardware hierarchy that exists for performance (SM, warp scheduler, warp, lane). The two overlap but are not the same thing. Most of the optimisations later come from understanding where they differ.
Kernel, thread, block, grid
Kernel
A function that runs on the GPU (the device). The CPU (the host) launches it with f<<<gridDim, blockDim>>>(args). The launch returns immediately; the GPU executes it asynchronously. One launch creates one grid.
Thread
The smallest unit. Every thread runs the same kernel code but with its own values of threadIdx and blockIdx, its own registers and its own program counter. Kernel code is written from the point of view of a single thread.
Block (thread block)
A group of threads that are guaranteed to run on the same SM at the same time. Threads in one block can share data through shared memory and can synchronise with __syncthreads(). Threads in different blocks cannot cheaply talk to each other.
Grid
All the blocks of one launch. Blocks are independent and the hardware may run them in any order, on any SM. Both gridDim and blockDim are 3D vectors (x, y, z), though most kernels use one or two dimensions.
Inside a kernel, a thread computes its global position from the built-ins. For a 2D layout the usual formula is x = blockIdx.x * blockDim.x + threadIdx.x and likewise for y. The value of threadIdx.x ranges from 0 to blockDim.x - 1; blockIdx.x ranges from 0 to gridDim.x - 1.
Warp: the unit the hardware actually executes
A warp is a group of 32 threads that the hardware executes together, in lockstep, one instruction at a time (this is the SIMT model, single instruction multiple threads). Warps do not appear anywhere in CUDA source code. They are formed by taking the threads of a block in order of their linear thread id and cutting every 32:
linearId = threadIdx.x + blockDim.x * (threadIdx.y + blockDim.y * threadIdx.z)
warpId = linearId / 32 // which warp in the block
laneId = linearId % 32 // position inside the warp, 0..31
Because threadIdx.x is the fastest-varying dimension, threads with consecutive threadIdx.x land in the same warp. This one fact decides which memory accesses get combined (kernel 2) and which shared memory accesses collide (kernels 7 and 8).
Three consequences of the warp being the execution unit:
- A memory instruction issued by a warp produces 32 addresses at once. The memory system services them as a set, and how many transactions it needs depends on how those 32 addresses are laid out. This is coalescing.
- If threads in a warp take different branches, the warp executes both paths with some lanes masked off. This is divergence, and it wastes issue slots.
- Scheduling, register allocation and latency hiding all happen per warp, not per thread.
Streaming multiprocessor (SM) and warp schedulers
The GPU is a collection of SMs. The RTX A6000 has 84 of them. Each SM holds a register file, an on-chip data store split between L1 cache and shared memory, and four warp schedulers. Every cycle each scheduler looks at the warps assigned to it, picks one that is ready (its operands have arrived), and issues that warp's next instruction to the execution units. When a warp is waiting on memory, the scheduler simply picks another warp. This is how GPUs hide latency: not with big caches and out-of-order execution like a CPU, but by keeping many warps in flight and switching between them for free.
Relevant limits on the A6000 (compute capability 8.6), used in the calculations later:
Occupancy is the number of warps resident on an SM divided by the maximum (48). It is limited by whichever of registers, shared memory or thread count runs out first. Higher occupancy gives the scheduler more warps to switch between, which helps hide latency, but it is not the only way to hide latency: a single warp with lots of independent instructions (instruction-level parallelism, ILP) also keeps the pipeline busy. The later kernels deliberately trade occupancy for ILP and register reuse.
Memory hierarchy
Registers
Scalar variables in kernel code live here. Fixed-size arrays with compile-time indices can too (float acc[8] becomes 8 registers). A thread can use up to 255. More registers per thread means fewer warps fit on the SM.
Shared memory (SMEM)
Programmer-managed scratchpad. Roughly 16x the bandwidth of global memory and far lower latency. Physically the same silicon as L1, split by configuration. Divided into 32 banks of 4-byte words; a warp instruction that hits the same bank at different addresses is serialised (a bank conflict).
L1 and L2 cache
L1 sits in each SM next to shared memory. L2 is shared by the whole chip (6 MB on the A6000). Both cache global memory transparently. The naive kernel already gets most of its data from L1, which is why the gain from shared memory alone (kernel 3) is smaller than you might expect.
Global memory (GMEM)
The 48 GB on the card. Highest latency (hundreds of cycles) and lowest bandwidth. Accessed in 32-byte sectors; a 128-byte line is four sectors. All inputs start here and the output must end here.
Row-major layout and strides
Matrices are stored as one flat array. Row-major means element (row, col) of an M x N matrix is at index row * N + col: consecutive columns in a row are neighbours in memory, consecutive rows are N floats apart. That gap (the stride) is why reading down a column is expensive and reading along a row is cheap.
FLOPs, arithmetic intensity and the roofline
FLOP and FMA
One floating point operation. A fused multiply-add d = a*b + c is one instruction (FFMA in SASS) but counts as two FLOPs. A matmul of size M x N x K does 2MNK FLOPs.
Arithmetic intensity (AI)
FLOPs performed per byte moved across some memory boundary. Raising it is the theme of kernels 3 to 6: do more math per byte loaded from global memory, then per byte loaded from shared memory.
Roofline
Attainable FLOP/s is at most the smaller of peak compute and (AI x memory bandwidth). At low AI the kernel is memory-bound and sits on the sloped part; at high AI it is compute-bound and sits under the flat roof.
PTX and SASS
PTX is NVIDIA's virtual assembly; SASS is the real machine code for one GPU generation. Instruction names used later: LDG (load global), STG (store global), LDS (load shared), FFMA (fp32 fused multiply-add). A .128 suffix means one instruction moves 128 bits (four floats).
The problem: SGEMM
SGEMM computes C = alpha * A @ B + beta * C in single precision, with A of size M x K, B of size K x N and C of size M x N. Each output element is a dot product of one row of A with one column of B, of length K. The article benchmarks M = N = K = 4092.
Napkin math for 4092 x 4092
So if a kernel moved only the minimum data it would be compute-bound by a factor of 13. Put differently: a kernel can afford to read global memory up to about 13 times more than the minimum and still be compute-bound. That is the budget the whole worklog is spent inside. cuBLAS reads about 500 MB, less than twice the minimum. The naive kernel, with zero caching, would read 548 GB.
Naive kernel
One thread per output element, straight from global memory.
The most direct mapping of the problem onto the thread hierarchy: launch a grid of 32 x 32 blocks, one block per 32 x 32 tile of C, one thread per element. Each thread runs the K-long dot product on its own and writes one value. Nothing is shared, so no synchronisation is needed.
// host side
dim3 grid(ceil_div(M, 32), ceil_div(N, 32));
dim3 block(32, 32); // 1024 threads
k1_naive<<<grid, block>>>(M, N, K, alpha, A, B, beta, C);
// device side, from one thread's point of view
__global__ void k1_naive(int M, int N, int K, float alpha,
const float *A, const float *B, float beta, float *C) {
int row = blockIdx.x * blockDim.x + threadIdx.x; // note: x picks the row
int col = blockIdx.y * blockDim.y + threadIdx.y;
if (row < M && col < N) {
float acc = 0.f;
for (int k = 0; k < K; ++k)
acc += A[row * K + k] * B[k * N + col];
C[row * N + col] = alpha * acc + beta * C[row * N + col];
}
}
Why it is slow
Look at which threads form a warp. With blockDim = (32, 32), a warp is 32 threads with the same threadIdx.y and threadIdx.x = 0..31. In this kernel threadIdx.x selects the row. So in a single iteration k, the 32 lanes of a warp read A[row*K + k] for 32 different rows: 32 floats that are each 4092 x 4 = 16368 bytes apart. The memory system has to fetch 32 separate 32-byte sectors to deliver 128 bytes of useful data. Meanwhile all 32 lanes read the same B[k*N + col], which is at least cheap (one broadcast).
The profiler shows the consequence: about 15 GB/s of global memory throughput on a card capable of 768 GB/s. The next kernel changes nothing except which thread owns which element.
Global memory coalescing
Make the 32 lanes of a warp read 32 neighbouring floats.
When a warp issues a load, the hardware collects the 32 addresses and merges those that fall in the same 32-byte sector (or 128-byte line) into one transaction. If the 32 lanes read 32 consecutive, aligned floats, that is 128 bytes served by a single 128-byte request. If they read 32 floats scattered across 32 different lines, it is 32 requests for the same amount of useful data. This merging is called coalescing and it is done at run time by the hardware, not by the compiler: the two kernels compile to identical SASS.
The fix is to flip which index the fast-varying thread id controls. Make the block one-dimensional (1024 threads) and derive row and column arithmetically so that consecutive threads get consecutive columns:
dim3 block(32 * 32); // same thread count, now 1D
int row = blockIdx.x * 32 + threadIdx.x / 32; // changes every 32 threads
int col = blockIdx.y * 32 + threadIdx.x % 32; // changes every thread
// dot product loop is unchanged
Now a warp is 32 threads with the same row and 32 consecutive columns. Per iteration k: every lane reads the same A[row*K + k] (one broadcast), and the lanes read B[k*N + col .. col+31], 128 contiguous bytes, one transaction.
Global memory throughput rises from 15 GB/s to 110 GB/s and the kernel gets 6.4x faster. Note that alignment also matters: the 32 addresses must start on a sector boundary for the minimum transaction count. Also note that lanes do not have to access addresses in lane order; any permutation of the same 32 consecutive addresses coalesces just as well.
Shared memory cache blocking
Load a chunk once per block, reuse it 32 times from on-chip memory.
A block computing a 32 x 32 tile of C needs a 32 x K slab of A and a K x 32 slab of B. Every element of that slab of A is used by 32 threads in the block (one per column of the tile). So instead of each thread reading from global memory independently, the block cooperates: it walks along K in chunks of BK = 32, and for each chunk all 1024 threads together copy a 32 x 32 piece of A and a 32 x 32 piece of B into shared memory, one element per thread. After a barrier, each thread does 32 multiply-adds against the shared copies, then the block moves to the next chunk.
__shared__ float As[32 * 32];
__shared__ float Bs[32 * 32];
int tRow = threadIdx.x / 32, tCol = threadIdx.x % 32; // tCol is the fast index
// move base pointers to this block's tile
A += blockIdx.x * 32 * K; B += blockIdx.y * 32; C += blockIdx.x * 32 * N + blockIdx.y * 32;
float acc = 0.f;
for (int k0 = 0; k0 < K; k0 += 32) {
As[tRow * 32 + tCol] = A[tRow * K + tCol]; // coalesced: tCol varies fastest
Bs[tRow * 32 + tCol] = B[tRow * N + tCol];
__syncthreads(); // wait until the whole chunk is in SMEM
A += 32; B += 32 * N; // advance to next chunk
for (int k = 0; k < 32; ++k)
acc += As[tRow * 32 + k] * Bs[k * 32 + tCol];
__syncthreads(); // nobody overwrites SMEM while others still read
}
C[tRow * N + tCol] = alpha * acc + beta * C[tRow * N + tCol];
Two barriers per chunk are needed. The first guarantees the chunk is fully written before anyone reads it. The second guarantees everyone has finished reading before a fast thread starts overwriting the buffer with the next chunk.
Occupancy check for this kernel
This kernel uses 8 KB of shared memory per block (two 32 x 32 float arrays), 37 registers per thread and 1024 threads per block. Feeding that into the SM limits:
- Shared memory: (8192 + 1024 bytes runtime overhead) per block into 102400 bytes per SM allows 11 blocks.
- Threads: 1024 per block into 1536 per SM allows only 1 block.
- Registers: 37 x 32 = 1184 per warp, rounded up to the 256 allocation unit gives 1280; x 32 warps = 40960 per block; 65536 / 40960 allows 1 block.
One block per SM, 32 warps out of 48 possible: 67% occupancy. That is decent, so occupancy is not the bottleneck. You can check other configurations in the calculator below.
The real bottleneck: shared memory instruction pressure
The inner loop compiles to two shared loads and one FMA:
ld.shared.f32 %f1, [As + ...];
ld.shared.f32 %f2, [Bs + ...];
fma.rn.f32 %f3, %f2, %f1, %f3;
Two memory instructions per math instruction. The profiler's warp-state sampling shows warps stalled on MIO throttle: the queue that feeds shared memory instructions is full. The kernel is compute-bound on paper but in practice it is bound by how fast it can issue LDS instructions. The fix is not more caching; it is fewer loads per FMA, which means each thread must produce more than one output so that loaded values can be reused from registers.
1D blocktiling
Each thread computes a column of 8 outputs and reuses one B value 8 times.
New tile shape: a block now owns a BM x BN = 64 x 64 tile of C and walks K in chunks of BK = 8. Each thread computes TM = 8 vertically adjacent outputs, so the block needs 64 x 64 / 8 = 512 threads. Shared memory holds a 64 x 8 chunk of A and an 8 x 64 chunk of B: 1024 floats, 4 KB.
The important change is the inner loop. For each k in the chunk, the thread loads one value from Bs (its column) into a register and then reuses it for all 8 of its rows, loading one As value per row:
float acc[TM] = {0}; // 8 accumulators, live in registers
for (int k0 = 0; k0 < K; k0 += BK) {
// cooperative GMEM -> SMEM copy of the 64x8 and 8x64 chunks (one float each)
As[aRow * BK + aCol] = A[aRow * K + aCol];
Bs[bRow * BN + bCol] = B[bRow * N + bCol];
__syncthreads();
A += BK; B += BK * N;
for (int k = 0; k < BK; ++k) {
float b = Bs[k * BN + tCol]; // loaded once ...
for (int i = 0; i < TM; ++i)
acc[i] += As[(tRow * TM + i) * BK + k] * b; // ... used 8 times
}
__syncthreads();
}
for (int i = 0; i < TM; ++i)
C[(tRow * TM + i) * N + tCol] = alpha * acc[i] + beta * C[(tRow * TM + i) * N + tCol];
Counting loads per output
Per k in the chunk a thread issues 1 + TM = 9 shared loads and does TM = 8 FMAs. That is 1.125 loads per FMA instead of 2. Per output element over the full K: K/32 global loads and 9K/8 shared loads, versus K/16 and 2K in kernel 3. The warp-stall profile confirms far fewer cycles lost to the memory pipeline, and throughput almost triples.
Sidenote: the compiler would have done this anyway
If the two inner loops are written in the natural order (outputs outer, k inner) with no explicit b register, the generated code is just as fast. Both loop counts are compile-time constants, so nvcc fully unrolls them, notices that the same Bs element is loaded 8 times, and keeps it in a register. With this loop order it also emits 128-bit LDS.128 for the As loads, because for one output row the eight k values sit side by side in memory. The Bs loads stay 32-bit, since consecutive k are BN floats apart. In kernel 5 the k loop moves outermost and the situation flips: the Bs slice a thread needs is contiguous and the As slice is strided, which is what kernel 6 fixes by transposing As.
2D blocktiling
Each thread computes an 8 x 8 square as an outer product held entirely in registers.
Tile constants become BM = BN = 128, BK = 8, TM = TN = 8. A block of 128 x 128 / 64 = 256 threads owns a 128 x 128 tile of C; each thread owns an 8 x 8 square of it, 64 accumulators in registers.
Loading the chunk into shared memory now takes several trips per thread: As is 128 x 8 = 1024 floats and Bs is 8 x 128 = 1024 floats, so each of the 256 threads copies 4 floats of each. The copy loop strides through the chunk so that, at every step, the 32 lanes of a warp still read 32 contiguous floats from global memory (coalesced).
The inner loop is the interesting part. For each k in the chunk, a thread copies the 8 As values it needs (a column slice) and the 8 Bs values it needs (a row slice) into registers, then does a rank-1 update of its 8 x 8 accumulator block: 64 FMAs from 16 loads.
float acc[TM * TN] = {0}; // 64 accumulators
float regA[TM], regB[TN]; // 16 operand registers
for (int k = 0; k < BK; ++k) {
for (int i = 0; i < TM; ++i) regA[i] = As[(tRow * TM + i) * BK + k];
for (int j = 0; j < TN; ++j) regB[j] = Bs[k * BN + tCol * TN + j];
for (int i = 0; i < TM; ++i)
for (int j = 0; j < TN; ++j)
acc[i * TN + j] += regA[i] * regB[j]; // outer product, all in registers
}
Why a square beats a column
A thread computing an r x c block of outputs needs r + c loads per k for r x c FMAs. For a column (c = 1) the ratio is (r+1)/r, which can never get below 1. For a square it is 2/r, which halves every time r doubles. That is the whole reason 2D tiling exists.
Counting again
Per output element over the full K: K/64 global loads and K/4 shared loads. Compared with kernel 4 that is another 2x fewer global loads and 4.5x fewer shared loads. Throughput doubles to 16 TFLOP/s. The block-level arithmetic intensity (FLOPs per byte moved from global memory into shared memory) is BM x BN / (2 (BM + BN)) = 32 FLOP per byte for a 128 x 128 tile, versus 8 for the 32 x 32 tile of kernel 3.
Vectorised memory access
Transpose As so both operand slices are contiguous, and move global data as float4.
Part 1: transpose As in shared memory
In kernel 5, the 8 values a thread needs from As for one k are As[(tRow*8 + i) * BK + k] for i = 0..7: eight floats spaced BK apart. A single load instruction cannot fetch them. If As is stored transposed, as As[k * BM + row], those eight values become As[k * BM + tRow*8 .. tRow*8+7]: contiguous, so two LDS.128 instructions replace eight LDS.32. The transpose is free: it is done by writing each element to the swapped position during the global-to-shared copy, which happens anyway.
Part 2: 128-bit global loads and stores
Each thread's share of the copy is four consecutive floats. Read them as one float4:
float4 a4 = reinterpret_cast<const float4*>(&A[aRow * K + aCol * 4])[0]; // LDG.E.128
As[(aCol * 4 + 0) * BM + aRow] = a4.x; // scatter into the transposed layout
As[(aCol * 4 + 1) * BM + aRow] = a4.y;
As[(aCol * 4 + 2) * BM + aRow] = a4.z;
As[(aCol * 4 + 3) * BM + aRow] = a4.w;
reinterpret_cast<float4*>(&Bs[bRow * BN + bCol * 4])[0] =
reinterpret_cast<const float4*>(&B[bRow * N + bCol * 4])[0]; // LDG.E.128 then STS.128
Why does the cast matter, when the compiler could have merged four adjacent 32-bit loads itself? Because a 128-bit load requires the address to be 16-byte aligned, and the compiler cannot prove that about a float* passed in as an argument. The reinterpret_cast to float4* is the programmer's promise that the pointer is aligned. Shared memory needs no such promise; the compiler owns its layout and vectorises those loads on its own.
Together the two changes bring another 2.3 TFLOP/s. The profiler now lists three things: shared memory bank conflicts, occupancy higher than needed, and no overlap between loading the next chunk and computing on the current one.
Bank conflicts (the two kernels that were dropped)
They removed the conflicts and still ran slower, so the article skips them. The concept still matters.
Shared memory is built from 32 banks, each 4 bytes wide, interleaved: byte address a lives in bank (a / 4) % 32. In one cycle each bank can serve one 32-bit word. When a warp issues a shared load, the 32 lanes' addresses are grouped by bank. If two lanes want different words from the same bank, the hardware replays the instruction for the second one. This is an n-way bank conflict, and it multiplies the cost of that instruction by n. Lanes reading the same word are fine (broadcast).
The outer-product kernels read As and Bs in patterns that depend on the tile sizes and on TM, TN, and those patterns produce conflicts. Kernels 7 and 8 rearranged the shared memory layout so that a warp's accesses land in 32 distinct banks. They succeeded at that but the extra index arithmetic and the changed access pattern cost more than the conflicts did, so the article moved on. cuBLAS does avoid conflicts, which is part of the remaining gap.
Autotuning
Same kernel, search the five tile parameters instead of guessing them.
By now the kernel has five template parameters. BM, BN, BK set how much of A and B is cached in shared memory per step. TM, TN set how much of that a thread pulls into registers. They interact with every hardware limit at once: shared memory capacity, register file size, occupancy, coalescing, the float4 copy loop, bank conflict patterns. Nobody can reason their way to the optimum, so the article does what every production library does: enumerate the sensible configurations and time them.
What "sensible" means
- The thread count is fixed by the tile: (BM / TM) x (BN / TN) threads, which must be a multiple of 32 and at most 1024.
- Each thread copies one float4 per pass of the load loop, so BM x BK and BN x BK must be divisible by 4 x threads.
- Shared memory (BM + BN) x BK x 4 bytes must fit, and registers per thread (TM x TN accumulators plus TM + TN operands plus addressing) must leave room for enough warps.
- TM and TN should be multiples of 4 so the register slices load as 128-bit vectors.
About 400 configurations survived those filters and were benchmarked with a script. On the A6000 the winner was BM = BN = 128, BK = 16, TM = TN = 8: only BK changed from kernel 6, for a gain of about 8% in the benchmark table. On an A100 the winner was BM = BN = 64, BK = 16, TM = TN = 4, and running the A6000's best configuration there would have left 6% on the table. The optimum is a property of the GPU, not of the algorithm, which is why compilers like Triton ship an autotuner and why cuBLAS ships hundreds of pre-tuned kernels.
Warptiling
Insert the warp as an explicit tiling level between block and thread.
So far the loop structure has two tiling levels: the block owns BM x BN and each thread owns TM x TN, with threads laid out in a plain row-major grid over the block tile. That layout ignores the fact that the 32 threads of a warp are scheduled together, share a register cache, and conflict with each other (and only each other) in shared memory. Warptiling adds a level: the block tile is divided among warps, each warp tile is divided among its 32 lanes, and each lane still computes TM x TN squares.
The new parameters
threadIdx.x / 32 and threadIdx.x % 32; the lane's row and column in the sub-tile are laneId / (WSUBN / TN) and laneId % (WSUBN / TN)Each thread therefore accumulates WMITER x WNITER squares of TM x TN. For each k in the chunk it loads WMITER x TM values from As and WNITER x TN values from Bs into registers, then updates all its squares with outer products. The k loop is kept outermost inside the chunk, so everything inside it is independent work the scheduler can overlap.
for (int k = 0; k < BK; ++k) {
for (int wm = 0; wm < WMITER; ++wm) // A slices for each sub-tile row
for (int i = 0; i < TM; ++i)
regA[wm * TM + i] = As[k * BM + warpRow * WM + wm * WSUBM + laneRow * TM + i];
for (int wn = 0; wn < WNITER; ++wn) // B slices for each sub-tile column
for (int j = 0; j < TN; ++j)
regB[wn * TN + j] = Bs[k * BN + warpCol * WN + wn * WSUBN + laneCol * TN + j];
for (int wm = 0; wm < WMITER; ++wm) // warp-level matmul
for (int wn = 0; wn < WNITER; ++wn)
for (int i = 0; i < TM; ++i)
for (int j = 0; j < TN; ++j)
acc[(wm * TM + i) * (WNITER * TN) + wn * TN + j] += regA[wm * TM + i] * regB[wn * TN + j];
}
Why it is faster
- Explicit parallelism at each hardware level. Blocks run in parallel on different SMs; warps run in parallel on the four schedulers of an SM (and interleave on one scheduler); the independent FMAs inside a thread overlap in the pipeline (ILP). Warptiling makes the middle level a deliberate choice instead of an accident of thread numbering.
- Register cache locality. Recent NVIDIA GPUs have a small operand cache in front of the register file. Tighter per-warp tiles mean the same As and Bs registers are reused in consecutive instructions.
- Bank conflicts are a per-warp phenomenon, so controlling which addresses a warp touches together is the lever for reducing them.
- It is the shape tensor cores want. A warp-level matrix multiply on a WSUBM x WSUBN sub-tile maps directly onto the warp-wide
mmainstructions that tensor cores execute. This kernel is one step from that.
After re-tuning the enlarged parameter space, throughput reaches 21.8 TFLOP/s, within 6% of cuBLAS. Something that did not help: thread block swizzling (remapping blockIdx to C tiles so that concurrently running blocks share rows of A in L2). L2 hit rate was already about 80%, and the swizzle produced no measurable gain, so it was removed.
The cuBLAS reference
Not one kernel but a library of hundreds, chosen at run time by shape.
Comparing kernel 10 with cuBLAS across matrix sizes shows two regimes. At 2048 and 4096 the hand-written kernel is within a few percent. At small sizes it loses badly. The reason is that cuBLAS is a dispatcher: it contains many SGEMM implementations (the article counted 16 distinct ones for square sizes up to 4096, out of a 500 MB binary) and picks one per call based on M, N, K, data type and GPU.
The trace at size 256 is instructive: cuBLAS launched a matmul kernel and a reduction kernel. That is split-K. A 256 x 256 output with 128 x 128 tiles is only 4 blocks, which cannot occupy 84 SMs. Splitting the K dimension across several blocks gives each SM something to do; each block produces a partial sum for its slice of K, and a second kernel adds the partials.
What is left (work in progress in the article)
The remaining 6%, and what a tensor core kernel would add.
- Double buffering. Right now each chunk is a strict sequence: copy from global memory, barrier, compute, barrier. With two shared memory buffers, the copy of chunk i + 1 can be issued before the compute on chunk i starts, so memory latency overlaps math. CUTLASS does this at both levels: global to shared, and shared to registers.
- Conflict-free shared memory layouts. Swizzled layouts that keep vectorised loads while spreading a warp's accesses over all 32 banks. Kernels 7 and 8 were the first attempt.
- Hardware copy paths. On Ampere,
cp.asyncmoves data from global memory straight into shared memory without passing through registers. On Hopper, TMA and warp specialisation (some warps only load, others only compute, with different register budgets) go further. - Tensor cores. Everything above is fp32 on CUDA cores. With TF32 or BF16 inputs, warp-level
mmainstructions raise peak throughput by roughly 3x on this card and turn the problem back into a memory-bound one, where the loading machinery above becomes the entire game.
Where Triton and the FLA kernels sit
The kernels above are written one level at a time: you decide which thread owns which element, you declare shared memory, you choose the register tile, you transpose by hand, you pipeline by hand. Triton flips the contract. You write a program at the block tile level (roughly kernel 3's viewpoint: one program owns a BM x BN tile and loops over K), and the compiler fills in kernels 2, 4 to 8, 10 and 11 from a handful of knobs. The flash-linear-attention library (fla-org/flash-linear-attention, the kernels behind the fla-hub models) is written entirely in Triton, so understanding what Triton takes off your hands is also understanding what FLA can and cannot control.
Vocabulary translation
tl.program_id(axis) is blockIdxthreadIdx anywhere in Triton sourceWhat you still own in Triton
- The grid decomposition: which program computes which tile, and in which order. This is kernel 1's decision and it still determines L2 locality and load balance.
- Block tile sizes and the number of warps. Too large a tile per warp and the compiler spills accumulators to local memory, which shows up as
LDLandSTLinstructions and a cliff in throughput. That is kernel 9's register budget, now enforced by the compiler instead of the occupancy calculator. - Dtypes at each boundary: FLA loads bf16, accumulates in fp32, and casts back before the next
tl.dot. With fp32 inputs Triton uses TF32 on tensor cores unless told otherwise. - Fusion and materialisation: which intermediate tiles go back to global memory and which stay in registers. The kernels above never had this choice because a GEMM has no intermediates; a linear attention chunk kernel has several.
- Masks for ragged edges (the tile quantisation problem from the primer) and variable-length batches.
What FLA builds on top of this
A causal linear attention layer with a decay (GLA, gated delta rule, KDA and their relatives) can be written as a recurrence over tokens with a state matrix S of size K x V per head. Computing it token by token is sequential and slow; computing it as one T x T attention matrix is quadratic. FLA's chunk kernels take the middle path: split the sequence into chunks of BT tokens (typically 64), handle everything inside a chunk as a small dense attention (a BT x BT matrix, causally masked), and pass everything between chunks through S.
Read the grid axes on that figure against the primer. program_id(2) runs over batch x heads: the outermost independent dimension, so it maps to different SMs. program_id(1) runs over chunks in the output kernel, or over tiles of the state in the state kernel. program_id(0) runs over tiles of the head dimension (BK or BV), which is what keeps the register tile of one program small enough. Each program is still a block of num_warps x 32 threads; the warp and thread tiling inside a tl.dot is chosen by the compiler.
Why FLA sits somewhere else on the roofline
- The matmuls are small. With BT = 64 and head dimension 128, each
tl.dotis a 64 x 64 x 128 or 64 x 128 x 64 product: about 1 MFLOP, versus 137 GFLOP for the article's GEMM. Block-level arithmetic intensity is bounded by the tile, so these kernels sit far to the left of the ridge point and are bound by memory traffic and latency, not FMA throughput. - The state pass is a scan. One program advances S through all T / BT chunks in order, so parallelism comes from batch x heads x state tiles, not from the sequence. For a small batch at long context this is the part that leaves SMs idle, the same shape of problem as cuBLAS's split-K case.
- S gets materialised. Writing the per-chunk state to global memory (T / BT x K x V floats per head) so the output kernel can read it back is exactly the kind of traffic kernels 3 to 6 spent all their effort avoiding. Fusing the two passes trades that traffic for a longer serial chain; FLA offers both fused and chunked paths for a reason.
- Delta-rule variants add a triangular solve inside each chunk (the
solve_trilkernels), which is why those kernels carry a sub-chunk size BC and autotune over it separately.
So the levers that matter for FLA are the ones Triton leaves in your hands: chunk size, which tensors are materialised, how the grid is cut so that batch x heads x tiles fills the GPU, bf16 operands with fp32 accumulation, and the autotune space over num_warps and num_stages. The thread-level work the article spends kernels 2 to 8 on is done by the compiler, and done about as well as a careful hand-written kernel for these tile shapes. Serving stacks such as vLLM vendor FLA's Triton ops for their gated delta rule models, so the same trade-offs carry over to inference.
Cheat sheet
| Kernel | Change | Tile (BM, BN, BK, TM, TN) | GMEM loads / output | SMEM loads / output | GFLOP/s | of cuBLAS |
|---|---|---|---|---|---|---|
| 1 | one thread per output, strided warp reads of A | 32, 32, K, 1, 1 | 2K | 0 | 309 | 1.3% |
| 2 | consecutive lanes read consecutive columns | 32, 32, K, 1, 1 | 2K (coalesced) | 0 | 1987 | 8.5% |
| 3 | block caches chunks of A and B in SMEM | 32, 32, 32, 1, 1 | K/16 | 2K | 2980 | 12.8% |
| 4 | thread owns a column of 8 outputs | 64, 64, 8, 8, 1 | K/32 | 9K/8 | 8475 | 36.5% |
| 5 | thread owns an 8 x 8 square, outer product in registers | 128, 128, 8, 8, 8 | K/64 | K/4 | 15972 | 68.7% |
| 6 | As transposed for LDS.128; float4 GMEM loads | 128, 128, 8, 8, 8 | K/64 (128-bit) | K/4 (128-bit) | 18237 | 78.4% |
| 9 | parameters searched by benchmark | 128, 128, 16, 8, 8 | K/64 | K/4 | 19721 | 84.8% |
| 10 | warp tile level between block and thread | + WM, WN, WMITER, WNITER | K/64 | below K/4 | 21779 | 93.7% |
| 0 | cuBLAS, shape-dependent dispatch | varies | 23250 | 100% |
Loads per output are counted for the full K loop, per thread, divided by the number of outputs the thread produces. For kernel 10 the shared loads per output depend on WMITER and WNITER: per k a thread loads (WMITER x TM + WNITER x TN) values for WMITER x WNITER x TM x TN FMAs.
The five rules the worklog teaches
- Make the 32 lanes of a warp touch 32 consecutive, aligned words. Everything else about global memory is secondary.
- Move data up the hierarchy once, then reuse it as many times as possible before moving more. Block tiles reuse through shared memory; thread tiles reuse through registers.
- Count instructions, not just bytes. Two shared loads per FMA is a bottleneck even when the bytes are on-chip.
- Give the compiler the information it cannot infer: compile-time loop bounds so it can unroll, and aligned vector types so it can emit 128-bit accesses.
- Organise around warps. They are the unit of scheduling, of coalescing, of bank conflicts and of tensor core instructions.