资讯动态

FlashInfer SM110 GQA Decode 实战指南:Jetson AGX Thor 上的精确 FP16 解码内核与 Prepared 持续解码

发布时间:2026/10/9 5:27:21 来源:尧图企业网站定制
大模型深度学习算子库后端高性能计算【免费下载链接】flashinferFlashInfer: Kernel Library for LLM Serving项目地址https://gitcode.com/gh_mirrors/fl/flashinfer点击查看免费下载FlashInfer 的sm110_gqa_decode是一套显式 opt-in 的实验性 FP16 GQAGrouped Query Attention解码路径专门面向 NVIDIA SM110Jetson AGX Thor的 compute capability 11.0 硬件。它以固定的 32 查询头 / 8 KV 头 / 128 头维度几何复用 SM110 的tcgen05张量核、TMA 与张量内存提供便捷 API 与 Prepared 持续解码 API 两套入口。读完本文你将掌握该内核的输入契约、双内核short/long路由机制、JIT 注册表与启动绑定原理并能用prepare_sm110_gqa_decodelaunch_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新一代张量核指令、TMATensor 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(featureSM110 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.0SM110即 Jetson AGX ThorCUDA 版本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与链接标志-lcudaTMA 描述符编码需要 CUDA driver API三个 Prepared 专用模块额外使用--use_fast_math编译标志见 jit.py 与gen_sm110_gqa_decode_module。输入契约固定的张量布局与几何该内核不接受任意形状输入布局是硬编码的参数形状dtype要求q[batch, 32, 128]FP16contiguousCUDA 设备kv[batch, 2, 8, capacity, 128]FP16contiguous索引 0 为 K、索引 1 为 Vsequence_lengths[batch]CUDA int32contiguous每个值必须落在闭区间[1, capacity]out可选[batch, 32, 128]FP16调用者自持contiguous不能与q别名其中kv的第 2 维是两个平面K 平面位于 index 0V 平面位于 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.pyCUDA 侧 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, dtypetorch.float16, devicecuda) kv torch.randn(1, 2, 8, 1024, 128, dtypetorch.float16, devicecuda) sequence_lengths torch.tensor([1024], dtypetorch.int32, devicecuda) out sm110_gqa_decode(q, kv, sequence_lengths)签名细节与 flashinfer/decode.py 一致参数默认值说明q/kv/sequence_lengths—见上文输入契约outNone传入时使用调用者自持输出缺省时内部torch.empty_like(q)分配。若与q共享存储会抛ValueErrorq_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.pycapacity 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.pyMODULES4 个可编译程序originalshortlong 两个内核与其 binding 的合并模块和三个 Prepared 专用模块n32_b4_direct、n32_disjoint_s10、n64_kvlast_s10每个模块列出其.cu源文件与编译标志ROUTES5 条启动路由记录“模块 FFI 入口 内核符号 分裂数”short→run_short/kernel_sm110_gqa_decode_shortnum_splits1long→run_long/kernel_sm110_gqa_decode_longnum_splits1n32_b4_direct→ split 1n32_disjoint_s10→ split 10n64_kvlast_s10→ split 10。PREPARED_ROUTESPrepared API 的精确形状默认路由表4:256→n32_b4_direct、1:1024→n64_kvlast_s10、1:4096→n32_disjoint_s10SHORT_CAPACITY_MAX 64short 内核的容量上界。注意ROUTES的注释num_splits大于 1 的路由会把 QK^T 分片并行计算各分片写出FP32 partials由最后一个 CTA 负责合并——这正是 Prepared API 需要一次性分配并保留 workspace 的原因。launch.cuhTMA 描述符与动态共享内存 opt-in所有 binding 都共享 sm110_gqa_decode_launch.cuhTensorMap6464 字节对齐的 128 字节载体承载CUtensorMaptensor 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 64SetMaxDynamicSharedMemory为每个接受该内核的设备做一次性动态共享内存上限 opt-inshort 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 KBKV/V 从偏移 17408 起占 16 KBshort 内核总 SMEM 50176 B张量核计算tcgen05.mma.cta_group::1.kind::f16在张量内存中执行 FP16 MMAscores 与输出分别位于 tensor memory 的TMEM_SCORES_OFFSET0与TMEM_OUTPUT_OFFSET128列区同步原语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 中可驻留全部 KVlong 内核则对更大容量做分块流水SMEM 总量略升至 50944 B 以容纳更多流水阶段。Prepared 持续解码为 CUDA Graph 与重复 launch 而生便捷 API 每次调用都会走一次模块加载缓存与参数组装且输出可自动分配。若你在做 serving 场景的持续 decode同一形状反复 launch、CUDA Graph 捕获重放应当使用 Prepared 三件套prepare_sm110_gqa_decodelaunch_sm110_gqa_decode_preparedREADME 原文示例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_launchprepared.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.pyq_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 longbatch : capacity默认路由num_splits内部特征源自 README4 : 256n32_b4_direct1N32 直接输出direct1 : 1024n64_kvlast_s1010N64、KV-last 布局1 : 4096n32_disjoint_s1010N32、ring-3 调度、disjoint 分裂其他capacity 64originallong1384 线程流水内核capacity 64originalshort1256 线程内核num_splits参数的完整规则prepared.py 与 decode.pyNone按上表默认选择形状专用路由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 完成编译/初始化/warmupcapture 流等待 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 被标记为 experimentalREADME 给出了直接理由它的张量布局与固定头几何是 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 以获得稳定可重放的启动行为。赞分享大模型深度学习算子库后端高性能计算【免费下载链接】flashinferFlashInfer: Kernel Library for LLM Serving项目地址https://gitcode.com/gh_mirrors/fl/flashinfer点击查看免费下载相关推荐FlashInfer 实验性 Balanced Paged GQA DecodeCake 后端SM100/SM103 上的片上负载均衡解码方案FlashInfer 实验性 Balanced Paged GQA DecodeCake 后端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 内核源码剖析 本大模型深度学习算子库后端高性能计算上一篇EMQX Message Streams 消息流功能实战基于 Topic Filter 的持久化消息集合与 $s/ 消费协议下一篇Zebraix 图工具详解基于序维 2 偏序集的 Jaywalk 图定义、测试与渲染能力创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑