CUDA 深入浅出(一):矩阵乘法为什么要分块
# CUDA 深入浅出(一):矩阵乘法为什么要分块
很多 CUDA 教程讲矩阵乘法时,很快就会扔出一句话:
each tile of A and B contains T² elements and is loaded once from slow memory. each element participates in T multiply-add operations during the tile multiply. So the data reuse factor is T.
这句话没有错,但它把太多东西塞进了一行。
如果你刚开始看 CUDA,很容易在这里卡住。T² 是什么,T³ 又从哪来,一个元素为什么会参与 T 次 multiply-add,最后怎么就得到 data reuse factor 等于 T?
我们慢一点,从一个小块矩阵乘法开始。
# 先看 T=4
假设 tile 大小是 T=4。
那么 A_tile 是一个 4×4 小矩阵,里面有 16 个元素。B_tile 也是一个 4×4 小矩阵,也有 16 个元素。
CUDA kernel 会先让一个 thread block 合作,把这两个小块从 global memory 搬到 shared memory:
global memory 可以理解成离计算单元很远的慢内存。shared memory 离一个 block 里的线程近得多,访问速度也高得多。
这就是原文里说的:
对于 T=4,每个 tile 有 4²=16 个元素。每个 tile 从慢内存里搬一次。
# 一个元素到底用了几次
现在看 A_tile × B_tile。
结果 C_tile 也是一个 4×4 小矩阵,一共有 16 个输出元素。
比如左上角:
这个输出元素需要 4 次乘法。
再看同一行的下一个输出:
注意 a11 又出现了。
继续往右看:
a11 一共参与了这四个输出:
所以 a11 被用了 4 次。
这不是巧合。A 里的一个元素 aij,行号固定,它会沿着 B 的不同列一路配对。B 有 T 列,它就参与 T 个输出。
所以对于一般的 T×T tile:
原文里的这句:
说的就是这个意思。一个元素从慢内存搬进来之后,没有用一次就丢掉,而是在这个小矩阵乘法里反复参与了 T 次计算。
# T³ 从哪里来
另一个常见卡点是计算量。
很多文章会直接写:
这一步看起来跳得很快。
矩阵乘法的结果里有 T×T 个输出元素。也就是:
每个输出元素都要做一条长度为 T 的点积:
也就是说,每个输出元素需要:
输出元素有 T² 个,所以总乘法数是:
总加法数严格算是:
HPC 里讨论 GEMM 性能时,通常把一次乘法和一次加法都算成一次浮点运算。于是总 FLOPs 更精确地写成:
当 T 比较大时,T² 这一项比 T³ 小很多。比如 T=32:
两者只差一点点。
所以大家会把一个 T×T tile multiply 的计算量近似写成:
这不是说加法真的有 T³ 次。严格数是 T³ - T²。写成 2T³ 是为了抓住主项,方便讨论性能趋势。
# 复用率为什么是 T
现在把访存量和计算量放到一起看。
一次 tile multiply 需要从慢内存加载:
这次小矩阵乘法完成的计算量约等于:
所以每读入一个元素,平均换来的计算量是:
这就是 data reuse factor 等于 T 的来源。
你也可以不用公式记它。一个 A 元素会被同一行的 T 个输出使用,一个 B 元素会被同一列的 T 个输出使用。它们从 slow memory 进到 shared memory 后,都被一个 block 里的线程榨出了 T 次价值。
# 如果不分块会怎样
最朴素的矩阵乘法 kernel 往往是:一个线程负责一个 C[i][j]。
这个线程会从 global memory 里读一行 A,再读一列 B,算完一个输出就结束。旁边的线程算 C[i][j+1] 时,又会重新读很多相同的 A 元素。
比如 a11。
C11 要用它,C12 也要用它,C13 还要用它。如果没有 shared memory 里的 tile,多个线程可能会反复从 global memory 读同一个 a11。
分块之后,线程先把 a11 搬进 shared memory。接下来同一个 block 里的多个线程都从 shared memory 里读它。
访问次数没有消失,但昂贵的 global memory 访问变少了。
CUDA 里很多优化都围绕同一个问题打转:你已经花钱把数据从远处搬来了,能不能在它离你近的时候多用几次?
# 这才是 Tiling 的核心
Tiling 没有减少矩阵乘法本身需要做的乘加次数。C=A×B 该算多少还是算多少。
Tiling 改变的是数据走路的方式。
没有分块时,你可能频繁去 global memory 拿同一批数据。分块之后,你把一小批数据搬到 shared memory,让 block 内的线程反复使用。
所以那句抽象的话可以翻译成工程语言:
这也是 GEMM 优化里最值得先记住的关系:
T 越大,单位访存对应的计算越多。GPU 的算力很强,但它讨厌线程一边等内存一边发呆。Tiling 的第一层意义,就是让数据少走远路,让线程多做计算。
