Back to blog

TokenWeave: Efficient Compute-Communication Overlap for Distributed LLM Inference

首个在低延迟推理场景下为张量并行LLM推理实现高效计算通信重叠的系统,通过融合AllReduce-RMSNorm核和波感知智能分割达成最高1.28×加速

TokenWeave: Efficient Compute-Communication Overlap for Distributed LLM Inference

一、论文概述

项目内容
标题TokenWeave: Efficient Compute-Communication Overlap for Distributed LLM Inference
作者Raja Gond, Nipun Kwatra, Ramachandran Ramjee
机构Microsoft Research India
论文arXiv:2505.11329
代码github.com/microsoft/tokenweave
发布2025-05-16 (v1), 2026-05-01 (v5, 最终版)
发表MLSys 2026
许可CC BY 4.0
领域Distributed, Parallel, and Cluster Computing (cs.DC); Machine Learning (cs.LG)

二、核心思想

问题定义

在分布式LLM推理中,即使通过NVLink连接的高性能GPU,张量并行(Tensor Parallelism)带来的通信开销仍高达端到端延迟的9-23%(Figure 1)。更令人惊讶的是,RMSNorm操作的开销也达到4-9%,此前一直被忽视。

现有的计算-通信重叠技术(如TileLink、NanoFlow、Flux)都无法在低延迟推理场景中有效工作,原因有二:

  1. 分解开销:将计算分解为更小的子任务以实现重叠,会导致GPU wave quantization效应,降低计算效率
  2. SM资源竞争:通信操作占用大量流式多处理器(SM),挤占了本可用于计算的资源

因此,尽管已有大量研究,vLLM、SGLang、TensorRT-LLM等主流推理系统默认并未开启任何通信重叠优化,张量并行推理仍承受高达20%的通信成本。

解决方案概述

TokenWeave是首个在低至1024 token的小批次下实现高效计算-通信重叠的系统,核心创新包括:

  1. 发现并优化RMSNorm:识别RMSNorm为关键瓶颈,设计融合的AllReduce-RMSNorm核
  2. 波感知智能分割(Wave-Aware Smart-Splitting):两路token分割,确保总wave数不增加
  3. NVSHARP/Multimem优化:仅用2-8个SM即可完成通信和归一化,释放大部分SM用于计算

三、技术架构

整体框架

TokenWeave架构

TokenWeave的三个核心技术:

┌─────────────────────────────────────────────────────────────┐
│                    TokenWeave Pipeline                       │
│                                                              │
│  输入批次 → Smart-Split (两路分割)                            │
│    ├─ Split 0 → Compute 0 (并行)                            │
│    │            + Fused-AR-RMSNorm 0 (与Split 1重叠)         │
│    └─ Split 1 → Compute 1 (并行)                            │
│              + Fused-AR-RMSNorm 1 (与Split 0重叠)            │
│                                                              │
│  关键: 通信(Split 0) 与 计算(Split 1) 交织执行               │
└─────────────────────────────────────────────────────────────┘

核心公式与设计要点

融合核优化原理:

标准做法中,vLLM在AllReduce之后执行RMSNorm,或者将AllReduce拆分为ReduceScatter+AllGather并在ReduceScatter后执行RMSNorm。后者虽然减少了每GPU的RMSNorm计算量(除以TP度),但引入了显著的额外开销(Figure 4):

序列长度RS+AG / AR 比率
642.03×
1K1.06×
8K1.01×
32K1.00×

对于小批次(≤8K),拆分开销达1-6%,抵消了RMSNorm节省。

融合核实现(CUDA Kernel伪代码,见附录A.1):

__global__ multimem_fused_allreduce_rmsnorm_kernel(...) {
    // 每个CTA处理一批token
    // 1. 通过multimem ld reduce add执行GPU间约简
    auto multimem_temp = multimem_ld_reduce_add<multimem_address_ptr>(offset);
    
    // 2. 累加残差并计算方差
    temp += residual_o[idx];
    variance[0] += temp.sum_squares();
    
    // 3. 归一化后直接写入输出
    vec_t temp = residual_o[idx] * s_variance * weight_v[idx];
    multimem_st<multimem_output_ptr>(offset, temp);
}

关键优化:消除中间HBM访问——约简、残差累加、方差计算、归一化均在共享内存中完成,只通过multimem指令进行最终的远程存储。

智能分割(Smart-Splitting):

传统等量分割:  Split 0 = T/2 tokens,  Split 1 = T/2 tokens
   
TokenWeave:     Split 0 = T/2 + offset tokens,  
                Split 1 = T/2 - offset tokens
                
其中 offset 通过离线分析使Split 0的最后wave满occupancy

wave量化效应示意(Figure 8):

假设kernel有10个CTA,GPU有4个SM:
Wave 1: CTA 0,1,2,3 → 4个SM满载
Wave 2: CTA 4,5,6,7 → 4个SM满载  
Wave 3: CTA 8,9     → 仅2个SM工作(浪费6个SM)

Smart-split确保Split 0的last wave满occupancy,
使两个split的总wave数 ≤ 原始kernel的wave数

Multimem SM效率(Figure 5, 10):

序列长度最优SM数说明
64-2562-4 SM极小规模
512-8K4-8 SM典型范围
16K-64K8-16 SM大规模

仅2-8个SM即可饱和通信带宽,释放绝大部分SM用于计算。

选择性启用策略(Figure 3)

if (num_tokens >= threshold) {
    // 启用完整TokenWeave: 智能分割 + 重叠 + 融合核
    TokenWeave(smart_split=true, overlap=true, fuse=true);
} else {
    // 仅启用融合核,避免分割开销
    TokenWeave(smart_split=false, overlap=false, fuse=true);
}

阈值设定(基于离线分析):

  • Dense模型(Llama, Qwen):1K tokens
  • MoE模型(Mixtral):4K tokens

模型组件

组件说明关键参数
Fused AllReduce-RMSNorm融合NVSHARP约简与RMSNorm,消除中间HBM访问4-8 SM, bf16
Smart-Splitting波感知的两路token分割,确保总wave数不增加offset通过离线profile确定
Selective Enable根据token数量动态决定是否启用分割/重叠Dense: 1K, MoE: 4K
NVSHARP/Multimem利用NVSwitch的in-network reductionPyTorch SymmetricMemory API

四、核心创新

创新点说明效果
RMSNorm识别与融合首次发现RMSNorm是TP通信优化的关键,设计融合AllReduce-RMSNorm核1.34-1.39×加速,消除中间HBM访问
波感知智能分割通过调整split偏移量使last wave满occupancy,消除wave量化开销在小序列长度下分割开销降至最低
NVSHARP/Multimem高效利用仅用2-8个SM完成通信+归一化,而非传统的16-20+ SM释放SM资源用于计算重叠
首个低延迟TP重叠系统在低至1024 token下实现有效重叠(之前方案需8K+)覆盖低延迟推理的实际工作负载
选择性启用机制根据批次大小动态选择最优策略避免小批次的分割惩罚

五、实验结果

评估设置

  • 硬件:8×H100 DGX(NVLink4/NVSHARP),以及8×B200 DGX验证
  • 模型:Llama-3.3-70B, Qwen2.5-72B, Mixtral-8x22B, Qwen3-235B-A22B
  • 基线:vLLM-Multimem(已优化AllReduce),vLLM-nocomm(理论下限)
  • 对比:TileLink, NanoFlow

端到端延迟性能

延迟增益

序列长度TokenWeave vs vLLM-Multimemvs vLLM-nocomm
64 tokens~1.04×接近
1K tokens1.20×接近/超越
4K tokens1.25×超越
8K tokens1.28×超越
32K tokens1.28×超越

关键发现:在≥4K序列长度时,TokenWeave不仅恢复所有通信开销,还因RMSNorm融合优化超越无通信的理论上限。

吞吐量增益

吞吐量增益

  • ShareGPT trace:1.19× 吞吐量提升
  • arXiv trace:1.15× 吞吐量提升
  • 在不同chunk size(1024-8192)下均保持稳定增益

与现有方案对比

vs TileLink(Figure 14):

序列长度TileLinkTokenWeave
256降速(overhead)1.16×
512降速1.22×
1K降速1.25×
2K1.16×1.27×
8K1.20×1.28×

TileLink在≤2K时反而降速,TokenWeave在整个序列长度范围内稳定加速。

vs NanoFlow(Figure 15):

NanoFlow在高吞吐设置下有优势,但在低延迟场景(小chunk size)下TokenWeave表现更好。

消融实验

消融实验

变体64 tokens1K4K32K
vLLM-Multimem1.00×1.00×1.00×1.00×
TokenWeave-fuseonly1.04×1.04×1.04×1.04×
TokenWeave (full)1.04×1.20×1.27×1.28×
  • fuseonly:仅融合核,在所有序列长度下提供~4%的稳定增益
  • full:融合核+智能分割+重叠,在小序列长度下增益主要来自融合核,大序列长度下重叠贡献显著

B200验证

在8×B200 DGX系统上的验证:

  • Decode延迟:Llama-3.3-70B获得1.01-1.05×提升
  • Prefill延迟:全TokenWeave实现最高1.22×加速
  • 融合核:在B200上达到1.24-1.38×加速(Table 2)

融合核性能详情(B200,Table 2)

序列长度AR(us)RMSNorm(us)AR+RMSNorm(us)Fused(us)加速比
3226.0814.4640.5430.461.33×
12832.2914.6746.9634.141.38×
1K60.2621.1281.3863.621.28×
4K166.6153.66220.27170.141.29×
64K2240.93654.512895.442236.021.29×

六、相关工作

工作方法局限TokenWeave优势
Flux/TileLink细粒度tile-level融合,在GEMM内核中嵌入通信仅适用于GEMM,小批次需8K+ tokens粗粒度token-level,低至1K tokens有效
NanoFlownanobatch调度,按kernel粒度分割依赖大批次,小批次开销大智能分割消除wave量化开销
DeepSeekexpert-parallel all-to-all重叠all-to-all开销50%+,易隐藏针对AllReduce(~20%开销)优化
GPipe训练中的pipeline parallelism不适用于推理的低延迟场景专为低延迟推理设计

七、总结

核心贡献

  1. 发现RMSNorm的关键作用:识别RMSNorm为TP通信优化的关键瓶颈(4-9%开销),设计融合AllReduce-RMSNorm核
  2. 波感知智能分割:两路token分割确保总wave数不增加,消除wave量化开销
  3. NVSHARP/Multimem高效利用:仅2-8个SM完成通信+归一化,释放SM用于重叠
  4. 首个低延迟TP重叠系统:在低至1024 token下实现有效重叠,覆盖实际推理场景
  5. 广泛验证:在H100和B200平台上,多种模型和trace下验证了1.28×延迟改善和1.19×吞吐提升

技术影响

  • TokenWeave已集成到vLLM-V1框架中,可直接惠及所有使用张量并行的LLM推理服务
  • 融合核优化在不分隔预填充和解码的场景中直接受益小decode批次,在分隔场景中大prefill批次可获得更大增益
  • 为NVSHARP/Multimem等新型硬件特性在ML推理中的高效利用提供了范例

局限性

  • 依赖NVSHARP/Multimem:当前实现依赖Hopper/Blackwell架构的NVSHARP支持,AMD GPU或其他架构需要适配
  • torch.compile限制:SymmetricMemory目前不被torch.compile支持,decode-only场景下需禁用torch.compile
  • CUDA Graph不支持:完整TokenWeave方案(智能分割+重叠)不支持CUDA Graph
  • B200上decode增益有限:由于上下文长度增加使RMSNorm占比下降,decode场景增益较小(1.01-1.05×)

八、参考资源