Skip to content

07-16 下午:分布式大模型训练与推理

最后更新于·约 13759 字

与后续实验的联系

Lab 5 的量化和端到端推理依赖张量形状、显存账本、键值缓存(KV cache)和批处理。分布式训练还需要理解前向、反向、梯度同步和参数更新之间的关系。PyTorch DistributedDataParallel(DDP,分布式数据并行)、FSDP(Fully Sharded Data Parallel,完全分片数据并行)和 Hugging Face 的缓存文档可作为实现参考。7 8 9

模型、batch(批次)或上下文长度超过单张 GPU 的显存和吞吐范围后,训练与推理需要在多张卡之间分工。多卡方案的名字不同,首先改变的都是数据和状态的放置方式。

数据并行(data parallelism,DP)沿 batch 切分样本。ZeRO(Zero Redundancy Optimizer,零冗余优化器)和 FSDP 把模型状态分片保存。流水线并行(pipeline parallelism,PP)按层切分模型;张量并行(tensor parallelism,TP)在同一层内切分矩阵;上下文并行沿序列切分。

理解任一方案时,分别记录每张卡保存什么、负责计算什么,以及何时需要交换数据。2 3

一次训练 step 在多卡上怎样进行

先从两张 GPU 的普通数据并行训练看起。设全局 batch 有 8 条样本,GPU 0 处理前 4 条,GPU 1 处理后 4 条;两张卡在 step 开始时都保存同一份参数 \(W\)

GPU 0:micro-batch 0..3  → forward → loss_0 → backward → 本地梯度 g_0
GPU 1:micro-batch 4..7  → forward → loss_1 → backward → 本地梯度 g_1

AllReduce(g_0, g_1)
        ↓
两张卡得到同一份全局梯度 g = (g_0 + g_1) / 2
        ↓
各自执行相同的 optimizer.step(),参数仍保持一致

这里的 rank 是分布式作业中一个独立进程的编号。micro-batch 是从全局 batch 切出的样本子集。每张卡先对自己的样本做前向计算,得到 logits(原始分数)和 loss(训练目标),再反向计算本地梯度。上图用等大小 micro-batch 的平均梯度作为例子。AllReduce 将对应参数的梯度规约并分发给所有参与 rank;框架随后按自己的约定求和或取平均,再更新参数。

可以把这一步写成接近框架执行顺序的伪代码。

# 每个 rank 都执行相同代码,但 batch_slice 不同
optimizer.zero_grad()
logits = model(batch_slice)
loss = cross_entropy(logits, labels_slice)
loss.backward()               # 得到本 rank 的本地梯度
all_reduce_gradients(model)   # 让各 rank 使用同一份全局梯度
optimizer.step()

数据并行只切 batch,模型参数仍完整复制在每张卡上。模型放不下时,FSDP 在使用某层前临时收集参数、使用后重新分片;张量并行在同一层内切开矩阵;流水线并行把连续层交给不同 stage;上下文并行沿序列切分。它们都保留前向、loss、反向、更新这条训练主线,只是把参数、activation 和通信安排到不同位置。

PyTorch DistributedDataParallel 文档的说明(译)

DistributedDataParallel 将模型复制到多个进程,每个进程处理不同输入。反向传播期间,模块会同步梯度,使模型副本在更新后保持一致。7

分布式训练的约束

训练扩展到多卡时,需要同时考虑模型状态、计算量和通信。参数、梯度、优化器状态和 activation 占用显存。矩阵乘法和 attention 消耗算力。梯度同步、模型分片和专家路由产生设备间通信。GPU 数量增加后,显存、计算和通信的比例也会改变。2 3

切分方向 每张卡主要保存 典型通信 先解决的约束
数据并行(DP) 完整模型副本、不同数据 AllReduce 梯度 增加 batch 吞吐
ZeRO(Zero Redundancy Optimizer)/ FSDP(Fully Sharded Data Parallel) 参数/梯度/优化器分片 AllGather、ReduceScatter 降低模型状态显存
流水线并行(PP) 连续若干层 相邻 stage 激活/梯度 沿深度拆层
张量并行(TP) 一层权重/激活的分片 层内 AllReduce/AllGather 拆大矩阵计算
上下文并行 序列或 K/V 分段 K/V 分段传递 降低长序列峰值

这些缩写描述数据如何放置,性能还取决于实际计算、通信和等待。每种切分都会引入 collective 或点对点通信。性能分析应把计算、通信和等待放在同一条时间线上,检查是否发生有效重叠。7 8

多卡训练缩写怎样区分
  • DP 复制模型、切分 batch。它主要增加训练吞吐。
  • ZeRO/FSDP 分片长期保存的参数、梯度和优化器状态。它主要减少单卡模型状态显存。
  • PP 按模型层切分。相邻 stage 之间传递 activation 和梯度。
  • TP 把同一层中的大矩阵切给多张卡。它通常要求较频繁的层内通信。
  • 上下文并行 按序列或 K/V 分段。它主要处理长序列带来的注意力计算和缓存峰值。

它们可以组合使用。组合时要把总 GPU 数写成各并行维度的乘积,再分别检查每个通信组传输什么张量。

训练与推理的显存对象不同

推理通常需要参数和 KV cache。训练还需要梯度、优化器状态以及为反向传播保留或重算的 activation。看到“7B(约 70 亿参数)模型的 BF16(bfloat16,16 位浮点格式)权重约 14 GB”时,这只是权重文件占用的量级。训练显存还要按优化器、精度、batch、序列长度和实现策略继续计算。

通信中的点对点和集合操作

点对点通信(point-to-point,P2P)由一个 rank send、另一个 rank recv 组成。集合通信(collective communication)由一个通信组内的所有 rank 共同参加,例如 4 张 GPU 同时执行一次 AllReduce,即先规约各卡的数据,再把同一份规约结果分发给所有参与者。NCCL(NVIDIA Collective Communications Library,NVIDIA 集合通信库)、MPI(Message Passing Interface,消息传递接口)和 Gloo 等通信库会依据通信器、拓扑和消息大小选择具体算法。6

集合操作要求参与者按相同顺序进入同一个 collective。某个 rank 没有进入、卡住或退出时,其他 rank 会等待,最后可能表现为 timeout。调试分布式程序时,应找到第一个没有到达同步点的 rank,而不只看报错的那张卡。

常见集合通信

设每个 rank 上有一个局部张量,collective 的名称描述数据在组内怎样变化。

操作 结束后数据在哪里 常见场景
Broadcast 根 rank 的同一份数据出现在所有 rank 初始化参数、同步控制信息
Reduce 各 rank 数据规约到根 rank 汇总和、最大值
AllReduce 规约结果出现在所有 rank 数据并行梯度同步
Scatter / Gather 根张量分块下发 / 各分块收回根 输入分发、结果汇总
AllGather 每个 rank 都得到所有分块的拼接 临时凑齐分片参数或激活
ReduceScatter 先规约,再让每个 rank 留一个分块 分片梯度与状态
All-to-All 每个 rank 给所有其他 rank 发送不同分块 MoE token 路由

用四个 rank 的 AllReduce 看语义。设局部向量分别是

rank0: [1, 0]    rank1: [0, 2]
rank2: [3, 0]    rank3: [0, 4]

NVIDIA NCCL User Guide 的 AllReduce 图。每个 rank 输入一块数据,归约后每个 rank 都获得同一份输出。图表达集合通信语义,具体使用环形还是树形算法由通信库和运行环境决定。

AllReduce 的输出复制到全部参与 rank,因此适合数据并行梯度同步。归约操作可以是 sum、max、min 等。训练中是否再除以 world size(通信组中的 rank 总数),取决于框架和优化器语义。 6

若做 sum AllReduce,四个 rank 结束后都得到 [4,6]。实现可以是 ring reduce-scatter 再 all-gather,也可以是树形规约。使用者首先关心结果语义和通信量,底层库还要考虑拓扑和算法。AllGather 的结果是拼接。若 rank0 有 [1,0]、rank1 有 [3,4],结束后两个 rank 都得到 [1,0,3,4]

集合通信自测

两个 rank 分别持有梯度分片 g0=[1,2,3]g1=[4,5,6]。执行 sum AllReduce 后各自得到什么?执行 AllGather 后各自得到什么?

详细答案

sum AllReduce 对相同位置相加,并把结果交给每个参与 rank。因此两个 rank 都得到 [1+4,2+5,3+6]=[5,7,9]

AllGather 按 rank 顺序拼接分片。因此两个 rank 都得到 [1,2,3,4,5,6]。数据并行需要每张卡拿到同一份全局梯度,适合 AllReduce。分片参数或激活需要临时凑回完整张量时,适合 AllGather。

Scatter 将一份输入分给多个进程。Gather 将各进程持有的结果收回。

Scatter 改变数据所有权,Gather 恢复集中表示。每个箭头都对应实际通信量。

图中对比 Allgather、Reduce-Scatter 与 Allreduce。AllGather 拼接各进程的 shard。后两者先规约,再让每个进程保留一个分块或得到完整结果。

三种集合通信都围绕每个进程持有一块数据展开。区别在于拼接结果和规约结果的放置方式。

AllGather 的结果是拼接后的原始数据。AllReduce 的结果是逐元素相加、平均等规约后的数据。训练中使用哪一种,要看需要的是完整参数、完整激活,还是一致的全局梯度。

思考题

数据并行同步梯度时,为什么常用 AllReduce 而不是 Reduce?

答案

数据并行的每张卡都要用同一份全局梯度继续更新自己的参数。Reduce 只把结果送到一个 rank,其他 rank 得不到结果。AllReduce 让所有 rank 结束后都拿到同一份规约结果。

通信延迟与带宽

课上用一个常见近似模型表示一次消息传输。

\[ T_{\mathrm{comm}} \approx \alpha + \frac{n}{\mathrm{BW}} = \alpha + n\beta \]

其中 \(T_{\mathrm{comm}}\) 是一次通信的估计时间,\(\alpha\) 是启动延迟,\(n\) 是消息字节数,\(\mathrm{BW}\) 是可用带宽,\(\beta=1/\mathrm{BW}\) 是每字节传输时间。这个式子省略拥塞和拓扑等因素,用于区分小消息的启动成本和大消息的搬运成本。

图中给出通信时间模型 \(\alpha + n\beta\)。\(\alpha\) 是建联或调度延迟,\(n/\mathrm{BW}\) 是传输时间。小消息主要受 \(\alpha\) 影响,大消息主要受 \(n\beta\) 影响。

通信时间由固定延迟和随消息大小增长的传输时间组成。这个模型可以帮助判断是否合并小消息,或用 ring 处理大消息。

例如 AllGather 可以用 ring 实现。把 \(p\) 个 rank 排成环,每个 rank 先持有自己的第 \(i\) 块。每一轮把当前块传给下家,同时从上家接收一块。经过 \(p-1\) 轮,每个 rank 都收到其余所有块。每轮只搬一块较小的数据,链路可以同时工作,因此对大消息常有较好的带宽利用率。Ring AllReduce 通常可理解为一段 ReduceScatter 加一段 AllGather。

图中的 Ring Algorithm。节点排成逻辑环,每轮每个节点向邻居传一块数据,经过多轮后每个节点都拿到完整结果。

ring 让所有链路的带宽同时被利用,适合大消息。多轮交换会带来启动和等待开销。

通信库会按消息大小选择算法。小张量可能承受不起 ring 的多轮启动开销,大张量则可能让一个根节点成为瓶颈。实际耗时还受 PCIe、NVLink/NVSwitch、跨节点网络和物理拓扑影响,不能只用 GPU 数量推断。

图中说明 NCCL 会根据 message size 自动选择通信算法。小消息优先减少启动开销,大消息优先利用带宽。用户一般不必手工指定。

通信库按消息大小在启动开销优先和带宽优先之间切换实现。这对应 \(\alpha\)\(n\beta\) 的相对影响。

思考题

AllReduce、AllGather、ReduceScatter 的结果分别落在哪里?各适合什么场景?

答案

ReduceScatter 把各 rank 数据归约后按片段分给所有 rank。AllGather 把各 rank 的片段收集成完整结果并分发给所有 rank。AllReduce 可以视为先 ReduceScatter 再 AllGather,因此每个 rank 都拿到同一份规约结果。数据并行的梯度平均常用 AllReduce,张量并行前向收集常用 AllGather。

思考题

四个 rank 约定执行 AllReduce,但 rank3 还在上一段计算中,前三个 rank 已经进入集合通信。系统可能表现成什么?调试时应先看哪里?

答案

前三个 rank 会在集合通信中等待,表面现象常是通信超时或进程挂起。真正的问题可能在 rank3 的前序计算、异常退出、条件分支或不同进程的通信顺序不一致。应找第一个没有到达同步点的 rank,并检查所有分支是否进入同一个 collective。

数据并行

数据并行(data parallelism, DP)沿 batch 维切分数据。假设有四张 GPU,GPU 0 处理 mini-batch 的第 0 份,GPU 1 处理第 1 份。每张卡都保留完整的模型参数 \(\theta\)。它们分别执行前向与反向,得到局部梯度 \(g_0,g_1,g_2,g_3\)

图中展示数据并行的前提。单张 GPU 能容纳完整的模型副本,各卡处理不同 batch,并通过通信保持参数一致。

每张卡都有一份完整模型,输入按 batch 切分。计算可以并行,参数通过通信保持一致。

要让四份模型继续保持一致,更新前需要计算全局平均梯度。

\[ g=\frac{1}{p}\sum_{i=0}^{p-1}g_i,\qquad \theta\leftarrow\theta-\eta g \]

这里 \(p\) 是参与同步的 GPU/rank 数,\(g_i\) 是第 \(i\) 张卡的本地梯度,\(g\) 是全局平均梯度,$ heta$ 是模型参数,\(\eta\) 是学习率。

AllReduce 可以同时完成求和和结果分发,所以每张 GPU 得到同一个 \(g\),再各自执行同样的 optimizer step,最终参数仍相同。PyTorch DistributedDataParallel 默认在 backward 中同步梯度。梯度除以 world size 的位置应以当前框架实现和训练配置为准。通信发生在反向传播产生梯度之后。缺少这一步时,四张卡会各自训练参数不同的模型。7

DP 的加速上限受通信约束。若单卡一个 step 的计算时间是 \(C\),AllReduce 梯度时间是 \(M\),理想多卡时间仍至少有 \(C+M\) 的一部分。GPU 数增加后,单卡 batch 变小、固定开销占比上升,通信量也会随梯度大小变化。因此吞吐曲线常常先近似线性增长,再逐渐弯曲,最后增加卡数不再划算。

较早的做法是 parameter server。多个 worker 把梯度交给服务器,服务器汇总、更新参数后再发回。它的逻辑直接,但服务器既是吞吐瓶颈也是单点故障。现代同步 DP 更常让 worker 直接集合通信。

图中的 Parameter Server。多个 Worker 把梯度发给集中的 Parameter Server,服务器汇总更新后把新参数发回各 Worker。

集中式参数服务器逻辑直观,但所有梯度都汇集到一点,因此会形成吞吐瓶颈和单点故障。

图中展示去中心化的 Allreduce DP。每个 worker 与相邻 worker 直接组成环或树做规约,不需要集中式服务器。

现代 DP 让 worker 直接通过 AllReduce 交换梯度,避免把通信集中到单个参数服务器。

Global Batch 与梯度累积

若每卡 batch 为 \(b\),数据并行大小为 \(p\),没有梯度累积时全局 batch 是 \(B_{\text{global}}=pb\)。显存不足而又希望维持大 batch 时,可以连续运行 \(k\) 个 micro-batch,只累加梯度,最后再同步和更新参数。此时全局 batch 为 \(pkb\)

batch 变大后,单步梯度的统计特性、每个 epoch 的更新次数和学习率设置都会变化。吞吐提升后,训练配置也要随并行规模一起重新检查,确认收敛速度和最终指标保持在预期范围。

AllReduce 的归约语义

数据并行的每个 rank 持有同一组参数、处理不同 mini-batch。若本地梯度为 \(g_r\),同步 SGD 需要得到

\[ g=\frac{1}{P}\sum_{r=0}^{P-1}g_r \]

\(P\) 是数据并行通信组的 rank 总数,\(r\) 是其中一个 rank 的编号,\(g_r\) 是该 rank 对自己 mini-batch 得到的梯度。

然后所有 rank 用同一个 \(g\) 更新同一份参数。AllReduce 完成的是求和,除以世界大小可能由框架或优化器完成。平均若执行两次,实际学习率会缩小 \(P\) 倍。梯度累积时,多个 micro-batch 的局部梯度应在正确边界再同步。no_sync 一类接口表示推迟 collective。

通信-计算重叠依赖梯度产生的顺序。反向传播从最后一层开始。当某个 bucket 的所有梯度准备好,就可在通信 stream 上启动 AllReduce,同时计算前面层的梯度。bucket 太小时,collective 启动次数较多。bucket 太大时,首个通信启动较晚。实际重叠需要用 profiler 时间线验证,还要检查网络、GPU DMA 和计算资源是否已经饱和。

思考题

数据并行为什么每个 GPU 都保存完整模型?这个特点带来什么通信和显存代价?

答案

数据并行的划分对象是小批量数据,每个 GPU 用自己的 micro batch 前向和反向,模型参数与优化器状态通常完整复制。反向后需要同步梯度,例如 AllReduce。它实现直接、扩展训练吞吐容易,但显存重复保存,模型变大后单卡装不下,于是需要 ZeRO、FSDP 或模型并行。

训练显存

模型状态与混合精度

\(\Psi\) 表示参数量。模型权重和梯度若用 BF16,各自约需 \(2\Psi\) 字节。使用 Adam 时,通常还有 FP32 主权重、一级动量和二级动量,各占约 \(4\Psi\) 字节。仅模型状态可以作如下粗略估算。

\[ 2\Psi\;\text{(BF16 参数)} + 2\Psi\;\text{(BF16 梯度)} + 4\Psi\;\text{(FP32 主参数)} + 4\Psi\;\text{(一阶动量)} + 4\Psi\;\text{(二阶动量)} = 16\Psi\;\text{bytes} \]

把这笔账落到 7B 模型上,可以逐项写出来。参数 14GB、梯度 14GB、FP32 主权重 28GB、一阶动量 28GB、二阶动量 28GB,合计约 112GB。若有 8 张 80GB A100,仅模型状态占用就低于总显存的一半,但 activation 会随 batch、序列长度和 checkpointing 策略继续增加。估算显存时宁可保守,因为 kernel 临时空间、通信 buffer、框架缓存和碎片都会消耗额外空间。

对一个 7B 参数模型,这一部分已约为 112 GB,尚未计入 activation、临时 buffer、通信 buffer 和框架开销。具体数值会随混合精度、优化器和是否保存主权重而变化。训练显存中占用很大的部分,往往来自训练附带保存的状态。

以一个 7B 参数模型为例,仅权重和训练附带状态就可能占据极大的显存。

图中以 LLaMA-2 7B 为例拆解参数量来源。embedding、attention 投影、FFN 等的参数相加,决定模型文件与显存的最小量级。

参数量由各层的权重矩阵之和决定。它是显存账本的起点,也是训练显存的一部分。

每个位置采用的精度不同。权重和梯度在计算过程中可以用 bf16 这类低精度,以节省显存和带宽。优化器状态和主权重通常需要更高的 fp32 精度,因为梯度需要反复累积和更新,低精度累加可能丢失很小的更新量。常见做法是以前向和反向使用 bf16,以 fp32 更新主权重,并以 fp32 保存优化器动量和方差。因此显存账本中会同时出现 bf16 计算副本和 fp32 副本。

Activation 与完整显存账本

激活是另一笔账。反向传播要用到前向阶段的中间结果,因此常规训练会保存多层 activation。sequence length、batch、hidden size 和层数增大后,activation 也可能成为主要瓶颈。activation checkpointing 通过少保存一部分中间结果、在反向时重算来换取显存。它处理的是激活开销。

activation checkpointing(激活检查点)可以用一条简单的时间线理解。普通训练在每个矩阵乘后保留中间结果,反向时直接使用。检查点方法只保留某些边界处的输入,反向到该段时重算中间结果并接着求梯度。它减少激活显存,增加前向计算次数,通常不减少优化器状态。ZeRO/FSDP 则减少参数、梯度和优化器状态的副本或分片。

训练与推理的账本不同。推理主要保存权重和 KV cache。训练还要额外保留梯度、优化器状态和 activation,因此两者的估算不能混用。

图中对比训练与推理的状态内存。训练要保留参数、梯度、优化器状态和 activation。推理则主要保存权重与 KV cache。

训练显存的大部分来自模型状态与 activation。推理显存的大部分来自权重与 KV cache。估算前应先分清负载类型。

checkpointing(activation recomputation)只保存若干边界 activation,反向时重跑中间前向以换取显存。它减少的是保存量,增加的是计算和可能的通信;不能把它与持久化检查点混为一谈。后者是故障恢复用的模型/优化器快照。选择并行策略前应分别列出参数、梯度、优化器状态、activation 与通信 buffer,而不是只报一个显存占用。

思考题

混合精度训练为什么通常仍要保存 FP32 权重副本?这和前向使用 FP16/BF16 是否矛盾?

答案

前向和部分计算用低精度提高速度、降低激活显存,但参数更新往往很小。若只在低精度权重上累加,小更新可能被舍入掉,长期训练不稳定。FP32 主权重保留更新精度,前向再转换成低精度,两者职责不同。

ZeRO / FSDP 分片模型状态

ZeRO 的三个阶段可以按切分对象记。ZeRO-1 切优化器状态,ZeRO-2 再切梯度,ZeRO-3 连参数也切。对于前述 Adam 的 16 bytes/参数,如果四张卡完全分片,每张卡的模型状态平均可降到约 4 bytes/参数;通信则从梯度 AllReduce 变为前向或反向按需 AllGather 参数,以及 ReduceScatter 梯度。

FSDP 可以理解为 PyTorch 中这类分片训练的工程实现。某一层计算前,各 rank 临时 AllGather 出该层完整参数;计算后释放完整副本,反向再按需收集,梯度用 ReduceScatter 分片写回。它和 ZeRO-3 的核心思想一致:每张卡长期保存自己负责的 shard(分片),当前层计算时再临时收集完整参数。

PyTorch FSDP 文档的说明(意译)

FSDP 将参数、梯度和优化器状态分片到数据并行工作进程中。在需要计算某层时,它临时收集完整参数;计算完成后再释放或重新分片,以降低常驻模型状态显存。8

普通 DP 中,每张卡都保存完整参数、完整梯度和完整优化器状态。数据已经切分,模型状态仍会重复保存。ZeRO(Zero Redundancy Optimizer)按阶段分片这些状态。FSDP(Fully Sharded Data Parallel)是一类围绕参数完全分片的实现。

  • ZeRO-1 分片优化器状态。每个 rank 只长期保存约 \(1/p\) 的 Adam 状态。
  • ZeRO-2 在 ZeRO-1 基础上再分片梯度。完成反向后用 ReduceScatter,让每张卡只留下自己负责的梯度块。
  • ZeRO-3 / FSDP 也分片模型参数。每张卡长期只持有参数、梯度和优化器状态的一个分片。

分得越彻底,单卡常驻状态越少,但运行时需要更多通信。ZeRO-3 在算某一层之前先 AllGather,把该层的完整参数临时凑齐;用完前向后释放不再需要的完整参数。反向到这一层时再 AllGather 参数,计算出梯度后用 ReduceScatter 把已规约的梯度留在对应 owner 上,由 owner 更新自己的参数分片。

图中展示 ZeRO-3/FSDP 的生命周期。前向前先 AllGather 各参数 shard 拼成完整参数,计算后释放。反向后用 ReduceScatter 把梯度留给各自 owner。

FSDP 在每次用到某层参数前临时 AllGather,计算后释放。它以通信换取较低的常驻显存。

图中以 GPT-3 175B 为例比较普通 DP 和各级 ZeRO 的单卡状态内存;分片级别越高,通信需求也越多。

ZeRO 依次分片优化器状态、梯度和参数。常驻显存下降,运行时 AllGather/ReduceScatter 增加。

ZeRO 论文的状态内存分解图。蓝色是参数,橙色是梯度,绿色是优化器状态。每一行对应 Baseline、只分片优化器状态、再分片梯度、再分片参数的不同阶段。

阅读这张图时横向比较同一张 GPU 上三种颜色的总高度。Baseline 中每张卡都长期保存完整三类状态。向下到 \(P_{os}\)\(P_{os+g}\)\(P_{os+g+p}\) 时,绿色、橙色、蓝色依次被分到不同 GPU,因此每张卡的常驻状态逐步下降。右侧的 7.5B 示例只用于说明趋势,实际显存还要加上 activation、通信 buffer 和框架开销。3

ZeRO 用通信换取更低的状态显存。模型因显存不足而无法训练时,这种交换很有价值。网络带宽不足、分片太细,或每层过小而频繁 AllGather 时,吞吐可能下降。参数状态过大时可考虑 ZeRO/FSDP。activation 过大时还要考虑 checkpoint、序列切分或减少 micro-batch。

思考题

ZeRO/FSDP 把参数和优化器状态切开后,为什么仍会引入通信?

答案

某一层需要完整参数时,各 rank 必须临时 AllGather 自己持有的分片;反向得到局部梯度后,又要通过 ReduceScatter 聚合并切回分片。显存占用下降,通信调度和重叠成为新的关键问题。

思考题

activation checkpointing 为什么能省显存,又为什么会让训练变慢?

答案

它不保存每个前向中间结果,只保存少量边界值,反向需要时重新执行一段前向来恢复激活。显存减少来自少存中间张量,时间增加来自重算。适合显存不足时换取更大的模型、batch 或序列长度,属于用计算换显存。

流水线并行

流水线并行(pipeline parallelism, PP)沿网络深度切分。若模型有很多层,可以让 GPU 0 负责前几层,GPU 1 负责中间层,GPU 2 负责后几层。前向时 activation 从第一个 stage 依次传到后面;反向时梯度再反向传回。每张卡只放本 stage 的层,因而可容纳单卡放不下的深模型。

最朴素的排程会产生 pipeline bubble。GPU 0 先忙,其他 stage 等待 activation。前向全部结束后,反向从最后一个 stage 才开始,前面的卡再次等待。卡的数量越多、每轮只有一个大 batch,空闲比例越明显。

图中画出 naive pipeline。一次前向和反向交替时,许多 stage 处在等待状态,空闲时间占比称为 bubble ratio。

朴素流水线中,一个 batch 串行穿过各 stage,前后期的 stage 大量空闲。这种空闲比例称为 bubble。

处理办法是把 batch 切为多个 micro-batch,形成填充流水线。GPU 0 算完 micro-batch 0 后就把它交给 GPU 1,自己立刻开始 micro-batch 1;GPU 1 接到 0 后开始计算,后续 stage 也依次被填满。等流水线稳定下来,多个 stage 可以同时处理不同 micro-batch 的前向或反向。bubble 没有消失,但被更多 micro-batch 摊薄。

图中展示 GPipe。输入 batch 被拆成多个 micro-batch,逐块沿流水线推进,让各 stage 尽量保持忙碌。

micro-batch 把大 batch 切小,流水线各 stage 因而能同时处理不同 micro-batch,从而摊薄 bubble。

PP 的通信主要是相邻 stage 间的 activation/gradient P2P 传输,延迟和负载均衡都很重要。如果某个 stage 的层更慢,整条流水线会被它卡住;切层时不能只按层数平均,还要考虑每层计算量、激活大小和显存。

流水线 bubble 可以用一个极小调度表看清楚。假设 4 个 stage、2 个 micro-batch,F 表示前向,B 表示反向,横轴是时间。若没有 micro-batch 交错,GPU 0 在第 1 个前向完成后要等后续 stage 逐步处理,反向又从 GPU 3 开始返回。画在时间线上,两端 GPU 有明显空闲,这些空闲就是 bubble。

GPU0: F0 F1 ........ B1 B0
GPU1: .... F0 F1 ... B1 B0
GPU2: ....... F0 F1 B1 B0
GPU3: .......... F0 B0 F1 B1

增加 micro-batch 数量、让不同 micro-batch 的前向和反向交错,可以填掉一部分空档,但通信次数和调度复杂度也随之增加。

思考题

流水线并行的 bubble 是怎么来的?

答案

前后 stage 之间有依赖,前面的 stage 要等激活送来,后面的 stage 在最早几个 micro-batch 时可能没有工作;反向阶段又会出现对称的空闲。把 batch 切成更多 micro-batch 可以降低 bubble 比例,但会增加调度和通信复杂度。

张量并行

PP 是按层切,张量并行(tensor parallelism, TP)则把一层内部的大矩阵乘法切到多张 GPU。Transformer 的 attention 投影和 FFN 都包含很大的矩阵乘法,这往往是计算最重的地方。

考虑线性层 \(Y=XW\)。若按 \(W\) 的列切分,计算过程如下。

\[ W=[W_0\;W_1\;\cdots\;W_{p-1}],\qquad Y=[XW_0\;XW_1\;\cdots\;XW_{p-1}] \]

\(X\) 的行是 token 或样本,列是输入 hidden 维度;\(W\) 的行对应输入维度,列对应输出维度。\(W_i\) 是第 \(i\) 张卡保存的一组输出列,因此 \(XW_i\) 正好是 \(Y\) 的同一组输出列。

每张 GPU 可独立算一个 \(XW_i\),得到 \(Y\) 的一部分列;若后续需要完整的 \(Y\),再做 AllGather。若按行切分,则每张卡先算局部乘积,最后用 AllReduce 把各部分相加。不同的切法决定通信发生在算子前还是算子后。

用数值说明列切和行切的区别。设

\[ X=\begin{bmatrix}1&2\end{bmatrix},\qquad W=\begin{bmatrix}1&1\\2&3\end{bmatrix} \]

完整结果是 \(XW=[5,7]\)。若按列切 \(W_0=[1,2]^T\)\(W_1=[1,3]^T\),GPU 0 算 \(XW_0=[5]\),GPU 1 算 \(XW_1=[7]\),需要时 AllGather 拼成 [5,7]。若按行切 \(W\),两块分别是 [[1,1]][[2,3]],GPU 0 得 [3],GPU 1 得 [4],必须相加才有 [7]。列切天然得到输出列分片,行切天然得到需要求和的部分和。

FFN 可以写成

\[ \operatorname{FFN}(X)=\phi(XW_1+b_1)W_2+b_2 \]

\(X\) 是 token 的隐藏表示,\(W_1\) 将最后一维扩展到中间维度,\(W_2\) 再投回原 hidden size;\(b_1,b_2\) 是偏置,\(\phi\) 是逐元素激活函数。

其中 \(\phi\) 是逐元素激活函数。第一层把 hidden size 扩到较大的中间维度,第二层再投回去。逐元素激活不会混合不同列,因此可在各 GPU 的局部结果上直接执行;恰当安排 \(W_1\) 的列切分与 \(W_2\) 的行切分,可把本来两次的全量通信压缩为较少的同步点。这就是 TP 在代码里常见列并行线性层和行并行线性层配对的原因。

TP 的通信频率很高,因为几乎每层都要在组内交换中间结果。因此通常把一个 TP group 放在同一台机器、使用 NVLink/NVSwitch 等高速互连;把频繁的 TP 通信跨慢网络,会非常不划算。

图中展示 Megatron-LM 的张量并行 MLP。FFN 的矩阵按列或行切到多卡,层内用多次 allreduce 同步中间结果。

张量并行把一层内的矩阵乘切开,卡间需要频繁 allreduce。因此 TP group 通常放在 NVLink 等高速互连内。

思考题

把一个 [H,4H] 权重按列切成两份后,两个 rank 各自得到的输出形状是什么?为什么还需要通信才能得到完整 FFN 结果?

答案

输入 [B,S,H] 乘每份 [H,2H],两个 rank 分别得到 [B,S,2H]。这只是中间结果的一部分,FFN 需要把这些片段合并成 [B,S,4H] 再过激活和第二层。因此张量并行在层内引入 AllGather 或 ReduceScatter,通信频率比数据并行高,通常放在高速互连内。

思考题

FFN 的第一层权重按列切,第二层权重按行切。这两种切法各自解决什么通信问题?

答案

第一层按列切后,每个 rank 得到输出特征的一段,不需要先通信输入的不同列。第二层按行切时,各 rank 局部输出可以直接按分块矩阵乘法相加,最后通过 AllReduce 得到完整结果。切分方式要顺着矩阵乘法的形状选择,避免为了凑形状引入多余通信。

多维并行与通信计算重叠

真实的大模型训练常把几种策略叠起来。节点内做 TP(Tensor Parallelism,张量并行),按层做 PP(Pipeline Parallelism,流水线并行),剩余维度做 DP(Data Parallelism,数据并行),再用 ZeRO/FSDP 分片状态。这常被称为 3D 并行。阅读一组并行度配置时,可以把总 GPU 数写成各维度的乘积,再逐维检查下列问题。

图中给出 6D 并行框架。经典的 3D 数据、张量和流水线并行之上,又叠加序列并行、上下文并行与专家并行,并附有一个 GPT-3 175B 的训练配置实例。

总 GPU 数等于各并行维度的乘积。例如 GPT-3 的配置 \(8\times 8\times 60=3840\) 块 A100。维度越多,切分越细,通信模式也越多样。

以 GPT-3 175B 为例,可把 8 块 A100 通过 NVLink 组成一个节点作为 TP 组,8 个节点按层做 PP,再复制 60 份做 DP,得到 3840 块 GPU 的集群配置。每个维度都对应一类通信。TP 在节点内频繁 AllReduce,PP 在相邻 stage 之间传 activation/gradient,DP 需要全局 AllReduce 梯度。SP(Sequence Parallelism,序列并行)、CP(Context Parallelism,上下文并行)沿序列维切分,EP(Expert Parallelism,专家并行)则按 expert 切分并触发 All-to-All。面对具体配置时,应先算清 GPU 如何分组、组间传什么数据。

检查并行配置时,把总卡数和乘积逐项写出,例如 3840 与 \(8\times8\times60\),随后给每个维度标注通信对象;TP 是节点内高频集合通信,PP 是相邻 stage 的点对点,DP 是全局梯度同步。最后核对每个组的大小是否匹配物理拓扑。TP 放在慢网络两端是最常见的错误之一,因为每层都可能等待通信。

  • DP 组里的卡拿不同数据,何时 AllReduce 梯度?
  • TP 组里的卡共享一层计算,哪个线性层后要 AllGather 或 AllReduce?
  • PP 的相邻 stage 之间传什么 activation 和 gradient?
  • ZeRO 分片后,哪一层前需要临时 AllGather 参数?

通信不一定必须停在计算之后才开始。反向传播从最后一层往前走,某层梯度就绪后可以立刻启动对应 bucket 的 AllReduce,同时 GPU 继续计算更前面的层。若通信在计算结束前完成,它的耗时被隐藏;若网络仍未完成,GPU 才需要等待。bucket 的大小过小会让启动延迟累积,过大又会推迟通信启动,因而是一个需要实测的折中。

图中用反向传播的时间线说明通信和计算重叠。反向算出一部分梯度后立刻发起通信,并与后续计算重叠,从而缩短可见的通信等待。

梯度按层或按 bucket 就绪后启动集合通信,不必等整个反向结束。重叠是否成立要看网络是否在计算完成前清空。

手动组合 TP、PP、DP、SP、CP、EP 并不容易,因此有了自动并行。它根据模型计算图和硬件拓扑,例如节点数量、节点内 NVLink、节点间 IB,选择每个算子的切分方式、流水线 stage 划分和 micro-batch 大小。

图中介绍 Alpa 的分层自动并行。先做算子间 stage 划分,再为 stage 内算子选择切分策略,最后把逻辑设备网格映射到物理拓扑。

自动并行将层间划分、层内切分和设备映射分别处理,再组合为整体方案。

这类工具(Alpa 之外,PyTorch DTensor、OneFlow SBP 也是同类思路)把设备网格抽象出来,统一描述张量在设备间如何分布,从而让同一份模型代码在不同并行度配置下运行。它的价值不在替代工程师手调,而在把可复现、可搜索的并行空间交给算法,减少每次换模型/换集群都要重写配置的负担。

思考题

总 GPU 数是 64,配置为 TP=4、PP=2、DP=8。这三组分别是什么?

答案

4 卡组成一个张量并行组,共同计算一层内被切开的矩阵;2 个流水线 stage 按层前后连接;8 份数据并行副本同时处理不同 micro-batch 并同步梯度。\(4\times2\times8=64\),读配置时先确认每组大小和物理拓扑是否匹配。

Transformer 层的算力估算

要判断长序列到底难在哪,先把一层 Transformer 的算力分清楚。以一层的 MLP 与 attention 为例,以 batch \(B\)、序列长 \(S\)、隐宽 \(H\) 计,各部分计算量大致如下。

图中给出每层算力估算表。Q/K/V/O 四个投影约 \(8BSH^2\),attention 的 \(QK^T\) 与 \(Softmax\times V\) 各约 \(2BS^2H\),FFN 的 SwiGLU 三个矩阵约 \(6BSH\)。

投影与 FFN 随 \(S\) 线性增长、随 \(H\) 平方增长。attention 的两个矩阵运算随 \(S\) 平方增长,只随 \(H\) 线性增长。

关键在 attention 的 \(S^2H\) 项。把这一层的总计算写成 \(8BSH^2+4BS^2H+6BSH^2\),当 \(S\ll H\)(短序列)时,\(SH^2\) 项占主导,FFN 与投影是瓶颈;当 \(S>H\)(长序列)时,\(S^2H\) 项爆炸,attention 变成瓶颈。

这里的系数来自矩阵乘的形状。一次 \([m,k]\times[k,n]\) 矩阵乘大约需要 \(2mkn\) 次乘加。QKV 和输出投影的形状都围绕 \(S\times H\)\(H\times H\),合起来给出 \(BSH^2\) 量级;\(QK^T\)\(S\times H\)\(H\times S\),乘 \(V\)\(S\times S\)\(S\times H\),所以各自都含 \(S^2H\)。做实验时,把每个矩阵的 shape 代入这个公式,就能解释 profiler 里 GEMM 和 attention kernel 的相对耗时。

图中画出 attention 与 MLP 计算量随序列长 \(S\) 变化的双对数曲线。attention \(\propto S^2H\) 增长更快,在 \(S\) 超过隐宽附近反超 MLP。

横轴是序列长 \(S\),纵轴是计算量。attention 的二次增长会让它在一段序列长度后反超 MLP,成为长上下文的主要成本。

序列并行、FlashAttention(按块计算并在线累计 softmax 的注意力实现)和 ring attention 都要处理 attention 的 \(S^2\) 项。它既带来较大的计算量,也让中间 attention 矩阵和 KV cache 难以完整物化或驻留单卡。可以先区分 \(S\ll H\)\(S>H\) 两种情况,再决定投影和 FFN 是否切分,或 attention 是否分块。

自回归语言模型推理

前面的章节讨论训练时怎样切分参数、梯度和计算。推理不再计算梯度,主要关心请求怎样进入模型、历史 token 怎样复用,以及多个请求怎样共享显存与 GPU 时间。

语言模型读入提示词(prompt)后,根据已有 token 预测下一个 token。新 token 加回输入,再预测下一个;直到遇到结束符或达到长度上限。这种逐步生成称为自回归。token 是分词器切出的处理单位,可以是词片、汉字、标点或字节片段。embedding 把 token 编号查成向量,常见输入形状写作 \([B,S,H]\)

Hugging Face 缓存文档的说明(意译)

自回归生成会反复使用前面 token 已经得到的 key 和 value。将它们保存在键值缓存中,可以避免每一步重复计算整段历史;缓存大小会随序列长度增长。9

维度 含义
\(B\) batch size,一次并行处理的请求数
\(S\) sequence length,当前序列的 token 数
\(H\) hidden size,每个 token 的隐藏向量维度

Transformer 的自回归注意力(attention)、前馈网络(feed-forward network,FFN)和因果 mask 的基本结构以原论文为准。1 例如,\([2,16,4096]\) 表示两个请求,每个请求当前有 16 个 token,每个 token 是 4096 个浮点数。只看其中第 1 个请求时,可以去掉 batch 维,得到 \([16,4096]\)。再取其中第 3 个 token,就得到长度为 4096 的向量。矩阵乘法中的许多形状错误,来自没有写明当前张量的各维含义。

张量读法自测

一个日志中出现的张量形状是 [3, 7, 1024]。这三个数分别表示什么?去掉 batch 维后形状是什么?取出第 2 个请求的第 5 个 token 后形状是什么?

详细答案

第一个维度 3 表示 batch 中有 3 个请求。第二个维度 7 表示每个请求当前有 7 个 token。最后一维 1024 表示每个 token 的隐藏向量长度。

取出一个请求后,batch 维消失,形状为 [7,1024]。再取出该请求的一个 token 后,序列维也消失,形状为 [1024]。读分布式日志时按 batch、序列、hidden 这个顺序理解,能判断当前操作是在请求级、序列级还是 token 级。

把一个 batch 送入 Transformer 后,可先按下面的顺序理解张量如何变化。

  1. embedding 查表将 token 编号变成 [B,S,H]
  2. attention 的线性投影把最后一维从 H 变成 Q、K、V 所需维度。
  3. \(QK^T\) 在每个 head 中产生 [B,heads,S,S] 的注意力分数。softmax 后再乘 V,张量形状回到 [B,heads,S,head_dim]
  4. FFN 对每个 token 独立做 [B,S,H] × [H,4H][B,S,4H] × [4H,H] 的矩阵乘。
  5. 最后的输出投影得到 [B,S,V] 的 logits。

embedding 只给 token 一份可计算的向量表示,本身不说明它排在序列第几位。位置编码或位置嵌入补上顺序信息。张量随后依次经过多个 Transformer block;每层包含 attention(注意力)与 FFN/MLP(前馈网络)等部分。最后的输出投影和 softmax 给出词表中每个 token 的概率。

图中展示 Transformer block 的因果注意力。每个位置只访问当前位置及其左侧 token,右侧位置由 mask 屏蔽。

因果(masked)注意力让每个位置只看它之前的 token。这使自回归语言模型能够逐 token 生成。

同一模型的两种时间尺度

处理已有 prompt 时,可以并行计算许多位置,矩阵乘法大而规则。逐 token 生成时,每一步依赖上一步采样结果,通常只能把不同请求拼在一起提高并行度。两种阶段使用同一套权重,却受不同资源限制。

训练比推理多出反向传播。前向传播得到预测分布,并与目标 token 计算交叉熵损失。反向传播从损失倒着求每个参数的梯度。优化器根据梯度更新参数。看到并行训练代码时,可先还原为 forward -> loss -> backward -> optimizer.step()。并行策略只是把其中的计算和数据拆到多张卡上。

训练语言模型时,loss 会在每个可预测位置上计算。输入序列 \([x_0,x_1,\ldots,x_{S-1}]\) 的目标是右移一位的 \([x_1,x_2,\ldots,x_S]\)。第 1 个位置的输入是 \(x_0\),模型要预测 \(x_1\)。第 2 个位置的输入包含 \(x_0,x_1\),要预测 \(x_2\)。因果 mask 保证每个位置只能用左侧信息。这样一次前向能同时提供 S 个训练样本。

Prefill、Decode 与推理调度

prefill 一次处理已有 prompt 的全部 token,并建立 KV cache。KV cache 按层保存已经计算出的 key 和 value,供之后的生成步骤复用。decode 每步输入一个新 token,同时读取历史 KV,因此更容易受显存带宽和单请求延迟影响。continuous batching 将不同请求的 decode step 合并成 batch,并在请求结束时补入新请求。运行时需要维护每个序列的长度、采样状态和 KV block 映射。

用一个短例子区分 prefill 和 decode。用户输入“今天 天气 很”,prefill 一次处理这三个 token,并保存每一层对它们的 K、V。第一步 decode 只把新 token“好”送入模型。attention 查询它和前四个 token 的 K/V,输出“。”。第二步只把“。”送入,前面四个 token 的 K/V 可以复用。没有 KV cache 时,每生成一个 token 都要把完整历史重新计算一次,序列越长,重复工作越多。

continuous batching 解决的是另一个问题。请求 A 已经生成到第 100 个 token,请求 B 刚到达,请求 C 在第 8 个 token 结束。传统静态 batch 会等最长请求结束。continuous batching 可以在每个 decode step 结束后移除 C、补入 B,让设备里的 token 数尽量稳定。相应地,运行时需要处理调度、变长张量和 KV 空间管理。

练习:判断该进 prefill 还是 decode 队列

  • 请求 D 刚收到 800 个 token 的长 prompt,尚未生成任何 token。
  • 请求 E 已生成 30 个 token,正在等第 31 个 token。
  • 请求 F 刚在第 17 个 token 满足停止条件。
  • 请求 G 的 prompt 有 3 个 token,刚加入系统。
详细答案
  • D 和 G 还没有生成第一个 token,需要先对完整 prompt 建立 K/V,因此进入 prefill 队列。
  • E 已经生成 30 个 token,只需要处理新 token 并读取历史 KV,因此留在 decode 队列。
  • F 已满足停止条件,应从活动 batch 中移除,并释放或标记其 KV block 供后续请求复用。

长 prefill 包含较大的矩阵计算,decode 更关注单 token 延迟。将两类请求放进同一静态队列时,长 prompt 可能推迟交互请求的下一 token;服务系统通常分开调度,或设置抢占与批处理策略。

图中画出 LLM 推理的两个阶段。Prefill 一次性处理全部输入 token 并建立 KV cache。之后的 decode 每步生成一个新 token,并读取已有 KV。

prefill 的矩阵计算较大且规则,通常更接近计算受限。decode 每步只生成一个 token,却要反复读取 KV,常受带宽和延迟影响。两者需要不同的调度取舍。

KV cache 是按层保存的 K/V 张量,其字节数近似为 layers × tokens × 2 × kv_heads × head_dim × bytes_per_element。采用 grouped-query attention 时,kv_heads 小于 query heads,缓存会显著下降。上下文接近显存上限时,调度器需要限制并发、换出、分页管理或拒绝请求。吞吐、首 token 延迟和每 token 延迟互相牵制,因此不能只报告 token/s。

系统通常同时关注首 token 时间(TTFT,Time To First Token)、后续 token 间延迟(TPOT,Time Per Output Token)和总体吞吐。为了吞吐而积累很大的 batch,可能伤害交互请求的首 token 延迟。连续批处理、分页式 KV cache 和分开调度 prefill 与 decode 都用于协调这些目标。

训练和推理也会用 TP、PP、序列并行,但关注点不同。训练更在意吞吐、全局 batch 和模型状态。推理更在意 KV cache 能否容纳、请求到达的动态性与用户实际等待时间。分析一个方案前,应先分清它要服务的负载。

思考题

prefill 和 decode 的计算模式有什么不同?为什么它们适合不同的优化方向?

答案

prefill 一次处理整个 prompt,可以在序列维度做批量矩阵乘,算力较容易喂满;decode 每步通常只生成一个 token,主要反复读取 KV Cache 和权重,更受显存带宽和调度影响。因此 prefill 常看 GEMM 与并行策略,decode 常看 KV Cache、batching 和访存。

思考题

训练语言模型时,为什么目标序列是输入序列右移一位?

答案

第 0 个位置看 \(x_0\) 要预测 \(x_1\),第 1 个位置看 \(x_0,x_1\) 要预测 \(x_2\)。因果 mask 保证每个位置只使用左侧 token,因此一次前向能同时得到多个“预测下一个 token”的训练样本。

思考题

KV cache 保存 K 和 V,为什么不保存每个历史 token 的 Q?

答案

生成第 \(t+1\) 个 token 时,新的 query 只来自新输入 token,历史 token 只需要作为 key 和 value 被查询。历史 Q 对后续注意力没有直接用处,保存它只会增加显存。反向传播需要重算或保存激活是训练时的问题,推理没有这条需求。

长上下文与序列并行

长上下文会让 attention 的中间量和 KV cache 迅速增大。即使参数已经由 TP 或 FSDP 分片,单张卡仍可能放不下完整的序列。序列并行(sequence parallelism)和上下文并行(context parallelism)于是沿 sequence 维分割 token 或 K/V 块。

以 ring attention 为例,每张卡保存一段 Q/K/V。它让 K/V 块在 rank 组成的逻辑环上逐轮传递,使每张卡的本地 Q 依次与完整上下文的 K/V 分块计算。数值稳定的分块 softmax 与 FlashAttention 的在线累积思想紧密相关。这样不必把所有 K/V 同时复制到每张卡,峰值内存随序列分片下降,代价是多轮环形通信。 4

上下文并行沿序列维切分 K/V,各 rank 沿环交换分段以覆盖完整上下文。

每张卡只持有部分 K/V 块,通过循环传递与本地 Q 逐步和全局 K/V 交互。通信次数随分片数上升,显存峰值随之下降。

这里有一个数值细节,softmax 不能简单把每一块各自归一化后再相加,因为全局分母依赖所有块。可维护运行中的最大值 \(m\)、归一化量 \(l\) 和加权结果,块到来时用新的最大值重新缩放旧累计量。这类 online softmax 保持数值稳定,使得 attention 可以按块流式计算;FlashAttention 等高效实现同样依赖分块计算、不物化巨大注意力矩阵的思想。

假设前一块 softmax 的分母是 \(l_1=2\),累计加权和是 \(u_1\);后一块的最大值更大,旧分数在全局 softmax 中必须整体缩小。online softmax 记录当前全局最大值 \(m\),遇到新块时先算 \(e^{m_{old}-m_{new}}\),用它缩放旧的 \(l\)\(u\),再累加新块。这样不需要先看到全部分数,也能得到和完整 softmax 相同的数学结果,同时避免大量 \(e^x\) 上溢或下溢。

思考题

online softmax 为什么不能把每块 softmax 的结果直接相加?

答案

全局 softmax 的分母依赖所有块的指数和。若新块的最大值更大,旧块分数的指数都必须按新的最大值重新缩放。online softmax 维护当前最大值、累计分母和累计加权和,遇到新块时先校正旧累计量再累加。

思考题

ring attention 已经把序列切到多张卡,为什么 attention 的总通信量通常不会消失?

答案

每个 Q 块最终仍要看到全局 K/V。若每个 rank 只保存一段 K/V,就必须通过环形传递让所有块相遇。显存峰值下降,因为不必同时在一张卡上保存完整序列,但通信轮数和传输量随分片数增加。它改变的是数据的放置和流动,不改变注意力需要全局交互的性质。

MoE 与 Expert Parallel

Mixture of Experts(MoE,混合专家)层由 router(路由器)为每个 token 选择少数 expert(专家网络)。模型可以拥有很多专家参数,而一次计算只激活其中少数几个,因此单个 token 的计算量不随总专家数等比例增加。条件计算和路由负载均衡也由此成为 MoE 系统的关键约束。5

MoE 的额外开销主要来自路由和网络。一个 GPU 上的 token 可能要送往另一张卡保存的 expert;专家计算完成后,结果还要按原 token 顺序送回,这常形成 All-to-All(全互换)通信。若 router 长期偏向少数专家,这些专家和所在 GPU 会排队,其他资源却闲置;容量限制和负载均衡损失用于缓解这种偏斜。性能瓶颈可能出在 token 重排、All-to-All、专家不均衡或跨节点拓扑。5

NVIDIA NCCL User Guide 的 All-to-All 图。每个 rank 的输入被切成多段,不同段分别送往不同 rank。每个 rank 最终接收来自所有参与者的一段数据。

MoE expert parallel 常用 All-to-All 交换按 expert 分组后的 token。通信量、分段布局和 expert 放置决定它能否与本地 expert GEMM 有效重叠。 6

一个 token 在 MoE 层里的路径可以这样读。设每层有 64 个 expert,每个 token 选 top-2。输入仍是 [B,S,H],router 对每行输出 64 个分数,softmax 后取最大两个,例如 expert 7 权重 0.6、expert 42 权重 0.4。系统把所有 token 按 expert 分组,expert 7 收到自己的 token 矩阵后执行 FFN,expert 42 同样执行;最后每个 token 的输出是 \(0.6E_7(x)+0.4E_{42}(x)\),再放回原序列位置。参数规模来自 64 个 expert,单个 token 计算量却只对应 2 个 expert。

思考题

MoE 的路由为什么必须考虑负载均衡?token 全都送往少数 expert 会发生什么?

答案

若路由把大部分 token 集中到少数 expert,这些 expert 所在 GPU 会成为瓶颈,其他计算资源闲置,还可能触发容量上限丢 token。负载均衡损失、容量因子和路由统计用于控制这种偏斜。MoE 的收益来自条件计算,前提是路由能同时兼顾质量和系统负载。


  1. A. Vaswani et al., Attention Is All You Need, https://arxiv.org/abs/1706.03762

  2. M. Shoeybi et al., Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism, https://arxiv.org/abs/1909.08053

  3. S. Rajbhandari et al., ZeRO: Memory Optimizations Toward Training Trillion Parameter Models, https://arxiv.org/abs/1910.02054

  4. T. Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness, https://arxiv.org/abs/2205.14135

  5. D. Lepikhin et al., GShard: Scaling Giant Models with Conditional Computation and Automatic Sharding, https://arxiv.org/abs/2006.16668

  6. NVIDIA, NCCL User Guide: Collective Operations, https://docs.nvidia.com/deeplearning/nccl/user-guide/docs/usage/collectives.html

  7. PyTorch, DistributedDataParallel, https://docs.pytorch.org/docs/stable/generated/torch.nn.parallel.DistributedDataParallel.html

  8. PyTorch, FullyShardedDataParallel, https://docs.pytorch.org/docs/stable/fsdp.html

  9. Hugging Face, KV cache strategies, https://huggingface.co/docs/transformers/main/cache_explanation

有用的话请给我个 star => Stars 本站总浏览