Skip to content

让深度学习跑得飞快:从第一性原理理解性能

原文标题:Making Deep Learning Go Brrrr From First Principles 作者:Horace He(PyTorch 团队) 原文链接:horace.io/brrr_intro.html

假设你想提升深度学习模型的性能,你会怎么做?很多人会从一袋碰巧能用的「技巧」里随手抓几个——也许是之前奏效过,也许是在推特上看到的:「用 in-place 操作!」「把梯度设为 None!」「装 PyTorch 1.10.0 但千万别装 1.10.1!」

用户之所以会采取这种临时拼凑的办法,是可以理解的:在现代系统(尤其是深度学习)上做性能优化,常常更像炼金术而非科学。话虽如此,从第一性原理出发仍然能排除掉一大片方案,让问题变得容易处理得多。

举个例子,在数据集上训练出好的性能也充满了猜测。但如果你的训练损失远低于测试损失,说明你处于「过拟合」区间,此时再去增加模型容量纯属浪费时间;反过来,如果训练损失与验证损失几乎相同,再去做正则化也是浪费时间。

同理,你可以把深度学习系统的效率理解为由三个不同部分构成:

  • 算力(Compute):GPU 实际执行浮点运算(FLOPS)所花的时间
  • 内存(Memory):GPU 内部搬运张量所花的时间
  • 开销(Overhead):其他所有时间

正如训练 ML 模型一样,知道当前处于哪个区间,就能帮你把优化聚焦在真正重要的地方。例如,如果你把所有时间都花在了内存搬运上(即处于内存带宽受限区间),那么提高 GPU 的 FLOPS 毫无帮助。另一方面,如果你把所有时间都花在大矩阵乘法上(即算力受限区间),那么把模型逻辑重写成 C++ 以减少开销也不会奏效。

所以,如果你想让你的 GPU 一直「brrrr」地跑下去,我们就来聊聊你的系统可能把时间花在的这三个部分——算力、内存带宽和开销。

让 GPU 保持高效运转的工程师们

在苦涩教训的背后,是一支让 GPU 高效运转的工程师大军。图片来源:Gwern

注意:本文大部分示例基于 GPU 与 PyTorch(因为我在 PyTorch 团队工作),但这些原理几乎可以推广到所有硬件和框架。

1. 算力(Compute)

优化深度学习系统的一个视角是:我们希望最大化处于算力受限区间的时间。你为那 312 teraflops 付了钱,理想情况下就该拿到这 312 teraflops。但为了让你昂贵的矩阵乘法物有所值,你需要减少花在其他部分上的时间。

为什么只强调最大化算力,而不说内存带宽?原因很简单——你可以减少开销或内存开销,但(大多数情况下)如果不改变实际执行的计算,你就无法减少所需的计算量。

让最大化算力利用更困难的是算力与内存带宽的增长速度差异。看看这张表,对比的是 CPU FLOPS 翻倍周期与内存带宽翻倍周期。[链接待更新:原文此链接拼写为 "bandwith"(少一个 d),正确拼写为 "bandwidth";两种写法当前均无法直接访问(前者连接失败,后者被 ACM 反爬返回 403)。]

算力增长 vs 内存带宽增长

可以把算力想象成一座工厂:我们向工厂发送指令(开销)、运入原材料(内存带宽),这一切都是为了让工厂高效运转(算力)。

工厂比喻

如果工厂的效率提升速度超过了我们向它供应原材料的速度,那么工厂要达到峰值效率就会越来越难。

工厂效率翻倍

即使工厂的规模(FLOPS)翻倍了,如果带宽跟不上,性能也不会跟着翻倍。

这种利用算力越来越难的状况,一方面意味着 ML 系统工程师永远有饭吃,另一方面也让我们理解瓶颈变得更重要了。

关于 FLOPS 还有一点补充:现代机器学习加速器都有专门做矩阵乘法的硬件,比如 NVIDIA 的「Tensor Cores」。

A100 规格

所以,如果你不做矩阵乘法,就只能达到 19.5 teraflops,而不是标称的 312。注意这不是 GPU 独有的——事实上,TPU 比 GPU 更不通用。

GPU 在矩阵乘法以外的事情上都慢得多,这初看似乎是个问题——那我们的 Layer Norm、激活函数等其他算子怎么办?其实,这些算子从 FLOPS 角度看只是四舍五入误差。比如看这张来自这篇论文的 BERT 各算子类型 FLOP 计数表,其中「Tensor Contraction」就是矩阵乘法。

BERT 各算子 FLOP 占比

可以看到,我们的非矩阵乘法算子总共只占了 FLOPS 的 0.2%,所以 GPU 算非矩阵乘法算子慢 15 倍根本无所谓。

不过在本例中,归一化和逐点算子的 FLOPS 分别只有矩阵乘法的 250 分之一和 700 分之一。

那么,为什么我们的非矩阵乘法算子会花比应有时间多得多的耗时?

回到我们的类比,罪魁祸首往往是往工厂搬运原材料的耗时。换句话说,就是内存带宽。

2. 带宽(Bandwidth)

带宽开销本质上就是把数据从一个地方搬到另一个地方所付出的代价。这可能包括把数据从 CPU 搬到 GPU、从一个节点搬到另一个节点,甚至是从 CUDA 全局内存搬到 CUDA 共享内存。这里我们重点关注的正是最后一种,它通常被称为「带宽开销」或「内存带宽开销」。

另外两种(通常分别称为「数据传输开销」和「网络开销」)当然也很重要,但深入讲分布式性能的话,这篇文章就永远写不完了。

要理解内存带宽开销,我们回到工厂类比。

虽然工厂是我们实际干活的地方,但它不适合当大批量仓库。很大程度上是因为既然这里在干实际工作,所有存储都是为「用起来快」而优化的(SRAM),而不是为了「存得多」。

那么实际的成果和原材料存哪里?典型做法是建一个仓库,大概选址在土地便宜、空间充足的地方(DRAM)。然后我们可以在工厂和仓库之间来回运货(内存带宽)。

工厂与仓库之间的搬运

这种在计算单元之间来回搬东西的代价,就是所谓的「内存带宽」开销。顺便说一句,你 GPU 的 DRAM 就是 nvidia-smi 里显示的那部分,也是你那些可爱的「CUDA Out of Memory」报错的主要来源。

需要注意的一点是:每次我们执行一个 GPU kernel,都需要把数据从 GPU 的 DRAM(即我们的仓库)搬出来,再搬回去。

现在想象一下执行 torch.cos 这种一元运算会发生什么:我们需要把数据从存储搬到仓库,对每一份数据做一点点计算,然后再把数据搬回去。搬运东西相当昂贵。结果就是,这里几乎所有时间都花在了搬运数据上,而不是实际的计算上。

既然我们所有时间都花在了内存带宽上,这种运算就被称为内存受限(memory-bound)运算,意味着我们没有在计算上花多少时间。

好吧,这不太理想。那能怎么办呢?我们来看看一系列算子是怎么执行的。

一串逐点算子的执行

这是一串逐点算子可能的执行方式。

嘿!这是个非常蠢的安排。为什么我们要反复把同一份数据送到全局内存再搬回计算单元?我们应该让数据待在工厂里,一次性做完所有计算,然后再搬回去!

算子融合

与其把我们的 cos 送回全局内存只为再读一遍,我们不如一口气把所有操作做完。

这就是算子融合(Operator Fusion)——深度学习编译器里最重要的优化。简单说,与其把数据写回全局内存再读一遍,不如一次性执行多个计算,省去额外的内存访问。

例如,执行 x.cos().cos() 通常需要 4 次全局读和写:

python
x1 = x.cos()  # 从全局内存读 x,写到 x1
x2 = x1.cos()  # 从全局内存读 x1,写到 x2

而有了算子融合,我们只需要 2 次全局内存读和写!所以算子融合能让它快 2 倍:

python
x2 = x.cos().cos()  # 从全局内存读 x,写到 x2

好多了。

这里有几个让事情变棘手的注意点。首先,GPU 在执行当前操作时需要知道接下来会发生什么。所以在 eager 模式下(PyTorch 一次只执行一个算子)你没法做这个优化。其次,我们实际上需要为它生成 CUDA 代码,这又打开了一个全新的潘多拉魔盒。

并不是所有算子融合都像逐点算子这么简单。你可以把逐点算子融合到归约上,或融合到矩阵乘法上。甚至矩阵乘法本身也可以看作「广播乘法 + 归约」的融合。

如果你对编写自定义 CUDA kernel 感兴趣,这里很可能就是你收益最大的地方。任意两个 PyTorch 算子之间都存在融合的机会,从而省去它们之间读写全局内存的带宽开销。此外,许多现有的编译器(如 NVFuser、Triton)都已经实现了这种融合。

最后,算子融合还会带来一些令人惊讶的后果。其一,融合后的 x.cos().cos() 与单独调用 x.cos() 几乎耗时完全相同。这就是为什么几乎所有激活函数耗时都差不多,尽管 gelu 显然比 relu 包含更多运算。

这个事实也给重计算/激活检查点(rematerialization/activation checkpointing)带来了一些有趣的推论。本质上,多做一点重计算可能会减少内存带宽,从而减少运行时间。因此,通过重计算,我们既能降低内存又能降低运行时间,我们正是利用这一点构建了「min-cut 最优重计算」

2.1 分析内存带宽开销

当需要判断你的操作是否内存带宽受限时,一个计算器就能帮上大忙。

对于简单算子,你可以直接推算内存带宽。例如,A100 有 1.5 TB/s 的全局内存带宽,每秒能执行 19.5 T flops 的计算。所以如果你用 32 位浮点数(即 4 字节),在一秒内可以加载 4000 亿个数……

所以……除非你在一元算子里做大约一百次运算,否则你花在内存访问上的时间会比实际计算多。

补充说明:这段是怎么算出来的?

这里的「大约一百次运算」不是拍脑袋,而是由 A100 的算力与带宽之比直接决定的,推导只有三步:

  1. 每秒能搬多少个数1.5×1012 字节/秒 ÷ 4 字节/个 ≈ 3.75×1011 个/秒 ≈ 4000 亿个。这就是原文那句「一秒内可以加载 4000 亿个数」的由来。
  2. 搬一个数的代价有多大:对 x.cos() 这种一元算子,每个元素要读一次(4 字节)、写一次(4 字节),共搬运 8 字节;而算一次 cos 只有 1 次浮点运算(1 FLOP)。
  3. 对比两者:搬运 8 字节耗时 =8÷(1.5×1012) 秒,计算 1 FLOP 耗时 =1÷(19.5×1012) 秒,两者之比 ≈ 104——搬一个数的时间,足够 GPU 算 104 次浮点运算。

反过来说,要让「计算时间追平内存时间」,你必须在每个元素上做 约 100 次运算。所以:

  • 只做 1 次运算(如 cosrelu)→ 内存时间是计算时间的上百倍 → 内存带宽受限(memory-bound),算力在闲置;
  • 只有当一元算子里塞进 超过约 100 次运算(即提高计算强度 compute intensity)→ 计算才追得上搬运 → 进入 算力受限(compute-bound)

这也正是为什么后面的微基准测试里 repeat < 32 饱和的是带宽、repeat > 64 才饱和算力,以及融合后的 x.cos().cos() 与单独 x.cos() 几乎一样快——因为两者都是内存受限,多算几下并不改变耗时。

(注:A100 官方标称的 HBM 带宽约为 1.6~2 TB/s,原文为行文方便用了 1.5 TB/s,推理逻辑不受影响。)

借助 NVFuser 这样的融合编译器,我们可以相当容易地自己实测这一点!代码在 Colab 里可以看到。

比如取一个这样的 PyTorch 函数:

python
def f(x: Tensor[N]):
    for _ in range(repeat):
        x = x * 2
    return x

用融合编译器对其做基准测试,我们就可以算出不同 repeat 值下的 FLOPS 和内存带宽。增大 repeat 是一种在不增加内存访问的情况下增加计算量的简单方法——这也叫提高计算强度(compute intensity)。

具体来说,假设我们对这段代码做基准测试,得到每秒执行的迭代次数。那么作为张量大小 N 的函数,我们会执行 2*N 次内存访问和 N * repeat 次 FLOP。因此,达到的内存带宽等于 bytes_per_elem * 2 * N * itrs_per_second,而 FLOPS 等于 N * repeat * itrs_per_second

现在,我们把运行时间、FLOPS 和达到的内存带宽作为计算强度的函数画出来。注意所有坐标轴都是对数-对数刻度。

微基准测试结果

首先注意到,直到执行 64 次乘法之前,运行时间几乎没有任何明显增加。这意味着在那之前,我们主要是内存带宽受限——计算大多处于闲置状态。

因此,一开始我们只达到了可怜的 0.2 teraflops。随着计算强度翻倍,这个数字近似线性增长,直到接近峰值 9.75 teraflops[^1]。一旦接近峰值 teraflops,我们就认为处于「算力受限」状态了。

最后你会看到,我们达到的内存带宽起初接近峰值,随着计算强度增加开始下降。这正是我们预期的:我们把越来越多的时间花在实际计算上,而不是访问内存。

在这个例子里,很容易看出何时算力受限、何时内存受限。对于 repeat < 32,我们饱和的是内存带宽,而算力没有被充分利用。反过来,一旦 repeat > 64,我们看到算力被饱和(即接近峰值 FLOPS),而内存带宽利用率在下降。

对于更大的系统,通常更难判断你是算力受限还是内存带宽受限,因为它们往往同时包含算力受限和内存受限的组件。

一个常见的衡量「你有多算力受限」的方法是,把你达到的 FLOPS 与峰值 FLOPS 相比,得到一个百分比。例如,如果你达到了峰值的 80%,那你就知道至少 80% 的时间是算力受限的——这已经很不错了!剩下的时间大概花在了内存访问或开销上。[^2]

不过,除了内存带宽开销,还有一件事可能会让你的 GPU 跑不起来。

3. 开销(Overhead)

开销是指你的代码花在「既不是搬运张量、也不是计算」上的时间。比如,花在 Python 解释器里的时间?是开销。花在 PyTorch 框架里的时间?是开销。花在启动 CUDA kernel(但还没执行它)的时间?也是……开销。

开销之所以如此顽固,首要原因是现代 GPU 真的太快了。A100 每秒能执行 312 万亿次浮点运算(312 teraflops)。相比之下,Python 真的非常非常慢。本地基准测试中,Python 一秒钟只能执行 3200 万次加法。

这意味着,在 Python 执行一次 FLOP 的时间里,A100 已经嚼完了 975 万次 FLOPS。

更糟的是,Python 解释器还不是唯一的开销来源——像 PyTorch 这样的框架在真正到达 kernel 之前还有好多层分发。如果你用 PyTorch 做同样的实验,只能做到每秒 28 万次运算。当然,这么小的张量不是典型用例。

比如,看这张 PyTorch 执行一次加法的 flamegraph 火焰图。看到那个小方块了吗?那才是实际执行计算的部分。其他一切都是纯粹的开销。

PyTorch 加法火焰图

看到这个,你可能震惊于居然还有人用 PyTorch。但请记住,现代深度学习模型执行的往往都是超大运算。而且,像 PyTorch 这样的框架是异步执行的。也就是说,当 PyTorch 正在运行一个 CUDA kernel 时,它可以继续排队更多工作。所以,只要 PyTorch 能「跑在」CUDA kernel 前面,框架的大部分开销就会被完全隐藏。

如果我们的 GPU 算子足够大,CPU 就可以跑到 GPU 前面(于是 CPU 开销就无关紧要了)。另一方面,如果 GPU 算子太小,那么 GPU 大部分时间就是一块昂贵的镇纸。

那怎么判断你处于这个区间呢?由于开销通常不随问题规模增长(而算力和内存会),最简单的办法就是直接增大数据规模。如果运行时间没有按比例增加,那你就是开销受限的。例如,如果你把 batch size 翻倍,但运行时间只增加了 10%,那你很可能处于开销受限区间。[^3]

另一个办法是使用 PyTorch profiler。在这里,粉色线条展示的是 CPU kernel 与 GPU kernel 的对应关系。

GPU 上大量空隙,说明它在等待 CPU 的开销

我们的 CPU 远远跑在 GPU 前面

顺便说一句,nvidia-smi 里的「GPU-Util」项(不是「Volatile GPU-Util」)基本上衡量的是最下面一行实际在运行 GPU kernel 的百分比。所以这也是目测开销的一个好办法。

这种开销存在的主要原因是 PyTorch 这类框架拥有的各种灵活性。本质上,需要花大量时间「搞清楚该做什么」。

这可能是来自 Python(查找属性或分发到正确的函数),也可能来自 PyTorch 内部的代码(PyTorch 的所有 dispatcher)。例如,当你执行 a + b 时,需要经过以下步骤:

  1. Python 需要查找 a 上的 __add__ 分发到什么。
  2. PyTorch 需要确定张量的许多属性(如 dtype、device、是否需要 autograd),以决定调用哪个 kernel。
  3. PyTorch 需要真正启动这个 kernel。

从根本上说,这种开销来自「每一步都能做不同事情」的灵活性。如果你不需要这种灵活性,一种解决方式是把它 trace 出来,比如用 jit.trace、FX 或 jax.jit。或者,你可以在更底层用 CUDA Graphs 之类的东西来做。

可惜,这要以失去灵活性为代价。我期待的一种做法是:在 VM 层面做内省,写一个更接近「真正的」JIT,从而两者兼得。更多内容见 TorchDynamo

4. 总结

如果你想让深度学习系统变快,最重要的事情是理解模型当前的瓶颈是什么。瓶颈决定了加速你的系统的正确方法。

我经常看到研究人员和想加速 PyTorch 代码的人,在不理解自己处于哪个区间的情况下盲目尝试各种东西。

性能区间可行的优化方案
开销受限Trace、算子融合、不用 Python、一个真正的 JIT :^)
带宽受限算子融合
算力受限用 Tensor Cores、给 NVIDIA 打钱

当然,可以说,用户居然需要思考这些东西,本身就反映框架的失败。PyTorch 的编译器或 profile API 一直都不是……最好用的,不过这是目前积极关注的方向。

无论如何,我觉得理解系统的基本原理几乎总是有用的——希望这篇文章对你有用。

PS:如果你喜欢这篇文章,我未来的大部分写作会放在 thonking.ai[链接待更新:thonking.ai 当前无法访问(连接失败)。]

致谢

感谢 Emily Shen、Qian Huang 以及 EleutherAI 的朋友们阅读本文的早期草稿并提供反馈。

参考文献与注释

[^1]: 这可能不是你在规格表上看到的 19.5 teraflops。原因是 GPU 有更专用的融合乘加(FMA)指令硬件。所以对于完全通用的计算,A100 实际上只能达到 9.75 teraflops。 [^2]: 关于 FLOPS 的计数方式有很多种,但现在在 PyTorch 里其实很容易就能优雅地做到——见 PyTorch FLOP counter。 [^3]: 这并非增大 batch size 可能不相应增加计算时间的唯一原因——在某些区间它还会提高计算强度。例如,在 MLP 里你通常做的是 [B, D] × [D, D] 的矩阵乘法。如果 B 小于 D(比如 batch size 为 1 而隐藏维度为 128),那么总内存带宽可能只增加一点点,而计算量却翻倍。

BibTeX 引用

bibtex
@article{he2022brrrrfromfirstprinciples,
  author={Horace He},
  title={Making Deep Learning Go Brrrr From First Principles},
  year={2022},
  url={https://horace.io/brrr_intro.html},
}

Maintained by Robin