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_layout是NHD,则为[kv_len, num_kv_heads, head_dim_qk];如果kv_layout是HND,则为[num_kv_heads, kv_len, head_dim_qk]。v (torch.Tensor) – 值张量,形状:如果
kv_layout是NHD,则为[kv_len, num_kv_heads, head_dim_vo];如果kv_layout是HND,则为[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]。掩码张量中的元素应为True或False,其中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 张量的布局,可以是
NHD或HND。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/fa2或fa3。默认为auto。如果设置为auto,该函数将根据设备架构和内核可用性自动选择后端。return_lse (bool) – 是否返回注意力 logits 的对数和指数值。
- 返回值:
如果
return_lse为False,则注意力输出的形状为:[qo_len, num_qo_heads, head_dim_vo]。如果return_lse为True,则返回一个包含两个张量的元组注意力输出,形状:
[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,该函数将使用 分组查询注意力。