flashinfer.decode.single_decode_with_kv_cache

flashinfer.decode.single_decode_with_kv_cache(q: Tensor, k: Tensor, v: Tensor, kv_layout: str = 'NHD', pos_encoding_mode: str = 'NONE', use_tensor_cores: bool = False, q_scale: float | None = None, k_scale: float | None = None, v_scale: float | None = None, 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, return_lse: Literal[False] = False) Tensor
flashinfer.decode.single_decode_with_kv_cache(q: Tensor, k: Tensor, v: Tensor, kv_layout: str = 'NHD', pos_encoding_mode: str = 'NONE', use_tensor_cores: bool = False, q_scale: float | None = None, k_scale: float | None = None, v_scale: float | None = None, 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, return_lse: Literal[True] = True) Tuple[Tensor, Tensor]

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

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

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

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

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

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

  • use_tensor_cores (bool) – 是否使用 tensor cores 进行计算。对于分组查询注意力中较大的组大小,会更快。默认为 False

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

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

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

  • 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 / sqrt(head_dim)

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

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

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

返回值:

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

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

  • log sum exp 值,形状:[num_qo_heads]

返回值类型:

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

示例

>>> import torch
>>> import flashinfer
>>> kv_len = 4096
>>> num_qo_heads = 32
>>> num_kv_heads = 32
>>> head_dim = 128
>>> q = torch.randn(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_decode_with_kv_cache(q, k, v)
>>> o.shape
torch.Size([32, 128])

注意

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