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_layout是HND,或者如果kv_layout是NHD,则形状为 [num_pages, 1 或 2, page_size, num_kv_heads, head_dim]。如果 kv_cache 是两个张量的元组,则它应该是一个形状为 [num_pages, num_kv_heads, page_size, head_dim] 的两个张量的元组,如果kv_layout是HND,或者如果kv_layout是NHD,则形状为 [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 张量的布局,可以是
NHD或HND。默认值为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/xqa或trtllm-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]