- 大模型
- 深度学习
- 算子库
- 后端
- 高性能计算
【免费下载链接】flashinfer
FlashInfer: Kernel Library for LLM Serving
本文以flashinfer/experimental/kimi_k3_latent_moe/csrc/README.md为骨架,讲解 FlashInfer 中 Kimi-K3 Stable LatentMoE(nvidia/Kimi-K3-NVFP4的KimiSparseMoeBlock)front/tail 两阶段密集投影的生成式 CUDA 内核:这些*.cu文件由 Cake 生成式程序导出流水线自动产出,配套的cake_jit.py注册表与cake_backend.py主机规划器共同完成按设备 SM 数、TP 度、token 数与路由 partial 数的运行时派发。读完本文,你将掌握该实验模块的目录结构、内核生成与 JIT 编译机制、四种路由(decode/prefill × front/tail)的调度规则、公开 Python API 的调用方式,以及它的硬件适配范围与已知限制。
一、csrc/README.md 说了什么:一份"生成源"声明
flashinfer/experimental/kimi_k3_latent_moe/csrc/README.md全文很短,核心信息有三条:
- 目录性质:
cake_kimi_k3_latent_moe/*.cu不是手写代码,而是由 Cake 仓库的生成式程序导出工具(tools/export-generated-programs run配合exports/kimi_k3_latent_moe/export.py)生成的,导出过程会同步填充上一级目录的../cake_jit.py(其中的MODULES、KERNELS、SPECIALIZATIONS三个注册表)。 - 文件组织:每个程序对应一对
_kernel.cu/_binding.cu,同一对源文件同时服务于 SM100a 与 SM103a 两种架构(两者编译同一份源码文本);所有程序共享的辅助 preamble 只保留一份,存放在cake_kimi_k3_latent_moe_device_common.cuh与cake_kimi_k3_latent_moe_host_common.cuh。 - 维护约束:该目录下没有任何手写内容,禁止手工编辑生成文件。
这篇 README 是理解整个模块的入口。真实的工程细节(计算语义、路由表、API、硬件限定)都记录在同目录的父级文档flashinfer/experimental/kimi_k3_latent_moe/README.md与源码中,下面逐层展开。
二、模块全景:目录结构与文件职责
flashinfer/experimental/kimi_k3_latent_moe/目录由四部分组成(已在仓库中确认存在):
| 文件/目录 | 职责 |
|---|---|
cake_backend.py | 主机规划器:decode/prefill 各路由的参数规划、算子合法性校验、启动绑定、runner 封装(约 1540 行) |
cake_jit.py | JIT 注册表:MODULES/KERNELS/SPECIALIZATIONS三个字典,由 Cake 导出流水线填充;gen_jit_spec生成编译规格、按架构选择 nvcc 参数 |
csrc/cake_kimi_k3_latent_moe/ | 生成的设备端与主机端绑定翻译单元(_kernel.cu+_binding.cu成对出现),另有_device_common.cuh/_host_common.cuh共享 preamble |
__init__.py | 包文档:说明这是 SM100/SM103 的实验性 Cake 后端,公共入口在flashinfer.kimi_k3_latent_moe |
配套的公共 API 层位于仓库根级模块flashinfer/kimi_k3_latent_moe.py(prepare_kimi_k3_latent_moe_front/kimi_k3_latent_moe_front/prepare_kimi_k3_latent_moe_tail/kimi_k3_latent_moe_tail四个函数),它只做一件事:强制backend="cake"后把调用转发到flashinfer.experimental.kimi_k3_latent_moe.cake_backend。也就是说,公开入口是实验性的薄封装,真正的规划与启动逻辑全部在实验包内。
生成文件的命名规律也值得注意:每个程序以 20 位十六进制 hash 命名(如cake_kimi_k3_latent_moe_a03cfa14a84509187ac0),这个 hash 在cake_jit.MODULES中与 per-arch 的 closure digest 一同记录,JIT 编译身份(identity)就取arch + closure 前 16 位拼接,用于缓存隔离。
三、计算语义:front/tail 两阶段到底算什么
Kimi-K3 的KimiSparseMoeBlock由"共享专家 + 路由专家"构成。本模块负责每个 rank、每一层包围路由专家的两组稠密投影,数学定义(摘自包内 README 与cake_backend.py模块 docstring,二者一致):
x [T, 7168] bf16 front: logits = FP32(x @ gate_weight.T) [T, 896] 路由器 logits latent = BF16(x @ down_weight.T) [T, 3584] 路由专家输入 shared_act = SiTU(x @ shared_gate.T, x @ shared_up.T) [T, 6144/TP] 共享专家中间激活 tail: y = KimiRMSNorm(sum of the P routed partials) [T, 3584] FP32 归一化 -> BF16 -> * bf16 权重, eps 1e-5 out = BF16(y[:, cols] @ up_weight[:, cols].T + shared_act @ shared_down.T) [T, 7168], fp32 累加只舍入一次 SiTU(g, u) = 4 tanh(g / 4) sigmoid(g) * 25 tanh(u / 25) (fp32 数学,一次 bf16 舍入)关键工程细节:
cols是当前 rank 的3584 / TP隐层切片(行并行 up 投影);当TP > 1时,调用方需要在 kernel 之后对out做 all-reduce。tail阶段输入是P个未归约的路由专家 partial([P, T, 3584]),在设备上以 FP32 求和、只舍入一次。- 权重直接使用 checkpoint 中
nn.Linear的[out, in]BF16 张量(serving 布局,vLLMKimiMoE/ SGLangKimiK3MoE):gate_weight、down_weight、norm_weight、up_weight为复制(replicated)权重;共享专家以MergedColumnParallel(gate/up 行切分,作为一份[2 * 6144/TP, 7168]拼接传入)+RowParallel(shared_down_weight[7168, 6144/TP])表示。不做任何拷贝或重打包——权重就是模型原布局。
从tests/experimental/test_cake_kimi_k3_latent_moe.py可以看到该语义的验收标准:BF16 输出容差 1e-2、FP32 logits 容差 1e-3(ATOL = RTOL = 1e-2,LOGITS_ATOL = LOGITS_RTOL = 1e-3),SiTU 的 beta 常数即源码中的SITU_BETA = 4.0、SITU_LINEAR_BETA = 25.0,与公式一一对应。
四、四种路由与内核选择:从 token 数到物理内核的映射
模块将调用分为四条路线(cake_backend.py中的规划函数decode_front_plan、decode_tail_plan、prefill_front_plan、front_split_plan、prefill_tail_plan、split_plan即 Cake 生产 launch 器的逐字节移植):
| 阶段 | T ≤ 128(decode) | T > 128(prefill) |
|---|---|---|
| front | 单次 launch 的 weight-streaming 交换 AB tcgen05 kernel(decode:*):128 行权重 tile 作为 MMA A 操作数,token 作为补齐后的 B 操作数(N ∈ {8, 16, 32, 64, 128});router / latent / shared SiTU tile 统一在一个 tile 空间;当2 × tiles能装进一个 wave 时用对齐的 2-CTA cluster 对(TP8),否则每 tile 一个 CTA(TP1) | 持久化 2-CTA tcgen05 GEMM(front:i<6144/TP>[e1];e1实例在 T ≤ 512 时以evict_first策略流式读权重);对 K 拼接的[gate; down; shared_gate; shared_up]行用 256×256 对 tile,带分类 epilogue;尾部 wave 的 tile 在能装进一个 wave 时以两个对齐 K 半块 + fp32 partial + 固定顺序 fix-up 运行(front_split_plan) |
| tail | 同一 streaming kernel,但把 KimiRMSNorm 融合进该 launch(decode:*_f*):epilogue warp 归一化路由行,TMA warp 同时流式读 shared-down 段;最小的 T 直接把归一化行写进常驻 smem B 操作数(_sb,staged rows_rs),更大的 T 走全局y_workspace协议 | 两段 launch:one-pass RMSNorm kernel(tail_norm:e<0|1>,late/early PDL trigger 由 GEMM grid 决定)+ 持久化 2-CTA GEMM(tail_gemm:tp<1|8>,对[up slice | shared down]做 trailing-wave stream-K),以 programmatic-dependent 方式启动 |
规划器如何工作:每个 plan 通过逻辑 key 命名其物理内核,key 在cake_jit.KERNELS中注册(每个 key 一个程序,同一份源码为 SM100a 与 SM103a 分别编译);cake_jit.SPECIALIZATIONS为共享同一程序的不同 key 附加编译期常量(例如 TP1/TP8 的 front/tail GEMM 差异)。从 1 到 16384 的每个 token 数、两种 TP 度、任意数量的路由 partial 都能解析到一个已注册程序。
按设备 SM 数派生的关键设计:每个 plan 都由设备自身 SM 数推导(torch.cuda.get_device_properties(...).multi_processor_count)。decode grid 是 tile 受限的(在合格部件上 94/112/131 个共驻 CTA),prefill GEMM 则以启动参数接收 resident-pair 窗口(num_items、full_items、sk_*),因此一个程序服务所有合格 SM 数。prepare_*只受理cake_backend.QUALIFIED_SM_COUNTS(148 和 152)列出的 SM 数,其他数量会抛出NotImplementedError并指名设备数量——因为 decode 程序依赖 grid 中每个 CTA 都驻留,而 plan 规则只在这些部件上验证过。
generated_program_available(device, stage, tp, num_tokens, num_partials)精确回答prepare_*会要求的路线是否可用(decode tail 实例依赖路由 partial 数 P)。plan 按 shape 做 memoize(_cached_plan装饰器,functools.lru_cache),每次 launch 不做规划、不做分配——per-device scratch buffer 在 prepare 阶段创建,因此准备好的 runner(或其 CUDA Graph 捕获)可以在绑定的 buffer 中写入新值后直接重放。
五、注册表机制:MODULES / KERNELS / SPECIALIZATIONS 的工程细节
cake_jit.py是理解生成式内核如何被 FlashInfer 的 JIT 管线拾起的关键。三个字典的分工如下:
MODULES:每个物理程序一条记录,包含role(全部为"kernel")、sources(相对csrc/的_kernel.cu+_binding.cu翻译单元)、compile_flags(统一--use_fast_math)、ffi_entry("run")、arg_plan(FFI 参数方案)、launch(block 尺寸、cluster 尺寸、cooperative、动态共享内存字节数)、arches(["sm_100a", "sm_103a"])与 per-arch closure digest。例如解码 tail 融合实例的arg_plan列出tma_buffer(A_R/A_L/A_S/A_2/B_1/B_2)、buffer(out_r/out_l/out_s/counters/routed/norm_w/y_out/tl)、parameter(num_tokens/k1_off/num_partials/eps)与grid_x/y/z;prefill tail GEMM 则额外携带m_tiles、k0_blocks、num_items、full_items、sk_ipc、sk_max_seg、sk_total等 stream-K 窗口参数。launch 规格显示 decode 实例 block 为 192 线程、prefill GEMM 为 320 线程,动态 smem 从约 197 KB 到 232 KB 不等,最大不超过cake_backend.MAX_DYN_SMEM = 232448。KERNELS:逻辑 key → 程序名。key 自带完整语义,例如decode:kimi_k3_latent_moe_decode_g112_n128_r7_d1_p0_w0_t0_56_0_k56_96_o7168_c2_f_la表示 grid 112、N_PAD 128、ring 7 段、TP8(cluster 2)且带融合 norm(_f)与 landing-zone alias(_la);tail_gemm:tp8e0f0s9n128表示 TP8、e0(无 evict_first)、f0(不融合 norm)、9 段 ring、128 宽 tile。cake_jit.kernel_module_name(arch, key)负责查表并校验目标架构,未注册的 key 会抛NotImplementedError并引用上游 issue 编号。SPECIALIZATIONS:key → 编译期#define。front 区分I_LOCAL(TP1=6144/TP8=768)与 tile 数(N_TILES/S_TILES);tail GEMM 区分 K 迭代计数(TP1:K1_ITERS=56, K2_ITERS=96, NUM_K_ITERS=152;TP8:7/12/19)。这些常量以-DNAME=value拼进 nvcc 命令行(gen_cake_kimi_k3_latent_moe_module),并用functools.cache保证每(程序、架构、define 集)只构建一次。
编译管线(gen_cake_kimi_k3_latent_moe_module)要点:源码根为包内csrc/,include 路径按"安装态还是 checkout 态"自动探测(_header_dirs()优先FLASHINFER_CSRC_DIR/FLASHINFER_INCLUDE_DIR,否则回落到仓库csrc/+include/),额外链接-lcuda,nvcc 参数取自sm100a_nvcc_flags/sm103a_nvcc_flags(flashinfer/jit/core.py提供),最终build_and_load()加载为可用模块。
六、公开 API 与可运行的调用示例
公共入口在flashinfer/kimi_k3_latent_moe.py,四个函数均带@flashinfer_experimental_api装饰器——显式调用即表示接受实验性特性,会触发 FlashInfer 的实验 API 警告,且没有自动后端选择(仅支持backend="cake")。包内 README 给出的最小示例:
import torch from flashinfer.kimi_k3_latent_moe import ( kimi_k3_latent_moe_front, kimi_k3_latent_moe_tail, prepare_kimi_k3_latent_moe_front, prepare_kimi_k3_latent_moe_tail, ) T, tp, rank = 16, 8, 0 i_local = 6144 // tp dev = torch.device("cuda") x = torch.randn(T, 7168, device=dev, dtype=torch.bfloat16) logits = torch.empty(T, 896, device=dev, dtype=torch.float32) latent = torch.empty(T, 3584, device=dev, dtype=torch.bfloat16) shared_act = torch.empty(T, i_local, device=dev, dtype=torch.bfloat16) # gate_weight [896, 7168], down_weight [3584, 7168], shared_gate_up [2 * i_local, 7168](gate 行在前、up 行在后) runner = prepare_kimi_k3_latent_moe_front(x, gate_weight, down_weight, shared_gate_up, logits, latent, shared_act) runner() # 或捕获进 CUDA Graph # tail: routed [P, T, 3584] partials, norm_weight [3584], up_weight [7168, 3584], shared_down [7168, i_local] out = torch.empty(T, 7168, device=dev, dtype=torch.bfloat16) y = torch.empty(T, 3584, device=dev, dtype=torch.bfloat16) kimi_k3_latent_moe_tail(routed, norm_weight, up_weight, shared_act, shared_down, out, tp=tp, rank=rank, y_workspace=y)各参数的形状与语义(依据flashinfer/kimi_k3_latent_moe.py的 docstring):
- front:
x为连续 BF16[T, 7168];gate_weight[896, 7168](复制);down_weight[3584, 7168](复制);shared_gate_up_weight[2 * 6144/TP, 7168](gate 行后接 up 行);输出logitsFP32[T, 896]、latentBF16[T, 3584]、shared_actBF16[T, 6144/TP],全部 caller-owned。 - tail:
routedBF16[P, T, 3584](P 个未归约 partial);norm_weightBF16[3584];up_weightBF16[7168, 3584](复制);shared_act为 front 阶段产物;shared_down_weightBF16[7168, 6144/TP];outBF16[T, 7168];y_workspaceBF16[T, 3584]接收归一化后的 latent,返回值即(y_workspace, out)。 - runner 语义:调用 runner 在当前流上启动 route,无 CUDA 分配、无主机同步,返回三个输出(front)或
(y_workspace, out)(tail);CUDA Graph 捕获需在捕获外 prepare。
注意T的上限约束:decode 程序只服务T ∈ {1, …, 128}(任意数量,按下一个 N 补齐),prefill 程序服务T > 128(任意数量,验证至 16384)。
七、测试验证:主机规划单测与 GPU 正确性
测试文件tests/experimental/test_cake_kimi_k3_latent_moe.py覆盖两层:
- 主机规划规则(CPU 单测):例如
test_decode_plan_rules验证 front TP1 流式 7+28+96=131 个 tile(每 tile 一个 CTA)、TP8 为 47 个 tile 映射成 94 个 cluster-pair CTA;front 在其 ring 保持 ≥4(单 CTA)/≥6(cluster 对)段时取 K depth 2(TP1 T≤64、TP8 T≤16),tail 恒为 depth 1 且 fused norm 恒开;test_split_plan_rules验证 stream-K 切分仅在余波 wave 填充 ≤55% cluster 时生效,且 152-SM 上 TP1 T=4096/8192 区域因切分收益过小而保持整 tile。这些断言直接印证了 README 中"decode grids are tile-bound"与 plan 规则从 Cake 移植的说法。 - GPU 正确性:对两阶段、TP 1/8、
T ∈ {1, 8, 16, 128, 256, 4096}(冒烟集还含 8192、16384)对照 FP32/BF16 torch 参考实现验证,要求 bit 级一致的重新 launch、CUDA Graph 重放,且y字节级一致(byte-exact)。
八、合格物理配置与限制清单
合格配置(README "Tested physical configurations",qualification 包括合约容差下的 GPU 正确性、bit-exact 重放与 CUDA-Graph 重放、对生产 launch 器与合约基线的完整调用 CUPTI cold-L2 基准):
| 设备 | 架构 | SM | 驱动 | 状态 |
|---|---|---|---|---|
| NVIDIA B200 | sm_100a | 148 | 580.82.07 | qualified(导出协议,每个分母行) |
| NVIDIA B300 SXM6 | sm_103a | 148 | 580.126.09 | qualified(导出协议,每个分母行) |
| NVIDIA GB300 (NVL72) | sm_103a | 152 | 580.159.03 | qualified(152-SM 配置每个分母行;逐行测量相对 torch/cuBLAS 链的完整调用加速) |
| NVIDIA GB200 (NVL72) | sm_100a | 152 | 未记录(无导出阶段) | admitted:152 SM 上每条路由的主机规划与程序注册已验证;GB200 硬件上包级 GPU 测试通过,完整调用合约行已对 torch/cuBLAS 链计时(front 每行更快;tail 见交付摘要);本次交付无导出阶段 |
限制:
- 仅 SM100(B200/GB200)与 SM103(B300/GB300):tcgen05/TMEM/TMA/2-CTA MMA 程序,且仅限上表的 SM 数。
- 仅 TP 1 与 TP 8(无专家并行);TP 12(7168/12 对 tail 切片不是 64 的倍数)超出范围。
- 仅模型布局权重;Cake kernel 的 chunk-major packed-weight 变体(decode 行 +3-8%,每 rank 多一次权重拷贝)未导出。
- 只支持验证行集的精确形状:decode 程序
T ∈ {1, …, 128}(任意数量,补齐到下一个 N),prefill 程序T > 128(验证至 16384)。
调试与扩展入口:想深入内核生成机制的读者可从flashinfer/experimental/kimi_k3_latent_moe/cake_backend.py(规划器常量如BLOCK_ROWS=128、CHUNK_K=64、MAX_DYN_SMEM=232448、LAND_ALIAS_N_PAD=128等全部在此)与flashinfer/experimental/kimi_k3_latent_moe/cake_jit.py(JIT 规格)入手;生成的.cu文件仅供阅读,切勿手改。
九、总结
Kimi-K3 LatentMoE 的 csrc README 虽然只有 9 行,却准确描述了 FlashInfer 中"生成式内核"工作流的核心原则:内核由工具生成而非手写、一份源码多架构复用、注册表驱动 JIT 编译、主机规划器按设备派生路由。结合包内 README、cake_backend.py、cake_jit.py与测试,可以还原出完整的调用链:flashinfer.kimi_k3_latent_moe薄封装 →cake_backend规划器(memoized plan)→cake_jit.KERNELS逻辑 key 解析 → per-arch nvcc JIT 编译 → 无分配、Graph 安全的一次性 launch。这套设计让 1~16384 个 token、TP1/TP8、任意路由 partial 数都能在 B200/B300/GB200/GB300 上解析到经过验证的物理内核,同时也通过QUALIFIED_SM_COUNTS与形状白名单将使用范围严格限定在验证过的硬件与形状之内——这正是实验性高算力 kernel 模块应有的工程严谨性。
- 大模型
- 深度学习
- 算子库
- 后端
- 高性能计算
【免费下载链接】flashinfer
FlashInfer: Kernel Library for LLM Serving
相关推荐
FlashInfer 实验性 Kimi-K3 Stable LatentMoE 前后投影指南:SM100/SM103 上的 tcgen05 生成式内核路线
FlashInfer 实验性 Kimi K3 Stable LatentMoE 前后投影指南:SM100/SM103 上的 tcgen05 生成式内核路线 本文
大模型深度学习算子库后端高性能计算FlashInfer Kimi-K3 Vision Tower 生成式源码体系:csrc 生成目录、Cake JIT 注册表与 SM100/SM103 程序路由解析
FlashInfer Kimi K3 Vision Tower 生成式源码体系:csrc 生成目录、Cake JIT 注册表与 SM100/SM103 程序路由
大模型深度学习算子库后端高性能计算FlashInfer Kimi-K3 FP8_PB_WO 投影 GEMM 实战指南:Cake 后端在 SM100/SM103 上的量化、调度与 CUDA Graph 用法
FlashInfer Kimi K3 FP8_PB_WO 投影 GEMM 实战指南:Cake 后端在 SM100/SM103 上的量化、调度与 CUDA Gra
大模型深度学习算子库后端高性能计算
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考