[spec] Dedup draft kv_indices sizing into spec_utils helpers (#27475)

This commit is contained in:
Liangsheng Yin
2026-06-07 00:42:29 -07:00
committed by GitHub
parent 5e2e0d5b49
commit eab2e02fa0
5 changed files with 82 additions and 32 deletions
@@ -27,7 +27,11 @@ from sglang.srt.layers.dp_attention import (
is_dp_attention_enabled, is_dp_attention_enabled,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.speculative.spec_utils import generate_draft_decode_kv_indices from sglang.srt.speculative.spec_utils import (
draft_kv_indices_buffer_width,
draft_kv_indices_used_len,
generate_draft_decode_kv_indices,
)
from sglang.srt.utils import is_gfx95_supported from sglang.srt.utils import is_gfx95_supported
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -2837,16 +2841,16 @@ class AiterMultiStepDraftBackend:
for i in range(self.speculative_num_steps - 1): for i in range(self.speculative_num_steps - 1):
forward_batch.spec_info.kv_indptr = self.kv_indptr[i, : bs + 1] forward_batch.spec_info.kv_indptr = self.kv_indptr[i, : bs + 1]
forward_batch.spec_info.kv_indices = kv_indices_buffer[i][ forward_batch.spec_info.kv_indices = kv_indices_buffer[i][
: seq_lens_sum * self.topk + bs * (i + 1) : draft_kv_indices_used_len(seq_lens_sum, self.topk, bs, i + 1)
] ]
call_fn(i, forward_batch) call_fn(i, forward_batch)
def init_forward_metadata(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch):
kv_indices_width = draft_kv_indices_buffer_width(
forward_batch.batch_size, self.topk, self.max_context_len
)
kv_indices = torch.empty( kv_indices = torch.empty(
( (self.speculative_num_steps, kv_indices_width),
self.speculative_num_steps,
forward_batch.batch_size * self.topk * self.max_context_len,
),
dtype=torch.int32, dtype=torch.int32,
device=self.device, device=self.device,
) )
@@ -2863,8 +2867,11 @@ class AiterMultiStepDraftBackend:
self.common_template(forward_batch, kv_indices, call_fn) self.common_template(forward_batch, kv_indices, call_fn)
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
kv_indices_width = draft_kv_indices_buffer_width(
max_bs, self.topk, self.max_context_len
)
self.cuda_graph_kv_indices = torch.zeros( self.cuda_graph_kv_indices = torch.zeros(
(self.speculative_num_steps, max_num_tokens * self.max_context_len), (self.speculative_num_steps, kv_indices_width),
dtype=torch.int32, dtype=torch.int32,
device=self.device, device=self.device,
) )
@@ -27,7 +27,11 @@ from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_info import SpecInput
from sglang.srt.speculative.spec_utils import generate_draft_decode_kv_indices from sglang.srt.speculative.spec_utils import (
draft_kv_indices_buffer_width,
draft_kv_indices_used_len,
generate_draft_decode_kv_indices,
)
from sglang.srt.utils import ( from sglang.srt.utils import (
get_int_env_var, get_int_env_var,
is_flashinfer_available, is_flashinfer_available,
@@ -1585,8 +1589,8 @@ class FlashInferMultiStepDraftBackend:
# Fail fast on an undersized kv_indices row: the kernel would otherwise write # Fail fast on an undersized kv_indices row: the kernel would otherwise write
# OOB and *silently* corrupt memory, only sometimes surfacing as a crash. # OOB and *silently* corrupt memory, only sometimes surfacing as a crash.
required_kv_indices_len = ( required_kv_indices_len = draft_kv_indices_used_len(
seq_lens_sum * self.topk + bs * self.speculative_num_steps seq_lens_sum, self.topk, bs, self.speculative_num_steps
) )
assert required_kv_indices_len <= kv_indices_buffer.shape[1], ( assert required_kv_indices_len <= kv_indices_buffer.shape[1], (
f"EAGLE draft kv_indices row too small: need {required_kv_indices_len} " f"EAGLE draft kv_indices row too small: need {required_kv_indices_len} "
@@ -1624,7 +1628,7 @@ class FlashInferMultiStepDraftBackend:
for i in range(self.speculative_num_steps - 1): for i in range(self.speculative_num_steps - 1):
forward_batch.spec_info.kv_indptr = self.kv_indptr[i, : bs + 1] forward_batch.spec_info.kv_indptr = self.kv_indptr[i, : bs + 1]
forward_batch.spec_info.kv_indices = kv_indices_buffer[i][ forward_batch.spec_info.kv_indices = kv_indices_buffer[i][
: seq_lens_sum * self.topk + bs * (i + 1) : draft_kv_indices_used_len(seq_lens_sum, self.topk, bs, i + 1)
] ]
global_override_indptr_cpu = indptr_cpu_whole[i] global_override_indptr_cpu = indptr_cpu_whole[i]
call_fn(i, forward_batch) call_fn(i, forward_batch)
@@ -1632,11 +1636,11 @@ class FlashInferMultiStepDraftBackend:
global_override_indptr_cpu = None global_override_indptr_cpu = None
def init_forward_metadata(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch):
kv_indices_width = draft_kv_indices_buffer_width(
forward_batch.batch_size, self.topk, self.max_context_len
)
kv_indices = torch.empty( kv_indices = torch.empty(
( (self.speculative_num_steps, kv_indices_width),
self.speculative_num_steps,
forward_batch.batch_size * self.topk * self.max_context_len,
),
dtype=torch.int32, dtype=torch.int32,
device="cuda", device="cuda",
) )
@@ -1656,8 +1660,11 @@ class FlashInferMultiStepDraftBackend:
# generate_draft_decode_kv_indices packs topk per-branch sequences per row, # generate_draft_decode_kv_indices packs topk per-branch sequences per row,
# so the row needs the topk factor -- same as the eager init_forward_metadata # so the row needs the topk factor -- same as the eager init_forward_metadata
# (batch_size * topk * max_context_len). Dropping it overflows the buffer. # (batch_size * topk * max_context_len). Dropping it overflows the buffer.
kv_indices_width = draft_kv_indices_buffer_width(
max_bs, self.topk, self.max_context_len
)
self.cuda_graph_kv_indices = torch.zeros( self.cuda_graph_kv_indices = torch.zeros(
(self.speculative_num_steps, max_bs * self.topk * self.max_context_len), (self.speculative_num_steps, kv_indices_width),
dtype=torch.int32, dtype=torch.int32,
device="cuda", device="cuda",
) )
@@ -25,7 +25,11 @@ from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_info import SpecInput
from sglang.srt.speculative.spec_utils import generate_draft_decode_kv_indices from sglang.srt.speculative.spec_utils import (
draft_kv_indices_buffer_width,
draft_kv_indices_used_len,
generate_draft_decode_kv_indices,
)
from sglang.srt.utils import ( from sglang.srt.utils import (
is_flashinfer_available, is_flashinfer_available,
is_sm100_supported, is_sm100_supported,
@@ -944,8 +948,8 @@ class FlashInferMLAMultiStepDraftBackend:
# Fail fast on an undersized kv_indices row: the kernel would otherwise # Fail fast on an undersized kv_indices row: the kernel would otherwise
# write OOB and silently corrupt memory. # write OOB and silently corrupt memory.
required_kv_indices_len = ( required_kv_indices_len = draft_kv_indices_used_len(
seq_lens_sum * self.topk + bs * self.speculative_num_steps seq_lens_sum, self.topk, bs, self.speculative_num_steps
) )
assert required_kv_indices_len <= kv_indices_buffer.shape[1], ( assert required_kv_indices_len <= kv_indices_buffer.shape[1], (
f"EAGLE draft kv_indices row too small: need {required_kv_indices_len} " f"EAGLE draft kv_indices row too small: need {required_kv_indices_len} "
@@ -979,16 +983,16 @@ class FlashInferMLAMultiStepDraftBackend:
for i in range(self.speculative_num_steps - 1): for i in range(self.speculative_num_steps - 1):
forward_batch.spec_info.kv_indptr = self.kv_indptr[i, : bs + 1] forward_batch.spec_info.kv_indptr = self.kv_indptr[i, : bs + 1]
forward_batch.spec_info.kv_indices = kv_indices_buffer[i][ forward_batch.spec_info.kv_indices = kv_indices_buffer[i][
: seq_lens_sum * self.topk + bs * (i + 1) : draft_kv_indices_used_len(seq_lens_sum, self.topk, bs, i + 1)
] ]
call_fn(i, forward_batch) call_fn(i, forward_batch)
def init_forward_metadata(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch):
kv_indices_width = draft_kv_indices_buffer_width(
forward_batch.batch_size, self.topk, self.max_context_len
)
kv_indices = torch.zeros( kv_indices = torch.zeros(
( (self.speculative_num_steps, kv_indices_width),
self.speculative_num_steps,
forward_batch.batch_size * self.topk * self.max_context_len,
),
dtype=torch.int32, dtype=torch.int32,
device="cuda", device="cuda",
) )
@@ -1007,8 +1011,11 @@ class FlashInferMLAMultiStepDraftBackend:
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
# Row holds topk per-branch sequences (generate_draft_decode_kv_indices), so # Row holds topk per-branch sequences (generate_draft_decode_kv_indices), so
# it needs the topk factor, matching the eager init_forward_metadata. # it needs the topk factor, matching the eager init_forward_metadata.
kv_indices_width = draft_kv_indices_buffer_width(
max_bs, self.topk, self.max_context_len
)
self.cuda_graph_kv_indices = torch.zeros( self.cuda_graph_kv_indices = torch.zeros(
(self.speculative_num_steps, max_bs * self.topk * self.max_context_len), (self.speculative_num_steps, kv_indices_width),
dtype=torch.int32, dtype=torch.int32,
device="cuda", device="cuda",
) )
@@ -16,7 +16,11 @@ from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.layers.radix_attention import AttentionType from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.speculative.spec_utils import generate_draft_decode_kv_indices from sglang.srt.speculative.spec_utils import (
draft_kv_indices_buffer_width,
draft_kv_indices_used_len,
generate_draft_decode_kv_indices,
)
from sglang.srt.utils import ( from sglang.srt.utils import (
get_bool_env_var, get_bool_env_var,
get_device_core_count, get_device_core_count,
@@ -1393,16 +1397,16 @@ class TritonMultiStepDraftBackend:
for i in range(self.speculative_num_steps - 1): for i in range(self.speculative_num_steps - 1):
forward_batch.spec_info.kv_indptr = self.kv_indptr[i, : bs + 1] forward_batch.spec_info.kv_indptr = self.kv_indptr[i, : bs + 1]
forward_batch.spec_info.kv_indices = kv_indices_buffer[i][ forward_batch.spec_info.kv_indices = kv_indices_buffer[i][
: seq_lens_sum * self.topk + bs * (i + 1) : draft_kv_indices_used_len(seq_lens_sum, self.topk, bs, i + 1)
] ]
call_fn(i, forward_batch) call_fn(i, forward_batch)
def init_forward_metadata(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch):
kv_indices_width = draft_kv_indices_buffer_width(
forward_batch.batch_size, self.topk, self.max_context_len
)
kv_indices = torch.empty( kv_indices = torch.empty(
( (self.speculative_num_steps, kv_indices_width),
self.speculative_num_steps,
forward_batch.batch_size * self.topk * self.max_context_len,
),
dtype=torch.int64, dtype=torch.int64,
device=self.device, device=self.device,
) )
@@ -1419,8 +1423,11 @@ class TritonMultiStepDraftBackend:
self.common_template(forward_batch, kv_indices, call_fn) self.common_template(forward_batch, kv_indices, call_fn)
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
kv_indices_width = draft_kv_indices_buffer_width(
max_bs, self.topk, self.max_context_len
)
self.cuda_graph_kv_indices = torch.zeros( self.cuda_graph_kv_indices = torch.zeros(
(self.speculative_num_steps, max_num_tokens * self.max_context_len), (self.speculative_num_steps, kv_indices_width),
dtype=torch.int64, dtype=torch.int64,
device=self.device, device=self.device,
) )
@@ -78,6 +78,28 @@ TREE_SPEC_KERNEL_AVAILABLE = (
) # This kernel is only available for CUDA and MUSA now ) # This kernel is only available for CUDA and MUSA now
def draft_kv_indices_buffer_width(
num_seqs: int, topk: int, max_context_len: int
) -> int:
"""Per-step row width of the EAGLE draft-decode kv_indices buffer.
num_seqs * topk branches each attend up to max_context_len KV slots; the topk
factor is mandatory -- dropping it under-allocates and overflows the row (#27338, #27460).
"""
return num_seqs * topk * max_context_len
def draft_kv_indices_used_len(
seq_lens_sum: int, topk: int, bs: int, num_steps: int
) -> int:
"""kv_indices length used through num_steps draft-decode steps.
bs = topk * num_seqs branches, one index appended per branch per step. Called with
num_steps = i + 1 (per-step slice) and speculative_num_steps (capacity assert).
"""
return seq_lens_sum * topk + bs * num_steps
def record_stream_each(tensors, stream): def record_stream_each(tensors, stream):
"""Call record_stream(stream) on each cuda tensor in `tensors`, skipping """Call record_stream(stream) on each cuda tensor in `tensors`, skipping
non-tensor / non-cuda entries. Tells the caching allocator that the non-tensor / non-cuda entries. Tells the caching allocator that the