- 大模型
- 深度学习
- 算子库
- 后端
- 高性能计算
【免费下载链接】flashinfer
FlashInfer: Kernel Library for LLM Serving
FlashInfer 的sm110_gqa_decode是一套显式 opt-in 的实验性 FP16 GQA(Grouped Query Attention)解码路径,专门面向 NVIDIA SM110(Jetson AGX Thor)的 compute capability 11.0 硬件。它以固定的 32 查询头 / 8 KV 头 / 128 头维度几何,复用 SM110 的tcgen05张量核、TMA 与张量内存,提供便捷 API 与 Prepared 持续解码 API 两套入口。读完本文,你将掌握该内核的输入契约、双内核(short/long)路由机制、JIT 注册表与启动绑定原理,并能用prepare_sm110_gqa_decode+launch_sm110_gqa_decode_prepared在 CUDA Graph 中实现零分配、无 host 同步的持续解码。
本文以 实验性包 README 为主体骨架,并结合 jit.py、backend.py、prepared.py 以及 csrc 启动头 等源码逐层展开。
为什么需要一条独立的 SM110 GQA Decode 路径
Jetson AGX Thor 搭载的 SM110 芯片属于与主流数据中心 GPU 不同的架构代际:它基于 Blackwell 架构的 compute capability 11.0,具备tcgen05新一代张量核指令、TMA(Tensor Memory Accelerator)异步批量拷贝和片上张量内存(tensor memory)。FlashInfer 现有的 decode 内核与 JIT 模板大多针对 SM80/SM90/SM100 系列,无法直接覆盖 SM110 的这一套执行模型,因此该包在flashinfer/experimental/下提供了精确 SM110 专用的 FP16 decode 实现:
- 固定几何:32 个 query 头、8 个 KV 头、头维度 128、每个请求仅 1 个 query token,即固定的 4:1 query-to-KV 头比例;
- 显式 opt-in:该后端只通过实验性 API 触达,不注册到 AOT 打包流程,也不参与自动 decode 路由(即常规
flashinfer.decode.*的按 shape 分发不会落到它头上); - 实验性生命周期:由维护者
@yyihuang负责,进度跟踪于 issue #5051,毕业(graduate)之前接口与内部布局都可能变化。
从源码看,顶层入口被@flashinfer_experimental_api(feature="SM110 GQA decode")装饰器标记,位于 flashinfer/decode.py,并由 flashinfer/init.py 导出sm110_gqa_decode、prepare_sm110_gqa_decode、launch_sm110_gqa_decode_prepared三个顶层名字。也就是说,你可以直接从from flashinfer import sm110_gqa_decode开始使用,但请把它当作有明确适用边界的实验功能,而非通用 decode 替代品。
硬件与编译前置条件
两条内核路径都使用 SM110 的tcgen05张量核指令、TMA 与张量内存,因此要求:
| 前置条件 | 说明 |
|---|---|
| 计算能力 | compute capability 11.0(SM110,即 Jetson AGX Thor) |
| CUDA 版本 | CUDA 13.0 或更新 |
| 运行设备 | 所有张量必须位于同一个 CUDA 设备上 |
这些条件在运行期会被强制校验。jit.py中的_check_exact_sm110a会读取设备计算能力并与is_sm110a_supported联合判断,不满足时抛出RuntimeError,提示 "SM110 GQA decode requires compute capability 11.0 and CUDA 13.0 or newer",并打印实际的计算能力与torch.version.cuda(见 jit.py)。
JIT 编译侧,所有模块都通过gen_jit_spec生成规范,公共部分追加sm110a_nvcc_flags与链接标志-lcuda(TMA 描述符编码需要 CUDA driver API);三个 Prepared 专用模块额外使用--use_fast_math编译标志(见 jit.py 与gen_sm110_gqa_decode_module)。
输入契约:固定的张量布局与几何
该内核不接受任意形状,输入布局是硬编码的:
| 参数 | 形状 | dtype | 要求 |
|---|---|---|---|
q | [batch, 32, 128] | FP16 | contiguous,CUDA 设备 |
kv | [batch, 2, 8, capacity, 128] | FP16 | contiguous;索引 0 为 K、索引 1 为 V |
sequence_lengths | [batch] | CUDA int32 | contiguous;每个值必须落在闭区间[1, capacity] |
out(可选) | [batch, 32, 128] | FP16 | 调用者自持,contiguous,不能与q别名 |
其中kv的第 2 维是两个平面:K 平面位于 index 0,V 平面位于 index 1,每平面形状为[batch, 8, capacity, 128]。capacity是 KV 缓存槽位总数(即kv.shape[-2])。
sequence_lengths的语义值得特别注意:长度只在设备端被读取,没有任何 API 调用会把它拷回 host。因此调用本身不会引入设备同步;但作为代价,[1, capacity]的区间约束是调用者的责任——如果你传入越界长度,内核行为未定义,且 host 无法在启动前发现。这一点在便捷 API 与 Prepared API 中是一致的(见 backend.py 与 prepared.py)。
Python 侧的_require_tensor会对形状、dtype、is_cuda、is_contiguous、设备一致性逐项校验(backend.py);CUDA 侧 binding 还会通过tvm_ffi_utils.h的check_dtype/check_contiguous/check_same_device/CheckGrid再做一遍防御(见 short binding)。
便捷 API 快速上手
README 给出的最小示例可以直接运行:
import torch from flashinfer import sm110_gqa_decode q = torch.randn(1, 32, 128, dtype=torch.float16, device="cuda") kv = torch.randn(1, 2, 8, 1024, 128, dtype=torch.float16, device="cuda") sequence_lengths = torch.tensor([1024], dtype=torch.int32, device="cuda") out = sm110_gqa_decode(q, kv, sequence_lengths)签名细节(与 flashinfer/decode.py 一致):
| 参数 | 默认值 | 说明 |
|---|---|---|
q/kv/sequence_lengths | — | 见上文输入契约 |
out | None | 传入时使用调用者自持输出;缺省时内部torch.empty_like(q)分配。若与q共享存储会抛ValueError |
q_scale | 1.0 | 在标准注意力缩放1/sqrt(128)之前额外施加的 query 缩放 |
关于q_scale的底层处理:Python 侧将q_scale * (1.0 / sqrt(128) / log(2))预计算为softmax_scale_log2传给内核(见 backend.py 与 launch 参数构造),因为 SM110 内核的 softmax 在log2 域用ex2.approx.ftz.f32指令计算指数,避免额外的乘法。如果你需要自定义注意力缩放,直接传q_scale即可,无需自己换算。
路由选择完全由capacity决定(backend.py):
capacity <= 64:走short 内核,256 线程;capacity > 64:走long 内核,384 线程、流水线化。
调用通过q.view(batch, 8, 4, 128).transpose(1, 2)将 query 重排为分组视图(每组 4 个 query 头对应 1 个 KV 头),然后一次性把 Q、K、V、O、lengths、log2 缩放、batch * 8的网格数等参数交给 FFI 入口。整个调用在torch.cuda.device(q.device)上下文中执行,异步入队到当前 PyTorch stream。
源码级原理:注册表、路由与启动绑定
jit.py:单一注册表
jit.py是包的“单一事实来源”,包含三张表(jit.py):
MODULES:4 个可编译程序,original(short+long 两个内核与其 binding 的合并模块)和三个 Prepared 专用模块n32_b4_direct、n32_disjoint_s10、n64_kvlast_s10,每个模块列出其.cu源文件与编译标志;ROUTES:5 条启动路由,记录“模块 + FFI 入口 + 内核符号 + 分裂数”:short→run_short/kernel_sm110_gqa_decode_short,num_splits=1;long→run_long/kernel_sm110_gqa_decode_long,num_splits=1;n32_b4_direct→ split 1;n32_disjoint_s10→ split 10;n64_kvlast_s10→ split 10。
PREPARED_ROUTES:Prepared API 的精确形状默认路由表,"4:256"→n32_b4_direct、"1:1024"→n64_kvlast_s10、"1:4096"→n32_disjoint_s10;SHORT_CAPACITY_MAX = 64:short 内核的容量上界。
注意ROUTES的注释:num_splits大于 1 的路由会把 QK^T 分片并行计算,各分片写出FP32 partials,由最后一个 CTA 负责合并——这正是 Prepared API 需要一次性分配并保留 workspace 的原因。
launch.cuh:TMA 描述符与动态共享内存 opt-in
所有 binding 都共享 sm110_gqa_decode_launch.cuh:
TensorMap64:64 字节对齐的 128 字节载体,承载CUtensorMap(tensor map 按值传给内核,static_assert保证 ABI 尺寸);EncodeQ:把 Q 视为 5D 张量[batch, 4, 8, 128](4 个头组 × 8 个 KV 头),编码 TMA 描述符,box 为(64, 64, 1, 1, 1),即一次搬运 64 个 head-dim 元素 × 64 行;EncodeKV:对 K/V 平面编码 4D 描述符[batch, 8, capacity, 128],box 为(64, box_tokens, 1, 1),short/long binding 均取box_tokens = 64;SetMaxDynamicSharedMemory:为每个接受该内核的设备做一次性动态共享内存上限 opt-in(short 50176 B、long 50944 B),函数内静态初始化保证只执行一次;Launch:封装cudaLaunchKernel并做错误检查。
编码前还会用CheckStridedFp16校验 TMA 源张量的最内层 stride 为 1、物理 stride 为正(contiguous 布局的硬性要求),以及CheckStride对解析出的全局 stride 做非负/非零校验——这就是为什么输入必须 contiguous。
binding 与内核:tcgen05 + mbarrier + exp2 softmax
每个 binding 都是薄启动器:run_short以dim3(256,1,1)块、run_long以dim3(384,1,1)块调用对应内核(见 short binding 与 long binding),内核签名以__grid_constant__方式接收三张 tensor map。
内核实现(short kernel)是典型的 SM110 编程模型组合:
- TMA 加载:
cp.async.bulk.tensor.4d/5d.shared::cta.global.mbarrier::complete_tx::bytes配合 mbarrier 的expect_tx机制,把 Q/K/V 异步搬入共享内存(SMEM 布局:Q 从偏移 1024 起占 16 KB,KV/V 从偏移 17408 起占 16 KB,short 内核总 SMEM 50176 B); - 张量核计算:
tcgen05.mma.cta_group::1.kind::f16在张量内存中执行 FP16 MMA,scores 与输出分别位于 tensor memory 的TMEM_SCORES_OFFSET=0与TMEM_OUTPUT_OFFSET=128列区; - 同步原语:
mbarrier.init/arrive/try_wait/arrive.expect_tx、elect.sync、tcgen05.commit全套 cluster/CTA 同步; - softmax 优化:
ex2.approx.ftz.f32计算exp2(scale_log2 * score)(与 log2 域缩放呼应),rcp.approx求倒数归一化,另有f32x2SIMD 打包的 FMA/加减/求最大辅助函数以及ex2_emulation_f32x2多项式模拟路径。
从源码结构可以推断:short 内核面向capacity <= 64的短前缀场景(SMEM 中可驻留全部 KV),long 内核则对更大容量做分块流水(SMEM 总量略升至 50944 B 以容纳更多流水阶段)。
Prepared 持续解码:为 CUDA Graph 与重复 launch 而生
便捷 API 每次调用都会走一次模块加载缓存与参数组装,且输出可自动分配。若你在做 serving 场景的持续 decode(同一形状反复 launch、CUDA Graph 捕获重放),应当使用 Prepared 三件套:prepare_sm110_gqa_decode+launch_sm110_gqa_decode_prepared(README 原文示例):
from flashinfer import prepare_sm110_gqa_decode, launch_sm110_gqa_decode_prepared output = torch.empty_like(q) prepared = prepare_sm110_gqa_decode( {"Q": q, "KV": kv, "O": output, "sequence_lengths": sequence_lengths} ) result = launch_sm110_gqa_decode_prepared(prepared) # result is output准备阶段做什么
prepare_for_launch(prepared.py)在Graph capture 之外完成:张量元数据校验(形状、dtype、contiguous、同设备)、O与任何输入 storage 的非别名检查、q_scale的有限且为正检查、路由选择、JIT 编译与 warmup,以及 split 路由所需 workspace 的一次性分配。返回的prepared是不透明字典(route、bindings、workspace、stages、launch_names、workspace_bytes 等字段),内部持有所有张量引用,调用方不得修改。
与便捷 API 的三点关键差异
- 必须提供调用者自持的
O,且其底层 storage 不得与Q、KV、sequence_lengths中任何一个共享(untyped_storage().data_ptr()逐一比对,见 prepared.py); q_scale必须有限且为正(默认 1.0),非法值直接抛ValueError;- 启动阶段零分配、零 host 读取:
launch_prepared仅把stages中的 FFI 调用按序执行(在tvm_ffi.use_torch_stream()上下文中使用当前 PyTorch stream),异步返回O。
生命周期与并发约束
- 保持
prepared对象与其内张量存活,直到所有异步工作完成、被捕获的 CUDA Graph 退役; - 每次 launch 前先更新张量内容(长度可写新值,只要仍在
[1, capacity]); - workspace 是可变的,因此并发执行的 stream 或 Graph 必须各自持有独立的 prepared 实例;有序 stream(含 event 同步交接)之间可以复用同一个 prepared(见
launch_prepared的 docstring)。
默认路由与 num_splits 语义
Prepared API 的默认路由针对三类精确形状做了专用内核(其余形状回退 original long):
| batch : capacity | 默认路由 | num_splits | 内部特征(源自 README) |
|---|---|---|---|
| 4 : 256 | n32_b4_direct | 1 | N32 直接输出(direct) |
| 1 : 1024 | n64_kvlast_s10 | 10 | N64、KV-last 布局 |
| 1 : 4096 | n32_disjoint_s10 | 10 | N32、ring-3 调度、disjoint 分裂 |
其他capacity > 64 | originallong | 1 | 384 线程流水内核 |
capacity <= 64 | originalshort | 1 | 256 线程内核 |
num_splits参数的完整规则(prepared.py 与 decode.py):
None:按上表默认选择形状专用路由;1:显式强制 original long(即使落在 B4/256 的 direct 路由上);10:两个专用 long 路由接受其固定分裂数;- 其余取值:接受集合为
{1, 2, 4, 8, 10, 16},但只有与所选路由固定分裂数一致才放行,否则抛ValueError(“exported fused route has a fixed split count” / “shape does not select an exported split tile”),不会为未知 split 现场生成新内核。
split > 1 时,准备阶段会分配 FP32 的partial_O [batch, 32, split, 128]、partial_max [batch, 32, split]、partial_sum(与 partial_max 同型)以及 uint32 的completed [batch * 8]完成计数器;最后一个 CTA 在合并后自行重置计数器,有序 launch 与 Graph 重放无需额外 reset kernel(见 prepared.py)。
端到端示例与基准
仓库提供了可直接运行的参考脚本:
- Prepared + CUDA Graph 示例:examples/experimental/sm110_gqa_decode_prepared.py。运行
python examples/experimental/sm110_gqa_decode_prepared.py --graph可看到完整的 Graph 流程:先用 side stream 完成编译/初始化/warmup(capture 流等待 warmup 结束),再在torch.cuda.graph(graph)上下文中捕获launch_sm110_gqa_decode_prepared(prepared),之后lengths.fill_(...)改写设备端长度(地址不变)并graph.replay()重放。不带--graph时则直接 launch 一次。默认--batch 1 --capacity 1024,恰好落在n64_kvlast_s10默认路由上; - 基准对比:benchmarks/bench_sm110_gqa_decode.py,在 SM110 设备上做 cold-L2 的 CUPTI 计时,与 PyTorch SDPA 对照(README 明确指出该基准为 cold-L2 对比,结论需在对应硬件上自行复现)。
实验性边界与毕业标准
为什么这个 API 被标记为 experimental?README 给出了直接理由:它的张量布局与固定头几何是 serving 工作负载特定的——32/8 头、head dim 128、每请求单 token、KV 平面堆叠格式,这些都不是通用 decode 的形状假设。因此:
- 它不参与 AOT 打包与自动路由,只能通过实验性 API 显式调用;
- 毕业(graduation)需要满足三个条件:更广泛的工作负载验证、稳定的打包覆盖、以及与FlashInfer 既有 decode API 的商定集成点(见 README)。
在毕业之前,建议把该 API 的使用范围限定在:已确认 SM110/Jetson AGX Thor 硬件、CUDA 13.0+ 环境、以及 4:1 GQA 固定几何的持续 decode 服务中,并始终通过prepare_*/launch_*_prepared路径复用 workspace 以获得稳定可重放的启动行为。
- 大模型
- 深度学习
- 算子库
- 后端
- 高性能计算
【免费下载链接】flashinfer
FlashInfer: Kernel Library for LLM Serving
相关推荐
FlashInfer 实验性 Balanced Paged GQA Decode(Cake 后端):SM100/SM103 上的片上负载均衡解码方案
FlashInfer 实验性 Balanced Paged GQA Decode(Cake 后端):SM100/SM103 上的片上负载均衡解码方案 本文深入解
大模型深度学习算子库后端高性能计算FlashInfer Gated Delta-Rule Decode API 完全指南:gdn_decode 系列内核的原理与实战
FlashInfer Gated Delta Rule Decode API 完全指南:gdn_decode 系列内核的原理与实战 本文导读 : docs/ap
大模型深度学习算子库后端高性能计算FlashInfer Cake GDN 非 CP 后端:SM100a/SM103a 上源自 Cake 生成的 Prefill 与 Decode 内核源码剖析
FlashInfer Cake GDN 非 CP 后端:SM100a/SM103a 上源自 Cake 生成的 Prefill 与 Decode 内核源码剖析 本
大模型深度学习算子库后端高性能计算
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考