CUDA 深入浅出(四):怎么高效访问 Shared Memory
# CUDA 深入浅出(四):怎么高效访问 Shared Memory
前三篇都在围绕 global memory 打转。
我们先讲了 Tiling 为什么能提高数据复用率,又讲了它为什么能减少 global memory 访问,接着讲了 coalesced load 让 warp 访问 global memory 时少浪费 transaction。
现在数据已经被搬进 shared memory 了。
接下来的问题变成:
shared memory 比 global memory 快很多,但它不是无限带宽的魔法盒。你读得太乱,warp 里的 thread 会互相撞。你读得太频繁,shared memory 也会变成瓶颈。
这一篇看两个东西:
前者讲 shared memory 的访问形状。后者讲一个线程读到数据以后,怎么在 register 里多用几次。
# shared memory 也有访问结构
你可以把 shared memory 想成一个 block 内部共享的小仓库。
但这个仓库不是一个大门。它分成很多个 bank。常见 CUDA 架构里,shared memory 有 32 个 bank。一个 warp 也有 32 个 thread。
理想情况下:
每个 thread 访问不同 bank,硬件可以并行服务这些访问。
如果多个 thread 访问同一个 bank 的不同地址,硬件就要分多轮处理。这个现象叫 bank conflict。
它和上一篇的 coalesced load 有点像。global memory 关心 warp 的地址能不能合成少量 transaction。shared memory 关心 warp 的地址会不会打到同一个 bank 上。
# bank 怎么决定
先用 float 看。一个 float 占 4 字节。
你可以先记一个简化模型:
如果一个 warp 里的 thread 访问连续的 float:
这很好。
如果 thread 跨步访问:
这些元素的下标对 32 取模都等于 0。
整个 warp 都打到 bank 0。硬件没法一口气处理完,只能拆成多轮。
这就是 shared memory 版本的“排队堵门”。
# 矩阵 tile 里最常见的坑
我们通常会这样声明 shared memory:
如果 thread 读同一行的连续列:
那么相邻 tx 访问的是:
这是一段连续地址。bank 分布也很自然。
麻烦常出在按列读:
如果 T=32,相邻 tx 访问的是:
在 row-major 布局里,每一行相隔 32 个 float。这些地址的元素下标差是 32。
对 bank 来说,它们可能都落在同一个 bank 上。
这就是一个典型的 32-way bank conflict。
# padding 为什么有用
常见修法是在 shared memory 里多加一列:
这样每一行不再正好隔 32 个 float,而是隔 33 个。
如果你按列访问:
对应的元素下标差变成 33。
对 32 个 bank 取模:
bank 被错开了。
padding 没有改变矩阵数学。它只是改变 shared memory 里的布局,让 warp 访问时不要挤进同一个 bank。
这招在 transpose、某些 tile 变换、或者需要按列读 shared memory 的 kernel 里很常见。
# 广播不是 bank conflict
有一种情况容易误判。
如果一个 warp 里的多个 thread 读取同一个 shared memory 地址,比如:
这通常可以走 broadcast。硬件把同一个值发给多个 thread。
bank conflict 说的是多个 thread 访问同一个 bank 里的不同地址。
所以判断 shared memory 访问时,要看两个问题:
同一个地址可以广播。不同地址就会排队。
# shared memory 读多了也会慢
解决 bank conflict 只处理了访问形状。
另一个问题是访问次数。
Tiled matmul 里,最朴素的内层循环通常长这样:
每个 thread 负责一个 C[row][col]。
每走一次 k,这个 thread 从 shared memory 读:
也就是两次 shared memory load。
对一个 T×T 的 C_tile 来说,线程数量是 T²。每个线程做 T 轮。shared memory load 数量大约是:
你已经把 global memory 压下去了,但 block 内部还在频繁读 shared memory。
shared memory 快,但没有 register 快。
# register tiling 在干什么
register tiling 的想法是:让一个 thread 不只算一个 C 元素,而是算一小块 C。
比如一个 thread 算 2×2 个输出:
它会在 register 里放 4 个 accumulator:
然后每一轮 k,它从 shared memory 读两个 A 值、两个 B 值:
接着在 register 里做四次 FMA:
这里的关键是复用。
a0 从 shared memory 读了一次,用在 c00 和 c01。
a1 从 shared memory 读了一次,用在 c10 和 c11。
b0 从 shared memory 读了一次,用在 c00 和 c10。
b1 从 shared memory 读了一次,用在 c01 和 c11。
同样是读 4 个 shared memory 值,这个 thread 做了 4 次乘加。
如果一个 thread 只算一个输出,它读 2 个值,做 1 次乘加。
register tiling 把 shared memory 里读出来的数据,放进 register 后多用几次。
# 从 1×1 到 2×2
对比一下。
普通 tiled matmul 里,一个 thread 每轮 k:
2×2 register tile 里,一个 thread 每轮 k:
shared memory load 从每次 FMA 平均 2 个值,变成每次 FMA 平均 1 个值。
如果一个 thread 算 4×4 小块,每轮 k:
平均到每次 FMA,shared memory load 变成:
这个比例会随着 register tile 变大继续下降。
这就是 register tiling 的价值:它减少 shared memory 到 register 的数据搬运,让每个 thread 拿到数据后多算几下。
# 一个小型代码形状
下面不是完整高性能 GEMM,只看 register tiling 的形状:
a、b、acc 都会尽量放在 register 里。
这段代码的核心不是数组写法,而是数据流:
shared memory 负责 block 内线程共享。register 负责单个 thread 内部复用。
# register tiling 也有代价
register 不是越多越好。
一个 thread 算 4×4 小块,就要至少 16 个 accumulator register。再加上 A fragment、B fragment、地址计算、循环变量,register 用量会涨得很快。
register 用太多,SM 上同时能驻留的 warp 数会下降。这个现象会影响 occupancy。
所以高性能 GEMM 里,tile size、warp tile、thread tile、register tile 都要一起调。
你可以先用一个直觉:
真正的工程优化就是在这两个方向之间找平衡。
# shared memory 优化的顺序
看一个 tiled GEMM,可以按这个顺序检查:
前三篇解决的是前两个问题。
这一篇解决第三和第四个问题。
shared memory 让一个 block 里的 thread 共享数据。register tiling 让一个 thread 在自己的小世界里继续复用数据。
你可以把 GEMM 的数据路径看成三段:
Tiling 把数据从 global memory 搬近。
Coalesced load 让这次搬运更整齐。
Bank conflict 处理 shared memory 内部的拥堵。
Register tiling 让 thread 拿到数据后多做几次计算。
CUDA GEMM 优化就是沿着这条路径,一层一层减少无效搬运。
