Chengshu@skadai · 2026.09.30
9,298 字 · 1,574 词 · 约 30 分钟

Stanford CS336 第五讲精读:GPU 硬件、访存瓶颈,以及 Flash Attention 是怎么被逼出来的

斯坦福 CS336(Language Modeling from Scratch, Spring 2025)第五讲完整讲义:为什么这一讲只讲单卡硬件、CPU 优化延迟而 GPU 优化吞吐、SM/warp/block 的执行模型与物理内存层级、同样重要的 TPU 支线,以及那张「波浪形」矩阵乘法性能图的三重谜底(roofline、tiling 对齐、wave quantization)——低精度、算子融合、重计算、访存合并、tiling 这五件提速法宝最终如何拼成 Flash Attention。

本篇属于系列 Stanford CS336 精读 · 第五讲

来源:YouTube 原视频(Stanford Online · CS336 Language Modeling from Scratch · Spring 2025 · Lecture 5: GPUs)

来源说明 这是斯坦福 CS336《Language Modeling from Scratch》2025 年春季第五讲的完整讲义,本讲由 Percy Liang 一人主讲。前四讲把数据、分词、架构、MoE 讲完之后,课程从这里进入「系统」部分——而系统部分的第一站不是并行策略,而是最底层的那块硬件:GPU。这一讲的结构非常清楚:先把 GPU 的物理构造和执行模型拆开(SM、warp、内存层级),再用一张「波浪形」的矩阵乘法性能图当谜面,逐条讲清低精度、算子融合、重计算、访存合并、tiling 五件提速法宝,最后用 Flash Attention 当谜底,把前面所有概念串成一个完整算法。文中 18 张配图均截取自视频对应时刻的幻灯片,并把该时刻的完整观点句(英文原句+中文翻译)拼合进图中。文中出现的时钟周期数、显存与算力倍率、SM 数量、tile 尺寸、提速百分比等数字,都是讲师当场引用的公开博客、论文、推文或他自己的估算,不是本文独立核实的事实;讲师口播中前后不一致或明确表示「不确定」的地方,下文会照实说明。

TL;DR

  • 这一讲是 CS336「系统部分」的地基。它刻意不讲并行策略(那是下一讲),只讲一块加速器内部:GPU 怎么执行、内存怎么摆放、以及为什么同一份数学在不同矩阵尺寸下能差出数倍性能。Percy 的目标很朴素——“让 CUDA 和 GPU 不再像魔法”。
  • CPU 优化延迟,GPU 优化吞吐。CPU 把芯片面积花在分支预测和控制逻辑上,让你那一个线程尽快跑完;GPU 反过来,用一点点控制逻辑去驱动成千上万个 ALU,单个任务慢一点没关系,整体吞吐高就行。
  • 一讲的主角其实是内存,不是计算。算力几十年间涨了 1000 倍到 10 万倍,显存带宽只涨了约 100 倍,互连涨得更慢。到了 H100 这一代,绝大多数算子的瓶颈已经从 FLOPs 变成了数据搬运——“memory movement is the bottleneck in all of this”。
  • 执行模型有三个粒度:block(一整块线程)被分配到 SM 上;block 里的线程按 32 个一组(warp)执行;同一 warp 内所有线程执行同一条指令、只是数据不同(SIMT)。这直接推出一个硬约束:warp 内的条件分支会串行化。
  • 五件提速法宝,全部围绕“少搬数据”:(1) 低精度(FP16 让同一条 ReLU 的访存从 8 字节/FLOP 降到 4 字节/FLOP);(2) 算子融合(把 sin²+cos² 的 5 次 kernel 启动压成 1 次);(3) 重计算(拿算力换访存,三个 sigmoid 的例子把访存从 8 次降到 5 次);(4) 访存合并(利用 DRAM 的 burst 模式,让一个 warp 的访问落进同一块);(5) tiling(把矩阵切成小块搬进共享内存,把全局访存减少 T 倍)。
  • 那张“波浪形”性能图有三个成因:左侧下滑是 roofline 的内存受限区;中段的锯齿来自 tiling 与矩阵尺寸的对齐(能不能整除 burst 段);某几个尺寸处的“悬崖”来自 wave quantization——1792 只需要 98 块 tile,1793 却要 120 块,而 A100 只有 108 个 SM,多出来的那一波只能低利用率地跑完。
  • Flash Attention = tiling + online softmax + 重计算。三个矩阵乘法本身是“教科书级的分块矩阵乘法”,真正难的是 softmax 这个全局操作;online softmax 让它变成可以逐块累计的形式,反向传播再用重计算避免存下 N² 的中间激活。

一、为什么这一讲只讲“单卡”

前四讲把语言模型从数据到架构讲了一遍,从这一讲开始,课程转向系统。有意思的是,系统部分的第一刀没有砍向并行策略,而是砍向了最底层的那块硅:GPU。

Percy 开场先交代了作业进度——作业一当晚截止,作业二马上放出,内容是让你用 Triton 实现 Flash Attention 2 的一部分——然后直接给出了本讲的两个学习目标:

  1. 讲完之后,你应该对 GPU 是怎么工作的感到“舒服”,不再觉得它神秘;
  2. 你应该有信心用 CUDA 去加速自己造出来的新架构、新算子。

他还特意点明了范围:今天只讲“单加速器”,也就是一块 GPU 内部的构造、执行模型和性能特征;把多卡、多机之间的并行留给下一讲。TPU 只会用一个“支线”简单带过,因为它在概念上和 GPU 高度相似。整个讲座因此分成三段:先理解硬件与执行模型,再理解性能(什么让它快、什么让它慢),最后拿 Flash Attention 当一次完整的动手演练——“看看这些东西是怎么拼到一起的”。

一个很诚实的态度贯穿始终:Percy 承认硬件不是他自己的研究领域,所以开头专门列了致谢——Horace He 的博客(那篇讲 GPU 性能的经典文章,里面有很多“反直觉的 GPU 冷知识”,比如“为什么全是 0 的矩阵乘法反而更快”)、CUDA Mode 社区、以及 Google 新出的那本 TPU 读物。他说这一讲的定位是“不深、但尽量完整”的覆盖,想深入的人应该顺着这些资源继续读。

而整场讲座的叙事钩子,从一开始就被埋下了:

演示实验里,随着方阵尺寸增大,GPU 的实际算力曲线会呈现出诡异的波浪形:某些尺寸特别快,相邻的某些尺寸又特别慢。这一讲结束时,Percy 承诺你会觉得“这图看起来完全正常”。

二、GPU 解剖:SM、warp 和内存层级

2.1 先看清大势:算力涨得比内存快得多

讲座先用几分钟铺背景。语言模型的性能由算力驱动(训练 scaling law,换成推理的曲线也一样成立),所以真正推动进步的,是更快的硬件、更高的利用率、更好的并行化。

问题在于,单线程的性能增长早就停了。早年 CPU 靠 Dennard 缩放(更小的晶体管、更高的频率、更低的功耗)一路变快,1980 到 2000 年代这条路就走到了尽头——晶体管密度还在涨,但单线程吞吐不再随之提升。于是整个行业转向并行扩展:不再是“让一次计算更快”,而是“让大量计算同时发生”。Percy 引用了 Bill Dally 主题演讲里那张曲线——从最早的 K20 到 H100,整数运算量呈现超指数增长——并且直言:语言模型想吃到这条曲线的红利,你就必须理解它是怎么来的。

2.2 CPU 与 GPU 是两种不同的生物

CPU 的设计目标是低延迟。程序里有大量的分支、条件判断和控制流,所以 CPU 把很大一部分芯片面积花在控制单元和分支预测上,让少数几个线程尽快跑完。

GPU 反其道而行:芯片上密密麻麻全是 ALU(计算单元),控制逻辑只占很小一块,用来指挥海量的并行计算。它的设计目标是高吞吐——单个任务慢一点无所谓,只要所有任务加起来完成得够快。而且 GPU 的线程极其轻量,可以随时停下来、随时唤醒,切换成本极低,这正是高利用率的前提。

2.3 执行模型:SM、block、warp、thread

GPU 的核心单元叫 SM(streaming multiprocessor,流式多处理器),内部又包含许多 SP(streaming processor,流式处理器)。你可以把 SM 理解成 CPU 里的一个“核”,一块 GPU 里有非常多个;在 Triton 这类编程模型里,你实际上就是以 SM 为单位思考问题的。以 A100 为例,讲师口播说有 128 个 SM(真值见 6.3 节,那里他用的是 108),每个 SM 里又有大量的 SP 和专门的矩阵乘法单元。

执行时有三个粒度:

  • block:一大批线程的集合,一个 block 会被分配到某一个 SM 上执行。SM 是自主单元,block 就是它的工作单位。
  • warp:block 内部的线程并不是逐个执行的,而是每 32 个连续编号的线程编成一组,这组就叫一个 warp。
  • thread:warp 里最小的工作单位,持有自己的寄存器和私有数据。

关键约束来自 SIMT(single instruction, multiple threads):同一个 warp 内的所有线程,永远在执行同一条指令,只是作用在不同的数据上。这意味着如果你在 warp 内写了一个 if/else,硬件没法让两拨线程同时执行两个分支——它只能先让一部分线程干活、把另一部分“睡觉”,再反过来。于是:

// 同一个 warp 里,两拨线程被迫串行执行
if (threadIdx.x < 4) {
    A();   // 前 4 个线程执行,其余 28 个被挂起
} else {
    X();   // 再让其余线程执行,前 4 个被挂起
}

这就是为什么“在超大规模并行的计算单元里塞条件分支”是个坏主意——它把你花钱买来的并行度直接打了对折。

2.4 内存层级:物理距离就是速度

如果说 GPU 是“计算的机器”,那么它同样是“搬数据的机器”,而且后者在当前这个时代可能更重要。要理解这一点,必须看物理布局——因为在这种速度下,内存离 SM 的物理距离,直接决定了访问它需要多少个时钟周期。

图 1|越靠近 SM 的内存越快:L1 与共享内存在 SM 内部,L2 在片上,全局内存在芯片之外

层级大致是:

  • 寄存器 / L1 / 共享内存(shared memory):位于 SM 内部,最快——讲师给的量级是约 20 个时钟周期;
  • L2 缓存:不在 SM 里,但仍在同一块芯片上,物理上紧挨着 SM,慢一个数量级——大约 200 到 300 个时钟周期;
  • 全局内存 / DRAM(HBM):在芯片之外,通过芯片边缘那些黄色连接器连过来,最慢。

“这个 10 倍的差距会狠狠地咬你。“如果一段计算必须频繁访问全局内存,SM 很可能算完了手头的工作就得空转等数据,利用率直接掉下去。讲座后面反复强调的每一件提速技巧,本质上都是在做同一件事:让数据尽量留在离 SM 更近的地方,把对全局内存的访问压到最少。

配合物理布局,还有一个逻辑内存模型用于编程:寄存器(单个数值,最快)、局部内存(local)、共享内存(shared)、全局内存(global),另外还有很少用到的常量内存(constant)。这里有一条非常重要的规则:每个线程可以访问自己的寄存器和共享内存,但跨 block 的信息传递必须经过全局内存。 这直接决定了后面所有 kernel 设计的形状——理想情况下,一个线程先把它需要的那一小片数据载入共享内存,然后所有线程都在这片快速内存上愉快地干活;反过来,如果每个线程都要去全局内存里东抓一把西抓一把,性能必然一塌糊涂。

2.5 GPU 为什么能赢,以及它本来不是为深度学习设计的

Percy 总结了 GPU 大获成功的三个原因:堆规模极其容易(想要更强,加 SM 就行,不必拉高频率、不必担心散热失控);编程模型相对好理解(每个 SM 内是一条指令作用于多份数据,尤其适合规整的矩阵运算);线程非常轻量(可以随时停止和启动,几乎没有状态负担,这让 SM 能一直保持高占用)。

而 GPU 最初是“图形处理器”,很长时间里并不做科学计算。后来研究者发现它可编程,就硬是用图形管线来跑矩阵乘法——讲座里引用了早期那篇“用图形硬件做快速矩阵乘法”的论文,讲他们怎么“黑”进纹理缓冲区(texture buffer)来让 GPU 做 matmul。再后来,NVIDIA 意识到矩阵乘法在深度学习里是被祝福过的操作(blessed operations):只要你在跑深度学习,绝大部分工作量就是矩阵乘法。于是在 V100 那一代引入了专门的张量核心(tensor cores)。

图 2|张量核心出现之后,矩阵乘法与非矩阵乘法算力拉开了一到两个数量级的差距

这张“不同代际 NVIDIA GPU 的 TFLOPs”图里,橙色线是矩阵乘法算力,蓝色线是非矩阵乘法算力。V100 之后两条线之间的鸿沟大到夸张,讲座里给出的量级是矩阵乘法比浮点运算快 10 倍以上。由此得到一条对架构设计者的硬性建议:你的网络里绝大部分计算必须落在矩阵乘法上,“如果你造了一个不以 matmul 为主的神经网络,你会陷入大麻烦”。

2.6 最该记住的一张图:算力在飞,带宽在爬

接下来是本讲最“该被记住”的一张图:把语言模型训练栈里不同组件的增长速度画在一起。

  • 蓝线:GPU 到主机(服务器)的连接带宽。PCIe、NVLink 这些互连当然也在涨,但涨得最慢。
  • 绿线:全局内存速度。从 GDDR 到 HBM2E,绝对值涨了约 100 倍——注意这是对数坐标。
  • 灰线:算力。同样是对数坐标,涨了大约 1000 倍到 100000 倍。

结论非常直白:早年你的瓶颈可能是 FLOPs(算力不够做矩阵乘法),但到了 H100 这一代,“你的瓶颈几乎注定是内存”。而且这个趋势不会逆转——DRAM 极难继续快速缩放,这条鸿沟只会越来越宽。所以设计任何“硬件友好”的算法,都必须越来越认真地思考访存。这也是他反复敲打的唯一主题。

2.7 这一节的事实清单

讲座在这里做了个小结,把“关于 GPU 你只要记住这些”压缩成三条:

  1. GPU 是大规模并行系统,同一条指令作用在海量线程上,里面有许多叫 SM 的“核”;
  2. 算力与矩阵乘法涨得远比内存快,这就是 GPU 的性能特征;
  3. 但也不是所有内存都慢——存在一个内存层级,有些内存极快。只要善用这个层级,你依然能拿到极好的性能。

三、支线:TPU 长什么样

讲完 GPU,Percy 用一个“side thread”专门聊了几分钟 TPU。之所以值得讲,是因为替代性加速器在概念上和 GPU 高度相似,理解了 GPU 就能迁移过去。而这一讲能补上这一节,本身也是因为去年资料太少、今年 Google 那本 TPU 读物出来了。

TPU 的结构可以这样对应过来理解:

  • tensor core 大致对应 GPU 的 SM,是一个可以独立操作数据的原子单元;
  • scalar unit 相当于控制单元,也能做类似 CPU 的任意操作;
  • vector unit 负责逐元素操作——你有一个向量要做 entry-wise 运算,交给它;
  • MXU 是一大块专门做矩阵乘法的硬件;
  • 向量内存 / 片上内存(SMEM) 极快,位于片上;HBM 位于片外,慢而大。

心智模型和 GPU 完全一致:外面是慢而大的内存,里面是快而小的内存,中间有一块专门做矩阵乘法的硬件。 差别主要在“多块加速器怎么连起来”,那是下一讲并行的话题。TPU 的 tensor core 反而比 SM 更“简单”,因为它只做矩阵乘法,不像 GPU 那样试图兼顾各种操作。

讲座里有个有趣的提问:它为什么叫 tensor?答案是“yes and no”——它操作的对象确实可以是任意张量,但 MXU 实际执行的运算永远是批量矩阵乘法,不会做更复杂的张量运算。

图 3|TPU 支线:tensor core 对应 SM,标量/向量单元对应控制与逐元素运算,MXU 专做矩阵乘法

四、谜面:一张“波浪形”的性能曲线

从第 25 分钟起,讲座进入第二段——理解性能。Percy 没有直接列技巧,而是先抛出一张图作为谜题。这张图的设定很简单:把两个方阵相乘,横轴是方阵边长 N,纵轴是实际达到的算力(可以理解为硬件利用率)。直觉是“N 越大、利用率越高”(因为固定开销被摊薄了),但实际曲线是一堆诡异的波浪:四条不同颜色的线,每条都起起伏伏,某些相邻尺寸之间会出现巨大的落差。

他甚至半开玩笑地承诺:“这一部分结束时,你会觉得这张图完全正常,是 GPU 的自然行为。”

第一层解释是系统课里的经典工具——roofline model(屋顶线模型):

图 4|roofline 模型:左侧是内存受限区,右侧是吞吐受限区

  • 曲线的左半部分(对角线)处于内存受限区:算力不足以喂饱计算单元,你被“每访问一个字节能换来多少 FLOPs”卡住;
  • 曲线的右半部分(平台)处于吞吐受限区:所有矩阵乘法单元都在满负荷工作,这才是理想状态;
  • 我们的目标就是把算法推向右侧,避免停留在左侧。

但 roofline 只解释了曲线的“大形状”,解释不了那些锯齿。为了拆清剩下的谜团,讲座列出了一张“技巧清单”,先排除唯一一个非内存因素(条件分支),然后逐个讲解五件与内存相关的法宝:低精度、算子融合、重计算、访存合并、tiling。最后再回到这张图,把每一处诡异的起伏归因到具体的机制上。

五、五件提速法宝

5.1 唯一的非内存问题:条件分支

前面已经讲过 SIMT 的后果。这里只补充一句:“你应该很清楚,不该在超大规模并行的计算单元里放条件判断。“因为 warp 内分歧会强制串行化,把并行度白白浪费掉。剩下的技巧,全都与内存有关。

5.2 法宝一:低精度

这是“最重要、也最应该经常用”的一件法宝。Percy 指出,前面那张让人兴奋的算力增长曲线里其实有个障眼法:真正驱动 GPU 进步的一大因素,是数值表示的变化——从 FP32 到 FP16、到 INT8,每降一次精度就带来数量级收益。

为什么?因为位数少了,要搬的东西就少了。即便你仍然要从全局内存里读,这些数据占用的字节数也大幅下降。用一个最简单的算例(ReLU,y = max(0, x),作用在长度 n 的向量上):

# FP32 版本
Memory access: 读 x (4 bytes) + 写 out (4 bytes) = 8 bytes
Operations:    1 次比较 + 1 次 FLOP
Intensity:     8 bytes / FLOP

# FP16 版本
Memory access: 读 x (2 bytes) + 写 out (2 bytes) = 4 bytes
Operations:    1 次比较 + 1 次 FLOP
Intensity:     4 bytes / FLOP

FLOPs 完全没变,访存直接减半——“某种意义上,你白拿了一倍的显存带宽”(前提是精度真的够用)。

图 5|低精度提升算术强度:同一个 ReLU,FP16 的访存强度只有 FP32 的一半

但低精度不是无脑把所有东西都降下去。真实做法是混合精度:矩阵乘法的输入用 16 位,但乘法与累加在 32 位里做(FP32 accumulator),因为部分和的累加需要更高精度,算完再降回 16 位。另外,像指数函数这类需要大动态范围的操作,可能需要 BF16 而不是 FP16,否则会爆掉或直接归零。低精度训练能稳不稳,是一堆精细工程问题;但只要能稳住,收益就是“把瓶颈吞吐翻倍”。

5.3 法宝二:算子融合

第二件法宝非常直观。讲座借用了 Horace He 的“工厂”比喻:工厂是你的计算单元,传送带是内存带宽。厂房再大、再多,传送带带宽是有限的,整体产出就被卡死在传送带上。而一个很隐蔽的浪费方式是这样的:

  • 坏模式:把方块从内存搬进计算单元 → 变成三角形 → 搬回内存 → 再搬进来 → 变成圆形 → 又搬回内存……数据在内存和计算单元之间来回横跳,每一次都付带宽成本。
  • 好模式:既然这些操作之间没有依赖,那就方块 → 三角形 → 圆形 → 矩形,全程留在计算单元里,只在最后把结果写回去一次。

这就是 kernel fusion(算子融合) 的心智模型。

具体例子是一个看起来很无辜的模块:输入 x,输出 sin²(x) + cos²(x)。如果你的代码是朴素写法,PyTorch 的计算图会启动一长串 CUDA kernel——算 sin、算 cos、算 sin²、算 cos²、再相加——每一步的结果都要写回全局内存再读出来,正是上面那个“坏模式”。但如果你用 torch.compile,编译器能看出这几个操作一共只用到很少的内存,把它们融合成一次 kernel 调用,中间结果根本不落地。

Percy 的建议很直接:“如果你还没用 torch.compile,你应该认真考虑在所有地方都用它。” 作业里也会让你体验。

5.4 法宝三:重计算

第三件法宝的思路是用算力换访存——既然算力过剩、带宽稀缺,那就拿多出来的算力去抵消访存。

回想反向传播:前向过程算出的激活值必须存下来(写全局内存),反向时再读回来,这种“存下来、读回去”本身就很贵。重计算(recomputation) 的做法是:前向时干脆不存激活值,反向时现场重算一遍。

讲座用一个“三层 sigmoid 叠起来”的小例子把账算得很清楚:

  • 朴素做法:前向读一次 X、写三次激活(S1、S2、out);反向读三次激活、写一次 dX,合计 8 次访存。而且这个运算里连一个矩阵乘法都没有,算术强度极低。
  • 重计算做法:前向不存激活,只读 X、写 out;反向读入 dout 和 X,在片上现场重新算出 S1、S2、out,再写出 dX。合计 5 次访存——“完全相同的计算,我只用了 5/8 的访存”。

代价是多算了三个 sigmoid。但如果你本来就因为访存受限而让计算单元空转,这笔交易简直太划算:用你过剩的资源,去买你稀缺的资源。

图 6|重计算:不存激活、反向重算,8 次访存压到 5 次

Percy 特意指出,这和梯度检查点(gradient checkpointing)是同一个技术,但目的不同:检查点是为了省显存、防止 OOM;这里是为了执行速度。

5.5 法宝四:DRAM 的 burst 模式与访存合并

这是 Percy 说自己“在真正研究硬件模型之前并不知道”的一件冷知识。GPU 里那块慢速全局内存(DRAM)之所以慢,是因为信号需要被搬到放大器(amplifier)里,这一步是瓶颈。硬件为了摊薄这个代价,做了一个优化——burst mode(突发模式):

当你去读一个地址时,你拿回来的不只是那个值,而是一整块。

图 7|DRAM 以 burst 为单位返回数据:读 1 个值,硬件顺带给你一整段

讲座给的示意是:地址空间被切成 burst 段(幻灯片上的基础例子是 16 字节地址空间、4 字节一段,也就是一次读回 4 个地址的数据;并注明现在常见的 burst 段已经到 128 字节甚至更大)。于是:

  • 如果你按随机顺序访问,每次都得付一次完整代价,吞吐低;
  • 如果你顺序访问,第一次读回 0–3,第二次读回 4–7……同样的带宽能换回 4 倍的有效吞吐。

这就引出访存合并(memory coalescing):如果同一个 warp 内 32 个线程的访问都落在同一段 burst 里,硬件就能把它们的请求合并成一次,一次拿回全部数据——“这会把你的访存吞吐直接提高 4 倍”。

图 8|访存合并:同一个 warp 的线程落在同一段 burst 里,请求才会被合并

这条规则会直接咬到矩阵乘法:把矩阵从全局内存读出来时,是按行遍历还是按列遍历,决定了访存能不能合并。幻灯片里那个“每个线程沿列走”的模式看起来更自然,实际上很慢——同一个时刻,不同线程读的是相隔很远的地址,落进不同的 burst 段,等于把每次读都放大成了一整块的无用功;反过来让相邻线程沿行前进,所有访问才能落进同一段。

5.6 法宝五(大头):tiling

最后一件、也是最重要的一件,是 tiling(分块):把访存成组地组织起来,让数据“搬进来一次、用够再走”,从而最小化对全局内存的访问。

讲座用矩阵乘法完整推了一遍。朴素的矩阵乘法里,每算一个输出元素,都要沿着 M 的一行和 N 的一列做内积,于是同一个矩阵元素会被反复从全局内存里读出来,而且访问既不合并、还高度重复——“这会非常慢”。

tiling 的理想形态是:花一次时间把一小块数据从全局内存搬进共享内存,然后在共享内存里做完所有能做的计算,再换下一块。 具体算法是:

# 把 M、N 都切成 T×T 的小块
for i in range(0, N, T):
    for j in range(0, N, T):
        acc = 0
        for k in range(0, N, T):
            load M[i:i+T, k:k+T] -> shared memory   # 全局 → 共享,一次搬一块
            load N[k:k+T, j:j+T] -> shared memory
            acc += M_tile @ N_tile                  # 所有计算都在共享内存里做
        P[i:i+T, j:j+T] = acc

图 9|tiling:把大矩阵切成小块,载入共享内存后反复复用

图 10|tiling 的实现:把 M、N 切成小块依次载入共享内存,在片上完成累加再写回

收益可以用一句“tiling 数学”概括:设矩阵边长为 N、tile 边长为 T。

  • 不做 tiling:每个输入元素都要从全局内存读 N 次;
  • 做 tiling:每个元素从全局内存只读 N/T 次,另外 T 次是从共享内存读的。

总读取量当然不可能减少(矩阵乘法的定义摆在那里),但读的来源从“慢的全局内存”变成了“快的共享内存”——全局访存减少了一个因子 T。tile 越大、共享内存越装得下,收益越大。

但 tiling 也带来一大堆新的复杂度,这正是 GPU 性能难以预测的根源:

  • 整除性:tile 尺寸 128 看起来很整,256×256 的矩阵刚好是 2×2 块;但 257 就糟了——你得用 6 块 tile 去覆盖,最右边那两块几乎是空的。而每块 tile 会被分配到一个 SM,意味着有 SM 在空转。必须挑 tile 尺寸,让矩阵维度尽量被整除。
  • 共享内存上限:tile 不能无限大,装不下就没意义。
  • 访存合并:装载 tile 时的遍历顺序也要能合并。
  • 与 burst 的对齐:这是最阴的一层。理想情况下每段 burst 刚好对齐一行 tile,读一块 tile 只需要取几段 burst;但只要矩阵尺寸“多加一个元素”,整行的对齐就被顶偏,一行数据横跨两段 burst,于是访存次数直接翻倍。

图 11|对齐问题:多加一个尾元素,burst 段与数据布局错位,访存翻倍

对应的解法是padding(填充):把矩阵尺寸补齐到让 burst 段与 tile 尺寸对齐的整数。讲座因此给出一条非常实用的经验:“小心 2 的幂”。

六、把谜题解开

有了上面这些工具,那张“波浪图”就可以逐层破解了。

6.1 先看几个真实的“民间优化”

tiling 与对齐的威力有多大?讲座举了一个真实例子:Andrej Karpathy 说 nanoGPT 有史以来最戏剧性的优化,是把词表大小从 50257 改成 50304——也就是最接近的 64 的倍数。加了 47 个“完全没用”的维度,却拿到了大约 25% 的加速,因为它让矩阵走上了占用率高得多的 kernel 路径。这条推文(2023 年 2 月)后来成了“形状对齐有多重要”的经典注脚。

图 12|Karpathy:把词表从 50257 改成 50304(64 的倍数),换来约 25% 加速

6.2 成因一:tiling 对齐

把那张性能图上的点按“矩阵尺寸能被几整除”重新着色,锯齿立刻有了规律:能被 32 整除的点高居紫色区域,能被 16 整除的仍在高处,K=8、K=2 依次下沉,完全不能整除(比如素数尺寸)的点跌到最底。原因正是上一节的 burst 对齐——一旦整除性变差,tile 就无法对齐 burst 段,访存被成倍放大。

图 13|按整除性着色后,锯齿立刻有了规律:K 越大,性能越好

6.3 成因二:wave quantization

但图上还有几处“悬崖”解释不了:从 1792 到 1793,只加了一个维度,性能为什么能塌掉一大截?讲座把账一步步算开:

  • 假设用 256×128 这个很自然的 tile 尺寸(之所以自然,是因为 GPU 的矩阵乘法单元本身就在大约 128 的尺度上工作);
  • 对 1792 而言,tile 数是 (1792/256) × (1792/128) = 7 × 14 = 98 块;
  • 尺寸变成 1793 后要向上取整,变成 8 × 15 = 120 块;
  • 而 A100 只有 108 个 SM。

于是 108 块 tile 一次性铺满所有 SM,剩下 12 块只能等第一波跑完、在极低的利用率下再跑一波。性能曲线上就表现为“先正常、然后掉下悬崖、再慢慢收尾”。这个现象叫 wave quantization(波量化)。

图 14|wave quantization:98 块 tile 全给 SM;120 块则要跑两波,第二波只有 12 块

一处诚实的标注:讲师在讲 GPU 执行模型时口播“A100 有 128 个 SM”,但在 wave quantization 这一节使用的是 108。A100 的真实规格是 108 个 SM,后者的计算方式与公开规格一致,前文那个 128 应属口误——本文按 108 理解。

6.4 第二部分小结

讲座把这一段的“提速法宝”收成一张清单:

  1. 减少总访存:算子融合,把多个算子合成一次 kernel;
  2. 访存合并(coalescing):顺序访问、让一个 warp 落进同一段 burst,白赚带宽;
  3. 把数据搬进更快的存储:tiling,让复用发生在共享内存里;
  4. 用别的资源换访存:低精度(换数值精度)、重计算(换算力)。

核心心法只有一句:你必须非常认真地对待内存在 GPU 性能中的角色。

七、把一切拼起来:Flash Attention

讲座最后一段是“动手演练”:用刚才学到的全部工具,把 Flash Attention 从头讲一遍。Percy 说,他不希望这些概念只是“一堆互不相关的 GPU 冷知识”,而希望大家看到它们如何组成现代高性能 Transformer 的一块基石。

Flash Attention 论文自己的说法是:用 tiling 和重计算两项成熟技术,实现对 HBM 访问的次二次方(sub-quadratic HBM accesses)。 注意措辞——计算量本身不可能降到次二次方(注意力就是要算那么多),能降到次二次方的是对全局内存的访问次数。这正好呼应本讲主题:既然内存是瓶颈,那就想办法让瓶颈项不要是二次方。

先看简单的一半。注意力由三个矩阵乘法(QKᵀ、softmax、乘以 V)加中间一个 softmax 组成。矩阵乘法部分就是前面讲的教科书式分块矩阵乘法:把 K、Q 切成小块,拷进 SRAM,相乘并累加。

图 15|Flash Attention 的第一半:KQV 就是一次标准的分块矩阵乘法

真正难的是 softmax。softmax 是一个全局操作——它按行工作,必须把整行的和(以及最大值)求出来才能归一化。这和“分块、永远不写回大矩阵”的理想直接冲突:如果你在逐块计算,你怎么知道后面还有没有更大的值、还有多少项要加?

答案是 online softmax(在线 softmax),出自 Milakov 与 Gimelshein 2018 年的工作。它的核心是维护两个可以增量更新的量:

# 标准数值稳定 softmax:需要先看完整行
#   m = max(x_1..x_n); y_i = exp(x_i - m) / sum_j exp(x_j - m)

# online softmax:只维护"到目前为止"的 max 和归一化和
m = -inf      # 目前为止见过的最大值
d = 0         # 目前为止的指数和(归一化项)
for x_j in stream:
    m_new = max(m, x_j)
    d = d * exp(m - m_new) + exp(x_j - m_new)   # 用新最大值校正旧和,再并入新项
    m = m_new
# 走完整条流之后,d 就是归一化项,可以直接算出 y_i

关键洞察是:当最大值从 m 更新到 m_new 时,旧的指数和可以按 exp(m - m_new) 这个因子整体缩放(幻灯片里管这叫 telescoping 式的推导),于是你不需要一开始就看到整行数据,只要有一个流式序列就能算出正确的归一化项。

图 16|online softmax:用可增量更新的 max 与归一化和,把 softmax 变成逐块可算

把这两半合起来,Flash Attention 的前向就是:

# Flash Attention 前向(示意)
for q_tile in Q_blocks:
    m, l, O = -inf, 0, 0          # 运行最大值、归一化和、输出累加器
    for k_tile, v_tile in zip(K_blocks, V_blocks):
        S = q_tile @ k_tile.T     # 分块做 QKᵀ
        m_new = max(m, rowmax(S))
        P = exp(S - m_new)        # 用当前最大值做数值稳定
        l = l * exp(m - m_new) + rowsum(P)   # 在线更新归一化项
        O = O * exp(m - m_new) + P @ v_tile  # 在线更新输出
        m = m_new
    O = O / l                     # 走完所有块,归一化项已就绪

图 17|把所有 tile 走一遍,就同时得到了输出和归一化项

讲座里有个同学追问了一个很关键的点:既然归一化项要看完所有 tile 才凑齐,那是不是还得回访一遍?Percy 的回答是:是的,你必须把 N² 个 tile 都走完一次才能输出 softmax;但走完的那一瞬间,你已经同时持有输出累加器和归一化项,可以直接除出结果,不需要为了 softmax 再回去重算。 真正省下来的是“不必物化那个 N² 的软最大化矩阵”——这正是论文所说的次二次方访存。

最后是反向传播。反向里如果老老实实存下 softmax 相关的激活,那又是 N² 大小的存储量,这正是“永远不要存 N² 东西”这条原则要避免的。所以 Flash Attention 在反向时逐块重计算这些中间量——再一次回到“用算力换访存”。Percy 没有展开推导,但指出其余部分是标准的梯度计算,只是按 tile 来做。

图 18|全讲收束:硬件缩放决定什么能被训练,而 GPU 的算力缩放方式在逼你重新思考数据搬运

我的笔记

  1. 这一讲的真正主角是内存,不是计算。 算力涨了 1000 到 100000 倍,带宽只涨了约 100 倍;理解 GPU 性能,先要接受“数据搬运才是瓶颈”这个前提。
  2. CPU 优化延迟,GPU 优化吞吐。 这个差别解释了从芯片面积分配、线程设计到“为什么可以容忍单个任务很慢”的一切。
  3. SIMT 让 warp 内的条件分支变成串行。 同一 warp 的 32 个线程永远执行同一条指令,“在并行计算单元里塞 if/else”是自废武功。
  4. 低精度是最划算的一招。 位数减半、访存减半,FLOPs 一点没变;但累加必须留在 FP32,需要大动态范围的地方要换 BF16——省的是带宽,赌的是数值稳定性。
  5. 算子融合、重计算、tiling 是同一件事的三种变体:都是想方设法让数据留在离 SM 更近的地方,少走那趟慢速全局内存。重计算尤其优雅——拿过剩的算力去买稀缺的带宽。
  6. DRAM 以 burst 为单位返回数据,所以顺序、对齐、可合并的访存能白赚数倍吞吐;反过来,“多加一个元素”就可能让访存翻倍、性能塌方。
  7. 那张波浪图的三个成因:roofline(内存受限区)、tiling/整除性对齐(锯齿)、wave quantization(98 块 vs 120 块 tile 对 108 个 SM 的错配)。矩阵形状不是审美问题,是性能问题。
  8. Flash Attention 不是魔法,而是 tiling + online softmax + 重计算的组合拳。 它把二次方的 HBM 访问压成了次二次方;最难的一步是让 softmax 这个全局操作变成可以逐块增量计算的形式。

附:课程信息与时间轴

时间 内容
0 开场:作业一晚截止,作业二是用 Triton 实现 Flash Attention 2
0 本讲目标:让 CUDA 和 GPU 不再神秘
0 待解之谜:为什么矩阵乘法性能呈波浪形
1 第二个目标:学会用 CUDA 加速自己的算法
2 致谢资源:Horace He 的博客、CUDA Mode、Google 的 TPU 读物
2 范围:只讲单卡硬件,并行留到下一讲
3 背景:算力越多,语言模型越好
4 Dennard 缩放与摩尔定律的终结
5 单线程性能见顶,转向并行扩展
5 Bill Dally 演讲:整数运算量的超指数增长
6 CPU vs GPU:控制单元 vs 计算单元
7 CPU 优化延迟,GPU 优化吞吐
8 GPU 解剖:SM 与 SP
9 A100 的 SM 数量(讲师此处口播 128,后文按 108 计算)
10 计算之外还有内存,后者可能更重要
11 内存层级:L1/共享内存 → L2 → 全局内存
12 20 个时钟周期 vs 200–300 个时钟周期
13 执行模型:block、warp、thread
14 warp = 32 个连续编号线程,SIMT
15 逻辑内存模型:寄存器、局部、共享、全局、常量
16 支线:TPU 长什么样
17 tensor core ≈ SM;标量单元、向量单元与 MXU
19 GPU 为何成功:易堆规模、编程模型清晰、线程轻量
20 早期研究者用图形硬件硬做矩阵乘法
21 张量核心让矩阵乘法与非矩阵乘法拉开数量级差距
22 相对缩放图:互连、显存、算力
23 H100 时代,瓶颈已经从 FLOPs 变成内存
24 回到那张谜题图
26 roofline 模型:内存受限区与吞吐受限区
27 技巧总览:条件分支 + 五件内存法宝
28 条件分支:warp 内分歧的代价
29 法宝一:低精度
30 FP32 → FP16 → INT8 带来数量级收益
30 ReLU 算例:8 bytes/FLOP vs 4 bytes/FLOP
31 混合精度:16 位输入、FP32 累加
32 需要大动态范围时改用 BF16
33 法宝二:算子融合
33 Horace He 的工厂比喻:传送带就是内存带宽
34 来回搬数据的坏模式 vs 一次算完的好模式
35 sin²+cos² 的朴素实现会启动一长串 kernel
37 用 torch.compile 自动完成融合
37 法宝三:重计算
38 三个 sigmoid 的例子:8 次访存 → 5 次
41 与梯度检查点是同一技术、不同目的
41 法宝四:DRAM 的 burst 模式
43 访存合并:一个 warp 落进同一段 burst
44 矩阵乘法的行/列遍历顺序决定能否合并
46 法宝五(大头):tiling
47 朴素矩阵乘法的重复访存问题
49 理想形态:搬一次、在共享内存里算够
51 tiling 的收益:全局访存减少 T 倍
53 tiling 的复杂性:整除性与空转的 tile
54 影响因素:访存合并、共享内存大小、维度整除
55 burst 对齐与 padding:小心 2 的幂
57 Karpathy 的 tweet:词表 50257 → 50304
58 回到谜题:波浪线的三重成因
59 成因一:按整除性着色,tiling 对齐
61 成因二:wave quantization(1792 vs 1793)
63 第二部分小结:四类提速手段
64 第三部分:Flash Attention
65 论文原话:用 tiling 与重计算实现次二次方的 HBM 访问
66 三个矩阵乘法 + 一个 softmax
67 softmax 是全局操作,与分块天然冲突
68 online softmax(Milakov & Gimelshein, 2018)
69 不必物化 N² 的软最大化矩阵
70 前向:逐块累加并校正最大值
71 走完所有 tile,归一化项就绪
71 反向传播:逐块重计算,避免存 N² 激活
73 全讲总结:内存搬运才是瓶颈

说明:本文是视频内容的整理、翻译与转述,观点均来自主讲人 Percy Liang;文中代码为讲座中算法与公式的整理版本,非官方作业代码。课程中引用的时钟周期、倍率、SM 数量、tile 尺寸、提速百分比与推文内容,多来自公开博客、论文、推文或讲师本人的估算,请自行核实;讲师口播前后不一致之处(如 A100 的 SM 数量)文中已照实标注。

讨论

这里是静态站点,没有内嵌评论区。如果这篇文章对你有用,欢迎通过 RSS 订阅后续更新。