flashinfer.prefill.cudnn_batch_prefill_with_kv_cache

flashinfer.prefill.cudnn_batch_prefill_with_kv_cache(q: Tensor, k_cache: Tensor, v_cache: Tensor, scale: float, workspace_buffer: Tensor, *, max_token_per_sequence: int, max_sequence_kv: int, actual_seq_lens_q: Tensor, actual_seq_lens_kv: Tensor, block_tables: Tensor | None = None, causal: bool, return_lse: bool, q_scale: Tensor | None = None, k_scale: Tensor | None = None, v_scale: Tensor | None = None, batch_offsets_q: Tensor | None = None, batch_offsets_o: Tensor | None = None, batch_offsets_k: Tensor | None = None, batch_offsets_v: Tensor | None = None, batch_offsets_stats: Tensor | None = None, out: Tensor | None = None, lse: Tensor | None = None, is_cuda_graph_compatible: bool = False, backend: str | None = None, o_data_type: dtype | None = None) tuple[Tensor, Tensor | None]

使用 cuDNN 执行具有分页 KV 缓存的批量预填充注意力。

参数:
  • q – 查询张量,形状为 (总 token 数,num_heads_qo,head_dim)

  • k_cache – 键缓存张量,如果启用分页 KV 缓存,形状为 (total_num_pages,num_heads_kv,page_size,head_dim),否则形状为 (总 KV 序列长度,num_heads_kv,d_qk)

  • v_cache – 值缓存张量,如果启用分页 KV 缓存,形状为 (total_num_pages,num_heads_kv,page_size,head_dim),否则形状为 (总 KV 序列长度,num_heads_kv,d_vo)

  • scale – 注意力分数的缩放因子,通常为 1/sqrt(head_dim)

  • workspace_buffer – cuDNN 操作的工作空间缓冲区。随批次大小缩放。大多数情况下,128 MB 应该足够

  • max_token_per_sequence – 每个查询序列的最大 token 数 (s_qo_max)

  • max_sequence_kv – 每个键/值序列的最大 token 数 (s_kv_max)

  • actual_seq_lens_q – 每个查询序列的实际 token 数,形状为 (batch_size,),位于 cpu 或设备上 (如果 cuda_graph 为 False,则位于 cpu 上)

  • actual_seq_lens_kv – 每个批次中键/值的实际序列长度,形状为 (batch_size,),位于 CPU 或设备上 (如果 cuda_graph 为 False,则位于 cpu 上)

  • block_tables – KV 缓存的分页表,形状为 (batch_size,num_pages_per_seq),位于 GPU 上

  • causal – 是否应用因果掩码

  • return_lse – 是否返回对数和指数值 (必须为 True)

  • out – 可选的预分配输出张量

  • lse – 如果 return_lse 为 True,则可选的预分配对数和指数值张量,否则返回 None

  • is_cuda_graph_compatible – 预填充操作是否与 CUDA 图兼容

  • q_scale – 查询张量的可选缩放张量,形状为 (1, 1, 1, 1),位于 GPU 上

  • k_scale – 键张量的可选缩放张量,形状为 (1, 1, 1, 1),位于 GPU 上

  • v_scale – 值张量的可选缩放张量,形状为 (1, 1, 1, 1),位于 GPU 上

  • batch_offsets_q – 查询张量的可选批次偏移量,形状为 (batch_size,),位于 GPU 上

  • batch_offsets_o – 输出张量的可选批次偏移量,形状为 (batch_size,),位于 GPU 上

  • batch_offsets_k – 键张量的可选批次偏移量,形状为 (batch_size,),位于 GPU 上

  • batch_offsets_v – 值张量的可选批次偏移量,形状为 (batch_size,),位于 GPU 上

  • o_data_type – 输出张量的可选数据类型

返回值:

输出张量,形状为 (batch_size * seq_len_q,num_heads_qo,head_dim)。如果 return_lse 为 True,则还返回形状为 (batch_size,seq_len_q,num_heads_qo) 的对数和指数值张量

注意

查询和 KV 头可以具有不同的尺寸 (num_heads_qo >= num_heads_kv)。在使用 cuda 图时,actual_seq_lens_q 和 actual_seq_lens_kv 必须位于与 q 相同的设备上。查询和键的头维度必须为 128 或 192。值和输出的头维度必须为 128