WangYu::Space

cat /dev/mind

大语言模型并行策略(五):上下文并行

分类:机器学习标签: LLM创建时间:2026-06-22 22:10:00

虽然序列并行可以降低超长序列的显存占用,但在计算 Attention 时,仍然需要拿到整个序列的上下文信息。Attention 计算需要保存整个序列对应的 Q、K、V 结果,即使使用了张量并行,单个 GPU 的显存也可能无法容纳整个序列的 Q、K、V,从而无法完成 Attention 的计算。上下文并行(Context Parallelism)正是用来解决超长序列的 Attention 计算问题的。

上下文并行的原理

要想理解上下文并行,首先需要了解 Attention 的计算原理。给定输入序列的表示矩阵 XX,我们通过线性变换得到查询矩阵 QQ、键矩阵 KK 和值矩阵 VV

Q=XWQ,K=XWK,V=XWVQ = XW_Q, \quad K = XW_K, \quad V = XW_V

接下来计算查询矩阵 QQ 与键矩阵 KK 的点积,得到注意力分数矩阵,最终通过 softmax 归一化后与值矩阵相乘,得到注意力输出:

Attention(Q,K,V)=softmax(QKTdk)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V

在 Attention 的计算中,查询矩阵 QQ 与键矩阵 KK 相乘后得到注意力矩阵,矩阵中的每一行表示一个 token 对整个序列的注意力分数,也就是每个 token 对其他 token 的关注程度。然后将注意力矩阵与值矩阵 VV 相乘,得到最终的注意力输出。

Attention 的流式计算

思考以上过程,可以发现 QQ 中的每一个 token 彼此之间是没有关系的,但是每个 token 需要依赖完整的 KKVV 来计算注意力输出。既然无法将 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,我们将其记为 QlocalQ_\text{local},但没有完整序列的 K、V,每次它只能从其他 GPU 上读取一个 token 对应的 K、V。

Attention 的计算公式为:

Attention(Q,K,V)=softmax(QKTdk)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V

首先令 sis_iQQKiK_i 的未经过 softmax 的注意力分数:

si=exp(QlocalKiTdk)s_i = \exp(\frac{Q_\text{local} K_i^T}{\sqrt{d_k}})

于是可以对 Attention 的计算公式进行如下改写:

Attention(Qlocal,K,V)=softmax(QlocalKTdk)V=isijsjVi=isiVijsj\begin{aligned} \text{Attention}(Q_\text{local}, K, V) &= \text{softmax} \left(\frac{Q_\text{local}K^T}{\sqrt{d_k}}\right)V \\ &= \sum_i \frac{s_i}{\sum_j s_j} V_i \\ &= \frac{\sum_i s_i V_i}{\sum_j s_j} \end{aligned}

我们可以在迭代过程中,不断累加分子和分母部分。当迭代完所有的 KiK_iViV_i 后,就可以得到 QlocalQ_\text{local} 对应的注意力输出。下面代码描述了计算过程:

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 发送和接收的总数据量为 2×GN×(N1)2 \times \frac{G}{N} \times (N-1),其中 GG 为整个序列的 K、V 的总数据量,NN 为 GPU 的数量。

总结

上下文并行通过将序列切分到不同的 GPU 上,并在每个 GPU 上进行流式的 Attention 计算,从而避免了在单个 GPU 上持有完整的 Q、K、V,以此降低了显存占用。Ring Attention 是上下文并行的一种实现,它通过环形通信的方式在 GPU 之间传递 K、V,实现了高效的注意力计算。

评论 评论内容仅博主可见,不会公开显示)