flashmla: sync-free spec via device-side draft-extend (#31090)

This commit is contained in:
Liangsheng Yin
2026-07-16 19:38:12 -07:00
committed by GitHub
parent 8432eafd3d
commit 9910ef8167
3 changed files with 148 additions and 50 deletions
@@ -37,19 +37,30 @@ class FlashMLADecodeMetadata:
flashmla_metadata: Optional[Tuple[torch.Tensor, torch.Tensor]] = None
num_splits: Optional[torch.Tensor] = None
block_kv_indices: Optional[torch.Tensor] = None
# K lens the kernel reads for draft-extend (window-aligned, device int32).
# Decode/verify compute theirs inline from forward_batch.seq_lens.
seq_lens_k: Optional[torch.Tensor] = None
def __init__(
self,
flashmla_metadata: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
num_splits: Optional[torch.Tensor] = None,
block_kv_indices: Optional[torch.Tensor] = None,
seq_lens_k: Optional[torch.Tensor] = None,
):
self.flashmla_metadata = flashmla_metadata
self.num_splits = num_splits
self.block_kv_indices = block_kv_indices
self.seq_lens_k = seq_lens_k
class FlashMLABackend(FlashInferMLAAttnBackend):
# Decode/verify/draft-extend metadata is built device-side and the
# tree-mask scratch is preallocated, so no seq_lens_cpu / seq_lens_sum
# D2H is needed. Prefill (EXTEND) goes through the FlashInferMLA parent,
# whose batches always carry the CPU mirror from the scheduler.
needs_cpu_seq_lens: bool = False
def __init__(
self,
model_runner: ModelRunner,
@@ -89,6 +100,14 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
self.cuda_graph_num_splits = None
self.cuda_graph_mla_metadata_view = None
self.cuda_graph_num_splits_view = None
# Static K-lens buffer bound by the draft-extend graph kernel.
self.cuda_graph_draft_extend_seq_lens_k = None
# Preallocated tree-mask scratch (see get_verify_buffers_to_fill_after_draft).
self.cuda_graph_custom_mask = None
self._eager_kv_indices_buf = None
# The worker fetches the tree-mask scratch from the target backend
# only; draft-side instances must not allocate it.
self.is_draft_runner = model_runner.is_draft_worker
# get dcp info
self.dcp_world_size = get_parallel().attn_dcp_size
@@ -100,7 +119,11 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
in_capture: bool = False,
):
forward_mode = forward_batch.forward_mode
if forward_mode.is_decode_or_idle() or forward_mode.is_target_verify():
if (
forward_mode.is_decode_or_idle()
or forward_mode.is_target_verify()
or forward_mode.is_draft_extend_v2()
):
self._apply_decode_target_verify_metadata(
bs=forward_batch.batch_size,
req_pool_indices=forward_batch.req_pool_indices,
@@ -113,18 +136,33 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
forward_batch, in_capture=in_capture
)
def _eager_block_kv_indices(self, bs: int, max_seqlen_pad: int) -> torch.Tensor:
"""Reused eager block-table scratch, grown to the running max shape.
The fill kernel rewrites each row up to its current K len and the
attention kernel reads no further, so stale tail content is unread.
"""
buf = self._eager_kv_indices_buf
if buf is None or buf.shape[0] < bs or buf.shape[1] < max_seqlen_pad:
rows = bs if buf is None else max(bs, buf.shape[0])
cols = max_seqlen_pad if buf is None else max(max_seqlen_pad, buf.shape[1])
buf = torch.full((rows, cols), -1, dtype=torch.int32, device=self.device)
self._eager_kv_indices_buf = buf
return buf[:bs, :max_seqlen_pad]
def init_forward_metadata(self, forward_batch: ForwardBatch):
bs = forward_batch.batch_size
# Host max only sizes the block table: CPU mirror when published,
# else the static bound (kernel reads are capped by seq_lens_k).
seq_lens_cpu = forward_batch.seq_lens_cpu
eager_max_k = (
seq_lens_cpu.max().item()
if seq_lens_cpu is not None
else self.max_context_len
)
if forward_batch.forward_mode.is_decode_or_idle():
max_seqlen_pad = triton.cdiv(
forward_batch.seq_lens_cpu.max().item(), PAGE_SIZE
)
block_kv_indices = torch.full(
(bs, max_seqlen_pad),
-1,
dtype=torch.int32,
device=forward_batch.seq_lens.device,
)
max_seqlen_pad = triton.cdiv(eager_max_k, PAGE_SIZE)
block_kv_indices = self._eager_block_kv_indices(bs, max_seqlen_pad)
create_flashmla_kv_indices_triton[
(bs, get_num_kv_index_blocks_flashmla(max_seqlen_pad, PAGE_SIZE))
](
@@ -134,7 +172,7 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
None,
block_kv_indices,
self.req_to_token.stride(0),
max_seqlen_pad,
block_kv_indices.stride(0),
)
mla_metadata, num_splits = get_mla_metadata(
forward_batch.seq_lens.to(torch.int32),
@@ -148,16 +186,10 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
block_kv_indices,
)
elif forward_batch.forward_mode.is_target_verify():
seq_lens_cpu = forward_batch.seq_lens_cpu + self.num_draft_tokens
seq_lens = forward_batch.seq_lens + self.num_draft_tokens
max_seqlen_pad = triton.cdiv(seq_lens_cpu.max().item(), PAGE_SIZE)
block_kv_indices = torch.full(
(bs, max_seqlen_pad),
-1,
dtype=torch.int32,
device=seq_lens.device,
)
max_seqlen_pad = triton.cdiv(eager_max_k + self.num_draft_tokens, PAGE_SIZE)
block_kv_indices = self._eager_block_kv_indices(bs, max_seqlen_pad)
create_flashmla_kv_indices_triton[
(bs, get_num_kv_index_blocks_flashmla(max_seqlen_pad, PAGE_SIZE))
](
@@ -167,7 +199,7 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
None,
block_kv_indices,
self.req_to_token.stride(0),
max_seqlen_pad,
block_kv_indices.stride(0),
)
mla_metadata, num_splits = get_mla_metadata(
seq_lens.to(torch.int32),
@@ -180,6 +212,39 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
num_splits,
block_kv_indices,
)
elif forward_batch.forward_mode.is_draft_extend_v2():
# Fixed-q draft-extend: every q window is num_draft_tokens wide
# (prepare_for_draft_extend pads the inputs); K lens window-aligned.
window = self.num_draft_tokens
seq_lens_k = (
forward_batch.seq_lens - forward_batch.extend_seq_lens + window
).to(torch.int32)
max_seqlen_pad = triton.cdiv(eager_max_k + window, PAGE_SIZE)
block_kv_indices = self._eager_block_kv_indices(bs, max_seqlen_pad)
create_flashmla_kv_indices_triton[
(bs, get_num_kv_index_blocks_flashmla(max_seqlen_pad, PAGE_SIZE))
](
self.req_to_token,
forward_batch.req_pool_indices,
seq_lens_k,
None,
block_kv_indices,
self.req_to_token.stride(0),
block_kv_indices.stride(0),
)
mla_metadata, num_splits = get_mla_metadata(
seq_lens_k,
window * self.num_q_heads,
1,
is_fp8_kvcache=self.is_fp8_kvcache,
)
self.forward_metadata = FlashMLADecodeMetadata(
mla_metadata,
num_splits,
block_kv_indices,
seq_lens_k,
)
else:
super().init_forward_metadata(forward_batch)
@@ -216,6 +281,22 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
self.cuda_graph_mla_metadata_view = None
self.cuda_graph_num_splits_view = None
if self.num_draft_tokens:
self.cuda_graph_draft_extend_seq_lens_k = torch.ones(
max_bs, dtype=torch.int32, device="cuda"
)
if not self.skip_prefill and not self.is_draft_runner:
# Worst-case FULL_MASK tree-mask scratch (bool); build_tree
# writes it in-place so the GPU-only path needs no seq_lens_sum.
self.cuda_graph_custom_mask = torch.zeros(
max_num_tokens * (self.max_context_len + self.num_draft_tokens),
dtype=torch.bool,
device="cuda",
)
def get_verify_buffers_to_fill_after_draft(self):
return [self.cuda_graph_custom_mask, None]
def _apply_decode_target_verify_metadata(
self,
bs: int,
@@ -224,12 +305,7 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
seq_lens_cpu: Optional[torch.Tensor],
forward_mode: ForwardMode,
):
"""Shared decode/target-verify capture+replay body.
Public entry: :py:meth:`init_forward_metadata_out_graph` (which routes
to this helper for decode/target-verify and falls back to the
FlashInferMLA parent for prefill/draft-extend).
"""
"""Shared decode/target-verify/draft-extend capture+replay body."""
if True:
seq_lens = seq_lens[:bs]
seq_lens_cpu = seq_lens_cpu[:bs] if seq_lens_cpu is not None else None
@@ -238,13 +314,15 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
seq_lens = seq_lens + self.num_draft_tokens
if seq_lens_cpu is not None:
seq_lens_cpu = seq_lens_cpu + self.num_draft_tokens
# draft_extend_v2 graph batches arrive with the padded q window
# already included in seq_lens; use them as-is.
seq_max = (
seq_lens_cpu.max().item()
if seq_lens_cpu is not None
else seq_lens.max().item()
)
max_seqlen_pad = triton.cdiv(seq_max, PAGE_SIZE)
# Tight block-table slice when the CPU mirror is free; static
# bound otherwise (no D2H; kernel reads are capped by seq_lens_k).
if seq_lens_cpu is not None:
max_seqlen_pad = triton.cdiv(seq_lens_cpu.max().item(), PAGE_SIZE)
else:
max_seqlen_pad = self.cuda_graph_kv_indices.shape[1]
create_flashmla_kv_indices_triton[
(
@@ -264,7 +342,9 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
)
q_head_mult = (
self.num_draft_tokens if forward_mode.is_target_verify() else 1
self.num_draft_tokens
if forward_mode.is_target_verify() or forward_mode.is_draft_extend_v2()
else 1
)
mla_metadata, num_splits = get_mla_metadata(
seq_lens.to(torch.int32),
@@ -299,10 +379,17 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
self.cuda_graph_mla_metadata[:actual_num_sm_parts].copy_(mla_metadata)
self.cuda_graph_num_splits[: bs + 1].copy_(num_splits)
seq_lens_k = None
if forward_mode.is_draft_extend_v2():
# The graph kernel binds this static buffer; refresh per replay.
self.cuda_graph_draft_extend_seq_lens_k[:bs].copy_(seq_lens)
seq_lens_k = self.cuda_graph_draft_extend_seq_lens_k[:bs]
self.forward_metadata = FlashMLADecodeMetadata(
self.cuda_graph_mla_metadata_view,
self.cuda_graph_num_splits_view,
self.cuda_graph_kv_indices[:bs, :max_seqlen_pad],
seq_lens_k,
)
def get_cuda_graph_seq_len_fill_value(self):
@@ -398,12 +485,10 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
forward_batch: ForwardBatch,
save_kv_cache: bool = True,
):
if forward_batch.forward_mode in (
ForwardMode.EXTEND,
ForwardMode.DRAFT_EXTEND_V2,
):
if forward_batch.forward_mode == ForwardMode.EXTEND:
return super().forward_extend(q, k, v, layer, forward_batch, save_kv_cache)
else:
# target_verify / draft_extend_v2: fixed-q decode-style kernel.
cache_loc = forward_batch.out_cache_loc
if k is not None:
@@ -414,7 +499,18 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
bs = forward_batch.batch_size
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
reshape_q = q.view(bs, -1, layer.tp_q_head_num, layer.head_dim)
if forward_batch.forward_mode.is_draft_extend_v2():
# prepare_for_draft_extend always emits the fixed q window.
window = self.num_draft_tokens
q_3d = q.view(-1, layer.tp_q_head_num, layer.head_dim)
assert q_3d.shape[0] == bs * window
reshape_q = q_3d.view(bs, window, *q_3d.shape[1:])
cache_seqlens = self.forward_metadata.seq_lens_k
else:
reshape_q = q.view(bs, -1, layer.tp_q_head_num, layer.head_dim)
cache_seqlens = (
forward_batch.seq_lens.to(torch.int32) + self.num_draft_tokens
)
if self.is_fp8_kvcache:
if layer.k_scale is not None:
q_scale = layer.k_scale
@@ -439,8 +535,7 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
q=reshape_q_fp8,
k_cache=k_cache.view(-1, PAGE_SIZE, 1, self.kv_cache_dim),
block_table=self.forward_metadata.block_kv_indices[:bs],
cache_seqlens=forward_batch.seq_lens.to(torch.int32)
+ self.num_draft_tokens,
cache_seqlens=cache_seqlens,
head_dim_v=self.kv_lora_rank,
tile_scheduler_metadata=self.forward_metadata.flashmla_metadata,
num_splits=self.forward_metadata.num_splits,
@@ -454,8 +549,7 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
q=reshape_q,
k_cache=k_cache.view(-1, PAGE_SIZE, 1, self.kv_cache_dim),
block_table=self.forward_metadata.block_kv_indices[:bs],
cache_seqlens=forward_batch.seq_lens.to(torch.int32)
+ self.num_draft_tokens,
cache_seqlens=cache_seqlens,
head_dim_v=self.kv_lora_rank,
tile_scheduler_metadata=self.forward_metadata.flashmla_metadata,
num_splits=self.forward_metadata.num_splits,
@@ -466,6 +560,9 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
class FlashMLAMultiStepDraftBackend:
# Read by decide_needs_cpu_seq_lens (getattr defaults missing flags to True).
needs_cpu_seq_lens: bool = False
def __init__(
self,
model_runner: ModelRunner,
+3 -8
View File
@@ -1,5 +1,3 @@
import logging
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.common import (
cpu_has_amx_support,
@@ -10,8 +8,6 @@ from sglang.srt.utils.common import (
is_npu,
)
logger = logging.getLogger(__name__)
class DraftBackendFactory:
def __init__(
@@ -374,10 +370,9 @@ class DraftBackendFactory:
return AscendAttnBackend(self.draft_model_runner)
def _create_flashmla_prefill_backend(self):
logger.warning(
"flashmla prefill backend is not yet supported for draft extend."
)
return None
from sglang.srt.layers.attention.flashmla_backend import FlashMLABackend
return FlashMLABackend(self.draft_model_runner, skip_prefill=False)
def _create_dsv4_prefill_backend(self):
# On NPU the "dsv4" backend resolves to the Ascend V4 subclass; its
@@ -449,6 +449,12 @@ class EagleDraftWorker(EagleDraftWorkerBase):
)
graph_supported_backend_types.append(DeepseekV4AttnBackend)
if _is_cuda:
# FlashMLA is CUDA-only; import lazily so CPU builds don't pull
# sgl_kernel.flash_mla at import time.
from sglang.srt.layers.attention.flashmla_backend import FlashMLABackend
graph_supported_backend_types.append(FlashMLABackend)
graph_supported_backend = isinstance(
self.draft_extend_attn_backend,