FlashInfer 注意力核

flashinfer.decode

单请求解码

single_decode_with_kv_cache()

使用 KV 缓存进行单请求解码,返回注意力输出。

批量解码

cudnn_batch_decode_with_kv_cache(q, k_cache, ...)

使用 cuDNN 执行带有分页 KV 缓存的批量解码注意力。

trtllm_batch_decode_with_kv_cache(query, ...)

class flashinfer.decode.BatchDecodeWithPagedKVCacheWrapper(float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, use_tensor_cores: bool = False, paged_kv_indptr_buffer: Tensor | None = None, paged_kv_indices_buffer: Tensor | None = None, paged_kv_last_page_len_buffer: Tensor | None = None, backend: str = 'auto', jit_args: List[Any] | None = None)

用于批量请求的分页 kv-cache 注意力解码的包装类(首次由 vLLM 提出)。

请查看 我们的教程 以了解页面表布局。

示例

>>> 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
>>> # allocate 128MB workspace buffer
>>> workspace_buffer = torch.zeros(128 * 1024 * 1024, dtype=torch.uint8, device="cuda:0")
>>> decode_wrapper = flashinfer.BatchDecodeWithPagedKVCacheWrapper(
...     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)
... ]
>>> # create auxiliary data structures for batch decode attention
>>> 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]
...     # compute batch decode attention, reuse auxiliary data structures for all layers
...     o = decode_wrapper.run(q, kv_cache)
...     outputs.append(o)
...
>>> outputs[0].shape
torch.Size([7, 64, 128])

注意

为了加速计算,FlashInfer 的批量解码注意力会创建一些辅助数据结构,这些数据结构可以在多个批量解码注意力调用之间重用(例如,不同的 Transformer 层)。这个包装类管理这些数据结构的生命周期。

__init__(float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, use_tensor_cores: bool = False, paged_kv_indptr_buffer: Tensor | None = None, paged_kv_indices_buffer: Tensor | None = None, paged_kv_last_page_len_buffer: Tensor | None = None, backend: str = 'auto', jit_args: List[Any] | None = None) None

BatchDecodeWithPagedKVCacheWrapper 的构造函数。

参数:
  • float_workspace_buffer (torch.Tensor. 必须在首次使用时初始化为 0。) – 用于在分割 k 算法中存储中间注意力结果的用户保留浮点工作区缓冲区。推荐大小为 128MB,工作区缓冲区的设备应与输入张量的设备相同。

  • kv_layout (str) – 输入 k/v 张量的布局,可以是 NHDHND

  • use_cuda_graph (bool) – 是否为批量解码注意力启用 CUDAGraph,如果启用,辅助数据结构将存储为提供的缓冲区。当启用 CUDAGraph 时,此包装器的生命周期内 batch_size 不能更改。

  • use_tensor_cores (bool) – 是否为计算使用张量核心。对于分组查询注意力中的大型组大小,速度会更快。默认值为 False

  • paged_kv_indptr_buffer (Optional[torch.Tensor]) – 用户在 GPU 上保留的缓冲区,用于存储分页 kv 缓存的 indptr,缓冲区的大小应为 [batch_size + 1]。仅当 use_cuda_graphTrue 时才需要。

  • paged_kv_indices_buffer (Optional[torch.Tensor]) – 用户在 GPU 上保留的缓冲区,用于存储分页 kv 缓存的页面索引,应足够大以存储生命周期内页索引的最大数量 (max_num_pages)。仅当 use_cuda_graphTrue 时才需要。

  • paged_kv_last_page_len_buffer (Optional[torch.Tensor]) – 用户在 GPU 上保留的缓冲区,用于存储最后一页中的条目数,缓冲区的大小应为 [batch_size]。仅当 use_cuda_graphTrue 时才需要。

  • backend (str) – 实现后端,可以是 auto/fa2/fa3trtllm-gen。默认值为 auto。如果设置为 auto,则包装器将根据设备架构和内核可用性自动选择后端。

  • jit_args (Optional[List[Any]]) – 如果提供,包装器将使用提供的参数创建 JIT 模块,否则,包装器将使用默认注意力实现。

plan(indptr: Tensor, indices: Tensor, last_page_len: Tensor, num_qo_heads: int, num_kv_heads: int, head_dim: int, page_size: int, pos_encoding_mode: str = 'NONE', window_left: int = -1, logits_soft_cap: float | None = None, q_data_type: str | dtype | None = 'float16', kv_data_type: str | dtype | None = None, o_data_type: str | dtype | None = None, data_type: str | dtype | None = None, sm_scale: float | None = None, rope_scale: float | None = None, rope_theta: float | None = None, non_blocking: bool = True, block_tables: Tensor | None = None, seq_lens: Tensor | None = None, fixed_split_size: int | None = None, disable_split_kv: bool = False) None

规划给定问题规范的批量解码。

参数:
  • indptr (torch.Tensor) – 分页 kv 缓存的 indptr,形状:[batch_size + 1],dtype:torch.int32

  • indices (torch.Tensor) – 分页 kv 缓存的页面索引,形状:[kv_indptr[-1]],dtype:torch.int32

  • last_page_len (torch.Tensor) – 分页 kv 缓存中每个请求的最后一页中的条目数,形状:[batch_size],dtype:torch.int32

  • num_qo_heads (int) – 查询/输出头的数量

  • num_kv_heads (int) – key/value 头的数量

  • head_dim (int) – 头部的维度

  • page_size (int) – 分页 kv 缓存的页面大小

  • pos_encoding_mode (str) – 注意力内核内部应用的 positional encoding,可以是 NONE/ROPE_LLAMA (LLAMA 风格的旋转嵌入) /ALIBI。默认为 NONE

  • window_left (int) – 注意力窗口的左(包含)窗口大小,设置为 -1 时,窗口大小将设置为序列的完整长度。默认值为 -1

  • logits_soft_cap (Optional[float]) – 注意力 logits 的软上限值(用于 Gemini、Grok 和 Gemma-2 等),如果未提供,将设置为 0。如果大于 0,logits 将根据公式进行截断:\(\texttt{logits_soft_cap} \times \mathrm{tanh}(x / \texttt{logits_soft_cap})\),其中 \(x\) 是输入 logits。

  • q_data_type (Optional[Union[str, torch.dtype]]) – 查询张量的的数据类型,默认为 torch.float16。

  • kv_data_type (Optional[Union[str, torch.dtype]]) – key/value 张量的数据类型。如果为 None,则设置为 q_data_type。默认为 None

  • o_data_type (Optional[Union[str, torch.dtype]]) – 输出张量的数据类型。如果为 None,则设置为 q_data_type。对于 FP8 输入,通常应设置为 torch.float16 或 torch.bfloat16。

  • data_type (Optional[Union[str, torch.dtype]]) – 查询和 key/value 张量的通用数据类型。默认为 torch.float16。data_type 已弃用,请改用 q_data_type 和 kv_data_type。

  • non_blocking (bool) – 是否异步将输入张量复制到设备,默认为 True

  • seq_lens (Optional[torch.Tensor]) – 一个 uint32 1D 张量,指示每个 prompt 的 kv 序列长度。形状:[batch_size]

  • block_tables (Optional[torch.Tensor]) – 一个 uint32 2D 张量,指示每个 prompt 的 block table。形状:[batch_size, max_num_blocks_per_seq]

  • fixed_split_size (Optional[int],) – FA2 split-kv 解码的固定分割大小,以页面为单位。目前仅受 tensor core 解码支持。建议设置为工作负载的平均序列长度。启用后,将导致 merge_states 内核中确定性的 softmax 分数减少,因此输出与批量大小无关。请参阅 https://thinkingmachines.ai/blog/defeating-nondeterminism-in-llm-inference/ 请注意,即使 bs 固定,kv 序列长度也可能发生变化,从而导致启动的 CTA 数量不同,因此与 CUDA 图的兼容性不能保证。

  • disable_split_kv (bool,) – 是否禁用 split-kv 以在 CUDA Graph 中实现确定性,默认为 False

注意

在调用任何 run()run_return_lse() 调用之前,应调用 plan() 方法,辅助数据结构将在调用期间创建并缓存以供多次运行调用使用。

num_qo_heads 必须是 num_kv_heads 的倍数。如果 num_qo_heads 不等于 num_kv_heads,该函数将使用 分组查询注意力

在 Cuda Graph 或 torch.compile 中无法使用 plan() 方法。

reset_workspace_buffer(float_workspace_buffer: Tensor, int_workspace_buffer: Tensor) None

重置工作区缓冲区。

参数:
  • float_workspace_buffer (torch.Tensor) – 新的 float 工作区缓冲区,该缓冲区的设备应与输入张量的设备相同。

  • int_workspace_buffer (torch.Tensor) – 新的 int 工作区缓冲区,该缓冲区的设备应与输入张量的设备相同。

run(q: Tensor, paged_kv_cache: Tensor | Tuple[Tensor, Tensor], *args, q_scale: float | None = None, k_scale: float | None = None, v_scale: float | None = None, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[False] = False, enable_pdl: bool | None = None, window_left: int | None = None) Tensor
run(q: Tensor, paged_kv_cache: Tensor | Tuple[Tensor, Tensor], *args, q_scale: float | None = None, k_scale: float | None = None, v_scale: float | None = None, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[True] = True, enable_pdl: bool | None = None, window_left: int | None = None) Tuple[Tensor, Tensor]

计算查询和分页 KV 缓存之间的批量解码注意力。

参数:
  • q (torch.Tensor) – 查询张量,形状:[batch_size, num_qo_heads, head_dim]

  • paged_kv_cache (Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]) –

    存储的分页 KV 缓存,作为张量元组或单个张量

    • 一个元组 (k_cache, v_cache),包含 4D 张量,每个张量的形状为:[max_num_pages, page_size, num_kv_heads, head_dim],如果 kv_layoutNHD,以及 [max_num_pages, num_kv_heads, page_size, head_dim],如果 kv_layoutHND

    • 一个 5D 张量,形状为:[max_num_pages, 2, page_size, num_kv_heads, head_dim],如果 kv_layoutNHD,以及 [max_num_pages, 2, num_kv_heads, page_size, head_dim],如果 kv_layoutHND。其中 paged_kv_cache[:, 0] 是 key 缓存,paged_kv_cache[:, 1] 是 value 缓存。

  • *args – 自定义内核的附加参数。

  • q_scale (Optional[float]) – 查询的校准比例,用于 fp8 输入,如果未提供,则设置为 1.0

  • k_scale (Optional[float]) – key 的校准比例,用于 fp8 输入,如果未提供,则设置为 1.0

  • v_scale (Optional[float]) – value 的校准比例,用于 fp8 输入,如果未提供,则设置为 1.0

  • out (Optional[torch.Tensor]) – 输出张量,如果未提供,则会在内部分配。

  • lse (Optional[torch.Tensor]) – 注意力 logits 的对数和指数,如果未提供,则会在内部分配。

  • return_lse (bool) – 是否返回注意力分数的对数和指数,默认为 False

  • enable_pdl (bool) – 是否启用程序依赖启动 (PDL)。请参阅 https://docs.nvda.net.cn/cuda/cuda-c-programming-guide/#programmatic-dependent-launch-and-synchronization,仅支持 >= sm90,并且当前仅支持 FA2 和 CUDA 核心解码。

  • q_len_per_req (int) – 每个请求的查询 token 数量,如果未提供,则设置为 1

返回值:

如果 return_lseFalse,则注意力输出,形状:[batch_size, num_qo_heads, head_dim]。如果 return_lseTrue,则为两个张量的元组

  • 注意力输出,形状:[batch_size, num_qo_heads, head_dim]

  • 注意力分数的对数和指数,形状:[batch_size, num_qo_heads]

返回值类型:

Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]

class flashinfer.decode.CUDAGraphBatchDecodeWithPagedKVCacheWrapper(workspace_buffer: Tensor, indptr_buffer: Tensor, indices_buffer: Tensor, last_page_len_buffer: Tensor, kv_layout: str = 'NHD', use_tensor_cores: bool = False)

CUDAGraph兼容的解码注意力包装器,带有分页kv缓存(首次在vLLM中提出),用于批量请求。

请注意,此包装器可能不如BatchDecodeWithPagedKVCacheWrapper高效,因为我们不会为了适应CUDAGraph的要求而为不同的批大小/序列长度/等分派到不同的内核。

请查看 我们的教程 以了解页面表布局。

注意

plan()方法无法被CUDAGraph捕获。

__init__(workspace_buffer: Tensor, indptr_buffer: Tensor, indices_buffer: Tensor, last_page_len_buffer: Tensor, kv_layout: str = 'NHD', use_tensor_cores: bool = False) None

BatchDecodeWithPagedKVCacheWrapper 的构造函数。

参数:
  • workspace_buffer (torch.Tensor) – 用户保留的GPU上的工作区缓冲区,用于存储辅助数据结构,建议大小为128MB,工作区缓冲区的设备应与输入张量的设备相同。

  • indptr_buffer (torch.Tensor) – 用户保留的GPU上的缓冲区,用于存储分页kv缓存的indptr,应足够大以存储最大批处理大小期间的indptr ([max_batch_size + 1]) 在此包装器的生命周期内。

  • indices_buffer (torch.Tensor) – 用户保留的GPU上的缓冲区,用于存储分页kv缓存的页面索引,应足够大以存储最大页面索引数 (max_num_pages) 在此包装器的生命周期内。

  • last_page_len_buffer (torch.Tensor) – 用户保留的GPU上的缓冲区,用于存储最后一页中的条目数,应足够大以存储最大批处理大小 ([max_batch_size]) 在此包装器的生命周期内。

  • use_tensor_cores (bool) – 是否为计算使用张量核心。对于分组查询注意力中的大型组大小,速度会更快。默认值为 False

  • kv_layout (str) – 输入 k/v 张量的布局,可以是 NHDHND

XQA

xqa(q, k_cache, v_cache, page_table, ...[, ...])

使用XQA内核应用带有分页KV缓存的注意力。 :param q: 查询张量,形状为 [batch_size, beam_width, num_q_heads, head_dim] 如果不使用推测解码,或者 [batch_size, beam_width, q_seq_len, num_q_heads, head_dim] 如果使用推测解码。 q_seq_len 是推测解码令牌的数量。 数据类型应为torch.float16或torch.bfloat16。 现在仅支持beam_width 1。 :type q: torch.Tensor :param k_cache: 分页K缓存张量,形状为 [num_pages, page_size, num_kv_heads, head_dim] 如果 kv_layoutNHD,或者 [num_pages, num_kv_heads, page_size, head_dim] 如果 kv_layoutHND。 数据类型应与查询张量匹配或为torch.float8_e4m3fn,在这种情况下xqa将运行fp8计算。 应与v_cache的数据类型相同。 :type k_cache: torch.Tensor :param v_cache: 分页V缓存张量,形状为 [num_pages, page_size, num_kv_heads, head_dim] 如果 kv_layoutNHD,或者 [num_pages, num_kv_heads, page_size, head_dim] 如果 kv_layoutHND。 数据类型应与查询张量匹配或为torch.float8_e4m3fn,在这种情况下xqa将运行fp8计算。 应与k_cache的数据类型相同。 :type v_cache: torch.Tensor :param page_table: 页面表张量,形状为 batch_size, nb_pages_per_seq。 数据类型应为torch.int32。 K和V共享相同的表。 :type page_table: torch.Tensor :param seq_lens: 序列长度张量,形状为 [batch_size, beam_width]。 数据类型应为torch.uint32。 :type seq_lens: torch.Tensor :param output: 输出张量,形状与查询张量匹配。 数据类型应与查询张量或kv张量匹配。 此张量将被就地修改。 :type output: torch.Tensor :param workspace_buffer: 用于临时计算的工作区缓冲区。 数据类型应为torch.uint8。 :type workspace_buffer: torch.Tensor :param semaphores: 用于同步的信号量缓冲区。 数据类型应为torch.uint32。 :type semaphores: torch.Tensor :param num_kv_heads: 注意力机制中的键值头数。 :type num_kv_heads: int :param page_size: 分页KV缓存中每个页面的大小。 必须是[16, 32, 64, 128]之一。 :type page_size: int :param sinks: 注意力sink值,形状为 [num_kv_heads, head_group_ratio]。 数据类型应为torch.float32。 如果为None,则不使用注意力sink。 :type sinks: Optional[torch.Tensor], default=None :param q_scale: 查询张量的缩放因子。 :type q_scale: Union[float, torch.Tensor], default=1.0 :param kv_scale: KV缓存的缩放因子。 :type kv_scale: Union[float, torch.Tensor], default=1.0 :param sliding_win_size: 注意力的滑动窗口大小。 如果为0,则不使用滑动窗口。 :type sliding_win_size: int, default=0 :param kv_layout: KV缓存的布局。 可以是 NHDHND。 :type kv_layout: str, default="NHD" :param sm_count: 要使用的流式多处理器的数量。 如果为None,将从设备推断。 :type sm_count: Optional[int], default=None :param enable_pdl: 是否启用PDL(持久数据加载器)优化。 如果为None,如果硬件支持,则设置为True。 :type enable_pdl: Optional[bool], default=None :param rcp_out_scale: 输出比例因子的倒数。 :type rcp_out_scale: float, default=1.0 :param q_seq_len: 查询序列长度。 当> 1时,启用推测解码模式。 :type q_seq_len: int, default=1 :param mask: 推测解码模式(当 q_seq_len > 1 时)的因果注意力掩码。 形状: [batch_size, q_seq_len, mask_size_per_row] 其中 mask_size_per_row = ((q_seq_len + 31) // 32) * 2。 数据类型应为torch.uint16(位打包格式,对齐到32位)。 :type mask: Optional[torch.Tensor], default=None。

xqa_mla(q, k_cache, v_cache, page_table, ...)

使用 XQA MLA(多头潜在注意力)内核应用带有分页 KV 缓存的注意力。 :param q: 查询张量,形状为 [batch_size, beam_width, num_q_heads, head_dim]。数据类型应为 torch.float8_e4m3fn。目前仅支持 beam_width 1。 :type q: torch.Tensor :param k_cache: 分页 K 缓存张量,形状为 [total_num_cache_heads, head_dim]。数据类型应为 torch.float8_e4m3fn :type k_cache: torch.Tensor :param v_cache: 分页 V 缓存张量,形状为 [total_num_cache_heads, head_dim]。数据类型应为 torch.float8_e4m3fn :type v_cache: torch.Tensor :param page_table: 页表张量,形状为 batch_size, nb_pages_per_seq。数据类型应为 torch.int32。K 和 V 共享相同的表。 :type page_table: torch.Tensor :param seq_lens: 序列长度张量,形状为 [batch_size, beam_width]。数据类型应为 torch.uint32。 :type seq_lens: torch.Tensor :param output: 输出张量,形状为 [batch_size, beam_width, num_q_heads, head_dim]。数据类型应为 torch.bfloat16。此张量将被就地修改。 :type output: torch.Tensor :param workspace_buffer: 用于临时计算的工作空间缓冲区。数据类型应为 torch.uint8。 :type workspace_buffer: torch.Tensor :param semaphores: 用于同步的信号量缓冲区。数据类型应为 torch.uint32。 :type semaphores: torch.Tensor :param page_size: 分页 KV 缓存中每个页面的大小。必须是 [16, 32, 64, 128] 中的一个。 :type page_size: int :param q_scale: 查询张量的缩放因子。 :type q_scale: Union[float, torch.Tensor], default=1.0 :param kv_scale: KV 缓存的缩放因子。 :type kv_scale: Union[float, torch.Tensor], default=1.0 :param sm_count: 要使用的流式多处理器数量。如果为 None,将从设备推断。 :type sm_count: Optional[int], default=None :param enable_pdl: 是否启用 PDL(持久数据加载器)优化。如果为 None,如果硬件支持,则设置为 True。 :type enable_pdl: Optional[bool], default=None.

flashinfer.prefill

用于单请求和批量服务设置中预填充和追加注意力的注意力内核。

单请求预填充/追加注意力

single_prefill_with_kv_cache()

使用 KV 缓存进行单个请求的预填充/追加注意力,返回注意力输出。

single_prefill_with_kv_cache_return_lse(q, k, v)

使用 KV 缓存进行单个请求的预填充/追加注意力,返回注意力输出。

批量预填充/追加注意力

cudnn_batch_prefill_with_kv_cache(q, ...[, ...])

使用 cuDNN 执行带有分页 KV 缓存的批量预填充注意力。

trtllm_batch_context_with_kv_cache(query, ...)

class flashinfer.prefill.BatchPrefillWithPagedKVCacheWrapper(float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, qo_indptr_buf: Tensor | None = None, paged_kv_indptr_buf: Tensor | None = None, paged_kv_indices_buf: Tensor | None = None, paged_kv_last_page_len_buf: Tensor | None = None, custom_mask_buf: Tensor | None = None, mask_indptr_buf: Tensor | None = None, backend: str = 'auto', jit_args: List[Any] | None = None, jit_kwargs: Dict[str, Any] | None = None)

用于批量请求的带有分页 kv 缓存的预填充/追加注意力的包装类。

请查看 我们的教程 以了解页面表布局。

示例

>>> import torch
>>> import flashinfer
>>> num_layers = 32
>>> num_qo_heads = 64
>>> num_kv_heads = 16
>>> head_dim = 128
>>> max_num_pages = 128
>>> page_size = 16
>>> # allocate 128MB workspace buffer
>>> workspace_buffer = torch.zeros(128 * 1024 * 1024, dtype=torch.uint8, device="cuda:0")
>>> prefill_wrapper = flashinfer.BatchPrefillWithPagedKVCacheWrapper(
...     workspace_buffer, "NHD"
... )
>>> batch_size = 7
>>> nnz_qo = 100
>>> qo_indptr = torch.tensor(
...     [0, 33, 44, 55, 66, 77, 88, nnz_qo], dtype=torch.int32, device="cuda:0"
... )
>>> paged_kv_indices = torch.arange(max_num_pages).int().to("cuda:0")
>>> paged_kv_indptr = torch.tensor(
...     [0, 17, 29, 44, 48, 66, 100, 128], dtype=torch.int32, device="cuda:0"
... )
>>> # 1 <= paged_kv_last_page_len <= page_size
>>> paged_kv_last_page_len = torch.tensor(
...     [1, 7, 14, 4, 3, 1, 16], dtype=torch.int32, device="cuda:0"
... )
>>> q_at_layer = torch.randn(num_layers, nnz_qo, num_qo_heads, head_dim).half().to("cuda:0")
>>> kv_cache_at_layer = torch.randn(
...     num_layers, max_num_pages, 2, page_size, num_kv_heads, head_dim, dtype=torch.float16, device="cuda:0"
... )
>>> # create auxiliary data structures for batch prefill attention
>>> prefill_wrapper.plan(
...     qo_indptr,
...     paged_kv_indptr,
...     paged_kv_indices,
...     paged_kv_last_page_len,
...     num_qo_heads,
...     num_kv_heads,
...     head_dim,
...     page_size,
...     causal=True,
... )
>>> outputs = []
>>> for i in range(num_layers):
...     q = q_at_layer[i]
...     kv_cache = kv_cache_at_layer[i]
...     # compute batch prefill attention, reuse auxiliary data structures
...     o = prefill_wrapper.run(q, kv_cache)
...     outputs.append(o)
...
>>> outputs[0].shape
torch.Size([100, 64, 128])
>>>
>>> # below is another example of creating custom mask for batch prefill attention
>>> mask_arr = []
>>> qo_len = (qo_indptr[1:] - qo_indptr[:-1]).cpu().tolist()
>>> kv_len = (page_size * (paged_kv_indptr[1:] - paged_kv_indptr[:-1] - 1) + paged_kv_last_page_len).cpu().tolist()
>>> for i in range(batch_size):
...     mask_i = torch.tril(
...         torch.full((qo_len[i], kv_len[i]), True, device="cuda:0"),
...         diagonal=(kv_len[i] - qo_len[i]),
...     )
...     mask_arr.append(mask_i.flatten())
...
>>> mask = torch.cat(mask_arr, dim=0)
>>> prefill_wrapper.plan(
...     qo_indptr,
...     paged_kv_indptr,
...     paged_kv_indices,
...     paged_kv_last_page_len,
...     num_qo_heads,
...     num_kv_heads,
...     head_dim,
...     page_size,
...     custom_mask=mask,
... )
>>> for i in range(num_layers):
...     q = q_at_layer[i]
...     kv_cache = kv_cache_at_layer[i]
...     # compute batch prefill attention, reuse auxiliary data structures
...     o_custom = prefill_wrapper.run(q, kv_cache)
...     assert torch.allclose(o_custom, outputs[i], rtol=1e-3, atol=1e-3)
...

注意

为了加速计算,FlashInfer 的批量预填充/追加注意力算子会创建一些辅助数据结构,这些数据结构可以在多个预填充/追加注意力调用之间重用(例如,不同的 Transformer 层)。此包装类管理这些数据结构的生命周期。

__init__(float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, qo_indptr_buf: Tensor | None = None, paged_kv_indptr_buf: Tensor | None = None, paged_kv_indices_buf: Tensor | None = None, paged_kv_last_page_len_buf: Tensor | None = None, custom_mask_buf: Tensor | None = None, mask_indptr_buf: Tensor | None = None, backend: str = 'auto', jit_args: List[Any] | None = None, jit_kwargs: Dict[str, Any] | None = None) None

BatchPrefillWithPagedKVCacheWrapper 的构造函数。

参数:
  • **float_workspace_buffer** (torch.Tensor) – 用于在 split-k 算法中存储中间注意力结果的用户保留的工作空间缓冲区。推荐大小为 128MB,工作空间缓冲区的设备应与输入张量的设备相同。

  • kv_layout (str) – 输入 k/v 张量的布局,可以是 NHDHND

  • use_cuda_graph (bool) – 是否启用 CUDA 图捕获,用于预填充内核。如果启用,辅助数据结构将存储在提供的缓冲区中。当启用 CUDAGraph 时,此包装器的生命周期内 batch_size 不能更改。

  • qo_indptr_buf (Optional[torch.Tensor]) – 用户预留的缓冲区,用于存储 qo_indptr 数组,缓冲区的大小应为 [batch_size + 1]。只有当 use_cuda_graphTrue 时,此参数才有效。

  • paged_kv_indptr_buf (Optional[torch.Tensor]) – 用户预留的缓冲区,用于存储 paged_kv_indptr 数组,此缓冲区的大小应为 [batch_size + 1]。只有当 use_cuda_graphTrue 时,此参数才有效。

  • paged_kv_indices_buf (Optional[torch.Tensor]) – 用户预留的缓冲区,用于存储 paged_kv_indices 数组,应足够大以存储包装器生命周期内 paged_kv_indices 数组的最大可能大小。只有当 use_cuda_graphTrue 时,此参数才有效。

  • paged_kv_last_page_len_buf (Optional[torch.Tensor]) – 用户预留的缓冲区,用于存储 paged_kv_last_page_len 数组,缓冲区的大小应为 [batch_size]。只有当 use_cuda_graphTrue 时,此参数才有效。

  • custom_mask_buf (Optional[torch.Tensor]) – 用户预留的缓冲区,用于存储自定义掩码张量,应足够大以存储包装器生命周期内打包的自定义掩码张量的最大可能大小。只有当 use_cuda_graph 设置为 True 并且在注意力计算中使用自定义掩码时,此参数才有效。

  • mask_indptr_buf (Optional[torch.Tensor]) – 用户预留的缓冲区,用于存储 mask_indptr 数组,缓冲区的大小应为 [batch_size + 1]。只有当 use_cuda_graphTrue 并且在注意力计算中使用自定义掩码时,此参数才有效。

  • backend (str) – 实现后端,可以是 auto/fa2/fa3/cudnntrtllm-gen。默认为 auto。如果设置为 auto,则包装器将根据设备架构和内核可用性自动选择后端。

  • jit_args (Optional[List[Any]]) – 如果提供,包装器将使用提供的参数创建 JIT 模块,否则,包装器将使用默认注意力实现。

  • jit_kwargs (Optional[Dict[str, Any]]) – 创建 JIT 模块的关键字参数,默认为 None。

plan(qo_indptr: Tensor, paged_kv_indptr: Tensor, paged_kv_indices: Tensor, paged_kv_last_page_len: Tensor, num_qo_heads: int, num_kv_heads: int, head_dim_qk: int, page_size: int, head_dim_vo: int | None = None, custom_mask: Tensor | None = None, packed_custom_mask: Tensor | None = None, causal: bool = False, pos_encoding_mode: str = 'NONE', use_fp16_qk_reduction: bool = False, sm_scale: float | None = None, window_left: int = -1, logits_soft_cap: float | None = None, rope_scale: float | None = None, rope_theta: float | None = None, q_data_type: str | dtype = 'float16', kv_data_type: str | dtype | None = None, o_data_type: str | dtype | None = None, non_blocking: bool = True, prefix_len_ptr: Tensor | None = None, token_pos_in_items_ptr: Tensor | None = None, token_pos_in_items_len: int = 0, max_item_len_ptr: Tensor | None = None, seq_lens: Tensor | None = None, seq_lens_q: Tensor | None = None, block_tables: Tensor | None = None, max_token_per_sequence: int | None = None, max_sequence_kv: int | None = None, fixed_split_size: int | None = None, disable_split_kv: bool = False) None

规划给定问题规范的 Paged KV-Cache 上批量预填充/追加注意力。

参数:
  • qo_indptr (torch.Tensor) – 查询/输出张量的 indptr,形状:[batch_size + 1]

  • paged_kv_indptr (torch.Tensor) – 分页 kv-cache 的 indptr,形状:[batch_size + 1]

  • paged_kv_indices (torch.Tensor) – 分页 kv-cache 的页索引,形状:[paged_kv_indptr[-1]]

  • paged_kv_last_page_len (torch.Tensor) – 分页 kv-cache 中每个请求的最后一页中的条目数,形状:[batch_size]

  • num_qo_heads (int) – 查询/输出头的数量。

  • num_kv_heads (int) – 键/值头的数量。

  • head_dim_qk (int) – 查询/键头的维度。

  • page_size (int) – 分页 kv-cache 中每个页面的大小。

  • head_dim_vo (Optional[int]) – 值/输出头的维度,如果未提供,将设置为 head_dim_qk

  • custom_mask (Optional[torch.Tensor]) –

    展平的布尔掩码张量,形状:(sum(q_len[i] * k_len[i] for i in range(batch_size))。掩码张量中的元素应为 TrueFalse,其中 False 表示注意力矩阵中相应元素将被屏蔽。

    有关掩码张量的展平布局的更多详细信息,请参阅 掩码布局

    当提供 custom_mask 且未提供 packed_custom_mask 时,该函数会将自定义掩码张量打包成 1D 压缩掩码张量,这会引入额外的开销。

  • packed_custom_mask (Optional[torch.Tensor]) – 1D 压缩的 uint8 掩码张量,如果提供,将忽略 plan() 中的 custom_mask。压缩掩码张量由 flashinfer.quantization.packbits() 生成。

  • causal (bool) – 是否将因果掩码应用于注意力矩阵。只有在 plan() 中未提供 custom_mask 时,此设置才有效。

  • pos_encoding_mode (str) – 注意力内核内部应用的 positional encoding,可以是 NONE/ROPE_LLAMA(LLAMA 风格的旋转嵌入)/ALIBI。默认值为 NONE

  • use_fp16_qk_reduction (bool) – 是否使用 fp16 进行 qk 缩减(速度更快,但精度略有损失)。

  • window_left (int) – 注意力窗口的左(包含)窗口大小,设置为 -1 时,窗口大小将设置为序列的完整长度。默认值为 -1

  • logits_soft_cap (Optional[float]) – 注意力 logits 的软上限值(用于 Gemini、Grok 和 Gemma-2 等),如果未提供,将设置为 0。如果大于 0,logits 将根据公式进行截断:\(\texttt{logits_soft_cap} \times \mathrm{tanh}(x / \texttt{logits_soft_cap})\),其中 \(x\) 是输入 logits。

  • sm_scale (Optional[float]) – softmax 中使用的缩放比例,如果未提供,则设置为 1.0 / sqrt(head_dim)

  • rope_scale (Optional[float]) – RoPE 插值中使用的缩放比例,如果未提供,则设置为 1.0

  • rope_theta (Optional[float]) – RoPE 中使用的 theta,如果未提供,则设置为 1e4

  • q_data_type (Union[str, torch.dtype]) – 查询张量的的数据类型,默认为 torch.float16。

  • kv_data_type (Optional[Union[str, torch.dtype]]) – key/value 张量的数据类型。如果为 None,则设置为 q_data_type

  • o_data_type (Optional[Union[str, torch.dtype]]) – 输出张量的数据类型。如果为 None,则设置为 q_data_type。对于 FP8 输入,通常应设置为 torch.float16 或 torch.bfloat16。

  • non_blocking (bool) – 是否异步将输入张量复制到设备,默认为 True

  • prefix_len_ptr (Optional[torch.Tensor]) – 前缀长度。一个 uint32 一维张量,指示每个 prompt 的前缀长度。张量大小等于 batch size。

  • token_pos_in_items_ptr (Optional[torch.Tensor]) – 一个 uint16 一维张量(在 flashinfer 中将被转换为 uint16),指示每个 item 的 token 位置,并从 0(分隔符)开始,针对每个 item。例如,如果对于此成员有 3 个长度为 3、2、4 的 item,则此向量将如下所示:[0, 1, 2, 3, 0, 1, 2, 0, 1, 2, 3, 4, 0],其中 4 个分隔符索引为 0。对于 batch size > 1,我们将它们连接成一个一维张量,并用零填充,以确保每个张量具有相同的长度,填充长度由 token_pos_in_items_len 减去每个 prompt 的原始 token_pos_in_items_ptr 的长度定义。

  • token_pos_in_items_len (int) – 用于 token_pos_in_items_ptr 的零填充长度,以更好地处理 bsz > 1 的情况。仍然使用上面的 3,2,4 示例。如果我们将 token_pos_in_items_len 设置为 20,它将是 [0, 1, 2, 3, 0, 1, 2, 0, 1, 2, 3, 4, 0, 0, 0, 0, 0, 0, 0, 0],其中有 7 个填充零。(请注意,末尾有 8 个零,其中第一个是 prompt 末尾的分隔符 token 0)

  • max_item_len_ptr (Optional[torch.Tensor]) – 一个 uint16 向量,包含每个 prompt 中所有 item 的最大 token 长度

  • seq_lens (Optional[torch.Tensor]) – 一个 uint32 1D 张量,指示每个 prompt 的 kv 序列长度。形状:[batch_size]

  • seq_lens_q (Optional[torch.Tensor]) – 一个 uint32 一维张量,指示每个 prompt 的 q 序列长度。形状:[batch_size]。如果未提供,则设置为与 seq_lens 相同的值。

  • block_tables (Optional[torch.Tensor]) – 一个 uint32 2D 张量,指示每个 prompt 的 block table。形状:[batch_size, max_num_blocks_per_seq]

  • max_token_per_sequence (Optional[int],) – 适用于 cudnn 后端。这是每个序列的最大 token 长度。

  • max_sequence_kv (Optional[int],) – 适用于 cudnn 后端。这是 kv 缓存中每个序列的最大序列长度。

  • fixed_split_size (Optional[int],) – FA2 split-kv prefill/decode 在页面中的固定分割大小。建议设置为工作负载的平均序列长度。启用后,将在 merge_states 内核中导致确定性的 softmax 分数归约,因此输出与 batch size 不变。请参阅 https://thinkingmachines.ai/blog/defeating-nondeterminism-in-llm-inference/ 请注意,由于即使 bs 固定,kv 序列长度也可能发生变化,从而导致启动的 CTA 数量不同,因此与 CUDA 图的兼容性不能保证。

  • disable_split_kv (bool,) – 是否禁用 split-kv 以在 CUDA Graph 中实现确定性,默认为 False

注意

在调用任何 run()run_return_lse() 之前,应调用 plan() 方法,辅助数据结构将在调用期间创建并缓存,以用于多次内核运行。

num_qo_heads 必须是 num_kv_heads 的倍数。如果 num_qo_heads 不等于 num_kv_heads,该函数将使用 分组查询注意力

在 Cuda Graph 或 torch.compile 中无法使用 plan() 方法。

reset_workspace_buffer(float_workspace_buffer: Tensor, int_workspace_buffer: Tensor) None

重置工作区缓冲区。

参数:
  • float_workspace_buffer (torch.Tensor) – 新的 float 工作区缓冲区,该缓冲区的设备应与输入张量的设备相同。

  • int_workspace_buffer (torch.Tensor) – 新的 int 工作区缓冲区,该缓冲区的设备应与输入张量的设备相同。

run(q: Tensor, paged_kv_cache: Tensor | Tuple[Tensor, Tensor], *args, k_scale: float | None = None, v_scale: float | None = None, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[False] = False, enable_pdl: bool | None = None, window_left: int | None = None) Tensor
run(q: Tensor, paged_kv_cache: Tensor | Tuple[Tensor, Tensor], *args, k_scale: float | None = None, v_scale: float | None = None, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[True] = True, enable_pdl: bool | None = None, window_left: int | None = None) Tuple[Tensor, Tensor]

计算查询和分页 kv 缓存之间的批量预填充/追加注意力。

参数:
  • q (torch.Tensor) – 查询张量,形状:[qo_indptr[-1], num_qo_heads, head_dim]

  • paged_kv_cache (Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]) –

    存储的分页 KV 缓存,作为张量元组或单个张量

    • 一个元组 (k_cache, v_cache),包含 4D 张量,每个张量的形状为:[max_num_pages, page_size, num_kv_heads, head_dim],如果 kv_layoutNHD,以及 [max_num_pages, num_kv_heads, page_size, head_dim],如果 kv_layoutHND

    • 一个 5D 张量,形状为:[max_num_pages, 2, page_size, num_kv_heads, head_dim],如果 kv_layoutNHD,以及 [max_num_pages, 2, num_kv_heads, page_size, head_dim],如果 kv_layoutHND。其中 paged_kv_cache[:, 0] 是 key 缓存,paged_kv_cache[:, 1] 是 value 缓存。

  • *args – 自定义内核的附加参数。

  • q_scale (Optional[Union[float, torch.Tensor]]) – fp8 输入的查询校准比例,如果未提供,则设置为 1.0

  • k_scale (Optional[Union[float, torch.Tensor]]) – 用于fp8输入的关键(key)的校准比例,如果未提供,则设置为 1.0

  • v_scale (Optional[Union[float, torch.Tensor]]) – 用于fp8输入的价值(value)的校准比例,如果未提供,则设置为 1.0

  • out (Optional[torch.Tensor]) – 输出张量,如果未提供,则会在内部分配。

  • lse (Optional[torch.Tensor]) – 注意力 logits 的对数和指数,如果未提供,则会在内部分配。

  • return_lse (bool) – 是否返回注意力输出的logsumexp

  • enable_pdl (bool) – 是否启用程序依赖启动 (PDL)。请参阅 https://docs.nvda.net.cn/cuda/cuda-c-programming-guide/#programmatic-dependent-launch-and-synchronization,仅支持 >= sm90,并且当前仅支持 FA2 和 CUDA 核心解码。

返回值:

如果 return_lseFalse,则注意力输出的形状为:[qo_indptr[-1], num_qo_heads, head_dim]。如果 return_lseTrue,则返回一个包含两个张量的元组

  • 注意力输出的形状为:[qo_indptr[-1], num_qo_heads, head_dim]

  • 注意力输出的logsumexp的形状为:[qo_indptr[-1], num_qo_heads]

返回值类型:

Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]

class flashinfer.prefill.BatchPrefillWithRaggedKVCacheWrapper(float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, qo_indptr_buf: Tensor | None = None, kv_indptr_buf: Tensor | None = None, custom_mask_buf: Tensor | None = None, mask_indptr_buf: Tensor | None = None, backend: str = 'auto', jit_args: List[Any] | None = None, jit_kwargs: Dict[str, Any] | None = None)

用于批量请求的带有不规则(tensor)kv-cache的预填充/追加注意力的包装类。

请查看 我们的教程 以了解不规则kv-cache布局。

示例

>>> import torch
>>> import flashinfer
>>> num_layers = 32
>>> num_qo_heads = 64
>>> num_kv_heads = 16
>>> head_dim = 128
>>> # allocate 128MB workspace buffer
>>> workspace_buffer = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device="cuda:0")
>>> prefill_wrapper = flashinfer.BatchPrefillWithRaggedKVCacheWrapper(
...     workspace_buffer, "NHD"
... )
>>> batch_size = 7
>>> nnz_kv = 100
>>> nnz_qo = 100
>>> qo_indptr = torch.tensor(
...     [0, 33, 44, 55, 66, 77, 88, nnz_qo], dtype=torch.int32, device="cuda:0"
... )
>>> kv_indptr = qo_indptr.clone()
>>> q_at_layer = torch.randn(num_layers, nnz_qo, num_qo_heads, head_dim).half().to("cuda:0")
>>> k_at_layer = torch.randn(num_layers, nnz_kv, num_kv_heads, head_dim).half().to("cuda:0")
>>> v_at_layer = torch.randn(num_layers, nnz_kv, num_kv_heads, head_dim).half().to("cuda:0")
>>> # create auxiliary data structures for batch prefill attention
>>> prefill_wrapper.plan(
...     qo_indptr,
...     kv_indptr,
...     num_qo_heads,
...     num_kv_heads,
...     head_dim,
...     causal=True,
... )
>>> outputs = []
>>> for i in range(num_layers):
...     q = q_at_layer[i]
...     k = k_at_layer[i]
...     v = v_at_layer[i]
...     # compute batch prefill attention, reuse auxiliary data structures
...     o = prefill_wrapper.run(q, k, v)
...     outputs.append(o)
...
>>> outputs[0].shape
torch.Size([100, 64, 128])
>>>
>>> # below is another example of creating custom mask for batch prefill attention
>>> mask_arr = []
>>> qo_len = (qo_indptr[1:] - qo_indptr[:-1]).cpu().tolist()
>>> kv_len = (kv_indptr[1:] - kv_indptr[:-1]).cpu().tolist()
>>> for i in range(batch_size):
...     mask_i = torch.tril(
...         torch.full((qo_len[i], kv_len[i]), True, device="cuda:0"),
...         diagonal=(kv_len[i] - qo_len[i]),
...     )
...     mask_arr.append(mask_i.flatten())
...
>>> mask = torch.cat(mask_arr, dim=0)
>>> prefill_wrapper.plan(
...     qo_indptr,
...     kv_indptr,
...     num_qo_heads,
...     num_kv_heads,
...     head_dim,
...     custom_mask=mask
... )
>>> outputs_custom_mask = []
>>> for i in range(num_layers):
...     q = q_at_layer[i]
...     k = k_at_layer[i]
...     v = v_at_layer[i]
...     # compute batch prefill attention, reuse auxiliary data structures
...     o_custom = prefill_wrapper.run(q, k, v)
...     assert torch.allclose(o_custom, outputs[i], rtol=1e-3, atol=1e-3)
...
>>> outputs_custom_mask[0].shape
torch.Size([100, 64, 128])

注意

为了加速计算,FlashInfer 的批量预填充/追加注意力算子会创建一些辅助数据结构,这些数据结构可以在多个预填充/追加注意力调用之间重用(例如,不同的 Transformer 层)。此包装类管理这些数据结构的生命周期。

__init__(float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, qo_indptr_buf: Tensor | None = None, kv_indptr_buf: Tensor | None = None, custom_mask_buf: Tensor | None = None, mask_indptr_buf: Tensor | None = None, backend: str = 'auto', jit_args: List[Any] | None = None, jit_kwargs: Dict[str, Any] | None = None) None

BatchPrefillWithRaggedKVCacheWrapper 的构造函数。

参数:
  • float_workspace_buffer (torch.Tensor) – 用于在split-k算法中存储中间注意力结果的用户保留的浮点工作区缓冲区。推荐大小为128MB,工作区缓冲区的设备应与输入张量的设备相同。

  • kv_layout (str) – 输入 k/v 张量的布局,可以是 NHDHND

  • use_cuda_graph (bool) – 是否为预填充内核启用CUDA图捕获,如果启用,辅助数据结构将存储为提供的缓冲区。

  • qo_indptr_buf (Optional[torch.Tensor]) – 用户保留的GPU缓冲区,用于存储 qo_indptr 数组,缓冲区的大小应为 [batch_size + 1]。只有当 use_cuda_graphTrue 时,此参数才有效。

  • kv_indptr_buf (Optional[torch.Tensor]) – 用户保留的GPU缓冲区,用于存储 kv_indptr 数组,缓冲区的大小应为 [batch_size + 1]。只有当 use_cuda_graphTrue 时,此参数才有效。

  • custom_mask_buf (Optional[torch.Tensor]) – 用户预留的 GPU 缓冲区,用于存储自定义掩码张量,应足够大以存储包装器生命周期内打包的自定义掩码张量的最大可能大小。当 use_cuda_graphTrue 且注意力计算中将使用自定义掩码时,此参数才有效。

  • mask_indptr_buf (Optional[torch.Tensor]) – 用户预留的 GPU 缓冲区,用于存储 mask_indptr 数组,缓冲区的大小应为 [batch_size]。当 use_cuda_graphTrue 且注意力计算中将使用自定义掩码时,此参数才有效。

  • backend (str) – 实现后端,可以是 auto/fa2/fa3/cudnncutlass。默认为 auto。如果设置为 auto,则包装器将根据设备架构和内核可用性自动选择后端。

  • jit_args (Optional[List[Any]]) – 如果提供,包装器将使用提供的参数创建 JIT 模块,否则,包装器将使用默认注意力实现。

  • jit_kwargs (Optional[Dict[str, Any]]) – 创建 JIT 模块的关键字参数,默认为 None。

plan(qo_indptr: Tensor, kv_indptr: Tensor, num_qo_heads: int, num_kv_heads: int, head_dim_qk: int, head_dim_vo: int | None = None, custom_mask: Tensor | None = None, packed_custom_mask: Tensor | None = None, causal: bool = False, pos_encoding_mode: str = 'NONE', use_fp16_qk_reduction: bool = False, window_left: int = -1, logits_soft_cap: float | None = None, sm_scale: float | None = None, rope_scale: float | None = None, rope_theta: float | None = None, q_data_type: str | dtype = 'float16', kv_data_type: str | dtype | None = None, o_data_type: str | dtype | None = None, non_blocking: bool = True, prefix_len_ptr: Tensor | None = None, token_pos_in_items_ptr: Tensor | None = None, token_pos_in_items_len: int = 0, max_item_len_ptr: Tensor | None = None, fixed_split_size: int | None = None, disable_split_kv: bool = False, seq_lens: Tensor | None = None, seq_lens_q: Tensor | None = None, max_token_per_sequence: int | None = None, max_sequence_kv: int | None = None, v_indptr: Tensor | None = None, o_indptr: Tensor | None = None) None

规划用于给定问题规范的 Ragged KV-Cache 的批量预填充/追加注意力。

参数:
  • qo_indptr (torch.Tensor) – 查询/输出张量的 indptr,形状:[batch_size + 1]

  • kv_indptr (torch.Tensor) – key/value 张量的 indptr,形状:[batch_size + 1]

  • num_qo_heads (int) – 查询/输出头的数量。

  • num_kv_heads (int) – 键/值头的数量。

  • head_dim_qk (int) – query/key 张量上的头的维度。

  • head_dim_vo (Optional[int]) – value/output 张量上的头的维度。如果未提供,将设置为 head_dim_qk

  • custom_mask (Optional[torch.Tensor]) –

    展平的布尔掩码张量,形状:(sum(q_len[i] * k_len[i] for i in range(batch_size))。掩码张量中的元素应为 TrueFalse,其中 False 表示注意力矩阵中相应元素将被屏蔽。

    有关掩码张量的展平布局的更多详细信息,请参阅 掩码布局

    当提供 custom_mask 且未提供 packed_custom_mask 时,该函数会将自定义掩码张量打包成 1D 压缩掩码张量,这会引入额外的开销。

  • packed_custom_mask (Optional[torch.Tensor]) –

    如果提供了 1D 压缩的 uint8 掩码张量,则会忽略 custom_mask。压缩的掩码张量由 flashinfer.quantization.packbits() 生成。

    如果提供,自定义掩码将在 softmax 和缩放之后添加到注意力矩阵中。掩码张量应与输入张量位于同一设备上。

  • causal (bool) – 是否将因果掩码应用于注意力矩阵。如果在 plan() 中提供了 mask,则忽略此参数。

  • pos_encoding_mode (str) – 注意力内核内部应用的 positional encoding,可以是 NONE/ROPE_LLAMA(LLAMA 风格的旋转嵌入)/ALIBI。默认值为 NONE

  • use_fp16_qk_reduction (bool) – 是否使用 fp16 进行 qk 缩减(速度更快,但精度略有损失)。

  • window_left (int) – 注意力窗口的左(包含)窗口大小,设置为 -1 时,窗口大小将设置为序列的完整长度。默认值为 -1

  • logits_soft_cap (Optional[float]) – 注意力 logits 的软上限值(用于 Gemini、Grok 和 Gemma-2 等),如果未提供,将设置为 0。如果大于 0,logits 将根据公式进行截断:\(\texttt{logits_soft_cap} \times \mathrm{tanh}(x / \texttt{logits_soft_cap})\),其中 \(x\) 是输入 logits。

  • sm_scale (Optional[float]) – softmax 中使用的缩放比例,如果未提供,将设置为 1.0 / sqrt(head_dim_qk)

  • rope_scale (Optional[float]) – RoPE 插值中使用的缩放比例,如果未提供,则设置为 1.0

  • rope_theta (Optional[float]) – RoPE 中使用的 theta,如果未提供,则设置为 1e4

  • q_data_type (Union[str, torch.dtype]) – 查询张量的数据类型,默认为 torch.float16。

  • kv_data_type (Optional[Union[str, torch.dtype]]) – key/value 张量的数据类型。如果为 None,则设置为 q_data_type

  • o_data_type (Optional[Union[str, torch.dtype]]) – 输出张量的数据类型。如果为 None,则设置为 q_data_type。对于 FP8 输入,通常应设置为 torch.float16 或 torch.bfloat16。

  • non_blocking (bool) – 是否异步将输入张量复制到设备,默认为 True

  • prefix_len_ptr (Optional[torch.Tensor]) – 前缀长度。一个 uint32 一维张量,指示每个 prompt 的前缀长度。张量大小等于 batch size。

  • token_pos_in_items_ptr (Optional[torch.Tensor]) – 一个 uint16 一维张量(在 flashinfer 中将被转换为 uint16),指示每个 item 的 token 位置,并从 0(分隔符)开始,针对每个 item。例如,如果对于此成员有 3 个长度为 3、2、4 的 item,则此向量将如下所示:[0, 1, 2, 3, 0, 1, 2, 0, 1, 2, 3, 4, 0],其中 4 个分隔符索引为 0。对于 batch size > 1,我们将它们连接成一个一维张量,并用零填充,以确保每个张量具有相同的长度,填充长度由 token_pos_in_items_len 减去每个 prompt 的原始 token_pos_in_items_ptr 的长度定义。

  • token_pos_in_items_len (int) – 用于 token_pos_in_items_ptr 的零填充长度,以更好地处理 bsz > 1 的情况。仍然使用上面的 3,2,4 示例。如果我们将 token_pos_in_items_len 设置为 20,它将是 [0, 1, 2, 3, 0, 1, 2, 0, 1, 2, 3, 4, 0, 0, 0, 0, 0, 0, 0, 0],其中有 7 个填充零。(请注意,末尾有 8 个零,其中第一个是 prompt 末尾的分隔符 token 0)

  • max_item_len_ptr (Optional[torch.Tensor]) – 一个 uint16 向量,包含每个 prompt 中所有 item 的最大 token 长度

  • fixed_split_size (Optional[int],) – 分割-kv FA2 预填充/解码的固定分割大小,以页为单位。建议设置为工作负载的平均序列长度。启用后,将导致合并状态内核中的确定性 softmax 分数降低,从而产生与批大小无关的输出。请参阅 https://thinkingmachines.ai/blog/defeating-nondeterminism-in-llm-inference/ 请注意,即使在固定 bs 的情况下,kv 序列长度也可能发生变化,从而导致启动的 CTA 数量不同,因此不保证与 CUDA 图的兼容性。

  • disable_split_kv (bool,) – 是否禁用 split-kv 以在 CUDA Graph 中实现确定性,默认为 False

  • seq_lens (Optional[torch.Tensor]) – 一个 uint32 1D 张量,指示每个 prompt 的 kv 序列长度。形状:[batch_size]

  • seq_lens_q (Optional[torch.Tensor]) – 一个 uint32 一维张量,指示每个 prompt 的 q 序列长度。形状:[batch_size]。如果未提供,则设置为与 seq_lens 相同的值。

  • max_token_per_sequence (Optional[int],) – 适用于 cudnn 后端。这是每个序列的最大 token 长度。

  • max_sequence_kv (Optional[int],) – 适用于 cudnn 后端。这是 kv 缓存中每个序列的最大序列长度。

  • v_indptr (Optional[torch.Tensor]) – 适用于 cudnn 后端。这是值张量的 indptr。

  • o_indptr (Optional[torch.Tensor]) – 适用于 cudnn 后端。这是输出张量的 indptr。

注意

在调用任何 run()run_return_lse() 调用之前,应调用 plan() 方法。辅助数据结构将在本次计划调用期间创建并缓存,以供多次内核运行使用。

num_qo_heads 必须是 num_kv_heads 的倍数。如果 num_qo_heads 不等于 num_kv_heads,该函数将使用 分组查询注意力

在 Cuda 图或 torch.compile 中无法使用 plan() 方法。

reset_workspace_buffer(float_workspace_buffer: Tensor, int_workspace_buffer) None

重置工作区缓冲区。

参数:
  • float_workspace_buffer (torch.Tensor) – 新的 float 工作区缓冲区,该缓冲区的设备应与输入张量的设备相同。

  • int_workspace_buffer (torch.Tensor) – 新的 int 工作区缓冲区,该缓冲区的设备应与输入张量的设备相同。

run(q: Tensor, k: Tensor, v: Tensor, *args, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[False] = False, enable_pdl: bool | None = None) Tensor
run(q: Tensor, k: Tensor, v: Tensor, *args, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[True] = True, enable_pdl: bool | None = None) Tuple[Tensor, Tensor]

计算查询和存储为不规则张量的 kv 缓存之间的批预填充/追加注意力。

参数:
  • q (torch.Tensor) – 查询张量,形状:[qo_indptr[-1], num_qo_heads, head_dim_qk]

  • k (torch.Tensor) – 键张量,形状:[kv_indptr[-1], num_kv_heads, head_dim_qk]

  • v (torch.Tensor) – 值张量,形状:[kv_indptr[-1], num_kv_heads, head_dim_vo]

  • *args – 自定义内核的附加参数。

  • q_scale (Optional[float]) – fp8 查询的校准比例,如果未提供,则设置为 1.0

  • k_scale (Optional[float]) – fp8 键的校准比例,如果未提供,则设置为 1.0

  • v_scale (Optional[float]) – fp8 值的校准比例,如果未提供,则设置为 1.0

  • o_scale (Optional[float]) – 输出的校准比例,如果未提供,则设置为 1.0

  • out (Optional[torch.Tensor]) – 输出张量,如果未提供,则会在内部分配。

  • lse (Optional[torch.Tensor]) – 注意力 logits 的对数和指数,如果未提供,则会在内部分配。

  • return_lse (bool) – 是否返回注意力输出的logsumexp

  • enable_pdl (bool) – 是否启用程序依赖启动 (PDL)。请参阅 https://docs.nvda.net.cn/cuda/cuda-c-programming-guide/#programmatic-dependent-launch-and-synchronization,仅支持 >= sm90,并且当前仅支持 FA2 和 CUDA 核心解码。

返回值:

如果 return_lseFalse,则注意力输出,形状:[qo_indptr[-1], num_qo_heads, head_dim_vo]。如果 return_lseTrue,则为两个张量的元组

  • 注意力输出,形状:[qo_indptr[-1], num_qo_heads, head_dim_vo]

  • 注意力输出的logsumexp的形状为:[qo_indptr[-1], num_qo_heads]

返回值类型:

Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]

flashinfer.mla

MLA(多头潜在注意力)是一种在 DeepSeek 系列模型中提出的注意力机制(DeepSeek-V2DeepSeek-V3DeepSeek-R1)。

MLA 的 PageAttention

class flashinfer.mla.BatchMLAPagedAttentionWrapper(float_workspace_buffer: Tensor, use_cuda_graph: bool = False, qo_indptr: Tensor | None = None, kv_indptr: Tensor | None = None, kv_indices: Tensor | None = None, kv_len_arr: Tensor | None = None, backend: str = 'auto')

DeepSeek 模型上 MLA(多头潜在注意力)PagedAttention 的包装类。此内核可用于解码、增量预填充,应与 矩阵吸收技巧 结合使用:其中 \(W_{UQ}\)\(W_{UK}\) 吸收,\(W_{UV}\)\(W_{O}\) 吸收。对于不使用矩阵吸收的 MLA 注意力(head_dim_qk=192head_dim_vo=128,这用于预填充自注意力阶段),请使用 flashinfer.prefill.BatchPrefillWithRaggedKVCacheWrapper

有关 MLA 中分页 KV-Cache 布局的更多信息,请参阅我们的教程 MLA 页面布局

有关 MLA 计算、矩阵吸收和 FlashInfer 的 MLA 实现的更多详细信息,请参阅我们的 博客文章

示例

>>> import torch
>>> import flashinfer
>>> num_local_heads = 128
>>> batch_size = 114
>>> head_dim_ckv = 512
>>> head_dim_kpe = 64
>>> page_size = 1
>>> mla_wrapper = flashinfer.mla.BatchMLAPagedAttentionWrapper(
...     torch.empty(128 * 1024 * 1024, dtype=torch.int8).to(0),
...     backend="fa2"
... )
>>> q_indptr = torch.arange(0, batch_size + 1).to(0).int() # for decode, each query length is 1
>>> kv_lens = torch.full((batch_size,), 999, dtype=torch.int32).to(0)
>>> kv_indptr = torch.arange(0, batch_size + 1).to(0).int() * 999
>>> kv_indices = torch.arange(0, batch_size * 999).to(0).int()
>>> q_nope = torch.randn(
...     batch_size * 1, num_local_heads, head_dim_ckv, dtype=torch.bfloat16, device="cuda"
... )
>>> q_pe = torch.zeros(
...     batch_size * 1, num_local_heads, head_dim_kpe, dtype=torch.bfloat16, device="cuda"
... )
>>> ckv = torch.randn(
...     batch_size * 999, 1, head_dim_ckv, dtype=torch.bfloat16, device="cuda"
... )
>>> kpe = torch.zeros(
...     batch_size * 999, 1, head_dim_kpe, dtype=torch.bfloat16, device="cuda"
... )
>>> sm_scale = 1.0 / ((128 + 64) ** 0.5)  # use head dimension before matrix absorption
>>> mla_wrapper.plan(
...     q_indptr,
...     kv_indptr,
...     kv_indices,
...     kv_lens,
...     num_local_heads,
...     head_dim_ckv,
...     head_dim_kpe,
...     page_size,
...     False,  # causal
...     sm_scale,
...     q_nope.dtype,
...     ckv.dtype,
... )
>>> o = mla_wrapper.run(q_nope, q_pe, ckv, kpe, return_lse=False)
>>> o.shape
torch.Size([114, 128, 512])
__init__(float_workspace_buffer: Tensor, use_cuda_graph: bool = False, qo_indptr: Tensor | None = None, kv_indptr: Tensor | None = None, kv_indices: Tensor | None = None, kv_len_arr: Tensor | None = None, backend: str = 'auto') None

BatchMLAPagedAttentionWrapper 的构造函数。

参数:
  • **float_workspace_buffer** (torch.Tensor) – 用于在 split-k 算法中存储中间注意力结果的用户保留的工作空间缓冲区。推荐大小为 128MB,工作空间缓冲区的设备应与输入张量的设备相同。

  • use_cuda_graph (bool, optional) – 是否为预填充内核启用 CUDA 图捕获,如果启用,辅助数据结构将存储在提供的缓冲区中。当启用 CUDAGraph 时,此包装器的生命周期内 batch_size 不能更改。

  • qo_indptr_buf (Optional[torch.Tensor]) – 用户预留的缓冲区,用于存储 qo_indptr 数组,缓冲区的大小应为 [batch_size + 1]。只有当 use_cuda_graphTrue 时,此参数才有效。

  • kv_indptr_buf (Optional[torch.Tensor]) – 用户保留的缓冲区,用于存储 kv_indptr 数组,缓冲区的大小应为 [batch_size + 1]。当 use_cuda_graphTrue 时,此参数才有效。

  • kv_indices_buf (Optional[torch.Tensor]) – 用户保留的缓冲区,用于存储 kv_indices 数组。当 use_cuda_graphTrue 时,此参数才有效。

  • kv_len_arr_buf (Optional[torch.Tensor]) – 用户保留的缓冲区,用于存储 kv_len_arr 数组,缓冲区的大小应为 [batch_size]。当 use_cuda_graphTrue 时,此参数才有效。

  • backend (str) – 实现后端,可以是 auto/fa2fa3。默认为 auto。如果设置为 auto,则该函数将根据设备架构和内核可用性自动选择后端。如果提供 cutlass,则 MLA 内核将由 CUTLASS 生成,并且只需要 float_workspace_buffer,其他参数将被忽略。

plan(qo_indptr: Tensor, kv_indptr: Tensor, kv_indices: Tensor, kv_len_arr: Tensor, num_heads: int, head_dim_ckv: int, head_dim_kpe: int, page_size: int, causal: bool, sm_scale: float, q_data_type: dtype, kv_data_type: dtype, use_profiler: bool = False) None

规划 MLA 注意力计算。

参数:
  • qo_indptr (torch.IntTensor) – 查询/输出张量的 indptr,形状:[batch_size + 1]。对于解码注意力,每个查询的长度为 1,张量的内容应为 [0, 1, 2, ..., batch_size]

  • kv_indptr (torch.IntTensor) – 分页 kv-cache 的 indptr,形状:[batch_size + 1]

  • kv_indices (torch.IntTensor) – 分页 kv-cache 的页面索引,形状:[kv_indptr[-1]] 或更大。

  • kv_len_arr (torch.IntTensor) – 每个请求的查询长度,形状:[batch_size]

  • num_heads (int) – 查询/输出张量中的头数。

  • head_dim_ckv (int) – 压缩 kv 的头维度。

  • head_dim_kpe (int) – rope k-cache 的头维度。

  • page_size (int) – 分页 kv-cache 的页面大小。

  • causal (bool) – 是否使用因果注意力。

  • sm_scale (float) – softmax 运算的缩放因子。

  • q_data_type (torch.dtype) – 查询张量的数据类型。

  • kv_data_type (torch.dtype) – kv-cache 张量的数据类型。

  • use_profiler (bool, optional) – 是否启用内核内分析器,默认值为 False。

run(q_nope: Tensor, q_pe: Tensor, ckv_cache: Tensor, kpe_cache: Tensor, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[False] = False, profiler_buffer: Tensor | None = None, kv_len: Tensor | None = None, page_table: Tensor | None = None, return_lse_base_on_e: bool = False) Tensor
run(q_nope: Tensor, q_pe: Tensor, ckv_cache: Tensor, kpe_cache: Tensor, out: Tensor | None = None, lse: Tensor | None = None, return_lse: Literal[True] = True, profiler_buffer: Tensor | None = None, kv_len: Tensor | None = None, page_table: Tensor | None = None, return_lse_base_on_e: bool = False) Tuple[Tensor, Tensor]

运行 MLA 注意力计算。

参数:
  • q_nope (torch.Tensor) – 不含 rope 的查询张量,形状:[batch_size, num_heads, head_dim_ckv]

  • q_pe (torch.Tensor) – 查询张量的 rope 部分,形状:[batch_size, num_heads, head_dim_kpe]

  • ckv_cache (torch.Tensor) – 压缩的 kv-cache 张量(不含 rope),形状:[num_pages, page_size, head_dim_ckv]head_dim_ckv 在 DeepSeek v2/v3 模型中为 512。

  • kpe_cache (torch.Tensor) – kv-cache 张量的 rope 部分,形状:[num_pages, page_size, head_dim_kpe]head_dim_kpe 在 DeepSeek v2/v3 模型中为 64。

  • out (Optional[torch.Tensor]) – 输出张量,如果未提供,则会在内部分配。

  • lse (Optional[torch.Tensor]) – 注意力 logits 的对数和指数,如果未提供,则会在内部分配。

  • return_lse (bool, optional) – 是否返回对数和指数值,默认值为 False。

  • profiler_buffer (Optional[torch.Tensor]) – 用于存储分析器数据的缓冲区。

  • kv_len (Optional[torch.Tensor]) – 每个请求的查询长度,形状:[batch_size]。 当 backendcutlass 时需要。

  • page_table (Optional[torch.Tensor]) – 分页 kv-cache 的页表,形状:[batch_size, num_pages]。 当 backendcutlass 时需要。