- 大模型
- 深度学习
- 算子库
- 后端
- 高性能计算
【免费下载链接】flashinfer
FlashInfer: Kernel Library for LLM Serving
POD(Prefill-On-Decode)是 FlashInfer 提供的一种融合注意力执行模式:在同一次 kernel launch中,同时执行单请求(或批量)的 prefill 注意力与批量的 decode 注意力,从而为"chunked prefill 与进行中的 decode 请求相互重叠"的 LLM 服务场景节省内核启动与调度开销。本文将基于 docs/api/pod.rst 定义的 API 骨架,结合 flashinfer/pod.py、csrc/pod.cu、csrc/batch_pod.cu 及测试代码,系统讲解 POD 的设计动机、两个 Wrapper 类的完整用法、plan/run生命周期与全部参数语义,以及 PDL(Programmatic Dependent Launch)、workspace、JIT 编译等底层实现细节。
POD 是什么:为什么需要一次 Launch 跑两种注意力
在 continuous batching(连续批处理)与 chunked prefill 架构下,服务端通常会遇到一个经典场景:一批正在逐 token 解码(decode)的请求尚未结束,此时又来了新的长序列请求需要 prefill。传统做法是分别调用 prefill kernel 与 decode kernel 两次 kernel launch:
- prefill 阶段:单请求、长序列、KV 长度大,属于计算密集型(compute-bound);
- decode 阶段:多请求、每请求仅 1 个新 token、KV 长度大,属于访存密集型(memory-bound)。
两次 launch 之间 kernel 无法共享调度资源,GPU 上也会出现"prefill 时 decode 空转、decode 时 prefill 空转"的利用率缝隙。POD 的思路(首次提出于 arxiv 2410.18038,该论文链接是 FlashInfer 官方文档注释中给出的原始出处)是把二者合并到一次 kernel 启动中,让 prefill 与 decode 两个任务在同一 launch 里并发执行、共享 SM 调度资源。
FlashInfer 在flashinfer.pod模块中提供两个 Python 层 Wrapper:
| Wrapper 类 | 覆盖场景 | 对应入口 |
|---|---|---|
PODWithPagedKVCacheWrapper | 单请求prefill(稠密 k/v 张量)+批量decode(paged kv-cache) | csrc/pod.cu |
BatchPODWithPagedKVCacheWrapper | 批量paged prefill +批量paged decode | csrc/batch_pod.cu |
两者的核心设计一致:plan阶段创建可复用的辅助数据结构,run阶段一次性完成融合计算并返回(out_p, out_d)两个输出。
单请求 POD:PODWithPagedKVCacheWrapper 完整示例
以下代码来自 flashinfer/pod.py 中PODWithPagedKVCacheWrapper类的 docstring 示例(L61-L125),它演示了在 32 层 Transformer 上逐层复用同一套 plan 辅助数据结构的标准用法:
import torch import flashinfer num_layers = 32 num_qo_heads = 64 num_kv_heads = 8 head_dim = 128 max_num_pages = 128 page_size = 16 # 分配 128MB workspace 缓冲区 workspace_buffer = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device="cuda:0") decode_wrapper = flashinfer.PODWithPagedKVCacheWrapper(workspace_buffer, "NHD") batch_size = 7 kv_page_indices = torch.arange(max_num_pages).int().to("cuda:0") kv_page_indptr = torch.tensor( [0, 17, 29, 44, 48, 66, 100, 128], dtype=torch.int32, device="cuda:0" ) # 1 <= kv_last_page_len <= page_size kv_last_page_len = torch.tensor( [1, 7, 14, 4, 3, 1, 16], dtype=torch.int32, device="cuda:0" ) kv_cache_at_layer = [ torch.randn( max_num_pages, 2, page_size, num_kv_heads, head_dim, dtype=torch.float16, device="cuda:0" ) for _ in range(num_layers) ] # 为 batch decode 注意力创建辅助数据结构 decode_wrapper.plan( kv_page_indptr, kv_page_indices, kv_last_page_len, num_qo_heads, num_kv_heads, head_dim, page_size, pos_encoding_mode="NONE", data_type=torch.float16 ) outputs = [] for i in range(num_layers): q = torch.randn(batch_size, num_qo_heads, head_dim).half().to("cuda:0") kv_cache = kv_cache_at_layer[i] # 计算 batch decode attention,所有层复用同一套辅助数据结构 # TODO_AK: DEMONSTRATE USAGE OF POD outputs.append(o) ... outputs[0].shape # torch.Size([7, 64, 128])需要说明的是,上述 docstring 示例中的run调用体本身尚是占位(源码中标注了TODO_AK),完整的run参数形态请参见下文"run 参数详解"。这个示例的价值在于完整呈现了 POD 的数据流骨架:plan只需要调用一次,run可跨层反复复用。
构造参数详解
PODWithPagedKVCacheWrapper.__init__的完整签名(flashinfer/pod.pyL127-L174):
def __init__( self, float_workspace_buffer: torch.Tensor, kv_layout: str = "NHD", use_cuda_graph: bool = False, paged_kv_indptr_buffer: Optional[torch.Tensor] = None, paged_kv_indices_buffer: Optional[torch.Tensor] = None, paged_kv_last_page_len_buffer: Optional[torch.Tensor] = None, jit_args: Optional[List[Any]] = None, ) -> None| 参数 | 类型 | 说明 |
|---|---|---|
float_workspace_buffer | torch.Tensor | 用户预留的 float workspace 缓冲区,用于存储 split-k 算法中的中间注意力结果;官方推荐大小128MB,其 device 必须与输入张量一致 |
kv_layout | str | 输入 k/v 张量的布局,"NHD"或"HND",默认"NHD" |
use_cuda_graph | bool | 是否启用 CUDAGraph 模式。启用后辅助数据结构将写入用户提供的缓冲区,且batch_size 在 wrapper 生命周期内不可变化(源码L218会从paged_kv_last_page_len_buffer长度固定_fixed_batch_size) |
paged_kv_indptr_buffer | Optional[torch.Tensor] | 仅在use_cuda_graph=True时需要,GPU 上存储 kv cache indptr 的预留缓冲区,大小为[batch_size + 1] |
paged_kv_indices_buffer | Optional[torch.Tensor] | 仅在use_cuda_graph=True时需要,需足够容纳生命周期内最大页索引数(max_num_pages) |
paged_kv_last_page_len_buffer | Optional[torch.Tensor] | 仅在use_cuda_graph=True时需要,大小[batch_size] |
jit_args | Optional[List[Any]] | 若提供,则用给定参数创建自定义 JIT 模块;否则使用默认注意力实现(当前源码中jit_args分支为注释掉的代码,实际固定走 tensor-core 路径,见L176-L189) |
从源码L188-L189可以看到一行关键注释:# Override options. Only tensor core version is performant.(仅 tensor core 版本才有性能),因此当前实现固定启用 tensor core,不再开放选择。
CUDAGraph 模式的运行时约束
构造函数中use_cuda_graph=True时(L205-L222):
- 三个 buffer 参数必须都是
torch.Tensor,否则直接抛ValueError; paged_kv_indptr_buffer长度必须等于batch_size + 1,否则抛错;- 同时
plan阶段(L352-L370)会校验运行期batch_size与初始化时固定的_fixed_batch_size一致,并校验indices长度不超过预留 buffer。
非 CUDAGraph 模式下(L371-L383),plan会以non_blocking方式把 indptr/indices/last_page_len 拷入设备端。
批量 POD:BatchPODWithPagedKVCacheWrapper 完整示例
BatchPODWithPagedKVCacheWrapper将 prefill 一侧也从"单请求稠密张量"扩展为"批量 paged kv-cache"。以下示例同样来自 flashinfer/pod.py docstring(L734-L823),展示了 2 条 prefill 请求(各 2048 token)与 128 条 decode 请求的融合:
import torch import flashinfer num_layers = 8 num_qo_heads = 64 num_kv_heads = 8 head_dim = 128 max_num_pages = 128 device = 0 page_block_size = 1 causal = True # 分配 128MB workspace 缓冲区(内部会均分为 prefill/decode 两份) workspace_buffer = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device="cuda:0") wrapper = flashinfer.BatchPODWithPagedKVCacheWrapper(workspace_buffer, "NHD") # Prefill 与 decode 参数 p_qo_lens = [2048] * 2 d_qo_lens = [1] * 128 p_kv_lens = [2048] * 2 d_kv_lens = [2048] * 128 # --- Prefill plan 输入 --- p_seq_lens_blocks = torch.ceil( torch.tensor(p_kv_lens, dtype=torch.int32) / page_block_size ).int() p_q_indptr = torch.cat( [torch.tensor([0]), torch.cumsum(torch.tensor(p_qo_lens), 0)], dim=0 ).int() p_kv_indptr = torch.cat( [torch.tensor([0]), torch.cumsum(p_seq_lens_blocks, 0)], dim=0 ).int() kv_indices_p = torch.arange(0, p_kv_indptr[-1], device=device, dtype=torch.int32) last_page_len_p = (p_seq_lens_blocks - 1) % page_block_size + 1 # --- Decode plan 输入 --- d_seq_lens_blocks = torch.ceil( torch.tensor(d_kv_lens, dtype=torch.int32) / page_block_size ).int() d_q_indptr = torch.cat( [torch.tensor([0]), torch.cumsum(torch.tensor(d_qo_lens), 0)], dim=0 ).int() d_kv_indptr = torch.cat( [torch.tensor([0]), torch.cumsum(d_seq_lens_blocks, 0)], dim=0 ).int() kv_indices_d = torch.arange(0, d_kv_indptr[-1], device=device, dtype=torch.int32) last_page_len_d = (d_seq_lens_blocks - 1) % page_block_size + 1 # 创建 batch prefill 与 decode 的辅助数据结构 wrapper.plan( p_q_indptr.to(device), p_kv_indptr.to(device), kv_indices_p.to(device), last_page_len_p, d_q_indptr.to(device), d_kv_indptr.to(device), kv_indices_d.to(device), last_page_len_d, num_qo_heads=num_qo_heads, num_kv_heads=num_kv_heads, head_dim=head_dim, page_size=page_block_size, q_data_type=torch.bfloat16, kv_data_type=torch.bfloat16, ) # Prefill 输入张量(批量 paged,形状为 [总页数, 2, page_size, num_kv_heads, head_dim]) q_p = torch.rand(p_q_indptr[-1].item(), num_qo_heads, head_dim).to(device, dtype=torch.bfloat16) kv_p = torch.randn(p_kv_indptr[-1], 2, page_block_size, num_kv_heads, head_dim).to( device, dtype=torch.bfloat16 ).unbind(1) # Decode 输入张量 q_d = torch.rand(d_q_indptr[-1].item(), num_qo_heads, head_dim).to(device, dtype=torch.bfloat16) kv_d = torch.randn(d_kv_indptr[-1], 2, page_block_size, num_kv_heads, head_dim).to( device, dtype=torch.bfloat16 ).unbind(1) for i in range(num_layers): o_p_batch, o_d_batch = wrapper.run( q_p, kv_p, q_d, kv_d, causal_p=causal, ) print(o_p_batch.shape, o_d_batch.shape) # torch.Size([4096, 64, 128]) torch.Size([128, 64, 128])与单请求版本的结构性差异
BatchPODWithPagedKVCacheWrapper的构造函数更简洁(flashinfer/pod.pyL833-L895),但内部有几个值得注意的实现细节:
- workspace 一分为二:
float_workspace_buffer在__init__中被torch.chunk(..., 2, dim=0)均分为_float_workspace_buffer_p与_float_workspace_buffer_d(L858-L862),prefill 与 decode 各用一半;int workspace 则各自独立分配 8MB(L864-L881)。 - SM 感知调度缓冲区:构造时分配
_sm_aware_sched,其大小是multi_processor_count + 2(L883-L887),用于在 kernel 内部做基于 SM 数量的任务分配,这是批量 POD 在同一个 launch 内协调 prefill/decode 两套工作量的关键设施。 - 不暴露 CUDAGraph:
BatchPODWithPagedKVCacheWrapper当前固定_use_cuda_graph = False,is_cuda_graph_enabled恒为False(L895-L899)。
plan/run 生命周期:辅助数据结构如何被复用
两个 Wrapper 都遵循plan一次、run多次的生命周期模型。plan与run的语义在 docstring 中有明确约定:
The
planmethod should be called before anyrunorrun_return_lsecalls, auxiliary data structures will be created during this call and cached for multiple run calls.
即:plan必须在任何run调用之前执行;plan期间创建辅助数据结构并被缓存,供后续多次run复用。这正是示例中"逐 Transformer 层复用同一套 plan 结果"能够成立的原因。同时文档明确两条限制:
num_qo_heads必须是num_kv_heads的倍数;若二者不相等则自动走 grouped query attention(GQA);plan不能在 CUDAGraph 或torch.compile中使用。
plan 参数总表(以 PODWithPagedKVCacheWrapper 为例)
plan的完整签名见 flashinfer/pod.pyL268-L287,参数语义如下:
| 参数 | 形状 / 类型 | 说明 |
|---|---|---|
indptr | [batch_size + 1] | paged kv-cache 的 indptr |
indices | [qo_indptr[-1]] | paged kv-cache 的页索引 |
last_page_len | [batch_size] | 每条请求最后一页的有效条目数,范围1..page_size |
num_qo_heads | int | query/output 头数 |
num_kv_heads | int | key/value 头数 |
head_dim | int | 头维度 |
page_size | int | paged kv-cache 页大小 |
pos_encoding_mode | str | 位置编码:"NONE"/"ROPE_LLAMA"(LLaMA 式旋转嵌入)/"ALIBI",默认"NONE" |
window_left | int | 注意力窗口左(含)边界;-1表示窗口为全序列长度,默认-1 |
q_data_type | Optional[Union[str, torch.dtype]] | query 张量数据类型,默认"float16" |
kv_data_type | Optional[Union[str, torch.dtype]] | key/value 数据类型,None时取q_data_type |
data_type | Optional[Union[str, torch.dtype]] | 同时设定 q/kv 类型;已废弃,请改用q_data_type/kv_data_type |
sm_scale | Optional[float] | softmax 缩放;None时默认1 / sqrt(head_dim),会在 wrapper 上缓存并在run时复用 |
rope_scale | Optional[float] | RoPE 插值缩放,仅在pos_encoding_mode != "NONE"时生效,默认1.0 |
rope_theta | Optional[float] | RoPE 频率基值,仅在pos_encoding_mode != "NONE"时生效,默认1e4 |
non_blocking | bool | 是否异步拷贝输入张量到设备,默认True |
BatchPODWithPagedKVCacheWrapper.plan在此基础上把输入扩展为 prefill / decode 两套qo_indptr、kv_indptr、kv_indices、last_page_len(共 8 个张量,见 flashinfer/pod.pyL901-L925),其余公共参数(heads、head_dim、page_size、pos_encoding、window、dtype、sm_scale、rope 等)语义一致。
plan 内部做了什么
从源码可以还原plan的核心流程(L385-L440):
- 将 indptr / last_page_len 拷回 host 端(
L385-L386),通过get_seq_lens(indptr_host, last_page_len_host, page_size)计算每条序列的实际 KV 长度(L401); - 规范化 dtype:
data_type若给定会回填q_data_type/kv_data_type,再经canonicalize_torch_dtype归一化(L388-L397); - 通过
get_batch_prefill_module("fa2", ...)获取底层 prefill 模块(L402-L417),其中以PosEncodingMode[pos_encoding_mode].value、window_left != -1(是否启用滑窗)、logits_soft_cap > 0(是否启用 logits soft cap,当前固定关闭)等作为模板参数; - 调用
self._cached_module.plan(...)完成 kernel 的plan,传入 float/int/pin-memory int 三个 workspace、indptr、KV 长度数组、batch size、head 配置、causal=False、window_left、fixed_split_size=-1、disable_split_kv=False、num_colocated_ctas=0、uniform_q_len=0(L419-L440),得到的_plan_info即为run阶段要复用的辅助数据结构。
Batch 版本 plan 的 colocated-CTA 调度细节
BatchPODWithPagedKVCacheWrapper.plan中有一个关键的调度决策(L1086-L1111):
num_colocated_ctas = self._plan_info_d[0] # Splitting small prefill causes unnecessary bandwidth contention if total_num_rows_p > 1536: num_colocated_ctas = 0 self._plan_info_p = self._cached_module.plan(..., num_colocated_ctas, 0)它先从 decode 侧 plan 结果中读出num_colocated_ctas,再决定 prefill 侧是否与 decode 共置 CTA;当 prefill 的总行数total_num_rows_p > 1536时强制关闭共置,因为拆分小规模 prefill 会造成不必要的带宽竞争——这是源码注释中明确写出的工程权衡。
run 参数详解:一次调用产出两个输出
PODWithPagedKVCacheWrapper.run的签名与完整语义见 flashinfer/pod.pyL452-L489:
def run( self, q_p: torch.Tensor, # [qo_len, num_qo_heads, head_dim] k_p: torch.Tensor, # 布局与 kv_layout_p 一致 v_p: torch.Tensor, # 布局与 kv_layout_p 一致 q_d: torch.Tensor, # [batch_size, num_qo_heads, head_dim] paged_kv_cache_d: Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]], # Prefill 选项 custom_mask_p=None, packed_custom_mask_p=None, causal_p=False, kv_layout_p="NHD", pos_encoding_mode_p="NONE", sm_scale_p=None, window_left_p=-1, rope_scale_p=None, rope_theta_p=None, return_lse_p=False, # Decode 选项 custom_mask_d=None, packed_custom_mask_d=None, causal_d=False, kv_layout_d="NHD", pos_encoding_mode_d="NONE", sm_scale_d=None, window_left_d=-1, rope_scale_d=None, rope_theta_d=None, q_scale=None, k_scale=None, v_scale=None, return_lse_d=False, use_fp16_qk_reduction=False, enable_pdl=None, *args, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]返回值恒为(out_p, out_d):prefill 输出形状[qo_len, num_qo_heads, head_dim],decode 输出形状[batch_size, num_qo_heads, head_dim]。
Prefill 侧参数
q_p/k_p/v_p:单请求 prefill 的 query/key/value 张量,形状如上表;custom_mask_p/packed_custom_mask_p:可选的自定义 mask(稠密 / 位打包两种形式)。当custom_mask_p非空而packed_custom_mask_p为空时,源码L601-L605会调用packbits(custom_mask_p.contiguous().view(-1), bitorder="little")自动打包;mask 布局可参考flashinfer.single_prefill_with_kv_cache;causal_p:是否对 prefill 侧施加因果 mask,默认False。注意 mask 优先级:有自定义 mask 时走MaskMode.CUSTOM,否则按causal_p走CAUSAL或NON_CAUSAL(L607-L613);kv_layout_p、pos_encoding_mode_p、sm_scale_p(默认1/sqrt(head_dim))、window_left_p、rope_scale_p(默认 1.0)、rope_theta_p(默认 1e4):prefill 侧独立生效的策略参数;return_lse_p:若为True,会为 prefill kernel 分配一个[qo_len, num_qo_heads]的 float32 LSE 缓冲区并交给 kernel 填充,但当前 Python 层并未把 LSE 透出返回值(docstring 明确标注这是当前 wrapper 的已知限制:kernel API 已具备,Python wrapper 尚未打通返回链路)。
Decode 侧参数与"plan 覆盖"行为
decode 侧存在一组当前被忽略的参数,这是使用 POD 时最容易踩的坑(flashinfer/pod.pyL539-L556的 docstring 与L629-L634的代码相互印证):
kv_layout_d:被忽略——decode 的 KV 布局永远取构造函数的kv_layout;pos_encoding_mode_d:被忽略——被plan缓存值self._pos_encoding_mode覆盖;sm_scale_d:被忽略——被plan缓存值self._sm_scale覆盖(其默认仍是1/sqrt(head_dim));window_left_d:被忽略——被plan缓存值self._window_left覆盖;rope_scale_d/rope_theta_d:被忽略——被plan缓存值覆盖。
换句话说,decode 侧的策略参数必须在plan阶段设置,run阶段传入的 decode 策略参数只是为与 prefill 侧签名对称而保留。源码L629-L634直接执行pos_encoding_mode_d = self._pos_encoding_mode等赋值,随后才校验_check_pos_encoding_mode。
decode 侧真正在run阶段生效的参数:
q_scale/k_scale/v_scale:FP8 校准缩放。q_scale、k_scale会以乘法方式折入 decode 的sm_scale(L642-L645),v_scale则在 kernel 返回后在 Python 层直接乘到out_d上(L722-L723);use_fp16_qk_reduction:是否用 FP16 累加 QK(精度更低、吞吐更高),默认False;enable_pdl:Programmatic Dependent Launch 开关。None时自动用device_support_pdl(q_p.device)探测当前设备是否支持(L583-L584);return_lse_d:与return_lse_p相同的分配但不返回的已知限制。
BatchPOD 的 run 差异
BatchPODWithPagedKVCacheWrapper.run(flashinfer/pod.pyL1122-L1145)把 prefill 输入也改为 paged 形式:q_p形状为[qo_indptr_p[-1], num_qo_heads, head_dim],paged_kv_cache_p与paged_kv_cache_d均为分页 KV cache((k_cache, v_cache)二元组或拼接张量),prefill 与 decode 的所有策略参数均取自plan缓存。其返回值有两种形态:
return_lse=False(默认):返回(out_p, out_d);return_lse=True:返回((out_p, lse_p), (out_d, lse_d))——注意与单请求版本不同,批量版本会把 LSE 真正返回给调用者。
run末尾(L1347-L1350)同样会把v_scale应用到 decode 输出。
底层原理:JIT 模块生成与一次 Launch 的实现
Python 层 JIT 入口
flashinfer/pod.py通过functools.cache缓存两个模块获取函数(L49-L58):
@functools.cache def get_pod_module(*args): module = gen_pod_module(*args).build_and_load() return SimpleNamespace(run_tensor=module.pod_with_kv_cache_tensor) @functools.cache def get_batch_pod_module(*args): module = gen_batch_pod_module(*args).build_and_load() return SimpleNamespace(run_tensor=module.batch_pod_with_kv_cache_tensor)模板参数由run时实际张量的 dtype、head_dim、两侧的PosEncodingMode值、use_sliding_window、use_logits_soft_cap、use_fp16_qk_reduction以及 indptr 的 dtype 决定(L658-L678/L1280-L1295),因此同一进程内不同配置会生成并缓存不同的 JIT 模块。
JIT 模块的模板实例化
gen_pod_module/gen_batch_pod_module定义在 flashinfer/jit/attention/modules.py(L619-L681/L684-L718)。它们会:
- 依据全部模板参数生成唯一 URI(
get_pod_uri),用于 JIT 缓存定位; - 声明额外张量
maybe_custom_mask(uint8_t)与maybe_alibi_slopes(float),以及额外标量logits_soft_cap、sm_scale、rope_rcp_scale、rope_rcp_theta(注意 kernel 侧接收的是 rope 参数的倒数); - 分别为 prefill 侧与 decode 侧实例化
DefaultAttention<use_custom_mask, use_sliding_window, use_logits_soft_cap, use_pos_encoding>变体,并包含flashinfer/attention/variants.cuh。
批处理版本只是在 URI 前加"batch_"前缀,其余模板逻辑相同。对应的 kernel 绑定实现在 csrc/pod.cu 与 csrc/batch_pod.cu,JIT 编译配置模板为 csrc/pod_customize_config.jinja、csrc/pod_kernel_inst.jinja 与 csrc/pod_jit_binding.cu。
run 的两次"设置"与一次 launch
以单请求版本run为例(L583-L720),调用前会做两套独立的参数准备:
- prefill 侧:分配 32MB 临时缓冲
tmp_p(_get_cache_buf("pod_with_kv_cache_tmp", 32 * 1024 * 1024, ...))、填充 mask/alibi/logits_soft_cap/sm_scale/rope 参数,out_p = torch.empty_like(q_p); - decode 侧:
_unpack_paged_kv_cache(paged_kv_cache_d, self._kv_layout)解包出 k/v,校验 q/kv dtype 与plan缓存一致(_check_cached_qkv_data_type),再套用 plan 缓存的策略参数; - 最终一次性调用
module_getter.run_tensor(...)(L679-L720),把 prefill 的(q_p, k_p, v_p, tmp_p, out_p, lse_p, mask, layout, window_left, sm_scale, rope 倒数, ...)与 decode 的(workspace, plan_info, q_d, k_cache_d, v_cache_d, indptr/indices/last_page_len, out_d, lse_d, ...)全部打包进同一次 kernel 调用,末尾传入enable_pdl。
PDL:让两个任务真正并发
enable_pdl默认通过device_support_pdl(q_p.device)自动探测。PDL(Programmatic Dependent Launch)允许 kernel 在程序内部依据运行期条件去启动依赖 kernel,这正是 POD 能在一次 launch 内让 prefill 与 decode 并发/协作执行的机制之一。在 csrc/pod.cu 中,enable_pdl被直接透传给底层 CUDA 内核入口(L56、L265附近),由底层决定采用 PDL 还是传统协作调度路径。
正确性验证:测试如何对标参考实现
仓库中的测试为 POD 提供了完整的正确性对标,可作为理解语义的辅助材料:
- tests/utils/test_pod_kernels.py:核心 kernel 测试。
test_pod_with_paged_kv_cache(L75-L107)对 prefill 长度(127/12288)、decode 批量(1/17/127)、KV 长度、page_size(1/16)、GQA 头数(8/32 头)等参数做笛卡尔积组合;prefill 参考结果来自flashinfer.prefill.single_prefill_with_kv_cache(L121-L127),decode 参考结果来自flashinfer.decode.BatchDecodeWithPagedKVCacheWrapper(L172-L179附近)——即以独立 prefill kernel 与独立 decode kernel 的输出作为融合 kernel 的 ground truth,逐元素对比验证融合后数值一致性; - tests/trace/test_pod_with_paged_kv_cache_run_reference_correctness.py 与 tests/trace/test_batch_pod_run_reference_correctness.py:验证
run在 trace 场景下的参考正确性,与 flashinfer/trace/templates/attention.py 中定义的pod_with_paged_kv_cache_run_trace、batch_pod_with_paged_kv_cache_run_trace追踪模板配套使用(run方法上的@flashinfer_api(trace=...)装饰器即是 trace 采集入口)。
使用建议与已知限制
综合 docstring 与源码,使用flashinfer.pod时应注意:
- decode 策略参数只在
plan生效:pos_encoding、window、sm_scale、rope 等 decode 侧配置必须在plan阶段传入,run阶段的同名参数会被 plan 缓存覆盖; - workspace 推荐 128MB,且 device 必须与输入一致;批量版本的 workspace 会被均分给 prefill/decode 两侧;
- CUDAGraph 模式仅单请求版本支持,且 batch_size 固定、indices 上限受预留 buffer 约束;批量版本当前不支持 CUDAGraph;
plan不能在 CUDAGraph 或torch.compile内调用,且num_qo_heads须为num_kv_heads的整数倍(否则走 GQA);- LSE 支持状态:单请求版本
return_lse_p/return_lse_d只分配缓冲区、暂不返回;批量版本return_lse=True会真正返回((out_p, lse_p), (out_d, lse_d)); - logits soft cap 当前不支持:源码固定
logits_soft_cap = 0.0(L347、L994),不要依赖该功能; - 模板缓存:POD kernel 经 JIT 按 dtype/head_dim/策略组合实例化并缓存,首次调用某组配置会有编译开销,可配合 FlashInfer 的 JIT 缓存机制复用。
POD 是 FlashInfer 面向 continuous batching 服务端优化的重要入口:它以"一次 kernel launch 同时服务 prefill 与 decode"的方式,把 chunked prefill 与 decode 重叠场景中的内核启动与调度开销压缩到最低。配合 flashinfer/page.py 的 paged kv-cache 数据结构与 flashinfer/trace/templates/attention.py 的 trace 能力,PODWithPagedKVCacheWrapper与BatchPODWithPagedKVCacheWrapper分别覆盖单请求与批量两种融合形态,可直接接入现有推理框架的解码主循环。
- 大模型
- 深度学习
- 算子库
- 后端
- 高性能计算
【免费下载链接】flashinfer
FlashInfer: Kernel Library for LLM Serving
相关推荐
Ray Serve LLM 分布式服务模式架构解析:数据并行注意力与 Prefill-Decode 解耦
Ray Serve LLM 分布式服务模式架构解析:数据并行注意力与 Prefill Decode 解耦 导读:本文基于 Ray Serve LLM 的架构文档
人工智能分布式训练强化学习任务调度模型推理服务后端MooncakeConnector 与 vLLM 解耦式 Prefill-Decode(PD)部署实战指南
MooncakeConnector 与 vLLM 解耦式 Prefill Decode(PD)部署实战指南 导读 本文是 Mooncake 项目中 vLLM 解
人工智能大模型模型推理服务后端exo 如何配置并运行 prefill/decode 分离基准测试?instance-links 与 prefill-decode.toml 实战
exo 如何配置并运行 prefill/decode 分离基准测试?instance links 与 prefill decode.toml 实战 如果你在 e
人工智能大模型本地部署模型推理服务
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考