flashinfer.xqa.xqa¶
- flashinfer.xqa.xqa(q: Tensor, k_cache: Tensor, v_cache: Tensor, page_table: Tensor, seq_lens: Tensor, output: Tensor, workspace_buffer: Tensor, semaphores: Tensor, num_kv_heads: int, page_size: int, sinks: Tensor | None = None, q_scale: float | Tensor = 1.0, kv_scale: float | Tensor = 1.0, sliding_win_size: int = 0, kv_layout: str = 'NHD', sm_count: int | None = None, enable_pdl: bool | None = None, rcp_out_scale: float = 1.0, q_seq_len: int = 1, mask: Tensor | None = None) None¶
使用 XQA 内核应用带有分页 KV 缓存的注意力。 :param q: 查询张量,形状为
[batch_size, beam_width, num_q_heads, head_dim](如果未使用推测解码),或者
[batch_size, beam_width, q_seq_len, num_q_heads, head_dim](如果使用推测解码)。q_seq_len是推测解码令牌的数量。数据类型应为 torch.float16 或 torch.bfloat16。现在仅支持 beam_width 1。- 参数:
k_cache (torch.Tensor) – 分页 K 缓存张量,形状为
[num_pages, page_size, num_kv_heads, head_dim](如果kv_layout为NHD),或者[num_pages, num_kv_heads, page_size, head_dim](如果kv_layout为HND)。数据类型应与查询张量匹配,或为 torch.float8_e4m3fn,在这种情况下 xqa 将运行 fp8 计算。应与 v_cache 具有相同的数据类型。v_cache (torch.Tensor) – 分页 V 缓存张量,形状为
[num_pages, page_size, num_kv_heads, head_dim](如果kv_layout为NHD),或者[num_pages, num_kv_heads, page_size, head_dim](如果kv_layout为HND)。数据类型应与查询张量匹配,或为 torch.float8_e4m3fn,在这种情况下 xqa 将运行 fp8 计算。应与 k_cache 具有相同的数据类型。page_table (torch.Tensor) – 页表张量,形状为
batch_size, nb_pages_per_seq。数据类型应为 torch.int32。K 和 V 共享相同的表。seq_lens (torch.Tensor) – 序列长度张量,形状为
[batch_size, beam_width]。数据类型应为 torch.uint32。output (torch.Tensor) – 输出张量,形状与查询张量匹配。数据类型应与查询张量或 kv 张量匹配。此张量将被就地修改。
workspace_buffer (torch.Tensor) – 用于临时计算的工作区缓冲区。数据类型应为 torch.uint8。
semaphores (torch.Tensor) – 用于同步的信号量缓冲区。数据类型应为 torch.uint32。
num_kv_heads (int) – 注意力机制中的键值头数。
page_size (int) – 分页 KV 缓存中每个页面的大小。必须是 [16, 32, 64, 128] 中的一个。
sinks (Optional[torch.Tensor], default=None) – 注意力 sink 值,形状为
[num_kv_heads, head_group_ratio]。数据类型应为 torch.float32。如果为 None,则不使用注意力 sink。q_scale (Union[float, torch.Tensor], default=1.0) – 查询张量的缩放因子。
kv_scale (Union[float, torch.Tensor], default=1.0) – KV 缓存的缩放因子。
sliding_win_size (int, default=0) – 注意力的滑动窗口大小。如果为 0,则不使用滑动窗口。
kv_layout (str, default="NHD") – KV 缓存的布局。可以是
NHD或HND。sm_count (Optional[int], default=None) – 要使用的流式多处理器的数量。如果为 None,将从设备推断。
enable_pdl (Optional[bool], default=None) – 是否启用 PDL(持久数据加载器)优化。如果为 None,如果硬件支持,则设置为 True。
rcp_out_scale (float, default=1.0) – 输出缩放因子的倒数。
q_seq_len (int, default=1) – 查询序列长度。当 > 1 时,启用推测解码模式。
mask (Optional[torch.Tensor], default=None) – 推测解码模式(当
q_seq_len > 1时)的因果注意力掩码。形状:[batch_size, q_seq_len, mask_size_per_row],其中mask_size_per_row = ((q_seq_len + 31) // 32) * 2。数据类型应为 torch.uint16(位打包格式,对齐到 32 位)。
注意
该函数会自动从张量形状推断几个参数:- batch_size 来自 q.shape[0] - num_q_heads 来自 q.shape[-2] - head_dim 来自 q.shape[-1] - input_dtype 来自 q.dtype - kv_cache_dtype 来自 k.dtype - head_group_ratio 来自 num_q_heads // num_kv_heads - max_seq_len 来自 page_table.shape[-1] * page_size