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
|
flashmla_metadata: Optional[Tuple[torch.Tensor, torch.Tensor]] = None
|
||||||
num_splits: Optional[torch.Tensor] = None
|
num_splits: Optional[torch.Tensor] = None
|
||||||
block_kv_indices: 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__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
flashmla_metadata: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
flashmla_metadata: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||||
num_splits: Optional[torch.Tensor] = None,
|
num_splits: Optional[torch.Tensor] = None,
|
||||||
block_kv_indices: Optional[torch.Tensor] = None,
|
block_kv_indices: Optional[torch.Tensor] = None,
|
||||||
|
seq_lens_k: Optional[torch.Tensor] = None,
|
||||||
):
|
):
|
||||||
self.flashmla_metadata = flashmla_metadata
|
self.flashmla_metadata = flashmla_metadata
|
||||||
self.num_splits = num_splits
|
self.num_splits = num_splits
|
||||||
self.block_kv_indices = block_kv_indices
|
self.block_kv_indices = block_kv_indices
|
||||||
|
self.seq_lens_k = seq_lens_k
|
||||||
|
|
||||||
|
|
||||||
class FlashMLABackend(FlashInferMLAAttnBackend):
|
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__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
model_runner: ModelRunner,
|
model_runner: ModelRunner,
|
||||||
@@ -89,6 +100,14 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
self.cuda_graph_num_splits = None
|
self.cuda_graph_num_splits = None
|
||||||
self.cuda_graph_mla_metadata_view = None
|
self.cuda_graph_mla_metadata_view = None
|
||||||
self.cuda_graph_num_splits_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
|
# get dcp info
|
||||||
self.dcp_world_size = get_parallel().attn_dcp_size
|
self.dcp_world_size = get_parallel().attn_dcp_size
|
||||||
@@ -100,7 +119,11 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
in_capture: bool = False,
|
in_capture: bool = False,
|
||||||
):
|
):
|
||||||
forward_mode = forward_batch.forward_mode
|
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(
|
self._apply_decode_target_verify_metadata(
|
||||||
bs=forward_batch.batch_size,
|
bs=forward_batch.batch_size,
|
||||||
req_pool_indices=forward_batch.req_pool_indices,
|
req_pool_indices=forward_batch.req_pool_indices,
|
||||||
@@ -113,18 +136,33 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
forward_batch, in_capture=in_capture
|
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):
|
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
||||||
bs = forward_batch.batch_size
|
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():
|
if forward_batch.forward_mode.is_decode_or_idle():
|
||||||
max_seqlen_pad = triton.cdiv(
|
max_seqlen_pad = triton.cdiv(eager_max_k, PAGE_SIZE)
|
||||||
forward_batch.seq_lens_cpu.max().item(), PAGE_SIZE
|
block_kv_indices = self._eager_block_kv_indices(bs, max_seqlen_pad)
|
||||||
)
|
|
||||||
block_kv_indices = torch.full(
|
|
||||||
(bs, max_seqlen_pad),
|
|
||||||
-1,
|
|
||||||
dtype=torch.int32,
|
|
||||||
device=forward_batch.seq_lens.device,
|
|
||||||
)
|
|
||||||
create_flashmla_kv_indices_triton[
|
create_flashmla_kv_indices_triton[
|
||||||
(bs, get_num_kv_index_blocks_flashmla(max_seqlen_pad, PAGE_SIZE))
|
(bs, get_num_kv_index_blocks_flashmla(max_seqlen_pad, PAGE_SIZE))
|
||||||
](
|
](
|
||||||
@@ -134,7 +172,7 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
None,
|
None,
|
||||||
block_kv_indices,
|
block_kv_indices,
|
||||||
self.req_to_token.stride(0),
|
self.req_to_token.stride(0),
|
||||||
max_seqlen_pad,
|
block_kv_indices.stride(0),
|
||||||
)
|
)
|
||||||
mla_metadata, num_splits = get_mla_metadata(
|
mla_metadata, num_splits = get_mla_metadata(
|
||||||
forward_batch.seq_lens.to(torch.int32),
|
forward_batch.seq_lens.to(torch.int32),
|
||||||
@@ -148,16 +186,10 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
block_kv_indices,
|
block_kv_indices,
|
||||||
)
|
)
|
||||||
elif forward_batch.forward_mode.is_target_verify():
|
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
|
seq_lens = forward_batch.seq_lens + self.num_draft_tokens
|
||||||
|
|
||||||
max_seqlen_pad = triton.cdiv(seq_lens_cpu.max().item(), PAGE_SIZE)
|
max_seqlen_pad = triton.cdiv(eager_max_k + self.num_draft_tokens, PAGE_SIZE)
|
||||||
block_kv_indices = torch.full(
|
block_kv_indices = self._eager_block_kv_indices(bs, max_seqlen_pad)
|
||||||
(bs, max_seqlen_pad),
|
|
||||||
-1,
|
|
||||||
dtype=torch.int32,
|
|
||||||
device=seq_lens.device,
|
|
||||||
)
|
|
||||||
create_flashmla_kv_indices_triton[
|
create_flashmla_kv_indices_triton[
|
||||||
(bs, get_num_kv_index_blocks_flashmla(max_seqlen_pad, PAGE_SIZE))
|
(bs, get_num_kv_index_blocks_flashmla(max_seqlen_pad, PAGE_SIZE))
|
||||||
](
|
](
|
||||||
@@ -167,7 +199,7 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
None,
|
None,
|
||||||
block_kv_indices,
|
block_kv_indices,
|
||||||
self.req_to_token.stride(0),
|
self.req_to_token.stride(0),
|
||||||
max_seqlen_pad,
|
block_kv_indices.stride(0),
|
||||||
)
|
)
|
||||||
mla_metadata, num_splits = get_mla_metadata(
|
mla_metadata, num_splits = get_mla_metadata(
|
||||||
seq_lens.to(torch.int32),
|
seq_lens.to(torch.int32),
|
||||||
@@ -180,6 +212,39 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
num_splits,
|
num_splits,
|
||||||
block_kv_indices,
|
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:
|
else:
|
||||||
super().init_forward_metadata(forward_batch)
|
super().init_forward_metadata(forward_batch)
|
||||||
|
|
||||||
@@ -216,6 +281,22 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
self.cuda_graph_mla_metadata_view = None
|
self.cuda_graph_mla_metadata_view = None
|
||||||
self.cuda_graph_num_splits_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(
|
def _apply_decode_target_verify_metadata(
|
||||||
self,
|
self,
|
||||||
bs: int,
|
bs: int,
|
||||||
@@ -224,12 +305,7 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
seq_lens_cpu: Optional[torch.Tensor],
|
seq_lens_cpu: Optional[torch.Tensor],
|
||||||
forward_mode: ForwardMode,
|
forward_mode: ForwardMode,
|
||||||
):
|
):
|
||||||
"""Shared decode/target-verify capture+replay body.
|
"""Shared decode/target-verify/draft-extend 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).
|
|
||||||
"""
|
|
||||||
if True:
|
if True:
|
||||||
seq_lens = seq_lens[:bs]
|
seq_lens = seq_lens[:bs]
|
||||||
seq_lens_cpu = seq_lens_cpu[:bs] if seq_lens_cpu is not None else None
|
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
|
seq_lens = seq_lens + self.num_draft_tokens
|
||||||
if seq_lens_cpu is not None:
|
if seq_lens_cpu is not None:
|
||||||
seq_lens_cpu = seq_lens_cpu + self.num_draft_tokens
|
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 = (
|
# Tight block-table slice when the CPU mirror is free; static
|
||||||
seq_lens_cpu.max().item()
|
# bound otherwise (no D2H; kernel reads are capped by seq_lens_k).
|
||||||
if seq_lens_cpu is not None
|
if seq_lens_cpu is not None:
|
||||||
else seq_lens.max().item()
|
max_seqlen_pad = triton.cdiv(seq_lens_cpu.max().item(), PAGE_SIZE)
|
||||||
)
|
else:
|
||||||
max_seqlen_pad = triton.cdiv(seq_max, PAGE_SIZE)
|
max_seqlen_pad = self.cuda_graph_kv_indices.shape[1]
|
||||||
|
|
||||||
create_flashmla_kv_indices_triton[
|
create_flashmla_kv_indices_triton[
|
||||||
(
|
(
|
||||||
@@ -264,7 +342,9 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
)
|
)
|
||||||
|
|
||||||
q_head_mult = (
|
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(
|
mla_metadata, num_splits = get_mla_metadata(
|
||||||
seq_lens.to(torch.int32),
|
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_mla_metadata[:actual_num_sm_parts].copy_(mla_metadata)
|
||||||
self.cuda_graph_num_splits[: bs + 1].copy_(num_splits)
|
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.forward_metadata = FlashMLADecodeMetadata(
|
||||||
self.cuda_graph_mla_metadata_view,
|
self.cuda_graph_mla_metadata_view,
|
||||||
self.cuda_graph_num_splits_view,
|
self.cuda_graph_num_splits_view,
|
||||||
self.cuda_graph_kv_indices[:bs, :max_seqlen_pad],
|
self.cuda_graph_kv_indices[:bs, :max_seqlen_pad],
|
||||||
|
seq_lens_k,
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_cuda_graph_seq_len_fill_value(self):
|
def get_cuda_graph_seq_len_fill_value(self):
|
||||||
@@ -398,12 +485,10 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
save_kv_cache: bool = True,
|
save_kv_cache: bool = True,
|
||||||
):
|
):
|
||||||
if forward_batch.forward_mode in (
|
if forward_batch.forward_mode == ForwardMode.EXTEND:
|
||||||
ForwardMode.EXTEND,
|
|
||||||
ForwardMode.DRAFT_EXTEND_V2,
|
|
||||||
):
|
|
||||||
return super().forward_extend(q, k, v, layer, forward_batch, save_kv_cache)
|
return super().forward_extend(q, k, v, layer, forward_batch, save_kv_cache)
|
||||||
else:
|
else:
|
||||||
|
# target_verify / draft_extend_v2: fixed-q decode-style kernel.
|
||||||
cache_loc = forward_batch.out_cache_loc
|
cache_loc = forward_batch.out_cache_loc
|
||||||
|
|
||||||
if k is not None:
|
if k is not None:
|
||||||
@@ -414,7 +499,18 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
bs = forward_batch.batch_size
|
bs = forward_batch.batch_size
|
||||||
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||||
|
|
||||||
|
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)
|
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 self.is_fp8_kvcache:
|
||||||
if layer.k_scale is not None:
|
if layer.k_scale is not None:
|
||||||
q_scale = layer.k_scale
|
q_scale = layer.k_scale
|
||||||
@@ -439,8 +535,7 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
q=reshape_q_fp8,
|
q=reshape_q_fp8,
|
||||||
k_cache=k_cache.view(-1, PAGE_SIZE, 1, self.kv_cache_dim),
|
k_cache=k_cache.view(-1, PAGE_SIZE, 1, self.kv_cache_dim),
|
||||||
block_table=self.forward_metadata.block_kv_indices[:bs],
|
block_table=self.forward_metadata.block_kv_indices[:bs],
|
||||||
cache_seqlens=forward_batch.seq_lens.to(torch.int32)
|
cache_seqlens=cache_seqlens,
|
||||||
+ self.num_draft_tokens,
|
|
||||||
head_dim_v=self.kv_lora_rank,
|
head_dim_v=self.kv_lora_rank,
|
||||||
tile_scheduler_metadata=self.forward_metadata.flashmla_metadata,
|
tile_scheduler_metadata=self.forward_metadata.flashmla_metadata,
|
||||||
num_splits=self.forward_metadata.num_splits,
|
num_splits=self.forward_metadata.num_splits,
|
||||||
@@ -454,8 +549,7 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
q=reshape_q,
|
q=reshape_q,
|
||||||
k_cache=k_cache.view(-1, PAGE_SIZE, 1, self.kv_cache_dim),
|
k_cache=k_cache.view(-1, PAGE_SIZE, 1, self.kv_cache_dim),
|
||||||
block_table=self.forward_metadata.block_kv_indices[:bs],
|
block_table=self.forward_metadata.block_kv_indices[:bs],
|
||||||
cache_seqlens=forward_batch.seq_lens.to(torch.int32)
|
cache_seqlens=cache_seqlens,
|
||||||
+ self.num_draft_tokens,
|
|
||||||
head_dim_v=self.kv_lora_rank,
|
head_dim_v=self.kv_lora_rank,
|
||||||
tile_scheduler_metadata=self.forward_metadata.flashmla_metadata,
|
tile_scheduler_metadata=self.forward_metadata.flashmla_metadata,
|
||||||
num_splits=self.forward_metadata.num_splits,
|
num_splits=self.forward_metadata.num_splits,
|
||||||
@@ -466,6 +560,9 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
|
|
||||||
|
|
||||||
class FlashMLAMultiStepDraftBackend:
|
class FlashMLAMultiStepDraftBackend:
|
||||||
|
# Read by decide_needs_cpu_seq_lens (getattr defaults missing flags to True).
|
||||||
|
needs_cpu_seq_lens: bool = False
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
model_runner: ModelRunner,
|
model_runner: ModelRunner,
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
import logging
|
|
||||||
|
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.utils.common import (
|
from sglang.srt.utils.common import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
@@ -10,8 +8,6 @@ from sglang.srt.utils.common import (
|
|||||||
is_npu,
|
is_npu,
|
||||||
)
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
class DraftBackendFactory:
|
class DraftBackendFactory:
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -374,10 +370,9 @@ class DraftBackendFactory:
|
|||||||
return AscendAttnBackend(self.draft_model_runner)
|
return AscendAttnBackend(self.draft_model_runner)
|
||||||
|
|
||||||
def _create_flashmla_prefill_backend(self):
|
def _create_flashmla_prefill_backend(self):
|
||||||
logger.warning(
|
from sglang.srt.layers.attention.flashmla_backend import FlashMLABackend
|
||||||
"flashmla prefill backend is not yet supported for draft extend."
|
|
||||||
)
|
return FlashMLABackend(self.draft_model_runner, skip_prefill=False)
|
||||||
return None
|
|
||||||
|
|
||||||
def _create_dsv4_prefill_backend(self):
|
def _create_dsv4_prefill_backend(self):
|
||||||
# On NPU the "dsv4" backend resolves to the Ascend V4 subclass; its
|
# 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)
|
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(
|
graph_supported_backend = isinstance(
|
||||||
self.draft_extend_attn_backend,
|
self.draft_extend_attn_backend,
|
||||||
|
|||||||
Reference in New Issue
Block a user