LLM Infra 的逻辑链

· 中文
目录

    本文面向关心 LLM 系统设计、希望建立整体认识的读者。从 LLM 本身出发,沿一条因果链贯穿训练与推理系统的主要设计,重点在每项设计的动机。每个主题依次说明遇到了什么问题、怎么解决、优化了什么、牺牲了什么;只保留理解设计所需的原理和数字,略去实现细节。数字都来自公开的硬件规格和论文,出处统一列在文末。

    LLM

    范围

    LLM(大语言模型)的训练与推理系统,以及其中主要设计的由来。

    要理解 LLM,先看它的 basic architecture 是什么

    1. 本文讨论 LLM 的训练与推理系统。
    2. 系统的开销由模型架构决定:分析开销之前,要先知道模型的架构。
    3. 当前的 LLM 几乎都采用同一种架构:Transformer。
    4. 因此先分析 Transformer 的架构,以及训练时 GPU 上的数据。

    Basic Architecture:Transformer

    架构

    每层两个主要模块:attention(让每个 token 汇总前文各个 token 的信息;内部分成若干个独立计算的 head)与 FFN(两层全连接网络,对每个 token 单独计算),$$L$$ 层堆叠;前后再加分词(把文本切成 token,即词或词的一部分)、embedding 等模块。大部分参数在这 $$L$$ 层的矩阵里。

    数据

    训练时 GPU 上的数据有四类:参数(上面这些矩阵本身,个数记为 $$\Phi$$)、梯度与优化器状态(每个参数配一个梯度,Adam 另配两个统计量),三者合称训练状态,大小只由模型决定;激活(每个模块前向的中间结果,反向时要用),大小随 batch(一次处理的序列数)和序列长度变化。输入的数据量很小,可以并入激活一起看。

    这些数据只有被存放、被计算、被搬运三种状态,每种状态对应一种开销

    1. 训练时 GPU 上的数据有四类:参数、梯度、优化器状态、激活。
    2. 任何一份数据,任何时刻都处于三种状态之一:正被存放、正被计算、正被搬运。
    3. 三种状态占用三种不同的硬件资源:存放占显存容量,计算占算力,搬运占带宽。三种资源各有一个上限,通常称为显存墙、算力墙、带宽墙。
    4. 三堵墙的单位分别是 Bytes、FLOP/s、Bytes/s。

    三种开销

    显存墙 · 算力墙 · 带宽墙
    Data Placement显存开销:数据放在哪、占多少显存显存墙 · BytesH100:80 GB
    Data Calculation计算开销:在数据上执行的浮点运算算力墙 · FLOP/sH100 BF16:989 TFLOP/s
    Data Movement搬运开销:机器之间、GPU 之间、HBM 与 SRAM 之间的数据搬运带宽墙 · Bytes/sH100:HBM 3.35 TB/s,NVLink 450 GB/s,跨机器网络 50 GB/s

    三种开销不同类:显存开销是容量,超过显存容量就无法运行;计算开销和搬运开销是时间,等于运算量 ÷ 实际达到的算力、字节数 ÷ 实际达到的带宽,运算单元空闲、带宽没有用满时,时间变长。HBM 即显存;SRAM 是芯片内容量小、速度快的存储。H100 按 SXM 规格:算力是不含稀疏的稠密值;NVLink 和跨机器网络按单个方向计,跨机器网络按每个 GPU 一块 400 Gb/s 网卡计。

    此后每项技术优化了什么、牺牲了什么,都用这三种开销描述。

    判断顺序:看到一项技术,按顺序问两个问题。① 它是不是去冗,即只去掉原本不起作用的部分?② 如果不是,它优化了哪种开销、牺牲了哪种开销?

    去冗 去掉原本不起作用的部分,例如重复存放的数据、运算单元的空闲时间:优化一种开销,不牺牲其他开销;实现上的少量额外开销不计。

    交换 优化一种开销,牺牲另一种。通常被优化的一项是当前场景的瓶颈,被牺牲的一项在这个场景下有富余。

    三种开销之外:每次启动 kernel(GPU 上执行的一个函数)或发起通信,另有微秒量级、与数据量无关的固定开销;少数技术牺牲的是数值误差、负载不均衡、延迟或调度的复杂度。卡片中以灰色标出。

    颜色表示变化的开销:显存开销 计算开销 搬运开销 三种开销之外

    有了三种开销就能分析交换,先从一个 GPU 内部开始

    \[ T = \frac{6\Phi D}{N \cdot R \cdot \eta} \qquad M = 16\Phi + A \]

    $$T$$:训练时间;$$D$$:训练的 token 数;$$N$$:GPU 数;$$R$$:每个 GPU 实际达到的算力(FLOP/s),只计模型本身的运算,重算不计入;$$\eta$$:并行效率。$$M$$:每个 GPU 的显存占用(byte),$$16\Phi$$ 是训练状态(参数、梯度和 Adam 的两个统计量,全用 FP32 时各 $$4\Phi$$ 字节),$$A$$ 是激活。

    1. 训练的总运算量约为 $$6\Phi D$$ FLOP:每个参数对每个 token,前向约 2 FLOP,反向约 4 FLOP。
    2. 分子由模型和数据决定,缩短训练时间只能增大分母的三个因子:单个 GPU 的实际算力 $$R$$(混合精度、算子优化),GPU 数 $$N$$(DDP 及之后的并行方法),并行效率 $$\eta$$(通信等待、GPU 空闲使它小于 1)。约束是每个 GPU 的显存占用 $$M$$ 不超过显存容量。
    3. 先改一个 GPU 上的两个默认设定,两者都是交换:用有富余的一种开销换瓶颈的一种,目的是提高 $$R$$ 或降低 $$M$$。
    4. 第一个默认设定是数值精度:每个数默认用 FP32 存,占 4 字节,矩阵乘不经过 Tensor Core(不启用 TF32 时)。H100 上 FP32 运算单元是 67 TFLOP/s,BF16 Tensor Core 是 989 TFLOP/s。
    5. 第二个是激活的保存方式:默认全部保存。1.3B 的模型(GPT-3 XL 的架构)一步处理 32 条 2048 token 的序列,不计 attention 的分数矩阵,激活按 BF16 计约 110 GB(计入约 500 GB),而训练状态只有 21 GB。
    6. 改变这两个设定的技术,分别是混合精度与 activation checkpointing。

    单 GPU 内的交换

    混合精度 · activation checkpointing
    问题

    全部用 FP32 时,矩阵乘的峰值只有 BF16 Tensor Core 的约 1/15,激活每个数占 4 字节;激活全部保存时,可能超过显存容量。

    方案
    混合精度交换

    矩阵乘的输入和激活用 BF16(每个数 2 字节),乘积在 FP32 中累加;softmax、归一化等归约和参数更新用 FP32。参数更新留在 FP32,因为更新量小于参数的约 1/256 时,在 BF16 中会被舍掉。训练状态是 BF16 的参数和梯度各 $$2\Phi$$,加 FP32 的参数和 Adam 的两个统计量 $$12\Phi$$,合计 $$16\Phi$$ 字节。

    优化计算开销:矩阵乘峰值约 15 倍优化显存开销:激活减半牺牲数值误差训练状态仍是 $$16\Phi$$:保留了 FP32 参数
    Activation checkpointing交换

    每隔约 $$\sqrt{L}$$ 层保存一层的激活;反向需要时,从最近的保存点重新做一遍前向。

    优化显存开销:激活从 $$L$$ 层降到约 $$2\sqrt{L}$$ 层牺牲计算开销:多一次前向,$$R$$ 降到约 3/4

    混合精度提高的是峰值算力,实际达到的却远低于峰值,差距由算子的实现决定

    \[ R = F \times u \]

    $$F$$:峰值算力;$$u$$:达成率。混合精度提高的是 $$F$$,$$u$$ 由算子的实现决定。

    1. BF16 Tensor Core 的峰值是 FP32 运算单元的约 15 倍,但这只是上限;实际达到的算力取决于每个算子。
    2. 差距有两个来源。一是受显存带宽限制的算子:RMSNorm 每读写 1 字节约做 1 FLOP,而 H100 每字节要做约 295 FLOP(989 TFLOP/s ÷ 3.35 TB/s)才能用满算力;它的时间由显存带宽决定,与峰值算力无关。
    3. 这类算子(norm、softmax、激活函数、逐元素运算)运算量少,占用的时间却多:用 PyTorch 在 V100 上训练 BERT-large 时,它们占 0.2% 的运算量、39% 的时间。
    4. 二是矩阵乘:芯片内放不下整个矩阵,运算单元要等数据从显存读进芯片。
    5. 判断方法:一个算子做 $$W$$ FLOP、读写显存 $$Q$$ 字节,时间至少是 $$W/F$$ 与 $$Q/B$$ 中的较大者($$B$$ 是显存带宽,推导见 roofline 一文)。arithmetic intensity $$W/Q$$ 低于 ridge point $$F/B$$ 时受显存带宽限制,高于时受算力限制。
    6. 两类算子分开优化:受显存带宽限制的,减少读写显存的字节(算子融合 → online softmax → FlashAttention);受算力限制的矩阵乘,减少运算单元等待数据的时间(分块、流水)。

    算子效率

    roofline · 算子融合 · FlashAttention · 分块与流水
    问题

    一部分算子的时间由显存带宽决定,与峰值算力无关;矩阵乘也要等数据从显存读进芯片。

    方案
    算子融合去冗

    把连续几个受显存带宽限制的算子合成一个 kernel,中间结果留在芯片内,不写回显存。

    优化搬运开销:读写显存的字节优化kernel 启动的固定开销
    FlashAttention交换

    attention 为每个 token 算出 query、key、value 三个向量,排成矩阵 $$Q$$、$$K$$、$$V$$,再依次算 $$S = QK^\top$$(token 两两之间的匹配分数)、$$P = \text{softmax}(S)$$、输出 $$O = PV$$。$$S$$、$$P$$ 都是 $$n \times n$$ 矩阵($$n$$ 是序列长度),$$n = 4096$$ 时一个 head 约 32 MB。softmax 要用整行的最大值与总和,标准实现分三个 kernel,把 $$S$$、$$P$$ 写回显存。

    FlashAttention 用 online softmax 逐块更新每行的最大值与总和,三步合进一个 kernel,在芯片内分块算完,$$S$$、$$P$$ 不写回显存;反向时用 $$Q$$、$$K$$ 和每行的这两个统计量重算 $$S$$、$$P$$。论文的例子(GPT-2 medium,序列长度 1024,A100)中,attention 的前向加反向时间从 41.7 ms 降到 7.3 ms。

    优化搬运开销:读写显存 40.3 → 4.4 GB优化显存开销:去掉激活中的 $$n^2$$ 项牺牲计算开销:反向重算 $$S$$、$$P$$,66.6 → 75.2 GFLOP
    分块与流水去冗

    $$n \times n$$ 的矩阵乘中每个数参与 $$n$$ 次乘加,每个数(BF16)只读写一次时 arithmetic intensity 是 $$n/3$$($$n = 4096$$ 时约 1365);但芯片内放不下整个矩阵。分块让读进芯片的数据多次使用;读下一块和算这一块同时进行。

    优化搬运开销:重复读显存的字节优化计算开销:运算单元等数据的时间

    峰值和达成率都提高后,一个 GPU 的算力仍有上限,只能增加 GPU 数

    6 × 405B × 15.6T token ≈ 3.8×10²⁵ FLOP;一个 H100 以 989 TFLOP/s 运行约 1200 年

    1. $$R = F \times u$$ 的两个因子都已提高:混合精度提高 $$F$$,算子优化提高 $$u$$(checkpointing 方向相反,用约 1/3 的额外计算换显存)。
    2. 分子 $$6\Phi D$$ 由模型和数据决定。Llama-3 405B:$$\Phi = 4.05 \times 10^{11}$$,$$D = 1.56 \times 10^{13}$$ token,$$6\Phi D \approx 3.8 \times 10^{25}$$ FLOP。
    3. 一个 H100 以峰值 989 TFLOP/s 持续运行,需要约 1200 年。
    4. 分子固定,$$R$$ 不超过峰值,$$\eta$$ 不超过 1,能继续增大的只有 GPU 数 $$N$$。

    DDP

    Data Parallelism
    问题

    一个 GPU 的算力有上限,前沿规模的训练在一个 GPU 上需要上千年。

    方案

    $$N$$ 个 GPU 各放一份完整模型,处理不同的数据;每步用 all-reduce(每个 GPU 出一份数据,结束后都得到它们的和)求梯度平均,所有副本做相同的更新。通信可以和反向计算同时进行,通信时间短于计算时间时 $$\eta$$ 接近 1。

    交换优化计算开销:每个 GPU 算 $$1/N$$ 的数据,训练时间约 $$1/N$$牺牲搬运开销:每步通信 $$2\Phi$$ 个元素显存开销不变:每个 GPU 仍存完整的 $$16\Phi$$,$$N$$ 份完全相同

    DDP 没有减少显存,每个 GPU 仍要放下完整的训练状态,放不下时怎么办?

    7B × 16 byte = 112 GB,超过 H100 的 80 GB

    1. DDP 的前提是每个 GPU 放得下完整的训练状态 $$16\Phi$$ 字节。
    2. $$16\Phi$$ 随参数量线性增长:7B 模型是 112 GB,超过 H100 的 80 GB,DDP 无法运行。
    3. 而 $$N$$ 个 GPU 上的 $$N$$ 份 $$16\Phi$$ 完全相同。
    4. 因此每个 GPU 只需存 $$1/N$$,用到时从其他 GPU 取回,信息不丢失。
    5. ZeRO 规定了分片的顺序,以及取回需要的通信。

    ZeRO / FSDP

    训练状态分片
    问题

    $$16\Phi$$ 超出一个 GPU 的显存,而 $$N$$ 个 GPU 存的是 $$N$$ 份相同的副本。

    方案

    按使用频率从低到高逐级分片。

    ZeRO-1去冗

    分片优化器状态($$12\Phi$$)。all-reduce 可以拆成 reduce-scatter(每个 GPU 得到 $$1/N$$ 的梯度和)和 all-gather(每个 GPU 把自己的 $$1/N$$ 发给所有 GPU)两步;ZeRO-1 让每个 GPU 在两步之间只更新自己那 $$1/N$$ 参数,再 all-gather 更新后的参数,所以每个 GPU 只需存 $$1/N$$ 的优化器状态,通信量不变。

    优化显存开销:$$16\Phi \to 4\Phi + 12\Phi/N$$搬运开销不变:每步通信 $$2\Phi$$
    ZeRO-2去冗

    再分片梯度($$2\Phi$$):reduce-scatter 之后,每个 GPU 只需要保留自己那 $$1/N$$ 的梯度和,其余的可以释放。

    优化显存开销:$$\to 2\Phi + 14\Phi/N$$搬运开销不变
    ZeRO-3交换

    再分片参数($$2\Phi$$):每层计算前取回完整参数,前向、反向各一次。PyTorch FSDP 的 FULL_SHARD 模式对应 ZeRO-3。

    优化显存开销:$$\to 16\Phi/N$$牺牲搬运开销:每步通信 $$2\Phi \to 3\Phi$$

    显存以字节计,通信量以元素个数计,沿用 ZeRO 论文的惯例。

    ZeRO 切分的是存储,每个 GPU 仍要算完整的模型,所以 GPU 多时同步参数的时间超过计算时间

    \[ \frac{t_\text{comm}}{t_\text{comp}} = \frac{3\Phi \times 2\ \text{byte} \,/\, B_\text{net}}{6\Phi \cdot (G/N) \,/\, R} = \frac{R \cdot N}{B_\text{net} \cdot G} \]

    $$G$$:每步的总 token 数(global batch);$$B_\text{net}$$:跨机器网络带宽。$$\Phi$$ 约掉了。

    1. ZeRO-3 每个 GPU 每步通信 $$3\Phi$$ 个元素,不随 $$N$$ 减小;计算量是 $$6\Phi$$ 乘每个 GPU 分到的 token 数。两者可以重叠,上式的比值小于 1 时,通信可以完全与计算重叠。
    2. 每步的总 token 数 $$G$$ 有上限:超过 critical batch size 后,继续加大 batch,减少的训练步数越来越少。
    3. $$G$$ 固定时,每个 GPU 分到 $$G/N$$ 个 token,通信与计算的时间比与 $$N$$ 成正比。
    4. Llama-3 405B 用 16384 个 H100,$$G$$ = 16M token,平均每个 GPU 只有约 1000 个 token。若这 16384 个 GPU 全部只用 ZeRO-3,按实测的每个 GPU 约 400 TFLOP/s 计,比值约为 8。
    5. 另外,每个 GPU 至少要处理一条完整的序列:序列长时,一条序列的激活就可能超过一个 GPU 的显存容量。
    6. 第一个限制来自每个 GPU 计算完整的模型,第二个来自每个 GPU 处理完整的序列。TP 把层内的矩阵乘切开,改由几个 GPU 共同计算同一层;CP 把序列切开。

    TP · SP · CP

    层内并行
    问题

    GPU 多时,同步参数的时间超过计算时间;序列长时,一条序列的激活超过一个 GPU 的显存容量。

    方案
    TP交换

    张量并行:把每个矩阵乘切成 $$t$$ 份,分给 $$t$$ 个 GPU;$$t$$ 个 GPU 共同算一份模型,上式的 $$N$$ 换成组数 $$N/t$$。

    优化显存开销:训练状态和 attention、FFN 内部的激活降到 $$1/t$$优化搬运开销:跨机器时每个 GPU 只同步 $$1/t$$ 的参数牺牲搬运开销:机器内每层前向、反向各 2 次 all-reduce,前向要等它完成
    SP去冗

    序列并行,这里指 Megatron-LM 的做法:TP 不切分 LayerNorm 和 Dropout,这部分激活每个 GPU 各存一份,$$t = 8$$ 时约占一层激活的 3/4(不计 attention 的分数矩阵)。SP 把它沿序列切成 $$t$$ 份,原来的 all-reduce 改写成等量的 reduce-scatter 和 all-gather。

    优化显存开销:这部分激活降到 $$1/t$$搬运开销不变
    CP交换

    context parallelism,用于长序列的另一类序列并行(Li et al. 称为 sequence parallelism,此后有 Ring Attention、DeepSpeed-Ulysses):整层的激活都沿序列切成 $$c$$ 份,attention 需要的其他位置的 K、V 通过通信获得。Ring Attention 沿环逐块传递,与计算重叠;DeepSpeed-Ulysses 用 all-to-all(每个 GPU 给其他每个 GPU 发不同的数据)改为按 head 切分;Llama-3 先 all-gather 全部 K、V。

    优化显存开销:整层的激活降到 $$1/c$$牺牲搬运开销:每层传 K、V

    TP 的通信限制了它自身的规模,跨机器需要通信量小的切法

    机器内 NVLink 450 GB/s,跨机器网络 50 GB/s(单向)

    1. TP 每层前向、反向各 2 次 all-reduce;前向要等 all-reduce 完成才能继续计算,所以难以和计算重叠。
    2. 这种通信要放在机器内的 NVLink 上,跨机器网络的带宽只有它的约 1/9。因此 TP 限于一台机器内,一台机器 8 个 GPU,$$t \le 8$$。
    3. 单靠 TP,8 个 GPU 放不下大模型:405B 的训练状态约 6.5 TB,8 个 H100 共 640 GB。与 ZeRO-3 组合时组数是 $$N/8$$,Llama-3 的规模下通信与计算的时间比只从约 8 降到约 1。
    4. 跨机器的切法需要通信量小:PP 按层切开,段之间只传边界处的激活,传输量远小于模型的参数量。

    PP

    Pipeline Parallelism
    问题

    TP 限于一台机器内;单靠 TP 放不下大模型,与 ZeRO-3 组合时,跨机器同步参数的时间仍约等于计算时间。

    方案

    把 $$L$$ 层分成 $$p$$ 段,段之间只传激活。batch 切成 $$m$$ 个 micro-batch 依次送入,各段同时处理不同的 micro-batch;开头和结尾部分 GPU 在等待,称为 bubble,占比 $$(p-1)/(m+p-1)$$。增大 $$m$$ 可以压低 bubble,但 micro-batch 太小时 arithmetic intensity 降低、固定开销占比增大。

    交换优化显存开销:每个 GPU 只存 $$L/p$$ 层优化搬运开销:跨机器时每个 GPU 只同步 $$1/(t \cdot p)$$ 的参数牺牲计算开销:bubble,$$p = 16$$、$$m = 64$$ 时约 19%

    三种并行单独用都不够,只能组合

    1. 三种并行各自的限制:TP 限于一台机器内;PP 有 bubble,micro-batch 不能太小;DP(数据并行,DDP 与 ZeRO 都属于这一类)的组数每增加一倍,每组分到的 token 减半,而 global batch 有上限(critical batch size)。
    2. 三个限制的来源各不相同:TP 缺的是跨机器的带宽,PP 缺的是足够多的 micro-batch,DP 缺的是 global batch 继续增大的余地。
    3. 来源不同,所以每种并行可以放在它的限制不起作用的位置。
    4. 于是组合方式基本确定:TP 在机器内,PP 跨机器,DP 在最外层,即 3D 并行。

    3D 并行

    TP × PP × DP
    问题

    TP 限于机器内,PP 有 bubble,DP 受 global batch 限制,单独使用都不够。

    方案

    $$N = t \times p \times d$$($$d$$ 是 DP 的组数)。Llama-3 405B 预训练的一种配置(序列长度 8K)用 $$8 \times 16 \times 128 = 16384$$ 个 H100,DP 用 FSDP;序列长度 128K 的阶段改为 TP 8、CP 16、PP 16、DP 8,报告称为 4D 并行。按上面的式子,通信时间 ÷ 计算时间 $$= R \cdot d/(B_\text{net} \cdot G)$$:

    切法组数 $$d$$通信时间 ÷ 计算时间
    全部 ZeRO-316384约 8
    加 TP($$t = 8$$)2048约 1
    再加 PP($$p = 16$$)128约 0.06
    优化搬运开销:跨机器通信从计算时间的约 8 倍降到约 6%

    Llama-3 的 FSDP 前向后不释放参数,每步通信 $$2\Phi$$ 个元素,梯度按 FP32 传,合计 $$6\Phi$$ 字节,与按 $$3\Phi \times 2$$ 字节计的结果相同。

    参数量继续增加时,每个 token 的运算量能否保持不变?

    1. 3D 并行解决了大模型的显存与训练时间问题。
    2. scaling law 给出继续增加参数的理由:同样数据下,参数越多,训练出的模型 loss(预测与正确答案的差距)越低。
    3. 但 dense 模型里每个 token 经过全部参数:参数翻倍,每个 token 的运算量也翻倍,训练和推理的成本都翻倍。
    4. 目标是让参数量与每个 token 的运算量解耦。
    5. 参数的大部分在 FFN(标准架构中约占每层的 2/3),而且每个 token 单独经过它:把 FFN 换成 $$E$$ 个结构相同的 FFN、每个 token 只经过其中 $$k$$ 个,就是 MoE。
    6. expert 多了,一个 GPU 放不下全部 expert;若用 ZeRO-3,每层要取回全部 $$E$$ 个 expert 的参数,而每个 token 只用其中 $$k$$ 个。改为参数不动、把 token 发到 expert 所在的 GPU,就是 EP。

    MoE · EP

    Mixture of Experts · Expert Parallelism
    问题

    dense 模型中,每个 token 的运算量与参数量成正比;改用 MoE 后,一个 GPU 放不下全部 expert,用 ZeRO-3 每层又要取回全部 expert 的参数。

    方案
    MoE交换

    每层的 FFN 换成 $$E$$ 个结构相同、参数各自训练的 FFN,叫 expert;router 为每个 token 选其中 $$k$$ 个。参数量随 $$E$$ 增长,每个 token 的运算量只随 $$k$$ 增长。DeepSeek-V3 共 671B 参数,每个 token 经过其中 37B。

    优化计算开销:$$6\Phi D$$ 中的 $$\Phi$$ 按每个 token 经过的参数计,671B → 37B牺牲负载不均衡:router 可能把 token 集中发给少数 expert,要额外均衡
    EP交换

    把 expert 分到不同的 GPU 上,每个 GPU 只存一部分 expert;token 发到它选中的 expert 所在的 GPU,算完再发回。与 ZeRO-3 相比,参数不动,传的是 token。DeepSeek-V3 每个 MoE 层有 256 个 routed expert(由 router 选择的 expert),训练时分到 64 个 GPU 上,每个 GPU 4 个;decode 时每个 GPU 只放 1 个 expert。

    优化显存开销:每个 GPU 只存一部分 expert牺牲搬运开销:每个 MoE 层前向、反向各 2 次 all-to-all

    这里与参数量相同的 dense 模型比较,两者的总显存相同;若与每个 token 运算量相同的 dense 模型比较,MoE 要多存约 634B 参数,由 EP 分到更多 GPU 上。

    模型部署后,按请求逐个生成 token

    不存中间结果时,生成 n 个 token 的总运算量至少随 n² 增长

    1. 模型训练完成后部署上线,按请求生成回答;推理只有前向计算。
    2. 推理分两个阶段:prefill 把整段输入一次送进模型;decode 逐个生成新 token,下一个 token 的计算要用到上一个。
    3. 生成第 $$n$$ 个 token 时,attention 要用前面每个位置的 key、value 向量(K、V)。
    4. 不存下来,每步都要重算整段前文。

    KV cache

    存下前文的 K、V
    问题

    生成第 $$n$$ 个 token 时,要把前面 $$n - 1$$ 个 token 重算一遍。

    方案

    把每层的 K、V 存下来,每步只算新 token,每个请求每步送进 1 个 token。

    交换优化计算开销:每步从整段前文降到 1 个 token牺牲显存开销:随长度增长,Llama 2 7B 每 token 512 KiB

    KV cache 之后,decode 每步只算 1 个 token,却要读一遍全部权重,瓶颈从算力变为显存带宽

    Llama 2 7B:每步读 13.5 GB 权重;batch 为 1 时,算力用到不足 1%

    \[ T_\text{step} \approx \max\!\left(\frac{Q_\text{w} + b \cdot n \cdot q_\text{KV}}{B},\ \frac{2\Phi \cdot b}{F}\right) \]

    $$Q_\text{w}$$:权重的字节数;$$b$$:一起算的请求数(batch);$$n$$:上下文长度;$$q_\text{KV}$$:每 token 的 KV cache 字节数。

    1. 训练时大量 token 一起经过模型,矩阵乘规模大,受算力限制。
    2. decode 每步读一遍全部 BF16 权重,每读 2 字节做 2 FLOP;$$b$$ 个请求一起算时,权重部分的 arithmetic intensity 约 $$b$$ FLOP/byte;$$b$$ 远小于 ridge point 295 时受显存带宽限制。
    3. batch 为 1 时,算力用到不足 1%(展开见 roofline 一文第 6 节)。瓶颈不同,训练侧的优化不能直接沿用。
    4. 按上式分两组优化:加大 $$b$$,使每次读权重生成更多 token;缩短每一步、减少步数。

    加大 batch

    第一组
    问题

    decode 每步读一遍全部权重,只为 $$b$$ 个请求各生成 1 个 token。

    上限

    $$b$$ 受显存容量限制:$$Q_\text{w} + b \cdot n \cdot q_\text{KV} \le$$ 显存容量。Llama 2 7B、上下文 4096 token、80 GB 的 H100:权重 13.5 GB,每个请求 KV cache 2.1 GB,$$b$$ 最多 30。这时每步读约 78 GB,其中 64 GB 是 KV cache;每个请求各读自己的 KV cache,加大 $$b$$ 不能分摊这部分读取。

    GQA(多个 query head 共用一组 K、V)和 MLA(把 K、V 压缩成短向量)都能减小 $$q_\text{KV}$$。

    方案

    四项技术中,PagedAttention 提高 continuous batching 能达到的 $$b$$;chunked prefill 与 prefix caching 各自独立。

    continuous batching去冗

    请求长短不一,整批一起开始、一起结束会留下空位。改为每一步都让完成的请求退出、新请求加入。

    优化搬运开销:读一遍权重生成更多有效 token优化新请求的排队时间
    PagedAttention去冗

    $$b$$ 的上限还取决于 KV cache 显存的利用率。按最大长度预留时,约 20% 存放有效数据(预先知道输出长度也只有约 38%)。改为按固定大小的块分配,块不要求连续。按块寻址让 attention kernel 慢约 20% 到 26%,但 $$b$$ 可以加大,端到端吞吐提高 2 到 4 倍,所以计为实现开销。

    优化显存开销:有效数据约 96%
    chunked prefill交换

    新请求的 prefill 插入时,正在 decode 的请求要等它算完。把 prefill 切成块,每步和 decode 一起算一块,用的是 decode 时有余量的那部分算力。

    优化decode 的停顿牺牲新请求第一个 token 的延迟牺牲搬运开销:每块都要重读前面各块的 KV cache
    prefix caching去冗

    很多请求的前缀相同(例如系统提示)。把算过的前缀的 KV cache 留在显存里复用;缓存只占空余的显存,需要加大 batch 时先淘汰最久未用的部分。

    优化计算开销:重复的 prefill优化显存开销:相同前缀只存一份

    第二组与加大 batch 互不依赖,作用是缩短每一步、减少步数

    一个请求的生成时间 ≈ 步数 × 每步读的字节 ÷ 实际达到的显存带宽

    1. 加大 batch 让每次读权重服务更多请求;第二组缩短单个请求的生成时间。
    2. 上式三个因子各有一项技术:Flash-Decoding 提高实际达到的带宽,量化减少每步读的字节,speculative decoding 减少原模型的步数。三项可以同时使用。

    缩短每步、减少步数

    第二组
    问题

    加大 batch 提高吞吐,但不缩短单个请求的生成时间:它由步数、每步读的字节和实际达到的显存带宽决定。

    方案
    Flash-Decoding去冗

    H100 有 132 个 SM(独立执行的运算单元组),同时读显存的 SM 越多,实际带宽越接近峰值。FlashAttention 按 batch、head 和 query 块把工作分给 SM;decode 时 query 只有 1 个位置,batch 小时分出的工作少于 SM 数。Flash-Decoding 再把 KV cache 沿长度切块,分给更多 SM,最后用一个小 kernel 合并各块的结果;Flash-Decoding 的博客报告,在 A100 上、长上下文时,attention 比 FlashAttention 最多快约 50 倍。

    优化搬运开销:空闲的 SM 也参与读 KV cache,实际带宽提高
    量化交换

    权重从 BF16 量化为 INT4 等低精度;只量化权重时,读进来先转回 BF16 再算。

    优化搬运开销:每步读的权重字节约 1/4牺牲计算开销:转回 BF16,用的是有余量的算力牺牲数值误差:GPTQ 等方法用校准数据减小
    speculative decoding交换

    生成每步得到 1 个 token,而验证多个给定的 token 只需一次前向。先用小模型生成 $$k$$ 个候选,原模型一次前向验证 $$k+1$$ 个位置,按拒绝采样规则决定接受到第几个,输出分布和逐个生成相同。收益取决于候选被接受的比例;batch 大时算力不再有余量,收益减小。

    优化搬运开销:原模型读权重的次数减少(小模型也要读权重,但小得多)牺牲计算开销:验证 $$k+1$$ 个位置

    prefill 和 decode 的瓶颈不同,放在同一组 GPU 上会互相影响

    1. prefill 一次处理整段输入,受算力限制;decode 受显存带宽限制。
    2. chunked prefill 让两者在同一组 GPU 上一起算,减少了 decode 的停顿;但每个 decode 步仍要等同批的 prefill 块算完,两者也只能用同一种并行方式和 batch。
    3. 和 3D 并行一样,按限制的来源分开放置。

    P/D 分离

    prefill 与 decode 分开部署
    问题

    prefill 受算力限制,decode 受显存带宽限制,共用一组 GPU 时互相影响。

    方案

    prefill 和 decode 放在两组 GPU 上,prefill 算完把 KV cache 传给 decode 组,每个请求只传一次。DistServe 报告,在 90% 的请求满足延迟要求的条件下,每个 GPU 每秒可服务的请求数最多是对比系统的 7.4 倍,或者能满足严格 12.6 倍的延迟要求。

    交换优化延迟:两个阶段不再互相影响,各自选择并行方式和 batch牺牲搬运开销:传输 KV cache

    最后,post-training 的 RL 把训练和推理两套系统放进同一个循环

    1. 至此,训练和推理各有一套系统,各自针对自己的瓶颈。
    2. RL 的一步是:模型生成回答(rollout),给回答打分,用打过分的样本更新参数。
    3. 生成是推理负载,大部分时间在 decode,受显存带宽限制;更新是训练负载,受算力限制。一个循环同时包含两种负载。
    4. 训练引擎缺少 KV cache 管理、continuous batching 等推理优化,用它生成很慢。
    5. 两套系统共用一份每步都在更新的参数,需要每步同步。

    RL post-training

    训练引擎与推理引擎
    问题

    训练引擎缺少 KV cache 管理、continuous batching 等推理优化,生成慢。

    方案

    rollout 用推理引擎(vLLM、SGLang),更新用训练引擎(FSDP、Megatron-LM),每步把新参数同步给推理引擎,并按推理引擎的切分方式重新切分。同步执行时,$$T_\text{RL} = T_\text{rollout} + T_\text{update} + T_\text{sync}$$。

    两种部署:共置,两个引擎在同一组 GPU 上轮流运行;分离,两个引擎在不同的 GPU 上。分离时还可以让下一批 rollout 用上一版参数提前开始,这时生成数据的参数落后一步或更多(off-policy),训练算法要能容忍这种差异。

    交换优化计算开销:生成时 KV cache 免去重算优化搬运开销:batch 更大,每次读权重生成更多 token牺牲搬运开销:同步参数,70B 模型每步 140 GB,分离部署时最多占一步时间的 36.4%牺牲调度两个引擎的复杂度

    汇总

    技术起因优化牺牲类型
    单 GPU 训练
    混合精度FP32 的矩阵乘不经过 Tensor Core,激活每个数占 4 字节计算显存:激活数值误差交换
    checkpointing激活超过显存容量显存:激活计算:多一次前向交换
    算子融合中间结果反复读写显存搬运kernel 启动—去冗
    FlashAttentionn × n 的 S、P 写回显存搬运显存计算:反向重算交换
    分块与流水矩阵乘等待数据搬运计算—去冗
    多 GPU 分布式训练
    DDP一个 GPU 的算力有上限计算:时间约 1/N搬运:同步梯度交换
    ZeRO-1/2N 个 GPU 存 N 份相同的训练状态显存—去冗
    ZeRO-3参数仍未分片显存搬运:2Φ → 3Φ交换
    TPGPU 多时同步参数的时间超过计算显存搬运:跨机器搬运:机器内交换
    SP(Megatron-LM)TP 不切分的激活每个 GPU 各存一份显存—去冗
    CP序列长,一条序列的激活放不下显存搬运:传 K、V交换
    PPTP 限于一台机器显存搬运:跨机器计算:bubble交换
    3D 并行三种并行各有限制搬运:跨机器—组合
    MoE每个 token 的运算量与参数量成正比计算负载不均衡交换
    EP一个 GPU 放不下全部 expert显存搬运:all-to-all交换
    推理
    KV cache每步重算整段前文计算显存交换
    continuous batching批内留下空位搬运排队时间—去冗
    PagedAttention按最大长度预留显存显存—去冗
    chunked prefill新请求的 prefill 让 decode 停顿decode 的停顿首 token 延迟搬运:重读 KV cache交换
    prefix caching相同的前缀重复 prefill计算显存—去冗
    Flash-Decodingbatch 小时读 KV cache 的 SM 少搬运—去冗
    量化每步读一遍全部权重搬运计算:转回 BF16数值误差交换
    speculative decoding每步只得到 1 个 token搬运:步数减少计算:验证交换
    P/D 分离prefill 与 decode 共用 GPU,互相影响延迟搬运:传 KV cache交换
    RL
    训练引擎 + 推理引擎训练引擎生成慢计算搬运搬运:同步参数调度复杂度交换

    出处

    1. H100 的算力、显存容量与带宽、NVLink:NVIDIA,H100 Tensor Core GPU 产品规格(SXM)
    2. H100 的 SM 数(132):NVIDIA,NVIDIA Hopper Architecture In-Depth,2022
    3. 混合精度:Micikevicius et al.,Mixed Precision Training,ICLR 2018
    4. BF16 训练:Kalamkar et al.,A Study of BFLOAT16 for Deep Learning Training,2019
    5. activation checkpointing:Chen et al.,Training Deep Nets with Sublinear Memory Cost,2016
    6. 激活的估算式、Megatron-LM 的 SP:Korthikanti et al.,Reducing Activation Recomputation in Large Transformer Models,MLSys 2023
    7. GPT-3 XL 的架构:Brown et al.,Language Models are Few-Shot Learners,NeurIPS 2020
    8. 各类算子的运算量与时间占比(BERT-large,V100):Ivanov et al.,Data Movement Is All You Need: A Case Study on Optimizing Transformers,MLSys 2021
    9. online softmax:Milakov 与 Gimelshein,Online Normalizer Calculation for Softmax,2018
    10. FlashAttention:Dao et al.,FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness,NeurIPS 2022
    11. Llama-3 405B 的参数、token 数、global batch、GPU 数、网络与并行配置:Llama Team,The Llama 3 Herd of Models,2024
    12. ZeRO:Rajbhandari et al.,ZeRO: Memory Optimizations Toward Training Trillion Parameter Models,SC 2020
    13. TP:Shoeybi et al.,Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism,2019
    14. CP(Li et al. 称为 sequence parallelism):Li et al.,Sequence Parallelism: Long Sequence Training from System Perspective,2021
    15. Ring Attention:Liu et al.,Ring Attention with Blockwise Transformers for Near-Infinite Context,2023
    16. DeepSpeed-Ulysses:Jacobs et al.,DeepSpeed Ulysses: System Optimizations for Enabling Training of Extreme Long Sequence Transformer Models,2023
    17. TP、PP、DP 的组合与通信分析:Narayanan et al.,Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM,SC 2021
    18. PP 与 bubble:Huang et al.,GPipe: Efficient Training of Giant Neural Networks using Pipeline Parallelism,NeurIPS 2019
    19. critical batch size:McCandlish et al.,An Empirical Model of Large-Batch Training,2018
    20. scaling law:Kaplan et al.,Scaling Laws for Neural Language Models,2020
    21. MoE:Shazeer et al.,Outrageously Large Neural Networks: The Sparsely-Gated Mixture-of-Experts Layer,ICLR 2017
    22. EP:Lepikhin et al.,GShard: Scaling Giant Models with Conditional Computation and Automatic Sharding,ICLR 2021
    23. DeepSeek-V3 的参数量与 EP 配置:DeepSeek-AI,DeepSeek-V3 Technical Report,2024
    24. Llama 2 7B 的层数与宽度(32 层,4096):Touvron et al.,LLaMA: Open and Efficient Foundation Language Models,2023
    25. Llama 2 7B 不用 GQA:Touvron et al.,Llama 2: Open Foundation and Fine-Tuned Chat Models,2023
    26. continuous batching:Yu et al.,Orca: A Distributed Serving System for Transformer-Based Generative Models,OSDI 2022
    27. PagedAttention:Kwon et al.,Efficient Memory Management for Large Language Model Serving with PagedAttention,SOSP 2023
    28. GQA:Ainslie et al.,GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints,EMNLP 2023
    29. MLA:DeepSeek-AI,DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model,2024
    30. chunked prefill:Agrawal et al.,Taming Throughput-Latency Tradeoff in LLM Inference with Sarathi-Serve,OSDI 2024
    31. prefix caching:Zheng et al.,SGLang: Efficient Execution of Structured Language Model Programs,NeurIPS 2024
    32. Flash-Decoding:Dao et al.,Flash-Decoding for long-context inference(PyTorch 博客),2023
    33. 量化:Frantar et al.,GPTQ: Accurate Post-Training Quantization for Generative Pre-trained Transformers,ICLR 2023
    34. speculative decoding:Leviathan et al.,Fast Inference from Transformers via Speculative Decoding,ICML 2023
    35. P/D 分离:Zhong et al.,DistServe: Disaggregating Prefill and Decoding for Goodput-optimized Large Language Model Serving,OSDI 2024
    36. RL 的两个引擎与部署:Sheng et al.,HybridFlow: A Flexible and Efficient RLHF Framework,EuroSys 2025
    37. rollout 提前开始的异步 RL:Noukhovitch et al.,Asynchronous RLHF: Faster and More Efficient Off-Policy RL for Language Models,ICLR 2025
    目录