flashinfer.cascade

合并注意力状态

merge_state(v_a, s_a, v_b, s_b)

合并两个 KV 分段的注意力输出 V 和 logsumexp 值 S

merge_state_in_place(v, s, v_other, s_other)

就地合并自注意力状态 (v, s) 与另一个状态 (v_other, s_other)

merge_states(v, s)

合并多个注意力状态 (v, s)。

级联注意力

级联注意力包装器类

class flashinfer.cascade.MultiLevelCascadeAttentionWrapper(num_levels, float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, qo_indptr_buf_arr: List[Tensor] | None = None, paged_kv_indptr_buf_arr: List[Tensor] | None = None, paged_kv_indices_buf_arr: List[Tensor] | None = None, paged_kv_last_page_len_buf_arr: List[Tensor] | None = None)

用于内存高效多级级联推理的注意力包装器,此 API 假定所有级别 KV 缓存都存储在统一的分页表中。

请参阅 多级级联推理数据布局 以了解级联推理中的数据布局。请注意,由于合并注意力结果的开销,增加级别数并不总是受益的。

级联推理的思想在我们的 博客文章 中介绍。

示例

>>> import torch
>>> import flashinfer
>>> num_layers = 32
>>> num_qo_heads = 64
>>> num_kv_heads = 8
>>> head_dim = 128
>>> page_size = 16
>>> # allocate 128MB workspace buffer
>>> workspace_buffer = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device="cuda:0")
>>> wrapper = flashinfer.MultiLevelCascadeAttentionWrapper(
...     2, workspace_buffer, "NHD"
... )
>>> batch_size = 7
>>> shared_kv_num_pages = 512
>>> unique_kv_num_pages = 128
>>> total_num_pages = shared_kv_num_pages + unique_kv_num_pages
>>> shared_kv_page_indices = torch.arange(shared_kv_num_pages).int().to("cuda:0")
>>> shared_kv_page_indptr = torch.tensor([0, shared_kv_num_pages], dtype=torch.int32, device="cuda:0")
>>> unique_kv_page_indices = torch.arange(shared_kv_num_pages, total_num_pages).int().to("cuda:0")
>>> unique_kv_page_indptr = torch.tensor(
...     [0, 17, 29, 44, 48, 66, 100, 128], dtype=torch.int32, device="cuda:0"
... )
>>> shared_kv_last_page_len = torch.tensor([page_size], dtype=torch.int32, device="cuda:0")
>>> # 1 <= kv_last_page_len <= page_size
>>> unique_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(
...         total_num_pages, 2, page_size, num_kv_heads, head_dim, dtype=torch.float16, device="cuda:0"
...     ) for _ in range(num_layers)
... ]
>>> qo_indptr_arr = [
...     torch.tensor([0, batch_size], dtype=torch.int32, device="cuda:0"),  # top-level for shared KV-Cache
...     torch.arange(batch_size + 1, dtype=torch.int32, device="cuda:0")    # bottom-level for unique KV-Cache
... ]
>>> # create auxiliary data structures for batch decode attention
>>> wrapper.plan(
...     qo_indptr_arr,
...     [shared_kv_page_indptr, unique_kv_page_indptr],
...     [shared_kv_page_indices, unique_kv_page_indices],
...     [shared_kv_last_page_len, unique_kv_last_page_len],
...     num_qo_heads,
...     num_kv_heads,
...     head_dim,
...     page_size,
... )
>>> outputs = []
>>> for i in range(num_layers):
...     q = torch.randn(batch_size, num_qo_heads, head_dim).half().to("cuda:0")
...     # compute batch decode attention, reuse auxiliary data structures for all layers
...     o = wrapper.run(q, kv_cache_at_layer[i])
...     outputs.append(o)
...
>>> outputs[0].shape
torch.Size([7, 64, 128])

参见

BatchPrefillWithPagedKVCacheWrapper

__init__(num_levels, float_workspace_buffer: Tensor, kv_layout: str = 'NHD', use_cuda_graph: bool = False, qo_indptr_buf_arr: List[Tensor] | None = None, paged_kv_indptr_buf_arr: List[Tensor] | None = None, paged_kv_indices_buf_arr: List[Tensor] | None = None, paged_kv_last_page_len_buf_arr: List[Tensor] | None = None) None

MultiLevelCascadeAttentionWrapper 的构造函数。

参数:
  • num_levels (int) – 级联注意力中的级别数。

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

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

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

  • qo_indptr_buf_arr (Optional[List[torch.Tensor]]) – 每个级别的 qo indptr 缓冲区数组,数组长度应等于级别数。每个张量的最后一个元素应该是查询/输出的总数。

  • paged_kv_indptr_buf_arr (Optional[List[torch.Tensor]]) – 每个级别的分页 kv 缓存 indptr 缓冲区数组,数组长度应等于级别数。

  • paged_kv_indices_buf_arr (Optional[List[torch.Tensor]]) – 每个级别的分页 kv 缓存索引缓冲区数组,数组长度应等于级别数。

  • paged_kv_last_page_len_buf_arr (Optional[List[torch.Tensor]]) – 每个级别的分页 kv 缓存最后一页长度缓冲区数组,数组长度应等于级别数。

plan(qo_indptr_arr: List[Tensor], paged_kv_indptr_arr: List[Tensor], paged_kv_indices_arr: List[Tensor], paged_kv_last_page_len: List[Tensor], num_qo_heads: int, num_kv_heads: int, head_dim: int, page_size: int, 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 = 'float16', kv_data_type: str | dtype | None = None)

为多级级联注意力创建辅助数据结构,用于在相同的解码步骤内的多次前向调用。请查看多级级联推理数据布局了解级联推理中的数据布局。

参数:
  • qo_indptr_arr (List[torch.Tensor]) – 每个级别的 qo indptr 张量数组,数组长度应等于级别数。每个张量的最后一个元素应为查询/输出的总数。

  • paged_kv_indptr_arr (List[torch.Tensor]) – 每个级别的分页 kv 缓存 indptr 张量数组,数组长度应等于级别数。

  • paged_kv_indices_arr (List[torch.Tensor]) – 每个级别的分页 kv 缓存索引张量数组,数组长度应等于级别数。

  • paged_kv_last_page_len (List[torch.Tensor]) – 每个级别的分页 kv 缓存最后一页长度张量数组,数组长度应等于级别数。

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

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

  • head_dim (int) – 头部的维度。

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

  • 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 (Optional[Union[str, torch.dtype]]) – 查询张量的数据类型。如果为 None,则设置为 torch.float16。

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

reset_workspace_buffer(float_workspace_buffer: Tensor, int_workspace_buffers: List[Tensor]) None

重置工作区缓冲区。

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

  • int_workspace_buffers (List[torch.Tensor]) – 新的 int 工作区缓冲区数组,新 int 工作区缓冲区的设备应与输入张量的设备相同。

run(q: Tensor, paged_kv_cache: Tensor)

计算多级级联注意力。

参数:
  • 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] 是键缓存,paged_kv_cache[:, 1] 是值缓存。

class flashinfer.cascade.BatchDecodeWithSharedPrefixPagedKVCacheWrapper(float_workspace_buffer: Tensor, kv_layout: str = 'NHD')

用于具有共享前缀分页 kv 缓存的批处理请求的解码注意力的包装器类。共享前缀 KV 缓存存储在独立的张量中,每个请求的唯一 KV 缓存存储在分页 KV 缓存数据结构中。

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

警告

此 API 将在未来弃用,请改用 MultiLevelCascadeAttentionWrapper

示例

>>> 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.empty(128 * 1024 * 1024, dtype=torch.uint8, device="cuda:0")
>>> wrapper = flashinfer.BatchDecodeWithSharedPrefixPagedKVCacheWrapper(
...     workspace_buffer, "NHD"
... )
>>> batch_size = 7
>>> shared_prefix_len = 8192
>>> unique_kv_page_indices = torch.arange(max_num_pages).int().to("cuda:0")
>>> unique_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
>>> unique_kv_last_page_len = torch.tensor(
...     [1, 7, 14, 4, 3, 1, 16], dtype=torch.int32, device="cuda:0"
... )
>>> unique_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)
... ]
>>> shared_k_data_at_layer = [
...     torch.randn(
...         shared_prefix_len, num_kv_heads, head_dim, dtype=torch.float16, device="cuda:0"
...     ) for _ in range(num_layers)
... ]
>>> shared_v_data_at_layer = [
...     torch.randn(
...         shared_prefix_len, num_kv_heads, head_dim, dtype=torch.float16, device="cuda:0"
...     ) for _ in range(num_layers)
... ]
>>> # create auxiliary data structures for batch decode attention
>>> wrapper.begin_forward(
...     unique_kv_page_indptr,
...     unique_kv_page_indices,
...     unique_kv_last_page_len,
...     num_qo_heads,
...     num_kv_heads,
...     head_dim,
...     page_size,
...     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")
...     k_shared = shared_k_data_at_layer[i]
...     v_shared = shared_v_data_at_layer[i]
...     unique_kv_cache = unique_kv_cache_at_layer[i]
...     # compute batch decode attention, reuse auxiliary data structures for all layers
...     o = wrapper.forward(q, k_shared, v_shared, unique_kv_cache)
...     outputs.append(o)
...
>>> outputs[0].shape
torch.Size([7, 64, 128])

注意

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

__init__(float_workspace_buffer: Tensor, kv_layout: str = 'NHD') None
begin_forward(unique_kv_indptr: Tensor, unique_kv_indices: Tensor, unique_kv_last_page_len: Tensor, num_qo_heads: int, num_kv_heads: int, head_dim: int, page_size: int, data_type: str = 'float16') None

为给定的问题规范规划共享前缀批量解码注意力。

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

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

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

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

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

  • head_dim (int) – 头部的维度

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

  • data_type (Union[str, torch.dtype]) – 分页 kv 缓存的数据类型

注意

在调用任何 forward()forward_return_lse() 调用之前,应调用 begin_forward() 方法,在此调用期间将创建辅助数据结构并缓存以供多次前向调用使用。

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

end_forward() None

警告:此函数已弃用,没有效果

forward(q: Tensor, k_shared: Tensor, v_shared: Tensor, unique_kv_cache: Tensor) Tensor

计算查询和共享前缀分页 kv 缓存之间的批量解码注意力。

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

  • k_shared (torch.Tensor) – 共享前缀键张量,形状:[shared_prefix_len, num_kv_heads, head_dim] 如果 kv_layoutNHD,或者 [num_kv_heads, shared_prefix_len, head_dim] 如果 kv_layoutHND

  • v_shared (torch.Tensor) – 共享前缀值张量,形状:[shared_prefix_len, num_kv_heads, head_dim] 如果 kv_layoutNHD,或者 [num_kv_heads, shared_prefix_len, head_dim] 如果 kv_layoutHND

  • unique_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] 是键缓存,paged_kv_cache[:, 1] 是值缓存。

返回值:

V – 注意力输出,形状:[batch_size, num_heads, head_dim]

返回值类型:

torch.Tensor

reset_workspace_buffer(float_workspace_buffer: Tensor, int_workspace_buffer) None

重置工作区缓冲区。

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

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

class flashinfer.cascade.BatchPrefillWithSharedPrefixPagedKVCacheWrapper(float_workspace_buffer: Tensor, kv_layout: str = 'NHD')

用于批量请求的分页 kv 缓存共享前缀预填充/追加注意力的包装器类。

请查看 我们的教程 以了解分页 kv 缓存布局。

警告

此 API 将在未来弃用,请改用 MultiLevelCascadeAttentionWrapper

示例

>>> 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.empty(128 * 1024 * 1024, dtype=torch.uint8, device="cuda:0")
>>> prefill_wrapper = flashinfer.BatchPrefillWithSharedPrefixPagedKVCacheWrapper(
...     workspace_buffer, "NHD"
... )
>>> batch_size = 7
>>> shared_prefix_len = 8192
>>> 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"
... )
>>> 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)
... ]
>>> shared_k_data_at_layer = [
...     torch.randn(
...         shared_prefix_len, num_kv_heads, head_dim, dtype=torch.float16, device="cuda:0"
...     ) for _ in range(num_layers)
... ]
>>> shared_v_data_at_layer = [
...     torch.randn(
...         shared_prefix_len, num_kv_heads, head_dim, dtype=torch.float16, device="cuda:0"
...     ) for _ in range(num_layers)
... ]
>>> # create auxiliary data structures for batch prefill attention
>>> prefill_wrapper.begin_forward(
...     qo_indptr,
...     paged_kv_indptr,
...     paged_kv_indices,
...     paged_kv_last_page_len,
...     num_qo_heads,
...     num_kv_heads,
...     head_dim,
...     page_size,
... )
>>> outputs = []
>>> for i in range(num_layers):
...     q = torch.randn(nnz_qo, num_qo_heads, head_dim).half().to("cuda:0")
...     kv_cache = kv_cache_at_layer[i]
...     k_shared = shared_k_data_at_layer[i]
...     v_shared = shared_v_data_at_layer[i]
...     # compute batch prefill attention, reuse auxiliary data structures
...     o = prefill_wrapper.forward(
...         q, k_shared, v_shared, kv_cache, causal=True
...     )
...     outputs.append(o)
...
s[0].shape>>> # clear auxiliary data structures
>>> prefill_wrapper.end_forward()
>>> outputs[0].shape
torch.Size([100, 64, 128])

注意

为了加速计算,FlashInfer 的共享前缀批量预填充/追加注意力运算符会创建一些辅助数据结构,这些数据结构可以在同一次预填充/追加步骤中的多次前向调用中重用(例如,不同的 Transformer 层)。这个包装器类管理这些数据结构的生命周期。

__init__(float_workspace_buffer: Tensor, kv_layout: str = 'NHD') None

BatchDecodeWithSharedPrefixPagedKVCacheWrapper 的构造函数。

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

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

begin_forward(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: int, page_size: int) None

为共享前缀批量预填充/追加注意力创建辅助数据结构,用于同一预填充/追加步骤中的多次前向调用。

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

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

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

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

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

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

  • head_dim (int) – 头部的维度。

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

注意

应在调用任何 begin_forward() 方法之前调用 forward()forward_return_lse() 调用,在此调用期间将创建辅助数据结构并缓存以供多次前向调用使用。

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

end_forward() None

警告:此函数已弃用,没有效果

forward(q: Tensor, k_shared: Tensor, v_shared: Tensor, unique_kv_cache: Tensor, causal: bool = False, use_fp16_qk_reduction: bool = False, sm_scale: float | None = None, rope_scale: float | None = None, rope_theta: float | None = None) Tensor

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

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

  • k_shared (torch.Tensor) – 共享前缀键张量,形状:[shared_prefix_len, num_kv_heads, head_dim] 如果 kv_layoutNHD,或者 [num_kv_heads, shared_prefix_len, head_dim] 如果 kv_layoutHND

  • torch.Tensor (v_shared ;) – 共享前缀值张量,形状:[shared_prefix_len, num_kv_heads, head_dim] 如果 kv_layoutNHD,或者 [num_kv_heads, shared_prefix_len, head_dim] 如果 kv_layoutHND

  • unique_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] 是键缓存,paged_kv_cache[:, 1] 是值缓存。

  • causal (bool) – 是否在注意力矩阵上应用因果掩码。

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

  • sm_scale (Optional[float]) – softmax的缩放比例,如果未提供,将设置为 1 / sqrt(head_dim)

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

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

返回值:

V – 注意力输出,形状:[qo_indptr[-1], num_heads, head_dim]

返回值类型:

torch.Tensor

reset_workspace_buffer(float_workspace_buffer: Tensor, int_workspace_buffer: Tensor) None

重置工作区缓冲区。

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

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