@ickma2311: Efficient AI 第19讲:分布式训练(第一部分)这一讲让我更清楚地了解了自注意力……

X AI KOLs Timeline 新闻

摘要

第19讲高效AI分布式训练总结,涵盖数据、流水线、张量和序列并行方法,并附有关内存和通信瓶颈的说明。

Efficient AI 第19讲:分布式训练(第一部分) 这一讲让我更清楚地了解了自注意力并行如何在GPU之间流动。 - 从Transformer注意力块开始:输入词元被投影为Q、K、V。 - 在张量并行中,QKV投影可以跨GPU拆分,通常按注意力头进行。每个设备只计算投影的一部分。 - 在局部注意力计算后,输出投影需要同步。不同GPU的部分输出通过通信原语(如All-Reduce)合并。 - 对于长序列,序列并行改变了轴:不再仅拆分头或隐藏维度,而是将词元跨设备拆分。 我的笔记: https://ickma2311.github.io/ML/HW-SW-codesign/efficient-ai-lecture-19-distributed-training-part-1.html…
查看原文
查看缓存全文

缓存时间: 2026/06/10 11:50

高效AI 第19讲:分布式训练(第一部分)

这节课让我更清晰地理解了自注意力并行化在GPU之间是如何流动的。

从Transformer注意力块开始:输入token被投影为Q、K、V。 在张量并行中,QKV投影可以按注意力头拆分到各个GPU上。每个设备只计算投影的一部分。 在局部注意力计算之后,输出投影需要同步。来自不同GPU的部分输出通过通信原语(如All-Reduce)合并。 对于长序列,序列并行改变了轴:我们不再只拆分头或隐藏维度,而是将token拆分到多个设备上。

我的笔记: https://ickma2311.github.io/ML/HW-SW-codesign/efficient-ai-lecture-19-distributed-training-part-1.html…


分布式训练(第一部分) – ∇ ickma.dev

来源: https://ickma2311.github.io/ML/HW-SW-codesign/efficient-ai-lecture-19-distributed-training-part-1.html 分布式训练是将模型训练拆分到多个设备上,同时控制两个瓶颈:

  • **内存:**模型、梯度、激活值和优化器状态能否装下?
  • **通信:**设备间能否以足够快的速度交换张量以保证计算不空闲?

主要技术要么拆分数据、要么拆分模型层、要么拆分层内的张量、要么拆分序列维度。

并行方法

数据并行

数据并行让每个设备都有一份相同的模型副本,但向每个设备发送不同的微批次。

每次迭代遵循以下模式:

  1. 复制: 每个设备以相同的模型权重开始。
  2. 前向: 每个设备在其本地微批次上计算损失。
  3. 反向: 每个设备计算本地梯度。
  4. 同步: 设备对梯度求平均,通常使用 All-Reduce
  5. 更新: 每个设备应用相同的平均梯度更新。

模型保持复制状态,而数据被分区。

流水线并行

流水线并行将模型层拆分为多个阶段。例如,一个100层的模型可以分配:

  • 第1-25层给GPU 1
  • 第26-50层给GPU 2
  • 第51-75层给GPU 3
  • 第76-100层给GPU 4

数据依次流经各个阶段。为减少空闲时间,批处理被拆分为微批次,这样不同GPU可以同时处理不同的微批次。

张量并行

张量并行将层内的张量进行分区。线性层和注意力块中的大型矩阵乘法被拆分到多个GPU上。

一个设备只计算矩阵乘法的本地切片。然后通过 All-ReduceAll-Gather 等集合通信合并部分结果。

序列并行

序列并行对序列维度进行分区。每个设备只处理一部分token。

这有助于长上下文训练,因为每个设备存储和计算的token更少。挑战在于注意力、归一化和softmax通常需要全局序列信息,因此设备必须交换部分结果。

数据并行

参数服务器

参数服务器架构将全局模型状态与工作节点的计算分开。

  • 参数服务器: 接收梯度,聚合并发送更新的模型权重返回。
  • 工作节点: 持有本地数据分片,计算本地梯度,并与服务器通信。
  • 全局状态: 参数服务器保持各工作节点模型的一致性。

单节点 vs. 分布式训练

单节点训练迭代简单:

\[ \text{样本} \rightarrow \text{计算梯度} \rightarrow \text{更新权重} \]

分布式训练增加了通信:

\[ \text{拉取权重} \rightarrow \text{样本} \rightarrow \text{计算梯度} \rightarrow \text{推送梯度} \rightarrow \text{全局更新} \]

好处是并行计算。代价是同步和通信开销。

通信原语

分布式训练依赖于通信集合。模型并行策略决定了哪个集合在关键路径上。

一对一

点对点通信将数据从一个进程发送到另一个特定进程。

示例:

节点0 -> 节点3

Scatter 和 Gather

Scatter 是一对多。源节点将张量拆分为块,并将每个块发送给一个工作节点。

Gather 是多对一。工作节点将本地结果发回源节点,源节点重建完整张量。

Reduce 和 Broadcast

Reduce 将来自多个工作节点的值聚合成一个结果。例如:

\[ [1] + [2] + [3] + [4] = [10] \]

Broadcast 将源节点的一个张量发送给所有其他节点。

All-Reduce 和 All-Gather

All-Reduce 结合了reduce和broadcast。每个工作节点贡献数据,每个工作节点都接收到相同的聚合结果。

这是同步数据并行训练的核心操作:

\[ g = \frac{1}{N} \sum_{i=1}^{N} g_i \]

其中 \(g_i\) 是来自工作节点 \(i\) 的梯度。

All-Gather 从所有工作节点收集数据,并将完整拼接后的结果分发回所有工作节点。

方法时间复杂度峰值节点带宽总带宽
参数服务器\(O(1)\)\(O(N)\)\(O(N)\)
All-Reduce: 顺序\(O(N)\)\(O(N)\)\(O(N)\)
All-Reduce: 环\(O(N)\)\(O(1)\)\(O(N)\)
All-Reduce: 并行\(O(1)\)\(O(N)\)\(O(N^2)\)

递归 All-Reduce

递归all-reduce使用倍增或减半的通信模式。

每一步,节点与偏移量递增的伙伴交换数据:

  1. 偏移量1
  2. 偏移量2
  3. 偏移量4
  4. 继续直到所有节点都贡献完毕

这在 \[ \log_2(N) \] 步通信后达到全局同步。

ZeRO 和 FSDP

训练期间的内存使用

以FP16权重为例:

  • 权重:每个参数2字节
  • 梯度:每个参数2字节
  • Adam优化器状态:每个参数约12字节

因此标准的复制式内存成本大致为:

\[ 2 + 2 + 12 = 16 \text{ 字节每参数} \]

在80 GB GPU上:

\[ \frac{80\text{ GB}}{16\text{ 字节}} \approx 5.0 \text{ 十亿参数} \]

这远低于大型语言模型的规模,因此在每个GPU上简单复制所有训练状态无法扩展。

ZeRO-1:分片优化器状态

ZeRO-1 将优化器状态拆分到 \(N\) 个GPU上。

权重和梯度仍被复制,但每个GPU只存储优化器状态的 \[ \frac{1}{N} \]。

每个参数近似内存:

\[ 2 + 2 + \frac{12}{N} \]

对于 \(N=64\):

\[ 2 + 2 + \frac{12}{64} \approx 4.2 \text{ 字节每参数} \]

这将模型容量提升到大约190亿参数(使用80 GB GPU)。

ZeRO-2:分片优化器状态和梯度

ZeRO-2 额外分片梯度。

每个参数近似内存:

\[ 2 + \frac{2}{N} + \frac{12}{N} \]

对于 \(N=64\):

\[ 2 + \frac{2}{64} + \frac{12}{64} \approx 2.2 \text{ 字节每参数} \]

(原文此处写36 billion,但计算应为约36.4 billion,保留原文近似数值)

这将容量提升到大约360亿参数(使用80 GB GPU)。

ZeRO-3:分片参数、梯度和优化器状态

ZeRO-3 分片所有三种主要训练状态:

  • 参数
  • 梯度
  • 优化器状态

每个参数近似内存:

\[ \frac{2}{N} + \frac{2}{N} + \frac{12}{N} = \frac{16}{N} \]

对于 \(N=64\):

\[ \frac{16}{64} = 0.25 \text{ 字节每参数} \]

使用80 GB GPU,这使训练数千亿参数的模型成为可能。

在PyTorch中,ZeRO-3风格的分片通过Fully Sharded Data ParallelFSDP)实现。

流水线并行

朴素流水线并行

朴素流水线并行将模型层拆分到多个GPU上,但每个阶段必须等待前一个阶段完成后才能运行。

这产生了流水线气泡

  • 早期的GPU在向前发送激活值后变为空闲
  • 后来的GPU在接收到工作之前等待
  • 硬件利用率不足

GPipe

GPipe 通过将大批次拆分为微批次来减少气泡。

例如:

\[ [16, 10, 512] \rightarrow 4 \times [4, 10, 512] \]

一旦GPU 0完成第一个微批次,它就将其发送给GPU 1,并立即开始下一个微批次。这样多个阶段可以同时处理不同的微批次。

目标是保持所有流水线阶段在大部分迭代时间内处于忙碌状态。

张量并行

张量并行切分大型权重张量,使得单个层操作可以分布到多个设备上。

FFN层

Transformer前馈网络通常形式为:

\[ X \rightarrow A \rightarrow \operatorname{GeLU} \rightarrow B \rightarrow Z \]

常见策略是:

  1. 按列拆分 \(A\)。
  2. 在每个本地激活切片上独立应用GeLU。
  3. 按行拆分 \(B\)。
  4. 最后使用 All-Reduce 将部分输出求和。

这避免了两次矩阵乘法之间的通信,只在第二次投影后支付通信开销。

QKV 投影

对于注意力机制,查询、键和值的投影通常按列拆分。

每个GPU只计算一部分注意力头:

\[ Q_i,\ K_i,\ V_i \]

输入 \(X\) 被广播到每个设备,每个设备独立计算其本地投影。

注意力和输出投影

在本地注意力头计算完毕后,输出投影通常是行并行的。

每个设备将其本地激活切片与输出矩阵的本地行分片相乘。然后部分输出通过 All-Reduce 求和:

\[ Z = \sum_i Z_i \]

序列并行

在注意力层中重新分区数据

序列并行将token拆分到多个GPU上。这增加了可处理的最大序列长度,因为每个GPU只拥有整个上下文的一部分。

缩放规则是:

\[ \#\text{GPU数} = \#\text{序列分区数} = \#\text{头分区数} \]

注意力仍需要全局信息,因此序列并行注意力通常需要 All-to-All 通信。

Ring Attention

Ring attention 将键和值分布到多个GPU上。

每个GPU持有:

  • 本地查询 \(Q\)
  • 键的一个分片 \(K\)
  • 值的一个分片 \(V\)

GPU 在一个环上传递 \(K\) 和 \(V\) 块。当GPU接收到一个新块时,它针对该块计算其本地查询的注意力分数。

关键思想是重叠:

\[ \text{通信} \quad \text{与} \quad \text{注意力计算} \]

这使得长上下文注意力比把所有键和值收集到每个设备上更具可扩展性。


来源:高效AI,第19讲:分布式训练第一部分。

相似文章