flashinfer.decode.cudnn_batch_decode_with_kv_cache

flashinfer.decode.cudnn_batch_decode_with_kv_cache(q: Tensor, k_cache: Tensor, v_cache: Tensor, scale: float, workspace_buffer: Tensor, *, max_sequence_kv: int, actual_seq_lens_kv: Tensor | None = None, block_tables: Tensor | None = None, is_cuda_graph_compatible: bool = False, batch_offsets_q: Tensor | None = None, batch_offsets_o: Tensor | None = None, batch_offsets_k: Tensor | None = None, batch_offsets_v: Tensor | None = None, out: Tensor | None = None) Tensor

使用 cuDNN 和分页 KV 缓存执行批量解码注意力。

参数:
  • q – 查询张量,形状为 (batch_size, num_heads_qo, head_dim),seq_len_q 是批处理中查询的最大序列长度

  • k_cache – 键缓存张量,形状为 (total_num_pages, num_heads_kv, page_size, head_dim)

  • v_cache – 值缓存张量,形状为 (total_num_pages, num_heads_kv, page_size, head_dim)

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

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

  • max_sequence_kv – 键/值序列的最大令牌数 (s_kv_max)

  • actual_seq_lens_kv – 批处理中键/值的实际序列长度,形状为 (batch_size,),位于 CPU 上

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

  • is_cuda_graph_compatible – 解码操作是否与 CUDA 图兼容

  • batch_offsets – 可选的批处理偏移张量,形状为 (batch_size,),位于 GPU 上

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

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

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

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

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

返回值:

输出张量,形状为 (batch_size, num_heads_qo, head_dim)

注意

目前仅支持因果注意力 (causal 必须为 True) 所有张量必须是连续的并且位于同一 CUDA 设备上 查询和 KV 头可以具有不同的大小 (num_heads_qo >= num_heads_kv)