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/xqatrtllm-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 张量核心操作支持此模式。

如果同时提供两者,将使用动态缩放比例张量。