大语言模型并行策略(五):上下文并行
虽然序列并行可以降低超长序列的显存占用,但在计算 Attention 时,仍然需要拿到整个序列的上下文信息。Attention 计算需要保存整个序列对应的 Q、K、V 结果,即使使用了张量并行,单个 GPU 的显存也可能无法容纳整个序列的 Q、K、V,从而无法完成 Attention 的计算。上下文并行(Context Parallelism)正是用来解决超长序列的 Attention 计算问题的。
上下文并行的原理
要想理解上下文并行,首先需要了解 Attention 的计算原理。给定输入序列的表示矩阵 ,我们通过线性变换得到查询矩阵 、键矩阵 和值矩阵 :
接下来计算查询矩阵 与键矩阵 的点积,得到注意力分数矩阵,最终通过 softmax 归一化后与值矩阵相乘,得到注意力输出:
在 Attention 的计算中,查询矩阵 与键矩阵 相乘后得到注意力矩阵,矩阵中的每一行表示一个 token 对整个序列的注意力分数,也就是每个 token 对其他 token 的关注程度。然后将注意力矩阵与值矩阵 相乘,得到最终的注意力输出。
Attention 的流式计算
思考以上过程,可以发现 中的每一个 token 彼此之间是没有关系的,但是每个 token 需要依赖完整的 和 来计算注意力输出。既然无法将 Q、K、V 完整地放置在单个 GPU 内,可以将序列中每个 token 对应的 Q、K、V 分别放置在不同的 GPU 上,每个 GPU 只持有一个 token 的 Q、K、V。当某个 GPU 需要其他 token 对应的 K、V 时,它可以通过网络从其他 GPU 上获取,这就将 Attention 计算转化成了如下问题:
def attention(q, k_iter, v_iter):
output = np.zeros_like(q)
for k, v in zip(k_iter, v_iter):
# ...
return output
在计算过程中,某个 GPU 持有本地 token 对应的 Q,我们将其记为 ,但没有完整序列的 K、V,每次它只能从其他 GPU 上读取一个 token 对应的 K、V。
Attention 的计算公式为:
首先令 为 对 的未经过 softmax 的注意力分数:
于是可以对 Attention 的计算公式进行如下改写:
我们可以在迭代过程中,不断累加分子和分母部分。当迭代完所有的 和 后,就可以得到 对应的注意力输出。下面代码描述了计算过程:
def attention(q, k_iter, v_iter):
d_k = q.shape[-1]
output = np.zeros_like(q)
expsum = 0.0
for k, v in zip(k_iter, v_iter):
# Q 对 K[i] 的未经过 softmax 的注意力分数
score = np.exp(np.dot(q, k) / np.sqrt(d_k))
expsum += score
output += score * v
return output / expsum
这里 Attention 的流式计算过程与 Flash Attention 的计算原理类似,上面的计算中我忽略了 safe softmax 的处理,更完善的计算过程可以参考文章 Flash Attention - 基本原理。
上面解释了如何在不持有完整上下文的情况下计算 Attention 输出,这就是上下文并行所依赖的核心原理。
Ring Attention
Ring Attention 是对上下文并行的一种实现。它的做法是将输入序列切分成子序列,然后将子序列分布到不同的 GPU 上进行计算。以 4 个 GPU 处理序列长度 128K 为例,Ring Attention 将序列切分为 4 个子序列,每个子序列长度为 32K,每个 GPU 只持有 32K 子序列对应的 Q、K、V。

在计算注意力时,4 个 GPU 之间互相传递自己持有的 K、V,当每个子序列的 K、V 经过所有 GPU 后,每个 GPU 就可以计算出自己持有的 Q 对应的注意力输出。

图片来源: https://coconut-mode.com/posts/ring-attention/
初始阶段,每个 GPU 持有本 GPU 对应子序列的 Q、K、V。后续每个 GPU 只需要将自己持有的 K、V 发送给下一个 GPU,同时接收上一个 GPU 发送过来的 K、V。
以上过程中,单个 GPU 始终只持有 32K 的 Q、K、V,而不需要持有完整的 128K 的 Q、K、V,从而降低了显存占用。
以 GPU 0 为例,整个注意力计算分为 4 个阶段,每个阶段它持有的数据如下:
Q[0-32K], K[0-32K], V[0-32K] # 阶段 1
Q[0-32K], K[32K-64K], V[32K-64K] # 阶段 2
Q[0-32K], K[64K-96K], V[64K-96K] # 阶段 3
Q[0-32K], K[96K-128K], V[96K-128K] # 阶段 4
通信模式
在 Ring Attention 中,GPU 之间的通信模式是环形的,每个 GPU 将自己持有的 K、V 发送给下一个 GPU,同时接收上一个 GPU 发送过来的 K、V。在实际的实现中,为了实现通信和计算的重叠,通常会使用双缓冲的方式,在计算当前阶段的注意力输出时,同时将当前阶段的 K、V 发送给下一个 GPU,并接收上一个 GPU 发送过来的 K、V。
每个 GPU 发送和接收的总数据量为 ,其中 为整个序列的 K、V 的总数据量, 为 GPU 的数量。
总结
上下文并行通过将序列切分到不同的 GPU 上,并在每个 GPU 上进行流式的 Attention 计算,从而避免了在单个 GPU 上持有完整的 Q、K、V,以此降低了显存占用。Ring Attention 是上下文并行的一种实现,它通过环形通信的方式在 GPU 之间传递 K、V,实现了高效的注意力计算。