CUDA 深入浅出(二):为什么分块可以减少 Global Memory 访问
# CUDA 深入浅出(二):为什么分块可以减少 Global Memory 访问
上一篇讲了一个结论:T×T 的 tile 乘法里,一个从 global memory 搬进来的元素,会被使用 T 次。
这篇换一个问法:
矩阵乘法本身还是那些乘加。你要算 C=A×B,每个 C[i][j] 还是要沿着 k 方向做点积。Tiling 省不掉数学运算。
Tiling 省的是重复搬运。
# naive kernel 怎么读数据
先看最朴素的 CUDA 写法。
一个 thread 负责一个输出元素:
这个 thread 要算 C[row][col]。
每走一次 k,它从 global memory 读:
所以一个输出元素需要:
也就是 2N 次 global memory load。
整个 C 有 N² 个输出元素。于是 naive kernel 对 A 和 B 的读取量大约是:
这里先不讨论 C 的写回。每个 kernel 都要把结果写回 C,这一项是 N² 级别。GEMM 的主要压力来自 A 和 B 在 k 循环里的反复读取。
# 重复读取藏在哪里
拿 T=4 看一个小例子。
假设一个 block 要计算 C_tile:
一共有 16 个输出元素。
如果每个 thread 自己从 global memory 读数据,那么算第一轮 k=1 时:
同一个 a11 被四个 thread 读了四次。
继续看第二行:
同一个 b11 也被四个 thread 读了四次。
问题不在于这些值参与了多次计算。矩阵乘法本来就需要它们参与多次计算。问题在于 naive kernel 让多个 thread 各自去 global memory 拿同一份数据。
你可以把它想成很多人排队去仓库拿同一把螺丝刀。工具确实要用很多次,但没必要每个人都跑一趟仓库。
# 分块之后怎么读
Tiling 让一个 block 先合作搬数据。
对于 T=4:
一个 block 先从 global memory 读取这 32 个元素,把它们放进 shared memory。
接下来,block 里的 thread 算 C_tile:
a11 还是被读了多次,但这些读取来自 shared memory。
global memory 只负责把 a11 搬进来一次。
这就是减少 global memory 访问的关键:Tiling 没有让数据少被使用,而是让重复使用发生在 shared memory 里。
# 对一个 C_tile 数数
现在用一般的 T×T 来数。
一个 block 要算一个 T×T 的 C_tile。在 k 方向的一段 tile multiply 里,它需要:
Tiled kernel 的 global memory load 是:
naive kernel 呢?
同一个 C_tile 里有 T² 个输出元素。每个输出元素在这一段 k tile 里需要做 T 次乘法。
每次乘法要读一个 A 和一个 B:
所以 naive kernel 的 global memory load 是:
现在差距出来了:
两者相除:
也就是说,对于这一段 tile multiply,分块把 A 和 B 的 global memory 读取量压低了 T 倍。
如果 T=16,这一段少读约 16 倍。
如果 T=32,这一段少读约 32 倍。
现实里还会受到 cache、coalescing、边界处理、shared memory bank conflict、occupancy 的影响。这个 T 倍是理想计数模型,但它抓住了 Tiling 的主要收益。
# 推广到整个 N×N 矩阵
假设 A、B、C 都是 N×N,并且 N 能被 T 整除。
naive kernel 对 A 和 B 的读取量大约是:
tiled kernel 怎么数?
整个 C 会被切成很多 T×T 的输出块:
每个 C_tile 沿着 k 方向要走:
每个 tile step 从 global memory 读取:
所以 tiled kernel 的读取量是:
整理一下:
和 naive 的 2N³ 比:
还是少了 T 倍。
这也是上一篇讲的 data reuse factor。站在复用的角度看,一个元素被用了 T 次。站在访存的角度看,你少做了 T 倍的 global memory load。
# 代码里的样子
一个典型 tiled matmul kernel 会长这样:
外层 tile 循环每走一步,block 从 global memory 搬一块 A 和一块 B。
内层 k 循环不再碰 global memory。它读的是:
这两个数组在 shared memory 里。
__syncthreads() 保证 block 里的 thread 都把数据搬完了,再开始算。算完这一段以后,block 再进入下一个 tile step,搬下一批数据。
这个结构把很多分散的 global load,变成了两件事:
# cache 不能替代 tiling 吗
GPU 有 L2 cache,也有每个 SM 附近的 cache。naive kernel 里,多个 thread 读同一个 A 或 B 元素时,cache 可能会命中一部分访问。
但 cache 不能让你明确控制数据复用。
shared memory 给了 kernel 一个可编程的暂存区。你自己决定:
cache 会根据硬件策略替你猜。Tiling 让你把复用关系写进程序结构。
这也是 CUDA 里 shared memory 难绕开的原因。你想让一个 block 内的 thread 复用同一批数据,就要给它们一个共同可见、延迟更低、带宽更高的地方。
# 减少的是远距离搬运
分块前,多个 thread 会从 global memory 里反复读取同一个 A 或 B 元素。
分块后,block 先把这批元素搬进 shared memory。多个 thread 仍然会读它们,但读的是 shared memory 里的副本。
所以这句话可以更精确地说:
矩阵乘法想跑快,光让线程变多不够。线程多了以后,如果每个线程都去远处搬同一份数据,GPU 会把大量时间花在等内存上。
Tiling 让一个 block 先把数据搬近,再让线程把这批数据反复用完。
这就是 2T³ 到 2T² 的意义。你没有少算,但你少跑了很多趟远路。
