diff --git a/python/sglang/srt/layers/attention/flashmla_backend.py b/python/sglang/srt/layers/attention/flashmla_backend.py index 19e420d2a..1abc64f84 100644 --- a/python/sglang/srt/layers/attention/flashmla_backend.py +++ b/python/sglang/srt/layers/attention/flashmla_backend.py @@ -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, diff --git a/python/sglang/srt/speculative/draft_utils.py b/python/sglang/srt/speculative/draft_utils.py index 5c254e9a2..bb7106eda 100644 --- a/python/sglang/srt/speculative/draft_utils.py +++ b/python/sglang/srt/speculative/draft_utils.py @@ -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 diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 1de74ecfb..e8493012c 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -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,