flashinfer.prefill.trtllm_batch_context_with_kv_cache¶
- flashinfer.prefill.trtllm_batch_context_with_kv_cache(query: Tensor, kv_cache: Tensor | Tuple[Tensor, Tensor], workspace_buffer: Tensor, block_tables: Tensor, seq_lens: Tensor, max_q_len: int, max_kv_len: int, bmm1_scale: float | Tensor, bmm2_scale: float | Tensor, batch_size: int, cum_seq_lens_q: Tensor, cum_seq_lens_kv: Tensor, 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, kv_layout: str = 'HND', enable_pdl: bool | None = None, sinks: List[Tensor] | None = None) Tensor | FP4Tensor¶
- 参数:
query (torch.Tensor) – query 张量,形状为 [num_tokens, num_heads, head_dim]
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 缓存,第二个张量是 value 缓存。workspace_buffer (torch.Tensor. 首次使用时必须初始化为 0。) – 工作区
block_tables (torch.Tensor) – kv 缓存的页表,[batch_size, num_pages]
seq_lens (torch.Tensor) – 一个 uint32 一维张量,指示每个 prompt 的 kv 序列长度。形状:
[batch_size]max_q_len (int) – query 的最大序列长度
max_kv_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。
batch_size (int) – 批大小
cum_seq_lens_q (torch.Tensor) – query 的累积序列长度。形状:
[batch_size + 1]cum_seq_lens_kv (torch.Tensor) – kv_cache 的累积序列长度。形状:
[batch_size + 1]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 输出张量比例因子的向量大小。
enable_pdl (Optional[bool] = None) – 是否启用程序依赖启动 (PDL)。请参阅 https://docs.nvda.net.cn/cuda/cuda-c-programming-guide/#programmatic-dependent-launch-and-synchronization 默认值为
None,这意味着如果设备支持 PDL,则将启用它。kv_layout (str = "HND") – kv-cache 的布局,可以是“HND”或“NHD”,默认值为“HND”。
sinks (Optional[List[torch.Tensor]] = None) – softmax 分母中的每个头的附加值。
- 返回值:
out – 输出 torch.Tensor 或 FP4Tensor。
- 返回值类型:
Union[torch.Tensor, FP4Tensor]