flashinfer.decode.trtllm_batch_decode_with_kv_cache

flashinfer.decode.trtllm_batch_decode_with_kv_cache(query: Tensor, kv_cache: Tensor | Tuple[Tensor, Tensor], workspace_buffer: Tensor, block_tables: Tensor, seq_lens: Tensor, max_seq_len: int, bmm1_scale: float | Tensor = 1.0, bmm2_scale: float | Tensor = 1.0, window_left: int = -1, out: Tensor | FP4Tensor | None = None, out_dtype: str | dtype | None = None, o_sf_scale: float | None = None, o_sf_vec_size: int | None = None, sinks: List[Tensor] | None = None, kv_layout: str = 'HND', enable_pdl: bool | None = None, backend: str = 'auto', q_len_per_req: int | None = 1, o_scale: float | None = 1.0, mask: Tensor | None = None, max_q_len: int | None = None, cum_seq_lens_q: Tensor | None = None) Tensor | FP4Tensor
参数:
  • query (torch.Tensor) – query 张量,形状为 [num_tokens, num_heads, head_dim],num_tokens = batch 中总的 query tokens 数量。

  • kv_cache (Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]) – 如果 kv_cache 是单个张量,则它应该是一个形状为 [num_pages, 1 或 2, num_kv_heads, page_size, head_dim] 的张量,如果 kv_layoutHND,或者如果 kv_layoutNHD,则形状为 [num_pages, 1 或 2, page_size, num_kv_heads, head_dim]。如果 kv_cache 是两个张量的元组,则它应该是一个形状为 [num_pages, num_kv_heads, page_size, head_dim] 的两个张量的元组,如果 kv_layoutHND,或者如果 kv_layoutNHD,则形状为 [num_pages, page_size, num_kv_heads, head_dim]。第一个张量是 key cache,第二个张量是 value cache。

  • workspace_buffer (torch.Tensor. 必须在首次使用时初始化为 0。) – 工作区

  • block_tables (torch.Tensor) – kv cache 的页表,[batch_size, num_pages]

  • seq_lens (torch.Tensor) – 一个 uint32 1D 张量,指示每个 prompt 的 kv 序列长度。形状:[batch_size]

  • max_seq_len (int) – kv_cache 的最大序列长度

  • bmm1_scale (Union[float, torch.Tensor]) – bmm1 输入的融合缩放。在使用 trtllm-gen 后端时,它可以是 dtype 为 torch.float32 的 torch.Tensor。

  • bmm2_scale (Union[float, torch.Tensor]) – bmm2 输入的融合缩放。在使用 trtllm-gen 后端时,它可以是 dtype 为 torch.float32 的 torch.Tensor。

  • window_left (int = -1) – 注意力窗口的左(包含)窗口大小,设置为 -1 时,窗口大小将设置为序列的完整长度。默认值为 -1

  • out (Optional[Union[torch.Tensor, FP4Tensor]] = None) – 输出张量,如果未提供,将使用 out_dtype 分配,如果未提供 out_dtype,将使用 query 的类型。

  • out_dtype (Optional[Union[torch.dtype, str]] = None) – 输出 dtype,如果未提供,将使用 out 的类型。对于 nvfp4,使用字符串 nvfp4

  • o_sf_scale (Optional[float] = None) – nvfp4 输出张量比例因子的缩放。

  • o_sf_vec_size (Optional[int] = None) – nvfp4 输出张量比例因子的向量大小。

  • sinks (Optional[List[torch.Tensor]] = None) – softmax 分母中的每个 head 的附加值。

  • kv_layout (str = "HND") – 输入 k/v 张量的布局,可以是 NHDHND。默认值为 HND

  • enable_pdl (Optional[bool] = None) – 是否启用程序依赖启动 (PDL)。请参阅 https://docs.nvda.net.cn/cuda/cuda-c-programming-guide/#programmatic-dependent-launch-and-synchronization。设置为 None 时,后端将基于设备架构和内核可用性进行选择。

  • backend (str = "auto") – 后端实现,可以是 auto/xqatrtllm-gen。默认为 auto。当设置为 auto 时,后端将根据设备架构和内核可用性进行选择。对于 sm_100 和 sm_103(blackwell 架构),auto 将选择 trtllm-gen 后端。对于 sm_90(hopper 架构)和 sm_120(blackwell 架构),auto 将选择 xqa 后端。

  • o_scale (Optional[float] = 1.0) – xqa fp8 输出的输出缩放因子。

  • mask (Optional[torch.Tensor] = None) – xqa 推测解码的因果注意力掩码。

  • max_q_len (Optional[int] = None) – 使用可变长度查询时,所有请求中的最大查询序列长度。仅由 trtllm-gen 后端支持。必须与 cum_seq_lens_q 一起提供。如果为 None,则所有请求使用由 q_len_per_req 指定的统一查询长度。

  • cum_seq_lens_q (Optional[torch.Tensor] = None) – 可变长度查询支持的累积查询序列长度,形状:[batch_size + 1],dtype:torch.int32。仅由 trtllm-gen 后端支持。必须与 max_q_len 一起提供。如果为 None,则所有请求使用由 q_len_per_req 指定的统一查询长度。

返回值:

out – 输出 torch.Tensor 或 FP4Tensor。

返回值类型:

Union[torch.Tensor, FP4Tensor]