flashinfer.prefill.single_prefill_with_kv_cache

flashinfer.prefill.single_prefill_with_kv_cache(q: Tensor, k: Tensor, v: Tensor, scale_q: Tensor | None = None, scale_k: Tensor | None = None, scale_v: Tensor | None = None, o_dtype: dtype | None = None, custom_mask: Tensor | None = None, packed_custom_mask: Tensor | None = None, causal: bool = False, kv_layout: str = 'NHD', 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, backend: str = 'auto', return_lse: Literal[False] = False) Tensor
flashinfer.prefill.single_prefill_with_kv_cache(q: Tensor, k: Tensor, v: Tensor, scale_q: Tensor | None = None, scale_k: Tensor | None = None, scale_v: Tensor | None = None, o_dtype: dtype | None = None, custom_mask: Tensor | None = None, packed_custom_mask: Tensor | None = None, causal: bool = False, kv_layout: str = 'NHD', 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, backend: str = 'auto', return_lse: Literal[True] = True) Tuple[Tensor, Tensor]

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

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

  • k (torch.Tensor) – 键张量,形状:如果 kv_layoutNHD,则为 [kv_len, num_kv_heads, head_dim_qk];如果 kv_layoutHND,则为 [num_kv_heads, kv_len, head_dim_qk]

  • v (torch.Tensor) – 值张量,形状:如果 kv_layoutNHD,则为 [kv_len, num_kv_heads, head_dim_vo];如果 kv_layoutHND,则为 [num_kv_heads, kv_len, head_dim_vo]

  • scale_q (Optional[torch.Tensor]) – 查询的缩放张量,每头量化,形状:[num_qo_heads]。用于 FP8 量化。如果未提供,将设置为 1.0

  • scale_k (Optional[torch.Tensor]) – 键的缩放张量,每头量化,形状:[num_kv_heads]。用于 FP8 量化。如果未提供,将设置为 1.0

  • scale_v (Optional[torch.Tensor]) – 值的缩放张量,每头量化,形状:[num_kv_heads]。用于 FP8 量化。如果未提供,将设置为 1.0

  • o_dtype (Optional[torch.dtype]) – 输出张量的数据类型,如果未提供,将设置为与 q 相同。这对于量化中的自动推断输出 dtype 来说是必要的。

  • custom_mask (Optional[torch.Tensor]) –

    自定义布尔掩码张量,形状:[qo_len, kv_len]。掩码张量中的元素应为 TrueFalse,其中 False 表示注意力矩阵中相应的元素将被屏蔽。

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

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

  • causal (bool) – 是否对注意力矩阵应用因果掩码。只有在未提供 custom_mask 时才有效。

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

  • 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。

  • backend (str) – 实现后端,可以是 auto/fa2fa3。默认为 auto。如果设置为 auto,该函数将根据设备架构和内核可用性自动选择后端。

  • return_lse (bool) – 是否返回注意力 logits 的对数和指数值。

返回值:

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

  • 注意力输出,形状:[qo_len, num_qo_heads, head_dim_vo]

  • 对数和指数值,形状:[qo_len, num_qo_heads]

返回值类型:

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

示例

>>> import torch
>>> import flashinfer
>>> qo_len = 128
>>> kv_len = 4096
>>> num_qo_heads = 32
>>> num_kv_heads = 4
>>> head_dim = 128
>>> q = torch.randn(qo_len, num_qo_heads, head_dim).half().to("cuda:0")
>>> k = torch.randn(kv_len, num_kv_heads, head_dim).half().to("cuda:0")
>>> v = torch.randn(kv_len, num_kv_heads, head_dim).half().to("cuda:0")
>>> o = flashinfer.single_prefill_with_kv_cache(q, k, v, causal=True,
        use_fp16_qk_reduction=True)
>>> o.shape
torch.Size([128, 32, 128])
>>> mask = torch.tril(
>>>     torch.full((qo_len, kv_len), True, device="cuda:0"),
>>>     diagonal=(kv_len - qo_len),
>>> )
>>> mask
tensor([[ True,  True,  True,  ..., False, False, False],
        [ True,  True,  True,  ..., False, False, False],
        [ True,  True,  True,  ..., False, False, False],
        ...,
        [ True,  True,  True,  ...,  True, False, False],
        [ True,  True,  True,  ...,  True,  True, False],
        [ True,  True,  True,  ...,  True,  True,  True]], device='cuda:0')
>>> o_custom = flashinfer.single_prefill_with_kv_cache(q, k, v, custom_mask=mask)
>>> torch.allclose(o, o_custom, rtol=1e-3, atol=1e-3)
True

注意

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