flashinfer.mla.trtllm_batch_decode_with_kv_cache_mla¶
- flashinfer.mla.trtllm_batch_decode_with_kv_cache_mla(query: Tensor, kv_cache: Tensor, workspace_buffer: Tensor, qk_nope_head_dim: int, kv_lora_rank: int, qk_rope_head_dim: int, block_tables: Tensor, seq_lens: Tensor, max_seq_len: int, sparse_mla_top_k: int = 0, out: Tensor | None = None, bmm1_scale: float | Tensor = 1.0, bmm2_scale: float | Tensor = 1.0, sinks: List[Tensor] | None = None, enable_pdl: bool = None, backend: str = 'auto') Tensor¶
- 参数:
query ([batch_size, q_len_per_request, num_heads, head_dim_qk], head_dim_qk = qk_nope_head_dim (kv_lora_rank) + qk_rope_head_dim, 应该拼接 q_nope + q_rope; q_len_per_request 是 MTP 查询长度。)
kv_cache ([num_pages, page_size, head_dim_ckv + head_dim_kpe] or [num_pages, 1, page_size, head_dim_ckv + head_dim_kpe], 应该拼接 ckv_cache + kpe_cache。为了向后兼容,支持 3D 和 4D 格式。)
workspace_buffer ([num_semaphores, 4], 用于多块模式。首次使用时必须初始化为 0。)
qk_nope_head_dim (qk_nope_head_dim, 必须为 128)
kv_lora_rank (kv_lora_rank, 必须为 512)
qk_rope_head_dim (qk_rope_head_dim, 必须为 64)
sparse_mla_top_k (稀疏 MLA top k, 对于非稀疏 MLA 必须为 0。)
block_tables (kv 缓存的page_table, [batch_size, num_pages])
seq_lens (query_len)
max_seq_len (kv_cache 的最大序列长度)
out (输出张量, 如果未提供, 将在内部分配)
bmm1_scale (mla bmm1 输入的融合缩放比例。) – 当使用 trtllm-gen 后端时,它可以是 dtype 为 torch.float32 的 torch.Tensor。
bmm2_scale (mla bmm2 输入的融合缩放比例。) – 当使用 trtllm-gen 后端时,它可以是 dtype 为 torch.float32 的 torch.Tensor。
sinks (softmax 分母中的每个 head 的附加值。)
backend (str = "auto") – 实现后端,可以是
auto/xqa或trtllm-gen。默认为auto。设置为auto时,后端将根据设备架构和内核可用性进行选择。对于 sm_100 和 sm_103(blackwell 架构),auto将选择trtllm-gen后端。对于 sm_120(blackwell 架构),auto将选择xqa后端。
注意
在 MLA 中,实际应用的 BMM1 和 BMM2 缩放比例将融合为:bmm1_scale = q_scale * k_scale * sm_scale / (head_dim_qk ** 0.5) bmm2_scale = v_scale * o_scale 或者,bmm1_scale = torch.Tensor([q_scale * k_scale * sm_scale / (head_dim_qk ** 0.5)) bmm2_scale = torch.Tensor([v_scale * o_scale])
对于 cuda 图捕获,这两个缩放比例因子应该是静态常量。应提供 (bmm1_scale, bmm2_scale) 或 (bmm1_scale_log2_tensor, bmm2_scale_tensor) 中的一个。
- 对于静态常量缩放比例因子,应将缩放比例因子作为浮点数提供。
(bmm1_scale, bmm2_scale)
- 对于设备上的融合缩放比例张量,这些张量可以动态变化,应将缩放比例因子作为 torch.Tensor 提供。
(bmm1_scale_log2_tensor, bmm2_scale_tensor)
目前,只有 fp8 张量核心操作支持此模式。
如果同时提供两者,将使用动态缩放比例张量。