<<<CUDA C++ 学习路线grid · block · warp · lane
阶段 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 最大的区别是:它由你显式管理。这既是负担也是武器——你比任何硬件预取器都更清楚自己的数据复用模式。

从朴素矩阵乘说起

04_matmul_naive.cu
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 │   │   累加器:全程待在寄存器里
├───┼───┼───┤ ├───┼───┤
│   │   │   │ │   │   │
└───┴───┴───┘ └───┴───┘
tiled matmul 的第 t 步
05_matmul_tiled.cu — 完整实现
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 个周期(最坏)
bank 编号 = (字节地址 / 4) % 32
打开 Bank Conflict 模拟器输入访问模式(stride、二维索引、padding 开关),直接看到 32 个 bank 的占用情况和冲突路数。强烈建议在这里亲手试一次「加 padding 前后」的对比。

转置里那个 +1 到底解决了什么

冲突是怎么发生的
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 周期
padding:一个字的代价换 32 倍速度
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 库用的都是这个思路。

XOR swizzle
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)];

自测

先自己在心里回答一遍,再展开对照。答不上来的说明这一段值得重读。