flashinfer.cascade¶
合并注意力状态¶
|
合并两个 KV 分段的注意力输出 |
|
就地合并自注意力状态 |
|
合并多个注意力状态 (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 张量的布局,可以是
NHD或HND。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_layout是NHD,以及[max_num_pages, num_kv_heads, page_size, head_dim],如果kv_layout是HND。单个 5D 张量,形状为:
[max_num_pages, 2, page_size, num_kv_heads, head_dim]如果kv_layout是NHD,并且[max_num_pages, 2, num_kv_heads, page_size, head_dim]如果kv_layout是HND。其中paged_kv_cache[:, 0]是键缓存,paged_kv_cache[:, 1]是值缓存。
用于具有共享前缀分页 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 层)中重用。这个包装器类管理这些数据结构的生命周期。
为给定的问题规范规划共享前缀批量解码注意力。
- 参数:
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,该函数将使用 分组查询注意力。
警告:此函数已弃用,没有效果
计算查询和共享前缀分页 kv 缓存之间的批量解码注意力。
- 参数:
q (torch.Tensor) – 查询张量,形状:
[batch_size, num_qo_heads, head_dim]。k_shared (torch.Tensor) – 共享前缀键张量,形状:
[shared_prefix_len, num_kv_heads, head_dim]如果kv_layout是NHD,或者[num_kv_heads, shared_prefix_len, head_dim]如果kv_layout是HND。v_shared (torch.Tensor) – 共享前缀值张量,形状:
[shared_prefix_len, num_kv_heads, head_dim]如果kv_layout是NHD,或者[num_kv_heads, shared_prefix_len, head_dim]如果kv_layout是HND。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_layout是NHD,以及[max_num_pages, num_kv_heads, page_size, head_dim],如果kv_layout是HND。单个 5D 张量,形状为:
[max_num_pages, 2, page_size, num_kv_heads, head_dim]如果kv_layout是NHD,并且[max_num_pages, 2, num_kv_heads, page_size, head_dim]如果kv_layout是HND。其中paged_kv_cache[:, 0]是键缓存,paged_kv_cache[:, 1]是值缓存。
- 返回值:
V – 注意力输出,形状:
[batch_size, num_heads, head_dim]- 返回值类型:
torch.Tensor
重置工作区缓冲区。
- 参数:
float_workspace_buffer (torch.Tensor) – 新的 float 工作区缓冲区,该缓冲区的设备应与输入张量的设备相同。
int_workspace_buffer (torch.Tensor) – 新的 int 工作区缓冲区,该缓冲区的设备应与输入张量的设备相同。
用于批量请求的分页 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 层)。这个包装器类管理这些数据结构的生命周期。
BatchDecodeWithSharedPrefixPagedKVCacheWrapper的构造函数。- 参数:
float_workspace_buffer (torch.Tensor) – 用于在split-k算法中存储中间注意力结果的用户保留的浮点工作区缓冲区。推荐大小为128MB,工作区缓冲区的设备应与输入张量的设备相同。
kv_layout (str) – 输入 k/v 张量的布局,可以是
NHD或HND。
为共享前缀批量预填充/追加注意力创建辅助数据结构,用于同一预填充/追加步骤中的多次前向调用。
- 参数:
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,该函数将使用 分组查询注意力。
警告:此函数已弃用,没有效果
计算查询和共享前缀分页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_layout是NHD,或者[num_kv_heads, shared_prefix_len, head_dim]如果kv_layout是HND。torch.Tensor (v_shared ;) – 共享前缀值张量,形状:
[shared_prefix_len, num_kv_heads, head_dim]如果kv_layout是NHD,或者[num_kv_heads, shared_prefix_len, head_dim]如果kv_layout是HND。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_layout是NHD,以及[max_num_pages, num_kv_heads, page_size, head_dim],如果kv_layout是HND。单个 5D 张量,形状为:
[max_num_pages, 2, page_size, num_kv_heads, head_dim]如果kv_layout是NHD,并且[max_num_pages, 2, num_kv_heads, page_size, head_dim]如果kv_layout是HND。其中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
重置工作区缓冲区。
- 参数:
float_workspace_buffer (torch.Tensor) – 新的 float 工作区缓冲区,该缓冲区的设备应与输入张量的设备相同。
int_workspace_buffer (torch.Tensor) – 新的 int 工作区缓冲区,该缓冲区的设备应与输入张量的设备相同。