先回顾:Naive 矩阵乘法长什么样
假设我们要算 C = A × B,矩阵都是 N×N。C 的每个元素 C[row][col] 是 A 的第 row 行和 B 的第 col 列做点积:
C[row][col] = A[row][0]*B[0][col] + A[row][1]*B[1][col] + ... + A[row][N-1]*B[N-1][col]
└─── 沿 K 维度求和 ──────────────────────────────────┘
Naive 版本每个线程算 C 的一个元素,从 Global Memory 读 A 的一整行和 B 的一整列。问题是 A 和 B 的数据被大量重复读取——计算 C[0][0] 读了 A 的第 0 行,计算 C[0][1] 又要再读一遍 A 的第 0 行。
Tiled 的核心思想:分块加载到 Shared Memory
把大矩阵切成小块(tile),每次只处理一小块,这样数据可以被 Block 内的线程复用。
画个图最直观。假设 N=8,TILE_SIZE=4:
矩阵 A (8×8) 矩阵 B (8×8) 矩阵 C (8×8)
┌────┬────┐ ┌────┬────┐ ┌────┬────┐
│A00 │A01 │ │B00 │B01 │ │C00 │C01 │
│4×4 │4×4 │ │4×4 │4×4 │ │4×4 │4×4 │
├────┼────┤ ├────┼────┤ ├────┼────┤
│A10 │A11 │ │B10 │B11 │ │C10 │C11 │
│4×4 │4×4 │ │4×4 │4×4 │ │4×4 │4×4 │
└────┴────┘ └────┴────┘ └────┴────┘
要算 C00(左上角的 4×4 子矩阵),需要:
C00 = A00 × B00 + A01 × B10
───────── ─────────
第 1 个 tile 第 2 个 tile
关键洞察:C00 这块 4×4 = 16 个元素,每个元素都需要读 A 的第 0-3 行和 B 的第 0-3 列。如果不用 Shared Memory,16 个线程各读各的,A 的同一行被读 4 次。用了 Shared Memory,整个 Block 协作把 A00 这块 4×4 一次性搬到 Shared Memory,然后 16 个线程都从 Shared Memory 读——Global Memory 访问量减少到 1/4。
逐行解读代码
__global__ void matMulShared(float *A, float *B, float *C, int N) {
// ① 声明 Shared Memory:每个 Block 分配两块 TILE_SIZE × TILE_SIZE 的共享空间
__shared__ float sA[TILE_SIZE][TILE_SIZE];
__shared__ float sB[TILE_SIZE][TILE_SIZE];
__shared__ 关键字告诉编译器这块内存在 SM 的 Shared Memory 上,Block 内所有线程共享。
// ② 计算当前线程负责 C 的哪个元素
int row = threadIdx.y + blockIdx.y * TILE_SIZE;
int col = threadIdx.x + blockIdx.x * TILE_SIZE;
假设 TILE_SIZE=4,blockIdx = (1, 2),threadIdx = (3, 1):
- row = 1 + 2 × 4 = 9 → 第 9 行
- col = 3 + 1 × 4 = 7 → 第 7 列
- 这个线程负责计算 C[9][7]
这里 threadIdx.x 对应列、threadIdx.y 对应行——正好满足上一课讲的合并访问原则(threadIdx.x 相邻 → col 相邻 → 内存连续)。
float sum = 0.0f;
// ③ 沿 K 维度(公共维度)分块遍历
for (int t = 0; t < N / TILE_SIZE; t++) {
t 是 tile 的编号。N=8, TILE_SIZE=4 时,t 从 0 到 1,循环两次——对应前面图里的 A00×B00 和 A01×B10。
// ④ 协作加载:每个线程从 Global Memory 读一个元素,存入 Shared Memory
sA[threadIdx.y][threadIdx.x] = A[row * N + (t * TILE_SIZE + threadIdx.x)];
sB[threadIdx.y][threadIdx.x] = B[(t * TILE_SIZE + threadIdx.y) * N + col];
这是最关键的一步。Block 内有 TILE_SIZE × TILE_SIZE 个线程(比如 4×4 = 16 个),每个线程负责搬一个元素:
线程 (ty=0, tx=0) 搬 A[row][t*4 + 0] → sA[0][0]
线程 (ty=0, tx=1) 搬 A[row][t*4 + 1] → sA[0][1]
...
线程 (ty=3, tx=3) 搬 A[row][t*4 + 3] → sA[3][3]
16 个线程各搬 1 个 → 一次协作就把 4×4 的 tile 从 Global Memory 搬到了 Shared Memory
注意 sA 和 sB 的索引方式:
sA[ty][tx]:ty 是行,tx 是列。每个线程搬 A 中当前行、第 t 块的一个元素sB[ty][tx]:ty 是行,tx 是列。每个线程搬 B 中第 t 块、当前列的一个元素
__syncthreads(); // ⑤ 屏障同步:等所有线程都搬完了再继续
为什么必须有这个? 如果没有,可能线程 0 已经搬完了开始计算,但线程 15 还没搬完——线程 0 就会读到 sA 中的垃圾数据。__syncthreads() 确保 Block 内所有线程都走到这里之后才一起继续。
// ⑥ 在 Shared Memory 内做计算
for (int k = 0; k < TILE_SIZE; k++) {
sum += sA[threadIdx.y][k] * sB[k][threadIdx.x];
}
现在数据在 Shared Memory 里了,做点积。对于计算 C[row][col] 的线程:
sum += sA[ty][0] * sB[0][tx] ← 从 Shared Memory 读,~20 cycles
sum += sA[ty][1] * sB[1][tx]
sum += sA[ty][2] * sB[2][tx]
sum += sA[ty][3] * sB[3][tx]
对比 naive 版直接从 Global Memory 读(~300+ cycles/次),每次访问快了 15 倍以上。
__syncthreads(); // ⑦ 确保所有线程算完了再加载下一个 tile(覆盖 sA/sB)
}
// ⑧ 写回结果
C[row * N + col] = sum;
}
第二个 __syncthreads() 也不能省:下一轮循环会用新的数据覆盖 sA 和 sB,如果还有线程没算完就被覆盖了,结果就错了。
一图总结整个流程
K 维度分块遍历
┌──────────┐
│ t = 0 │
└────┬─────┘
▼
┌────────────────────────────────────┐
│ 16 个线程协作加载 A_tile, B_tile │ ← Global → Shared
│ __syncthreads() │
│ 每个线程在 Shared Memory 内做部分求和 │ ← 快速读取
│ __syncthreads() │
└────────────────┬───────────────────┘
▼
┌──────────┐
│ t = 1 │ ← 加载下一块 tile,覆盖 sA/sB
└────┬─────┘
▼
... 重复 ...
▼
所有 tile 算完
▼
写 sum 到 Global Memory
核心收益:Naive 版本中 A 的每一行被读 N 次(C 的每一列都要用),tiled 版本中只从 Global Memory 读 N/TILE_SIZE 次(每次搬一个 tile,Block 内复用 TILE_SIZE 次)。当 TILE_SIZE=32 时,Global Memory 带宽需求降低到 1/32。

235

被折叠的 条评论
为什么被折叠?



