Why two loops?
In a naive kernel, 1 thread calculates 1 output element. It loads 1 value from A and 1 from B for every multiplication.
With Register Tiling, 1 thread calculates TM=4 output elements! To do this efficiently:
- The
dotIdxloop iterates through the K dimension (size BK=4). - We load one value from
Bsinto a fast register (regB). - The
resultIndexloop multiplies that singleregBvalue against TM=4 different values fromAs.
Result: We reused the memory loaded into regB 4 times! This drastically reduces memory bandwidth bottlenecks.
Kernel Code
Hover for info
for (int bkIdx = 0; bkIdx < K; bkIdx += BK) {// 1. Fetch Global to Shared (As, Bs)__syncthreads();for (int dotIdx = 0; dotIdx < BK; ++dotIdx) {float regB = Bs[dotIdx][threadCol];#pragma unrollfor (int resultIndex = 0; resultIndex < TM; ++resultIndex) {threadResults[resultIndex] +=As[threadRow*TM + resultIndex][dotIdx] * regB;}}__syncthreads();}
Step 1 of X
Initialization
1
X
Global Memory (Original Matrices)
Matrix A [8][8]
Matrix B [8][8]
Cols 4-7 are ignored by this specific block
Shared Memory
As [8][4]
Rows 0-3 are computed by Thread(0,0)
Shared Memory
Bs [4][4]
Col 0 is for Threads(*,0)
Registers
regB
-
Registers
threadResults[TM]
Global Memory
Block Matrix C [8][4]
Thread(0,0) computes the highlighted 4x1 section.