[Perf] FlashInfer MLA: remove blocking D2H in spec-decode plan (#27689)

This commit is contained in:
elvischenv
2026-08-12 12:29:03 -07:00
committed by GitHub
parent 3974b00359
commit 4f883636a2
@@ -314,6 +314,18 @@ class FlashInferMLAAttnBackend(AttentionBackend):
self.decode_cuda_graph_metadata = {}
self.prefill_cuda_graph_metadata = {} # For verify
# Pinned host buffers for the fast prefill path plan
if not skip_prefill:
self.fast_plan_qo_indptr_cpu = torch.zeros(
(max_bs + 1,), dtype=torch.int32, device="cpu", pin_memory=True
)
self.fast_plan_kv_indptr_cpu = torch.zeros(
(max_bs + 1,), dtype=torch.int32, device="cpu", pin_memory=True
)
self.fast_plan_kv_len_arr_cpu = torch.zeros(
(max_bs,), dtype=torch.int32, device="cpu", pin_memory=True
)
def init_forward_metadata_out_graph(
self,
forward_batch: ForwardBatch,
@@ -379,6 +391,9 @@ class FlashInferMLAAttnBackend(AttentionBackend):
spec_info=spec_info,
seq_lens_cpu=seq_lens_cpu,
)
if forward_mode.is_target_verify():
# use sync-free fast_mla_prefill_plan for replay
prefill_wrapper.plan = partial(fast_mla_prefill_plan, prefill_wrapper)
else:
self._apply_cuda_graph_metadata(
bs=bs,
@@ -424,6 +439,28 @@ class FlashInferMLAAttnBackend(AttentionBackend):
and not is_in_breakable_cuda_graph()
)
# build host indptr/len arrays for eager DRAFT_EXTEND_V2 fast plan path
qo_indptr_cpu = kv_indptr_cpu = kv_len_arr_cpu = None
spec_info = forward_batch.spec_info
if (
not use_ragged
and forward_batch.forward_mode.is_draft_extend_v2()
and forward_batch.seq_lens_cpu is not None
and spec_info is not None
):
ndt = spec_info.num_tokens_per_req
bs = forward_batch.batch_size
self.fast_plan_qo_indptr_cpu[: bs + 1] = torch.arange(
0, (bs + 1) * ndt, ndt, dtype=torch.int32
)
self.fast_plan_kv_len_arr_cpu[:bs] = forward_batch.seq_lens_cpu[:bs]
self.fast_plan_kv_indptr_cpu[1 : bs + 1] = torch.cumsum(
self.fast_plan_kv_len_arr_cpu[:bs], dim=0
)
qo_indptr_cpu = self.fast_plan_qo_indptr_cpu[: bs + 1]
kv_indptr_cpu = self.fast_plan_kv_indptr_cpu[: bs + 1]
kv_len_arr_cpu = self.fast_plan_kv_len_arr_cpu[:bs]
self.indices_updater_prefill.update(
forward_batch.req_pool_indices,
forward_batch.seq_lens,
@@ -432,6 +469,9 @@ class FlashInferMLAAttnBackend(AttentionBackend):
prefill_wrapper_paged=self.prefill_wrapper_paged,
use_ragged=use_ragged,
attn_dcp_metadata=forward_batch.attn_dcp_metadata,
qo_indptr_cpu=qo_indptr_cpu,
kv_indptr_cpu=kv_indptr_cpu,
kv_len_arr_cpu=kv_len_arr_cpu,
)
self.forward_metadata = PrefillMetadata(
self.prefill_wrapper_paged, use_ragged
@@ -516,6 +556,18 @@ class FlashInferMLAAttnBackend(AttentionBackend):
**self.fast_decode_kwargs,
)
elif forward_mode.is_target_verify():
# build host indptr/len arrays for target-verify fast plan path
assert (
seq_lens_cpu is not None and spec_info is not None
), "target-verify cuda-graph replay requires host-resident seq_lens_cpu"
ndt = spec_info.draft_token_num
self.fast_plan_qo_indptr_cpu[: bs + 1] = torch.arange(
0, (bs + 1) * ndt, ndt, dtype=torch.int32
)
self.fast_plan_kv_len_arr_cpu[:bs] = seq_lens_cpu[:bs] + ndt
self.fast_plan_kv_indptr_cpu[1 : bs + 1] = torch.cumsum(
self.fast_plan_kv_len_arr_cpu[:bs], dim=0
)
self.indices_updater_prefill.update(
req_pool_indices[:bs],
seq_lens[:bs],
@@ -526,6 +578,9 @@ class FlashInferMLAAttnBackend(AttentionBackend):
],
use_ragged=False,
spec_info=spec_info,
qo_indptr_cpu=self.fast_plan_qo_indptr_cpu[: bs + 1],
kv_indptr_cpu=self.fast_plan_kv_indptr_cpu[: bs + 1],
kv_len_arr_cpu=self.fast_plan_kv_len_arr_cpu[:bs],
)
else:
raise ValueError(f"Invalid forward mode: {forward_mode=}")
@@ -853,7 +908,7 @@ class FlashInferMLAIndicesUpdaterPrefill:
def update(
self,
req_pool_indices: torch.Tnesor,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
seq_lens_sum: int,
prefix_lens: torch.Tensor,
@@ -861,6 +916,9 @@ class FlashInferMLAIndicesUpdaterPrefill:
use_ragged: bool,
spec_info: Optional[SpecInput] = None,
attn_dcp_metadata: Optional[DecodeContextParallelMetadata] = None,
qo_indptr_cpu: Optional[torch.Tensor] = None,
kv_indptr_cpu: Optional[torch.Tensor] = None,
kv_len_arr_cpu: Optional[torch.Tensor] = None,
):
if use_ragged:
paged_kernel_lens = prefix_lens
@@ -882,6 +940,9 @@ class FlashInferMLAIndicesUpdaterPrefill:
use_ragged,
spec_info,
attn_dcp_metadata=attn_dcp_metadata,
qo_indptr_cpu=qo_indptr_cpu,
kv_indptr_cpu=kv_indptr_cpu,
kv_len_arr_cpu=kv_len_arr_cpu,
)
def call_begin_forward(
@@ -898,6 +959,9 @@ class FlashInferMLAIndicesUpdaterPrefill:
use_ragged: bool,
spec_info: Optional[SpecInput] = None,
attn_dcp_metadata: Optional[DecodeContextParallelMetadata] = None,
qo_indptr_cpu: Optional[torch.Tensor] = None,
kv_indptr_cpu: Optional[torch.Tensor] = None,
kv_len_arr_cpu: Optional[torch.Tensor] = None,
):
bs = len(seq_lens)
sm_scale = self.scaling
@@ -960,12 +1024,23 @@ class FlashInferMLAIndicesUpdaterPrefill:
kv_indptr = attn_dcp_metadata.dcp_kv_indptr
if attn_dcp_metadata.dcp_kv_indices is not None:
kv_indices = attn_dcp_metadata.dcp_kv_indices
kv_len_arr = kv_indptr[1:] - kv_indptr[:-1]
# DCP splits kv across CP ranks; the host-side fast-plan arrays are
# built from full seq_lens and don't match, so fall back to GPU.
qo_indptr_cpu = kv_indptr_cpu = kv_len_arr_cpu = None
plan_qo_indptr = qo_indptr if qo_indptr_cpu is None else qo_indptr_cpu
plan_kv_indptr = kv_indptr if kv_indptr_cpu is None else kv_indptr_cpu
plan_kv_len_arr = (
kv_indptr[1:] - kv_indptr[:-1]
if kv_len_arr_cpu is None
else kv_len_arr_cpu
)
wrapper_paged.plan(
qo_indptr,
kv_indptr,
plan_qo_indptr,
plan_kv_indptr,
kv_indices,
kv_len_arr,
plan_kv_len_arr,
self.num_local_heads,
self.kv_lora_rank,
self.qk_rope_head_dim,
@@ -1178,3 +1253,50 @@ def fast_mla_decode_plan(
)
except Exception as e:
raise RuntimeError(f"Error in alternate MLA plan: {e}")
def fast_mla_prefill_plan(
self,
qo_indptr_cpu: torch.Tensor,
kv_indptr_cpu: torch.Tensor,
kv_indices: torch.Tensor,
kv_len_arr_cpu: torch.Tensor,
num_heads: int,
head_dim_ckv: int,
head_dim_kpe: int,
page_size: int,
causal: bool,
sm_scale: float,
q_data_type: torch.dtype,
kv_data_type: torch.dtype,
) -> None:
"""Sync-free BatchMLAPagedAttentionWrapper.plan for the target-verify CUDA
graph replay. Like fast_mla_decode_plan it hands host-known qo/kv indptr +
lengths straight to _cached_module.plan (no per-replay device-to-host copy).
Decode's indices updater writes the cuda-graph buffers in place so its fast
plan can skip them; verify metadata is freshly built each step, so refresh
the bound buffers here exactly as stock plan()'s use_cuda_graph branch does
(host->device / device->device, non-blocking).
"""
self._causal = causal
self._page_size = page_size
self._sm_scale = sm_scale
self._qo_indptr_buf.copy_(qo_indptr_cpu, non_blocking=True)
self._kv_indptr_buf.copy_(kv_indptr_cpu, non_blocking=True)
self._kv_indices_buf[: len(kv_indices)].copy_(kv_indices, non_blocking=True)
self._kv_len_arr_buf.copy_(kv_len_arr_cpu, non_blocking=True)
try:
self._cached_module.plan(
self._float_workspace_buffer,
self._int_workspace_buffer,
self._pin_memory_int_workspace_buffer,
qo_indptr_cpu,
kv_indptr_cpu,
kv_len_arr_cpu,
num_heads,
head_dim_ckv,
causal,
)
except Exception as e:
raise RuntimeError(f"Error in alternate MLA prefill plan: {e}")