☰
FlashInfer SM110 GQA Decode 实战指南:Jetson AGX Thor 上的精确 FP16 解码内核与 Prepared 持续解码
2026/10/9 5:27:07 网站建设 项目流程
  • 大模型
  • 深度学习
  • 算子库
  • 后端
  • 高性能计算

【免费下载链接】flashinfer

FlashInfer: Kernel Library for LLM Serving

项目地址:https://gitcode.com/gh_mirrors/fl/flashinfer
点击查看免费下载

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]FP16contiguous,CUDA 设备
kv[batch, 2, 8, capacity, 128]FP16contiguous;索引 0 为 K、索引 1 为 V
sequence_lengths[batch]CUDA int32contiguous;每个值必须落在闭区间[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—见上文输入契约
outNone传入时使用调用者自持输出;缺省时内部torch.empty_like(q)分配。若与q共享存储会抛ValueError
q_scale1.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 的三点关键差异

  1. 必须提供调用者自持的O,且其底层 storage 不得与Q、KV、sequence_lengths中任何一个共享(untyped_storage().data_ptr()逐一比对,见 prepared.py);
  2. q_scale必须有限且为正(默认 1.0),非法值直接抛ValueError;
  3. 启动阶段零分配、零 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 : 256n32_b4_direct1N32 直接输出(direct)
1 : 1024n64_kvlast_s1010N64、KV-last 布局
1 : 4096n32_disjoint_s1010N32、ring-3 调度、disjoint 分裂
其他capacity > 64originallong1384 线程流水内核
capacity <= 64originalshort1256 线程内核

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

项目地址:https://gitcode.com/gh_mirrors/fl/flashinfer
点击查看免费下载

相关推荐

上一篇:EMQX Message Streams 消息流功能实战:基于 Topic Filter 的持久化消息集合与 `$s/` 消费协议
下一篇:Zebraix 图工具详解:基于序维 2 偏序集的 Jaywalk 图定义、测试与渲染能力

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询