Back to blog

FastTree: Optimizing Attention Kernel and Runtime for Tree-Structured LLM Inference

优化树结构LLM推理的注意力内核和运行时

FastTree: Optimizing Attention Kernel and Runtime for Tree-Structured LLM Inference

一、论文概述

项目内容
标题FastTree: Optimizing Attention Kernel and Runtime for Tree-Structured LLM Inference
作者Zaifeng Pan, Yitong Ding, Yue Guan, Zheng Wang, Zhongkai Yu, Xulong Tang, Yida Wang, Yufei Ding
机构University of California San Diego, University of Pittsburgh, AWS
会议MLSys 2025
论文OpenReview
代码GitHub
许可Apache-2.0
领域LLM推理优化、注意力计算、GPU内核、树结构推理

二、核心思想

问题定义

在现代LLM应用中,**树结构前缀共享(tree-structured prefix sharing)**非常普遍,包括few-shot learning、文档QA、tree of thoughts等场景。现有的LLM服务系统使用基数树(radix tree)来组织全局KV缓存,促进不同查询之间的缓存重用,从而减少不必要的内存使用。

然而,这些系统仍然依赖传统的计算模式进行注意力操作,导致严重的性能问题:

  1. 冗余内存加载:共享的KV缓存需要从慢速全局内存为每个查询重复加载
  2. GPU张量核心利用率不足:解码阶段的GEMV操作难以有效利用高性能张量核心
  3. 共享内存使用不充分:不同查询的线程块无法直接访问彼此的共享内存进行数据重用

现有系统的计算模式

Llama-2-7B执行时间分解

关键观察:

  • 当序列长度为2048时,注意力计算可占模型前向时间的65%
  • FlashAttention在填充Q矩阵后,有效计算不到1%
  • 批处理查询只能帮助增加projection和MLP层的计算强度,无法直接 benefit 注意力

解决方案概述

FastTree提出为高效处理通过基数树共享上下文的查询定制GPU内核:

  1. 树结构自适应运行时优化:使用贪心启发式算法划分树以最小化开销
  2. 高效树结构注意力内核:针对树结构查询优化的GPU内核
  3. 长上下文分割:分割冗长的上下文以缓解尾部效应

三、技术架构

整体框架图

FastTree系统概述

FastTree作为SGLang的插件实现,加速给定基数树的KV缓存注意力计算。系统工作流程:

  1. 客户端发送查询请求到LLM服务系统
  2. 系统组织全局KV缓存为基数树
  3. FastTree运行时生成上下文-查询分组计划
  4. 高效注意力内核执行优化的注意力计算

核心公式

自注意力计算:

Attention⁡(Q,K,V)=softmax⁡(QKTd)V\operatorname{Attention} (Q, K, V) = \operatorname{softmax} (\frac {Q K ^ {T}}{\sqrt {d}}) V

其中d是头维度。

问题形式化:

给定运行时的KV缓存树 T=(V,E)T = (V, E),目标是确定赋值函数 f:E→{0,1}f : E \to \{ 0, 1 \} 以最小化GPU内核执行延迟:

min⁡f Latency ( Groups (VT(T,f(E))))\min _{f} \quad \text { Latency } (\text { Groups } (\mathrm{VT} (T, f (E))))

 s.t. f(e)∈{0,1},∀e∈E\text { s.t. } \quad f (e) \in \{0, 1 \}, \quad \forall e \in E

其中:

  • VT⁡(T,f(E))\operatorname { V T } ( T , f ( E ) ):基于二元边赋值将原始树转换为虚拟树
  • Groups(⋅)\mathrm { G r o u p s } ( \cdot ):将虚拟树映射为分组计划
  •  Latency (⋅)\text { Latency } ( \cdot ):评估GPU内核延迟

核心组件

1. 树结构自适应运行时优化

上下文-查询分组:

给定KV缓存树,存在多种可能的分组计划。例如,查询 Q1Q_1 到 QMQ_M 共享前缀A1和B1,而查询 QM+1Q_{M+1} 到 QM+NQ_{M+N} 共享A1和B2。

分组计划示例

不同分组计划的性能差异:

  • 案例1(M=N=8,tile大小=16):计划2需要填充,导致性能下降
  • 案例2(M=N=128):计划1引入更多中间结果和归约步骤

二元边赋值:

边赋值与虚拟树

将问题形式化为二元边赋值任务:

  • 赋值1:连接两个节点
  • 赋值0:保持分离
  • 生成虚拟树指导分组策略

2. 贪心启发式算法

算法1:贪心二元边赋值

输入:基数树 T = (V, E)
1: 初始化 A 为每个节点的关联查询数
2: 初始化 L 为每个节点的上下文长度
3: for 节点 v in BFS(V) do
4:   nQ_curr = A[v]  ▷ 当前聚合查询数
5:   len_v = L[v]    ▷ 累积上下文长度
6:   for l in leaves(v) do
7:     nQ_l = A[l], len_l = L[l]
8:     C_0 = SplitKVCost(nQ_curr, nQ_l, len_l, len_v)
9:     C_1 = SplitQCost(nQ_curr, nQ_l, len_l, len_v)
10:    if C_0 >= C_1 then
11:      赋值1给边 v → l
12:      nQ_curr = nQ_curr - nQ_l
13:      L[l] = len_l + len_v
14:    else
15:      赋值0给边 v → l
16:    end if
17:  end for
18: end for

开销模型:

填充开销函数:

CP,q(nQ,len)=Pad⁡(TSq,nQ)⋅len⋅dC _ {P, q} (n Q, l e n) = \operatorname{Pad} (T S _ {q}, n Q) \cdot l e n \cdot d

CP,c(nQ, len )=nQ⋅Pad⁡(TSc,min⁡(len,TSc))⋅dC _ {P, c} (n Q, \text { len }) = n Q \cdot \operatorname{Pad} (T S _ {c}, \min (l e n, T S _ {c})) \cdot d

其中填充函数:

Pad⁡(TS,N)=TS−((N−1)%TS+1)\operatorname{Pad} (T S, N) = T S - ((N - 1) \% T S + 1)

中间结果开销:

SplitKVCost⁡R=γ⋅nQl⋅d\operatorname{SplitKVCost} _ {R} = \gamma \cdot n Q _ {l} \cdot d

3. GPU高效的长上下文分割

问题:GPU利用率不足的两种情况:

  1. 块级并行不足:启动的线程块数小于GPU容量
  2. 长尾效应:部分节点具有极长上下文,导致最后几波活跃块数很少

GPU利用率不足问题

解决方案:在运行时分割超过阈值的长上下文节点,虽然引入中间结果归约开销,但性能提升完全掩盖开销。

高效树结构注意力内核设计

内核架构:

注意力内核设计

关键设计:

  1. 单内核处理所有组:不同上下文-查询组分派到不同的线程块集
  2. FlashAttention风格处理:逐tile处理,利用在线softmax技术
  3. 多阶段tiling:自适应选择查询维度的tile大小
    • 树根附近:更大tile大小,最大化KV重用
    • 叶子附近:更小tile大小,避免共享内存浪费

优化效果:

  • 减少慢速全局内存事务
  • 将GEMV操作转换为GEMM操作
  • 有效利用张量核心

四、核心创新

创新点说明理论/实验依据
树结构注意力内核针对树结构查询定制的GPU内核减少冗余内存加载,提升张量核心利用率
贪心启发式分组线性时间复杂度的最优分组搜索平衡填充开销和中间结果开销
长上下文分割解决GPU利用率不足问题最高1.9×加速
多阶段tiling自适应选择tile大小平衡并行性和数据重用
SGLang集成作为SGLang插件实现无缝集成现有系统

五、实验结果

实验设置

  • 硬件:NVIDIA H100 GPU (80GB),CUDA 12.2
  • 实现:基于Triton实现内核,作为SGLang v0.2.13插件
  • 基线:
    • FlashAttention v2.6.3
    • FlashInfer v0.1.6
    • SGLang Triton内核 v0.2.13
    • DeFT
    • Multi-Level Cascade Attention (CascadeAttn)
  • 模型:Llama-2-7B (GQA=1),Mistral-7B (GQA=4)
  • 基准测试:
    • (A) 多级系统提示
    • (B) 多few-shot学习
    • (C) 多链推理
    • (D) 多文档QA

内核基准测试

内核基准测试结果

关键结果:

基线平均加速比
FlashAttention5.1×
SGLang Triton9.2×
FlashInfer4.2×
DeFT10.6×
CascadeAttn2.1×

树配置说明:

  • N:每层的节点数
  • C:每层的每节点上下文长度
  • GQA比率:1, 4, 16

关键发现:

  • GQA比率低时,FastTree相比FlashAttention和FlashInfer加速更显著
  • 即使GQA=16,FastTree仍优于查询分离的内核
  • 贪心启发式在复杂树结构中可带来最高2.2×加速

端到端性能

端到端性能对比

关键结果:

模型基线平均加速比
Llama-2-7BSGLang Triton2.4×
Llama-2-7BSGLang FlashInfer1.6×
Mistral-7BSGLang Triton3.1×
Mistral-7BSGLang FlashInfer1.9×

整体提升:FastTree将SGLang的吞吐量提升高达 2.2×(相比FlashInfer后端)

执行时间分解

执行时间分解对比

关键发现:

  • FastTree显著减少解码延迟
  • CPU预处理开销可忽略不计
  • 解码加速:相比SGLang-Triton 2.3×,相比SGLang-FlashInfer 1.9×

GPU内核分解

GPU内核分解

关键发现:

  • 未优化的注意力操作占总内核执行时间的很大比例
  • FastTree显著减少注意力计算时间
  • 其他内核(element-wise和reduction)只占很小比例

消融实验

贪心优化效果:

  • 简单树结构:贪心启发式倾向于直接聚合
  • 复杂树结构:贪心启发式可带来最高2.2×加速

长上下文分割效果:

  • 配置N=1,10,C=4000,400:最高1.9×加速
  • 解决GPU SM利用率不足问题

六、相关工作

方法特点与FastTree的区别
SGLang基数树KV缓存管理FastTree优化注意力内核
vLLM连续批处理、分页KV缓存未针对树结构优化
FlashAttentionIO感知的注意力内核不支持树结构共享
FlashInfer优化的CUDA内核库查询分离计算
DeFT树结构优化固定tile大小,额外masking
CascadeAttn多级级联注意力未解决分组挑战

七、总结

核心贡献

  1. 树结构注意力内核:首次为树结构LLM推理定制GPU内核
  2. 运行时优化:提出树结构自适应的贪心启发式运行时优化策略
  3. 长上下文分割:解决GPU利用率不足问题
  4. 系统集成:作为SGLang插件实现,易于部署
  5. 显著性能提升:内核加速最高10.6×,端到端吞吐量提升最高2.2×

技术影响

  • 性能提升:显著提升树结构LLM推理的吞吐量
  • 内存效率:减少冗余的KV缓存加载
  • GPU利用率:通过查询聚合提升张量核心利用率
  • 通用性:适用于各种树结构配置和GQA比率
  • 实用性:作为SGLang插件,易于集成到现有系统

局限性

  1. 硬件依赖:超参数仅针对H100调优,未利用Hopper特定特性(如TMA)
  2. 树结构限制:主要针对基数树结构优化
  3. 实现复杂度:需要定制化GPU内核
  4. 预填充优化:当前仅优化解码阶段

八、参考资源

  • 论文:OpenReview
  • 代码:GitHub
  • 相关项目:SGLang, FlashAttention, FlashInfer, DeFT