SGEMM Register Tiling: Thread(0,0) View

Understanding the dotIdx and resultIndex loops.
Parameters: M=8, N=8, K=8 | BM=8, BN=4, BK=4, TM=4.

Select Thread Matrix threadIdx.(y, x)

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 dotIdx loop iterates through the K dimension (size BK=4).
  • We load one value from Bs into a fast register (regB).
  • The resultIndex loop multiplies that single regB value against TM=4 different values from As.

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 unroll
for (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.