SGEMM: Shared Memory Tiling Visualizer

Watch threads cooperatively load data to fast On-Chip SRAM before computing.

Global DRAM Reads
0
Shared SRAM Reads
0

Simulation Engine

Init Block
Problem Size
M=6, N=6, K=9, TILE=3

Use small sizes so the grids stay readable.

Selected Thread
T(0,0)

Pick a thread to follow across the full tile-local dot product.

Speed600ms
Current Focus: Block (0, 0) Skip to Next Block →
// BlockDim=(2,2), TILE_SIZE=2
__shared__ float tileA[2][2];
__shared__ float tileB[2][2]; // no padding needed
int globalRow = blockIdx.y * 2 + threadIdx.y;
int globalCol = blockIdx.x * 2 + threadIdx.x;
for (int tileIdx = 0; tileIdx < ceil(K / 2); tileIdx++) {
tileA[ty][tx] = (globalRow < M && tileColA < K) ? A[globalRow * K + tileColA] : 0;
tileB[ty][tx] = (tileRowB < K && globalCol < N) ? B[tileRowB * N + globalCol] : 0;
__syncthreads(); // Barrier
for (int k = 0; k < 2; k++) {
partialSum += tileA[ty][k] * tileB[k][tx];
}
__syncthreads(); // Barrier
}
C[globalRow*N + globalCol] = partialSum;

Global Memory (DRAM)

Matrix A (4 × 8)
Matrix B (8 × 4)

Shared Memory (SRAM)

tileA [2][2]
tileB [2][2]

Registers & Matrix C

Matrix C (4 × 4)

Block Thread Monitor (BlockDim = 2×2)

Current Phase Details
Selected Thread Lens
T(0,0) -> C(0,0)

tileA Row
tileB Column
Dot Product Timeline
Thread (tx,ty) Global Target Action Local partialSum

Why Tiling Works (The Math)

Without Shared Memory: Each thread reads an entire row of A and column of B from slow global DRAM.

\(\text{DRAM Reads} = 2 \times M \times N \times K\)

With Shared Memory: Threads cooperatively load a tile into fast SRAM, then reuse it \(\text{TILE_SIZE}\) times.

\(\text{DRAM Reads} = \frac{2 \times M \times N \times K}{\text{TILE_SIZE}}\)

Current example: \(2 \times 4 \times 4 \times 8 / 2 = 128\) ideal DRAM reads.

Why exactly \(2 \times M \times N \times K\) for the naive approach?

  1. Total Number of Threads (\(M \times N\)): We launch one thread for every single element in the output Matrix C (which has \(M\) rows and \(N\) columns).
  2. What Each Thread Does (\(K\) steps): To compute its single target value, a thread calculates the dot product of one row from Matrix A and one column from Matrix B. Both have a length of \(K\).
  3. Memory Reads per Thread (\(2 \times K\)): In every iteration of that \(K\)-length loop, the thread reads one value from A and one from B (2 reads).
  4. The Final Math: Total Threads × Reads per Thread = \((M \times N) \times (2 \times K) = 2 \times M \times N \times K\).

Why Tiling Fixes This: Neighboring threads need the same rows of A or columns of B. The naive approach fetches the exact same data from slow DRAM repeatedly. Tiling brings a chunk of that data into fast shared memory once, then threads read from the shared memory, drastically cutting down that massive number!

Why no padding? Neither tile needs padding. Shared memory is divided into 32 banks, and a bank conflict happens only when threads of the same warp access different addresses in the same bank in the same instruction. A warp is 32 threads with consecutive threadIdx.x. At a fixed k, the warp reads tileA[ty][k] (one address for everyone: a broadcast) and tileB[k][tx] (32 consecutive words of one row: 32 different banks). Each thread does walk down a column of tileB over the k loop, but those reads are separate instructions, so they cannot conflict. Padding matters when one warp reads down a column at once, as in the matrix transpose example.