大语言模型并行策略(二):数据并行
在训练大模型时,最简单的做法是使用一个 GPU,使用一批训练数据,在单台机器上,从头到尾完成前向计算、反向求梯度,然后更新权重的整个流程。对于小模型而言,这种方式完全够用,而且实现上很容易。但在训练大模型时,通常都需要使用百亿甚至上千亿的 token,如果仅使用单个 GPU 来训练大模型,训练速度会非常慢。另外,为了保证训练的稳定性,需要使用较大的 batch size 来计算出更加稳定的梯度。为了处理完指定 batch size 的训练样本,完成一次权重的更新,使用单个 GPU 也需要消耗大量的时间。
数据并行的原理
数据并行的思想非常简单,它将模型复制多份,将训练数据切分为多个子集,然后将这些子集分发到不同的模型实例上进行训练。多张 GPU 同时持有同一份完全相同的模型权重,各自负责使用一小部分数据进行训练。各 GPU 独立完成了前向和反向计算之后,可以得到由不同训练数据计算得到的梯度。

然后将多个 GPU 上计算出来的梯度进行汇总,各个 GPU 得到汇总后的平均梯度后,使用相同的平均梯度来更新模型参数。这样就等效于使用超大的 batch 来进行训练。下面是数据并行的训练流程:
数据并行的核心目的是提高训练的吞吐量,它利用更多的硬件资源,将训练任务分布到多个 GPU 上进行计算,从而在相同的时间内处理更多的训练数据,提升整体的训练效率。
数据并行的优化:计算和通信重叠
数据并行整体上可以分为 4 个阶段:
- 前向计算:每个 GPU 使用自己负责的训练数据,独立完成前向计算,得到损失值。
- 反向计算:每个 GPU 使用自己负责的训练数据,独立完成反向计算,得到模型参数的梯度。
- 梯度汇总:多个 GPU 之间进行通信,将各自计算出来的梯度进行汇总,得到最终的平均梯度。
- 参数更新:每个 GPU 使用相同的平均梯度来更新模型参数。
上述过程中,每个 GPU 先独立完成前向计算与反向传播,得到本地梯度,随后进行 All-Reduce 同步并汇总梯度,只有梯度汇总完成后,才能进行参数更新。在同步梯度时,所有计算都需要等待通信完成,通信和计算是串行的。可以在反向传播过程中逐层进行梯度汇总,这样就可以将通信和计算重叠起来,从而减少 GPU 等待的时间,提高训练速度。

在反向传播阶段,不需要等待所有层的梯度计算完成后再进行通信,而是可以在每一层的梯度计算完成后,立即将该层的梯度发送给其他 GPU 进行汇总,这样就可以将通信和计算重叠起来,从而减少通信开销,提高训练速度。
数据并行中的通信模式
在执行完前向和反向计算后,每个 GPU 都会得到模型参数的梯度,每个 GPU 的梯度由不同的训练数据计算得到,多个 GPU 上的梯度需要求和取平均,然后各个 GPU 使用相同的平均梯度来更新模型参数。
每个 GPU 的梯度量与模型中可训练的参数量相同,设为 G 字节。假设有 N 个 GPU,那么每个 GPU 需要向其他的 N-1 个 GPU 发送自己的梯度,同时也需要从其他的 N-1 个 GPU 接收梯度。似乎每个 GPU 需要发送和接收的总数据量为 ,而且通信量会随着 GPU 数量的增加而线性增加,感觉通信开销会非常大。
在同步梯度时,会使用一种叫做 all-reduce 的通信模式,这个通信模式将多个 GPU 上的梯度进行汇总,然后将结果分发回每个 GPU。在实际实现中,all-reduce 通常会使用 reduce-scatter + all-gather 的方式来实现。
All-Reduce 的工作原理
Reduce-scatter 的工作原理是将待发送数据分为 N 个小块,每个 GPU 负责其中的一个小块,GPU 之间通过多次通信,将每个 GPU 的小块数据进行汇总,最终每个 GPU 都会得到自己负责的小块的汇总结果。
上图中,每个 GPU 最初都持有一个向量,这些向量被划分为 4 个小块,执行完 reduce-scatter 操作后,每个 GPU 都会得到自己负责的小块的汇总结果。Reduce-Scatter 的原理如下图所示:
首先 GPU 0/2 和 GPU 1/3 交换 1/2 的数据,结果如上面子图二所示。然后 GPU 0/1 和 GPU 2/3 交换 1/4 的数据,最终每个 GPU 都会得到自己负责的小块的汇总结果。整个过程中,每个 GPU 只需要发送和接收 字节的数据,通信量不会随着 GPU 数量的增加而线性增加。
之后可以使用 all-gather 操作将其他 GPU 上的梯度块收集回来,最终每个 GPU 都会得到完整的汇总后的数据。下面是 all-gather 的示意图:
All-Gather 的原理如下图所示:
在 All-Gather 操作中,每个 GPU 发送与接收的数据量为 字节,通信量同样不会随着 GPU 数量的增加而线性增加。
回到数据并行的讨论中,假如有 N 个 GPU,单个 GPU 的梯度量为 G 字节,使用 All-Reduce 的通信模式,那么每个 GPU 需要发送和接收的数据量为:
当 N 较大时,通信开销接近于 2G,这个通信开销是可以接受的,远远小于最初估算的 。
总结
数据并行并没有减少单个 GPU 的显存占用量,它只是一种提升吞吐量的手段。其核心思想是将模型复制多份,将训练数据切分为多个子集,分发到不同的模型实例上进行训练,然后多个 GPU 之间通过通信将各自计算出来的梯度汇总为平均梯度,每个 GPU 再使用相同的平均梯度来更新模型参数。
数据并行的主要优点是实现简单,易于扩展。它的主要目的实际上是提升吞吐量,在训练过程中,为了保持训练的稳定性,通常需要使用较大的 batch size 来计算出更加稳定的梯度。但是较大的 batch size 可能导致激活值过多,进而导致 OOM。一种方式是使用梯度累加的方式来模拟大 batch size 的训练,但这种方式会增加训练时间。而使用数据并行之后,在多个 GPU 上并行地处理小的 batch,汇总梯度后,这实际上等价于使用更大的 batch size 来训练。