TokenWeave 策略详解:以大模型推理为例

1. 背景:张量并行中的计算冗余

在大型 Transformer 模型(如 LLaMA-70B)的张量并行推理中,每个 GPU 负责隐藏维度的 1/N(N 为张量并行度)。以 N=2hidden_size=4096seq_len=2048 为例:

  • 前向计算(FFN 或 Attention)后,每个 GPU 持有一个局部张量:

    • GPU0:[2048, 2048](前半部分特征)
    • GPU1:[2048, 2048](后半部分特征)
  • AllReduce(求和) 将两个 GPU 的局部结果合并为完整特征张量,每个 GPU 均获得:

    • [2048, 4096](所有 token 的所有特征维度)
  • 随后,每个 GPU 独立执行残差加法与 RMSNorm

    output = RMSNorm(residual + allreduce_output) * weight
    

问题:RMSNorm 是逐 token 操作,每个 token 的均方根统计量基于其全部 4096 维特征计算。当 N 个 GPU 都持有完全相同的完整张量时,RMSNorm 被重复计算 N 次,造成严重的算力浪费。


2. TokenWeave 核心思想:操作重排序

TokenWeave 提出:将 RMSNorm 提前到 AllReduce 的中间阶段执行,消除冗余计算

传统顺序:

AllReduce → 残差加法 → RMSNorm

TokenWeave 顺序:

ReduceScatter → 残差加法 → RMSNorm → AllGather

关键约束:ReduceScatter 必须沿 token 边界切分,确保每个 GPU 获得若干个完整 token(即特征维度完整),而非特征维度的子集。这样 RMSNorm 可在部分 token 上独立、无冗余地执行。


3. 详细过程举例(N=2,seq_len=2048)

3.1 初始状态(FFN 输出后)
  • GPU0:[2048, 2048] # token0~2047,特征前半
  • GPU1:[2048, 2048] # token0~2047,特征后半
3.2 ReduceScatter(按 token 切分)

将 token 均匀分配给各 GPU,每个 GPU 接收完整特征的对应 token 子集。

  • 通信操作:跨 GPU 规约并分发
    对于 GPU0,它需要从 GPU1 获取 token0~1023 的后半特征,加上自己的前半特征,得到 token0~1023 的完整张量。
    对于 GPU1,类似地得到 token1024~2047 的完整张量。

  • 结果:

    • GPU0:[1024, 4096] # token0~1023 完整特征
    • GPU1:[1024, 4096] # token1024~2047 完整特征

此时,每个 GPU 拥有不同 token 的完整 embedding,且彼此不重叠。

3.3 残差加法与 RMSNorm

每个 GPU 独立处理自己负责的 token 子集:

  • 残差读取:残差张量通常每个 GPU 都有完整副本,但这里只需读取对应 token 部分。
  • RMSNorm:对 [1024, 4096] 逐 token 计算均方根,归一化,乘以权重。

计算量:原来每个 GPU 需处理 2048 个 token,现在只需处理 1024 个 token,计算量减少为 1/N(此处 N=2)。

3.4 AllGather

将各 GPU 的归一化结果拼接回完整序列:

  • 通信操作:GPU0 将自己的 [1024, 4096] 发给 GPU1,GPU1 将自己的 [1024, 4096] 发给 GPU0。
  • 结果:两个 GPU 都获得 [2048, 4096] 的完整归一化张量,与原始输出等价,但 RMSNorm 总计算量减半

4. 为什么简单的重排序反而可能变慢?

尽管计算量降低,但通信模式的变化可能导致整体性能倒退:

方面 传统 AllReduce 简单分解为 ReduceScatter + AllGather
通信原语 单次 AllReduce(通常硬件优化,如 NVLink 直接规约广播) 两次独立通信:ReduceScatter + AllGather
内核启动 1 个通信内核 2 个通信内核(额外调度开销)
同步开销 流水化良好,通信可隐藏 两次全局同步,GPU 空闲等待
内存访问 规约结果直接写入 HBM,RMSNorm 独立内核读取 中间结果需写 HBM,再读入 RMSNorm 内核,再写 HBM 供 AllGather
通信量 约 2×(N-1)/N × 数据量 相同(但分两次,延迟累积)

关键瓶颈:分解后的 RMSNorm 内核与通信内核串行执行,且中间数据在 HBM 往返,通信延迟无法被计算掩盖,甚至可能因额外内核启动而劣化。


5. 融合内核:TokenWeave 的真正加速器

为解决上述问题,TokenWeave 提出了融合的 ReduceScatter + RMSNorm + AllGather 内核(即代码清单 1)。其核心设计:

  • 一次内核启动,完成所有操作。
  • 网络内规约:通过 multimem_ld_reduce_add 直接在加载时完成跨 GPU 求和,无需中间 HBM 存储
  • 就地计算:在寄存器中融合残差、累加方差,避免多余读写
  • 远程直接存储:归一化结果通过 multimem_st 直接写入 AllGather 目标地址,省去显式 AllGather 内核

收益

  • 通信量不变,但通信与计算完全流水,延迟被隐藏。
  • HBM 访问次数从 5 次降为 2 次。
  • RMSNorm 计算量减少为 1/N,且无额外通信惩罚。

6. 总结:大模型分布式推理的优化方向

TokenWeave 示例揭示了一个通用原则:在分布式系统中,通过重新排序计算与通信,并借助融合内核消除冗余,可实现“计算换通信”的反直觉优化。对于百亿/千亿参数模型,此类融合可将单层延迟降低 20% 以上,是继张量并行、序列并行后的又一重要优化维度。

为什么 AllReduce 后每个 GPU 都得到完整张量?

在张量并行(Tensor Parallelism)的列切分(Column-wise)模式下,每个 GPU 只负责计算输出特征维度的 1/N。但是,后续的残差连接、层归一化等操作需要完整的特征向量。因此,必须通过通信将分散在各 GPU 的部分结果合并成完整的张量

关键问题

  • AllReduce 的语义是“规约”(如求和)并将结果广播给所有参与者。
  • 列切分下,每个 GPU 持有的是不同部分的特征,不是相同特征的不同副本。如何用 AllReduce 把“不同部分”合并成“完整”

实现方法:零填充 + AllReduce 求和

这种模式的实现很简单:

  1. 每个 GPU 预先分配一个完整大小的输出缓冲区(形状 [seq_len, hidden_size]),初始化为零。
  2. 在计算 FFN 或 Attention 时,每个 GPU 只填充自己负责的那 1/N 列,其余列保持为 0。
    • 例如:GPU 0 填充第 0~2047 列,GPU 1 填充第 2048~4095 列(hidden_size=4096,N=2)。
  3. 执行 AllReduce(求和)
    • GPU 0 将自己的部分(第 0~2047 列有值,第 2048~4095 列为 0)广播并求和;
    • GPU 1 将自己的部分(第 2048~4095 列有值,第 0~2047 列为 0)广播并求和。
    • 求和后,GPU 0 得到:第 0~2047 列 = 自己的值 + GPU 1 的 0;第 2048~4095 列 = 自己的 0 + GPU 1 的值 → 完整张量
    • GPU 1 同理。

结果:AllReduce 结束后,每个 GPU 都拥有完整的输出张量


数值示例(N=2,hidden=4,简化)

假设:

  • hidden_size = 4,切分给 GPU0: 列0~1,GPU1: 列2~3。
  • 单个 token,实际计算结果:
    • GPU0 局部输出:[a, b, 0, 0]
    • GPU1 局部输出:[0, 0, c, d]

AllReduce(求和)

  • GPU0 收到 GPU1 的 [0,0,c,d],加到本地 → [a, b, c, d]
  • GPU1 收到 GPU0 的 [a,b,0,0],加到本地 → [a, b, c, d]

结果:两个 GPU 都持有 [a,b,c,d]


为什么选择 AllReduce 而不是 AllGather?

通信原语 操作 通信量(每个 GPU) 优点 缺点
AllGather 收集所有 GPU 的部分并拼接 发送 size/N,接收 size×(N-1)/N 无冗余计算,内存不重复 需要两次通信(若后续要规约)
AllReduce (求和) 广播完整张量,每个 GPU 贡献非零部分 发送 size,接收 size 一次通信完成聚合+广播 内存浪费(完整缓冲区),额外加法

在 Transformer 层中,残差连接和 RMSNorm 需要完整张量,且之后通常还有 AllGather 用于序列并行。AllReduce 方案可以一次通信既完成聚合又完成广播,代码简洁且与后续操作容易融合(如本文的融合内核)。虽然通信量略大(发送完整 size),但在 NVLink 高速互联下,额外开销可被计算掩盖。


融合内核中的体现

在代码清单 1 中,multimem_ld_reduce_add<16> 正是这种思想的极致体现:

  • 每个线程加载自己负责的向量位置,但地址是远程 GPU 对应位置的地址
  • 硬件在 NVLink 路径上自动完成远程加载值的求和,并直接返回给线程。
  • 线程不需要预先零填充,因为远程地址指向对端 GPU 已经计算好的局部值。
  • 结果:线程拿到的是全局规约后的完整向量,AllReduce 被内嵌在加载指令中,无显式通信内核。

总结

AllReduce 后每个 GPU 得到完整张量,是因为:

  1. 每个 GPU 只计算完整张量的一个切片,其余部分补零。
  2. AllReduce 的求和语义将这些零和有效值相加,使得每个 GPU 获得所有切片的拼接结果。
  3. 这是一种实现简单、易于融合的通信模式,尤其适合 Transformer 中需要完整张量的后续计算。

这种模式与“先 AllGather 再本地处理”在数学上等价,但通信模式不同,为深度融合提供了机会(如 TokenWeave 的 ReduceScatter+RMSNorm+AllGather 融合)。

Logo

脑启社区是一个专注类脑智能领域的开发者社区。欢迎加入社区,共建类脑智能生态。社区为开发者提供了丰富的开源类脑工具软件、类脑算法模型及数据集、类脑知识库、类脑技术培训课程以及类脑应用案例等资源。

更多推荐