阶段 2 · 执行模型与内存07 / 15约 26 分钟cuda/05_matmul_tiled.cu
共享内存与 Bank Conflict:片上存储的正确用法
tiled 矩阵乘完整推导,以及那个神秘的 [TILE][TILE + 1]
学完这一课你会
- 理解共享内存作为「程序员管理的缓存」的价值
- 能从零推导 tiled 矩阵乘并说清每一步为什么这么写
- 彻底搞懂 32 个 bank 的机制和 conflict 的成因
- 掌握 padding 与 swizzle 两种消除 conflict 的手段
共享内存是每个 SM 上的一块高速 SRAM,延迟约为全局内存的 1/20,由 block 内所有线程共享。它和 CPU 的 L1 最大的区别是:它由你显式管理。这既是负担也是武器——你比任何硬件预取器都更清楚自己的数据复用模式。
从朴素矩阵乘说起
CUDA C++
1__global__ void matmulNaive(const float* A, const float* B, float* C, int N) {2 int col = blockIdx.x * blockDim.x + threadIdx.x;3 int row = blockIdx.y * blockDim.y + threadIdx.y;4 if (row >= N || col >= N) return;5 6 float acc = 0.0f;7 for (int k = 0; k < N; ++k) {8 acc += A[row * N + k] * B[k * N + col]; // 每轮 2 次全局读,1 次乘加9 }10 C[row * N + col] = acc;11}算一下这个 kernel 的账:内层循环每次迭代读 2 个 float(8 字节)只做 2 次浮点运算,计算强度 0.25 FLOP/Byte,是硬件拐点的 1/50。更浪费的是,A 的同一行被同一行上的 N 个线程各读了一遍,B 的同一列也被读了 N 遍——总共 2N³ 次全局访存,而实际只有 2N² 个不同的数据。
B 每一步: ┌───┬───┬───┐ 1. 各线程协作把 A 的一块、B 的一块搬进共享内存 │ │Bt │ │ 2. __syncthreads() 等所有人搬完 ├───┼───┼───┤ 3. 在共享内存里做 TILE 次乘加,累加到寄存器 │ │ │ │ 4. __syncthreads() 等所有人算完再进下一块 └───┴───┴───┘ A 全局读次数:2 * N^3 / TILE ┌───┬───┬───┐ ┌───┬───┐ 共享内存读:2 * N^3(但快 20 倍) │At │ │ │ │ C │ │ 累加器:全程待在寄存器里 ├───┼───┼───┤ ├───┼───┤ │ │ │ │ │ │ │ └───┴───┴───┘ └───┴───┘
CUDA C++
1#define TILE 322 3__global__ void matmulTiled(const float* A, const float* B, float* C, int N) {4 __shared__ float sA[TILE][TILE];5 __shared__ float sB[TILE][TILE];6 7 int tx = threadIdx.x, ty = threadIdx.y;8 int row = blockIdx.y * TILE + ty;9 int col = blockIdx.x * TILE + tx;10 11 float acc = 0.0f;12 13 for (int t = 0; t < (N + TILE - 1) / TILE; ++t) {14 // 协作加载:每个线程负责一个元素,越界补零省掉内层判断15 int aCol = t * TILE + tx;16 int bRow = t * TILE + ty;17 sA[ty][tx] = (row < N && aCol < N) ? A[row * N + aCol] : 0.0f;18 sB[ty][tx] = (bRow < N && col < N) ? B[bRow * N + col] : 0.0f;19 20 __syncthreads(); // 屏障 1:确保整块都搬完了21 22 #pragma unroll23 for (int k = 0; k < TILE; ++k) {24 acc += sA[ty][k] * sB[k][tx];25 }26 27 __syncthreads(); // 屏障 2:确保都算完了,才能覆盖共享内存28 }29 30 if (row < N && col < N) C[row * N + col] = acc;31}32 个 Bank:共享内存的并行结构
共享内存被划分成 32 个 bank,每个 bank 宽 4 字节,按地址交错排列:地址 0 在 bank 0,地址 4 在 bank 1……地址 128 又回到 bank 0。32 个 bank 可以在一个周期内同时服务 32 个请求。当一个 warp 内多个线程访问同一个 bank 的不同地址时,硬件只能串行处理,这就是 bank conflict。
float 下标: 0 1 2 ... 31 32 33 ...
bank: 0 1 2 ... 31 0 1 ...
└──────── 一轮 32 个 bank ────────┘
无冲突:32 个线程访问 32 个不同 bank → 1 个周期
广播 :32 个线程访问同一个 bank 的同一地址 → 1 个周期(硬件广播,不算冲突)
2 路冲突:每 2 个线程撞同一个 bank 的不同地址 → 2 个周期
32 路冲突:全部撞进同一个 bank 的不同地址 → 32 个周期(最坏)转置里那个 +1 到底解决了什么
CUDA C++
1__shared__ float tile[32][32];2 3// 按列访问:threadIdx.x = 0..31,行不同、列相同4float v = tile[threadIdx.x][0];5 6// 线程 i 访问的 float 下标 = i * 32 + 07// bank = (i * 32) % 32 = 0 ← 32 个线程全撞在 bank 0!8// 结果:32 路冲突,本该 1 周期的访问变成 32 周期CUDA C++
1__shared__ float tile[32][33]; // 每行多一个 float2 3float v = tile[threadIdx.x][0];4 5// 线程 i 访问的 float 下标 = i * 33 + 06// bank = (i * 33) % 32 = i ← 33 与 32 互质,正好打散到 32 个 bank7// 结果:零冲突。代价仅 32 * 4 = 128 字节额外共享内存另一条路:swizzle
padding 会浪费共享内存,在 tile 很大时可能挤压占用率。更高级的做法是 swizzle——不改变存储大小,而是用异或打乱列索引,让同列的不同行落在不同 bank 上。CUTLASS 和各种高性能 GEMM 库用的都是这个思路。
CUDA C++
1__shared__ float tile[32][32]; // 不加 padding,不浪费一个字节2 3__device__ inline int swz(int row, int col) {4 return col ^ (row & 31); // 每一行的列顺序被异或打乱5}6 7// 写入8tile[ty][swz(ty, tx)] = value;9 10// 按列读取时,不同 row 的同一逻辑 col 被映射到了不同物理 bank11float v = tile[threadIdx.x][swz(threadIdx.x, 0)];自测
先自己在心里回答一遍,再展开对照。答不上来的说明这一段值得重读。