How to Implement Sequence Parallelism in DeltaNet

发表时间: 2025-11 · Blog post by YyWang (yywangcs.notion.site)

原文: https://yywangcs.notion.site/DeltaNet-2a9fc9f5d8058013a498f34e0b25bd52

速读

一句话结论 本文针对 DeltaNet 状态更新的循环依赖瓶颈,提出了一种基于增量计算与后期校正的序列并行算法,成功将严格串行的状态计算解耦,实现了长序列在多个计算单元上的高效并行处理。

要解决什么问题 DeltaNet 的前向计算过程主要由三部分组成:WY 表示计算、基于循环方式的状态更新以及最终输出的计算。在处理长序列时,通常会将序列切分为多个数据块进行分块计算。其中,WY 表示计算和输出计算天然不具备循环依赖,可以轻松实现各个数据块之间的完全并行处理。然而,真正的卡点在于核心的状态更新阶段。在这个阶段,后一个数据块的初始状态严格依赖于前一个数据块计算结束时的最终状态。这种固有的循环依赖性导致状态更新只能以串行方式向前推进。当面对极长序列时,无法将其有效切分并分配到单个 GPU 内部的不同流式多处理器(SM)或多个 GPU 上进行并行加速,导致计算单元的并行计算资源无法被充分利用,这成为了制约 DeltaNet 扩展到更长上下文和更高并行度的主要计算瓶颈。

怎么做的 为了绕开状态更新的严格串行卡点,本文提出了一种基于“增量计算与后期校正”的序列并行策略。核心思路是将原本必须等待真实初始状态才能开始的计算,拆分为可以提前进行的局部增量计算,以及后续一次性完成的全局状态校正。具体而言,算法将长序列划分为多个子序列,并分配给不同的计算单元。第一个计算单元按照原始的循环方式串行处理首个子序列,计算出真实的最终状态。对于后续的计算单元,由于在计算开始时前置依赖尚未完成,真实的初始状态是未知的,因此它们会使用一个全零矩阵作为占位符,提前开始计算临时的“增量状态”。在 DeltaNet 中,第 $i$ 个数据块的真实状态 $\mathbf{S}_{[i+1]}$ 更新公式为: $$ \mathbf{S}_{[i+1]} = \mathbf{S}_{[i]} + \left(\mathbf{U}_{[i]} - \mathbf{W}_{[i]}\mathbf{S}_{[i]}^\top\right)^\top \mathbf{K}_{[i]} $$ 使用零矩阵作为初始状态时,计算单元内部循环计算出的增量状态 $\mathbf{S}_{[i+1]}^{\Delta}$ 同样遵循该更新逻辑。将真实状态与增量状态的更新公式相减并进行递推,可以发现两者之间的差值可以通过一个累积的“缩放矩阵” $\mathbf{M}$ 来桥接。因此,后续的计算单元在计算增量状态的同时,还会初始化一个单位矩阵 $\mathbf{M}_{\text{start}}=\mathbf{I}$,并同步循环更新这个缩放矩阵: $$ \mathbf M_{[i+1]}=\mathbf M_{[i]}(\mathbf{I}-\mathbf{W}_{[i]}^\top \mathbf{K}_{[i]}) $$ 在计算过程中,每个时间步的增量状态和缩放矩阵都会被实时存储到全局内存中。一旦第一个计算单元完成了首个子序列的计算,它就会将真实的最终状态(即第二个子序列的真实初始状态 $\mathbf S_{\text{start}}$)广播出去。此时,后续计算单元利用之前存下的增量状态和缩放矩阵,通过以下非循环的校正公式,一次性还原出所有正确的真实状态: $$ \mathbf S_{[i]}=\mathbf S_{[i]}^{\Delta} +\mathbf S_{\text{start}}\mathbf M_{i} $$ 这个关键的校正步骤彻底摆脱了时间步之间的循环依赖,不仅可以在子序列内部对所有时间步进行完全并行计算,还能直接与最终输出的计算过程相融合,从而大幅减少额外的内存访问开销。通过这种设计,算法成功将原本漫长的串行等待时间,转化为了高度并行的增量预计算和高效的后期校正。

效果如何 由于本文档为理论算法设计与推导笔记,原文本中并未提供具体的模型规模、训练数据量、硬件配置等实验环境搭建细节,也未列出用于对比的基线方法和具体的量化性能指标。就方法本身的机制而言,该并行算法的主要代价在于引入了额外的内存读写开销:在增量计算阶段,需要将每个时间步生成的缩放矩阵持续写出到全局内存中,并在随后的校正阶段再次将其读入。此外,在并行度较高(例如划分为 4 个或更多子序列)的场景下,状态校正过程本身存在一条串行的依赖链。即第二个子序列必须等待第一个子序列的真实状态来完成校正,第三个子序列又必须等待第二个子序列校正完毕后产出的真实状态,以此类推。这意味着虽然增量计算是完全并行的,但校正阶段的启动时间会随着子序列的顺位依次向后延迟,这构成了该算法在极高并行度下需要面对的固有局限性。

DeltaNet如何做序列并行

A1 主要贡献

本文旨在探讨如何对DeltaNet模型实施序列并行化,特别是在单个GPU内部不同流式多处理器(SM)之间的并行计算,其原理同样适用于多GPU间的并行。

通过这种方式,算法成功地将原本严格串行的状态更新过程,分解为大部分可并行的“增量计算”阶段和一个同样可并行的“校正”阶段,从而有效利用了并行计算资源。

A2 方法细节

DeltaNet的核心计算流程

DeltaNet的计算过程。对于DeltaNet模型,其主要的计算流程可以分解为以下三个步骤:

  1. WY表示计算。
  2. 通过循环(Recurrent)方式更新状态 $\mathbf S$。
  3. 计算最终输出 $\mathbf O$。
    在这些步骤中,步骤1和步骤3天然不具备循环依赖,序列中的不同数据块(chunk)之间可以完全并行处理。真正的循环依赖存在于步骤2,即状态 $\mathbf S$ 的更新过程中。

状态更新的循环公式

状态S的更新公式。根据博客文章【1,DeltaNet Explained (Part II) (2024) https://sustcsonglin.github.io/blog/2024/deltanet-2/】中 的“Chunkwise Parallel Form for DeltaNet”一节,步骤2中状态更新的计算公式如下:

$$\begin{aligned} \begin{align*} \mathbf{S}_{[i+1]} &= \mathbf{S}_{[i]} (\mathbf{I}-\mathbf{W}_{[i]}^\top \mathbf{K}_{[i]}) + \mathbf{U}_{[i]}^\top \mathbf{K}_{[i]} \\ &= \mathbf{S}_{[i]} + \left(\mathbf{U}_{[i]} - \mathbf{W}_{[i]}\mathbf{S}_{[i]}^\top\right)^\top \mathbf{K}_{[i]} && \in \mathbb{R}^{d\times d} \end{align*} \end{aligned}$$

在此公式中,$\mathbf{S}_{[i]} := \mathbf{S}_{iC} \in \mathbb{R}^{d \times d}$ 代表第 $i$ 个chunk的初始状态。

序列并行的基本思想与挑战

序列并行的挑战与示例。为了具体说明如何对DeltaNet进行序列并行,我们考虑一个实例:一个长度为8192的序列,使用大小为64的chunk进行分块计算,总共会切分出128个chunk。这些chunk对应的初始状态序列为 $\mathbf{S}_{[0]}, \mathbf{S}_{[1]}, \mathbf{S}_{[2]}, \dots, \mathbf{S}_{[127]}$。在非并行模式下,这128个状态会依据上述公式在同一个SM上循环执行,即通过$\mathbf{S}_{[0]}$ 计算出 $\mathbf{S}_{[1]}$,然后依次向前推进计算。

CP=2序列并行下的依赖问题。现在,假设我们要实现 CP=2 的序列并行,即在2个SM上执行计算。这意味着需要将序列拆分为两个子序列,每个子序列长度为4096。其中,状态 $\mathbf{S}_{[0]}, \mathbf{S}_{[1]}, \dots, \mathbf{S}_{[63]}$ 在第一个SM上计算,而状态 $\mathbf{S}_{[64]}, \mathbf{S}_{[65]}, \dots, \mathbf{S}_{[127]}$ 在第二个SM上计算。这里的主要问题是,第二个SM上所有状态的计算都依赖于真实的初始状态 $\mathbf{S}_{[64]}$。然而,在并行计算开始时,由于第一个SM尚未完成计算,真实的 $\mathbf{S}_{[64]}$ 是不可用的。因此,第二个SM只能使用一个不完整的或临时的初始状态,记为 $\mathbf{S}_{[64]}^{\Delta}$。在实际操作中,这个$\mathbf{S}_{[64]}^{\Delta}$ 通常被初始化为一个全零的状态矩阵。

并行算法的数学推导

真实状态与增量状态的更新关系。我们分别写出真实状态和使用临时初始状态(增量状态)时的更新公式:

$$\begin{aligned} \begin{align*} \mathbf{S}_{[i+1]} &= \mathbf{S}_{[i]} (\mathbf{I}-\mathbf{W}_{[i]}^\top \mathbf{K}_{[i]}) + \mathbf{U}_{[i]}^\top \mathbf{K}_{[i]} \\ \mathbf{S}_{[i+1]}^{\Delta} &= \mathbf{S}_{[i]}^{\Delta} (\mathbf{I}-\mathbf{W}_{[i]}^\top \mathbf{K}_{[i]}) + \mathbf{U}_{[i]}^\top \mathbf{K}_{[i]} \end{align*} \end{aligned}$$


将两式相减,可以轻易得到它们之间的差值关系:$\mathbf{S}_{[i+1]} - \mathbf{S}_{[i+1]}^{\Delta} = (\mathbf{S}_{[i]} - \mathbf{S}_{[i]}^{\Delta}) (\mathbf{I}-\mathbf{W}_{[i]}^\top \mathbf{K}_{[i]})$。通过递推,我们可以得到一个更通用的公式,它描述了从某个起始点 $i$ 开始,经过 $k$ 步后,真实状态与增量状态之间的差值:

$$ \mathbf{S}_{[i+k]} - \mathbf{S}_{[i+k]}^{\Delta} = (\mathbf{S}_{[i]} - \mathbf{S}_{[i]}^{\Delta}) \prod_{j=0}^{k-1}(\mathbf{I}-\mathbf{W}_{[i+j]}^\top \mathbf{K}_{[i+j]}) $$
这个公式是实现序列并行的关键。

CP=2 序列并行算法实现

CP=2并行算法步骤。基于上述推导,可以设计出如下的序列并行算法:
1. 第一个SM的任务:负责计算第一个子序列(chunks 0-63)。其计算逻辑与原始的串行方式完全一致,通过循环计算最终得到真实的 $\mathbf{S}_{[64]}$。
2. 第二个SM的任务:负责计算第二个子序列(chunks 64-127)。
- 首先,初始化其初始状态为零矩阵,即 $\mathbf{S}_{[64]}^{\Delta}=0$。
- 然后,循环地计算增量状态 $\mathbf{S}_{[i+1]}^{\Delta}=\mathbf{S}_{[i]}^{\Delta} + \left(\mathbf{U}_{[i]} - \mathbf{W}_{[i]}\mathbf{S}_{[i]}^\top\right)^\top \mathbf{K}_{[i]}$。
- 同时,初始化一个缩放矩阵 $\mathbf M_{[64]}=\mathbf I$(单位矩阵),并循环地更新它:$\mathbf M_{[i+1]}=\mathbf M_{[i]}(\mathbf{I}-\mathbf{W}_{[i]}^\top \mathbf{K}_{[i]})$。
- 在计算过程中,将每一时间步的增量状态 $\mathbf{S}_{[i+1]}^{\Delta}$ 和缩放矩阵 $\mathbf M_{[i+1]}$ 存储到全局内存(global memory)中。

  1. 状态校正:当第一个SM计算完成后,真实的 $\mathbf{S}_{[64]}$ 就绪。此时,可以利用之前存储的增量信息,通过校正公式 $\mathbf S_{[i]}=\mathbf S_{[i]}^{\Delta} +\mathbf S_{[64]}\mathbf M_{i}$ 来计算出第二个子序列所有真实的最终状态 $\mathbf S_{[i]}$。值得注意的是,这个校正步骤本身不是循环的,可以并行地对所有时间步进行计算,并且能够与输出 $\mathbf O$ 的计算过程融合(fuse),以减少访存开销。

该算法的主要开销在于,会额外增加缩放矩阵 $\mathbf M_{[i]}$ 从全局内存的写出和读入操作。

CP=4 序列并行算法实现

CP=4并行算法步骤。对于 CP=4 的情况,计算被分配到4个SM上,每个SM处理32个chunk,其原理与CP=2类似:
1. 第一个SM的任务:负责计算第一个子序列(chunks 0-31),计算逻辑保持不变,最终得到真实的 $\mathbf{S}_{[32]}$。
2. 其他三个SM的任务:分别负责计算后三个子序列。
- 它们的初始状态均被设置为零矩阵:$\mathbf{S}_{[32]}^{\Delta}=\mathbf{S}_{[64]}^{\Delta}=\mathbf{S}_{[96]}^{\Delta}=0$。
- 在每个SM内部,以循环方式计算各自的增量状态 $\mathbf{S}_{[i+1]}^{\Delta}$ 和缩放矩阵 $\mathbf M_{[i+1]}$,其中缩放矩阵的初始值 $\mathbf M_{[32]}=\mathbf M_{[64]}=\mathbf M_{[96]}=\mathbf I$。
- 计算出的 $\mathbf{S}_{[i+1]}^{\Delta}$ 和 $\mathbf M_{[i+1]}$ 被存储到全局内存中。

  1. 状态校正:此过程存在一个串行的依赖链。
    • 当第一个SM计算完成并得到真实的 $\mathbf{S}_{[32]}$ 后,可以利用它和第二个SM计算的增量信息,通过校正公式计算出第二个子序列的真实状态,并得到真实的 $\mathbf{S}_{[64]}$。
    • 接着,使用刚计算出的真实 $\mathbf{S}_{[64]}$,结合第三个SM的增量信息,计算出第三个子序列的真实状态,并得到真实的 $\mathbf{S}_{[96]}$。
    • 最后,使用真实的 $\mathbf{S}_{[96]}$ 来校正第四个子序列的状态。
    • 具体的校正公式为:对于第二个SM上的状态,使用 $\mathbf S{[i]}=\mathbf S_{[i]}^{\Delta} +\mathbf S_{[32]}\mathbf M_{i}$;对于第三个SM上的状态,使用 $\mathbf S_{[i]}=\mathbf S_{[i]}^{\Delta}+\mathbf S_{[64]}\mathbf M_{i}$;对于第四个SM上的状态,使用 $\mathbf S_{[i]}=\mathbf S_{[i]}^{\Delta} +\mathbf S_{[96]}\mathbf M_{i}$。
    • 同样地,这个校正步骤在每个子序列内部是并行的,并且可以和输出 $\mathbf O$ 的计算融合在一起。

A4 参考文献

[1] DeltaNet Explained (Part II). (2024). Retrieved from https://sustcsonglin.github.io/blog/2024/deltanet-2/
- 引用位置: 方法细节 - 状态更新的循环公式
- 引用内容: 引用了该博客文章中关于DeltaNet分块并行形式(Chunkwise Parallel Form)的状态更新公式。