flashmla: sync-free spec via device-side draft-extend (#31090)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user