CUDA 深入浅出(五):Vectorized Load 和 Shared Memory Layout
# CUDA 深入浅出(五):Vectorized Load 和 Shared Memory Layout
前面几篇已经把数据从 global memory 搬到了 shared memory,又从 shared memory 讲到了 register tiling。
现在补上中间一个很工程的细节:
这对应两个常见优化:
它们经常一起出现。vectorized load 负责把 global memory 的连续数据成组搬进来。shared memory layout 负责把这些数据摆成后续计算更爱读的形状。
# 一个 thread 一次只搬一个 float 有点浪费
先看普通 tile load:
每个 thread 搬一个 float。
如果一个 block 要搬 T×T 个元素,就需要很多 thread 或很多轮 load。即使这些 load 已经 coalesced,指令数量仍然不少。
CUDA 里可以让一个 thread 一次搬多个连续元素。比如一次搬 4 个 float:
float4 是 16 字节。一个 thread 发一条 vectorized load,就能拿到 4 个连续的 float。
这不会让你少搬字节。
你本来要搬 16 字节,现在还是搬 16 字节。变化在于:
指令数量少了,地址计算少了,memory pipeline 也更容易吃满。
# vectorized load 需要连续和对齐
float4 不是随便 cast 就能用。
你要满足几个条件:
比如 row-major 矩阵里,同一行的连续列很适合做 float4 load:
它们在内存里挨着。
如果你沿着列搬:
这些地址之间隔着 N 个元素。它们不是一段连续的 16 字节。你不能把这种访问伪装成一个 float4。
所以 vectorized load 第一条规则很朴素:
这和上一篇 coalesced load 的方向一致。warp 里的相邻 thread 读相邻地址,一个 thread 内部也读相邻地址。
# vectorized load 和 coalescing 的关系
coalescing 看的是 warp 里多个 thread 的地址。
vectorized load 看的是单个 thread 一次读几个连续元素。
二者可以叠在一起。
假设一个 warp 里 32 个 thread,每个 thread 读一个 float4:
整个 warp 覆盖了一大段连续内存。
这是一种很舒服的访问形状:
如果每个 thread 都读 float4,但 thread 之间地址乱跳,硬件仍然会发出很多 transaction。vectorized load 不是 coalescing 的替代品。
你需要同时满足:
# 搬进 shared memory 后为什么不直接原样放
假设 global memory 里 B_tile 是 row-major。
为了 coalesced + vectorized load,我们希望 thread 从 global memory 里按行连续读取:
这对 global memory 很友好。
但计算阶段不一定想按这个形状读。
在矩阵乘法内层:
register tiling 里,一个 thread 或一组 thread 往往需要在每个 k 上拿到一小段 B 的列方向数据,或者让不同 lane 以某种固定模式读 B fragment。
如果你把 B_tile 原样存进 shared memory,后面读 shared memory 时可能会出现:
于是很多 GEMM kernel 会在 global -> shared 这一步顺手改变布局。
常见做法包括:
我们先看 transpose。它最容易理解。
# shared memory transpose 的基本想法
假设你从 global memory 里按行读 B:
这是连续读,适合 vectorized load。
读到 register 后,你不一定要按同样位置写进 shared memory。
你可以转置写入:
global memory 侧还是连续读。
shared memory 侧变成转置布局。
后面计算阶段如果更适合按 Bs[col][k] 读,就能得到更顺的 shared memory 访问。
这就是很多优化代码看起来“绕”的原因:global memory 的最佳读法和 shared memory 的最佳读法,不一定是同一种二维布局。
# 搬运和计算想要的布局不同
搬运阶段关心:
计算阶段关心:
如果你只照顾搬运,shared memory 里可能摆得不好。
如果你只照顾计算,global memory load 可能变成 stride 访问。
成熟的 tiled GEMM 会把这两件事拆开:
中间靠 register 做一次短暂停留。
# 一个简化的 B tile 搬运例子
下面的代码不是完整 kernel,只看数据形状。
假设一个 thread 负责从 global memory 读取 B 的连续 4 个元素:
这次 load 读的是:
global memory 很满意。
接着写 shared memory 时,可以原样写:
也可以转置写:
第一种 layout 保留 global 里的 row-major 形状。
第二种 layout 把 B_tile 变成 shared memory 里的转置形状。
选哪个,取决于后面的 compute loop 怎么读 Bs。
# transpose 不是为了炫技
写高性能 GEMM 时,你经常会看到类似这样的 shared memory 定义:
或者:
维度顺序看起来和数学矩阵不一样。
原因通常有两个:
比如计算阶段每个 thread 每轮 k 都要取一组 B 值:
如果 shared memory 里这些值正好连续,thread 就能用更少的地址计算、更好的 bank 分布读出来。
如果它们跨步分布,计算阶段会不断为糟糕 layout 付费。
你在搬运阶段做一次 transpose,是为了让后面几百上千次 FMA 的读取更舒服。
# padding 和 transpose 经常一起用
上一篇讲过,按列读 T×T shared tile 时,如果 T=32,很容易出现 bank conflict。
转置可能改变读写方向,但它不自动消除所有 bank conflict。
所以你经常会看到:
多出来的 +1 是 padding。
它让每一行的 stride 从 32 变成 33。按列读时,地址对 32 个 bank 取模会错开。
transpose 解决“后面想按什么形状读”的问题。
padding 解决“这种读法会不会撞 bank”的问题。
layout optimization 往往就是这两类动作叠在一起。
# vectorized store 到 shared memory 要小心
vectorized load 从 global memory 读出 float4 后,你可能想把它也用 float4 直接写进 shared memory:
这在 layout 原样保存、地址对齐、没有越界时可以工作。
但如果你要 transpose,就不能把 float4 原封不动写进去。v.x、v.y、v.z、v.w 会落到不同的 shared memory 位置。
这时你会看到代码把一个 vector 拆开写:
这看起来比一次 vector store 麻烦,但它换来了后面 compute loop 的整齐访问。
优化不是只看搬运阶段哪一行代码短。你要看整条路径:
如果拆开写 shared memory 能让后续读少冲突、少地址计算、少重排,它就值得。
# 边界处理会破坏漂亮形状
float4 load 最怕边界。
如果矩阵宽度不是 4 的倍数,或者 tile 落在矩阵右边缘,你可能只剩 1 到 3 个有效元素。
这时不能盲目读一个 float4。
常见处理方式有两种:
高性能 kernel 通常会尽量让主路径没有太多 if。边界 tile 单独处理,或者用 predicate 把无效元素写成 0。
比如:
这段比 float4 慢,但它只发生在边界。
主体区域保持 vectorized load,性能才不会被大量边界判断拖住。
# 对齐也要从数据分配开始
float4 要求地址按 16 字节对齐。
如果矩阵起始地址对齐,但每一行的 stride 不是 4 的倍数,某些行的起始位置可能不再 16 字节对齐。
所以工程里经常会关心 leading dimension:
它们不只是矩阵宽度。它们还决定每一行在内存里的起点。
如果你希望每一行都能做 aligned vectorized load,lda 或 ldb 最好满足相应对齐要求。
这也是很多库会使用 padded leading dimension 的原因。数学矩阵是 M×N,内存里的每一行可能会补齐到更适合硬件访问的长度。
# 这一篇放进 GEMM 路径里看
到现在,我们已经有了一条更完整的数据路径:
每一步都有自己的目标。
global memory load 阶段:
shared memory store 阶段:
shared memory load 阶段:
compute 阶段:
这就是为什么优化后的 GEMM kernel 看起来不像课本里的矩阵乘法。
课本公式只写 C=A×B。
GPU kernel 要回答一串更具体的问题:
vectorized load 和 shared memory layout optimization 就卡在这串问题中间。
它们不改变计算结果。它们改变数据在硬件里的走法。
