[cuda]TokenWeave
TokenWeave 策略详解:以大模型推理为例
1. 背景:张量并行中的计算冗余
在大型 Transformer 模型(如 LLaMA-70B)的张量并行推理中,每个 GPU 负责隐藏维度的 1/N(N 为张量并行度)。以 N=2、hidden_size=4096、seq_len=2048 为例:
-
前向计算(FFN 或 Attention)后,每个 GPU 持有一个局部张量:
- GPU0:
[2048, 2048](前半部分特征) - GPU1:
[2048, 2048](后半部分特征)
- GPU0:
-
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 完整特征
- GPU0:
此时,每个 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 求和
这种模式的实现很简单:
- 每个 GPU 预先分配一个完整大小的输出缓冲区(形状
[seq_len, hidden_size]),初始化为零。 - 在计算 FFN 或 Attention 时,每个 GPU 只填充自己负责的那 1/N 列,其余列保持为 0。
- 例如:GPU 0 填充第 0~2047 列,GPU 1 填充第 2048~4095 列(hidden_size=4096,N=2)。
- 执行 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]
- GPU0 局部输出:
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 得到完整张量,是因为:
- 每个 GPU 只计算完整张量的一个切片,其余部分补零。
- AllReduce 的求和语义将这些零和有效值相加,使得每个 GPU 获得所有切片的拼接结果。
- 这是一种实现简单、易于融合的通信模式,尤其适合 Transformer 中需要完整张量的后续计算。
这种模式与“先 AllGather 再本地处理”在数学上等价,但通信模式不同,为深度融合提供了机会(如 TokenWeave 的 ReduceScatter+RMSNorm+AllGather 融合)。
更多推荐
所有评论(0)