[Refactor] Share the page-aligned decode alloc lens between EAGLE and DFLASH (#35382)

This commit is contained in:
Liangsheng Yin
2026-08-18 14:39:41 -07:00
committed by GitHub
parent 7b5410c999
commit 87a09494fa
3 changed files with 47 additions and 40 deletions
@@ -47,6 +47,29 @@ def get_alloc_reserve_per_decode() -> int:
return 2 * get_alloc_len_per_decode() return 2 * get_alloc_len_per_decode()
def page_aligned_decode_alloc_lens(
reqs,
*,
reserve: int,
page_size: int,
):
"""Whole-page decode alloc lens: nxt rounds committed up to page so allocated
== recorded (unaligned tails leak at ps>1)."""
cur_kv_lens = [0] * len(reqs)
nxt_kv_lens = [0] * len(reqs)
num_needed_tokens = 0
for i, r in enumerate(reqs):
cur = r.kv.kv_allocated_len
nxt = max(
cur,
(r.kv_committed_len + reserve + page_size - 1) // page_size * page_size,
)
cur_kv_lens[i] = cur
nxt_kv_lens[i] = nxt
num_needed_tokens += nxt - cur
return cur_kv_lens, nxt_kv_lens, num_needed_tokens
def get_req_to_token_extra_context_len() -> int: def get_req_to_token_extra_context_len() -> int:
"""req_to_token row headroom beyond the model context length. """req_to_token row headroom beyond the model context length.
+14 -19
View File
@@ -9,6 +9,7 @@ import torch
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.mem_cache.allocation import alloc_for_spec_decode from sglang.srt.mem_cache.allocation import alloc_for_spec_decode
from sglang.srt.mem_cache.allocation_sizing import page_aligned_decode_alloc_lens
from sglang.srt.runtime_context import get_spec from sglang.srt.runtime_context import get_spec
from sglang.srt.speculative.spec_info import SpecInput, SpecInputType from sglang.srt.speculative.spec_info import SpecInput, SpecInputType
from sglang.srt.utils.common import is_pin_memory_available from sglang.srt.utils.common import is_pin_memory_available
@@ -135,6 +136,7 @@ class DFlashDraftInputV2(SpecInput):
assert self._prepare_nxt_kv_lens_gpu_buf is not None assert self._prepare_nxt_kv_lens_gpu_buf is not None
batch_seq_lens_cpu_t = self._prepare_batch_seq_lens_cpu_buf[:bs] batch_seq_lens_cpu_t = self._prepare_batch_seq_lens_cpu_buf[:bs]
cur_kv_lens_cpu_t = self._prepare_cur_kv_lens_cpu_buf[:bs] cur_kv_lens_cpu_t = self._prepare_cur_kv_lens_cpu_buf[:bs]
nxt_kv_lens_cpu_t = self._prepare_nxt_kv_lens_cpu_buf[:bs]
# For DFLASH, each decode step needs a fixed-size verify block. # For DFLASH, each decode step needs a fixed-size verify block.
block_size = int(get_spec().speculative_num_draft_tokens) block_size = int(get_spec().speculative_num_draft_tokens)
@@ -142,37 +144,30 @@ class DFlashDraftInputV2(SpecInput):
raise ValueError( raise ValueError(
f"DFLASH invalid speculative_num_draft_tokens={block_size}." f"DFLASH invalid speculative_num_draft_tokens={block_size}."
) )
reserve = 2 * block_size
page_size = batch.token_to_kv_pool_allocator.page_size page_size = batch.token_to_kv_pool_allocator.page_size
nxt_kv_lens_cpu_t = self._prepare_nxt_kv_lens_cpu_buf[:bs]
committed_seq_lens_sum = 0 cur_kv_lens, nxt_kv_lens, num_needed_tokens = page_aligned_decode_alloc_lens(
nxt_kv_lens_sum = 0 batch.reqs,
num_needed_tokens = 0 reserve=reserve,
page_size=page_size,
)
max_top_k = 1 max_top_k = 1
uniform_top_k_value = None uniform_top_k_value = None
uniform_top_k = True uniform_top_k = True
for i, req in enumerate(batch.reqs): nxt_kv_lens_sum = 0
committed_seq_lens_sum = 0
for i, (req, cur, nxt) in enumerate(zip(batch.reqs, cur_kv_lens, nxt_kv_lens)):
committed_len = int(req.kv_committed_len) committed_len = int(req.kv_committed_len)
# Read the allocation watermark from the req object like EAGLE. committed_seq_lens_sum += committed_len
cur = int(req.kv.kv_allocated_len)
# Whole-page accounting (same as eagle_prepare_for_decode): the
# paged allocator hands out full pages, so an unaligned reserve
# strands the tail of the last page -- allocated but never recorded.
nxt = max(
cur,
(committed_len + 2 * block_size + page_size - 1)
// page_size
* page_size,
)
top_k = int(req.sampling_params.top_k) top_k = int(req.sampling_params.top_k)
batch_seq_lens_cpu_t[i] = committed_len batch_seq_lens_cpu_t[i] = committed_len
cur_kv_lens_cpu_t[i] = cur cur_kv_lens_cpu_t[i] = cur
nxt_kv_lens_cpu_t[i] = nxt nxt_kv_lens_cpu_t[i] = nxt
committed_seq_lens_sum += committed_len
nxt_kv_lens_sum += nxt nxt_kv_lens_sum += nxt
num_needed_tokens += nxt - cur
if top_k > max_top_k: if top_k > max_top_k:
max_top_k = top_k max_top_k = top_k
if i == 0: if i == 0:
+10 -21
View File
@@ -15,7 +15,10 @@ from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
maybe_build_dsv4_verify_bundle, maybe_build_dsv4_verify_bundle,
) )
from sglang.srt.mem_cache.allocation import alloc_for_spec_decode from sglang.srt.mem_cache.allocation import alloc_for_spec_decode
from sglang.srt.mem_cache.allocation_sizing import get_alloc_reserve_per_decode from sglang.srt.mem_cache.allocation_sizing import (
get_alloc_reserve_per_decode,
page_aligned_decode_alloc_lens,
)
from sglang.srt.runtime_context import get_parallel, get_spec from sglang.srt.runtime_context import get_parallel, get_spec
from sglang.srt.utils import ( from sglang.srt.utils import (
is_cpu, is_cpu,
@@ -907,26 +910,12 @@ def eagle_prepare_for_decode(batch: ScheduleBatch):
page_size = batch.token_to_kv_pool_allocator.page_size page_size = batch.token_to_kv_pool_allocator.page_size
double_alloc = get_alloc_reserve_per_decode() double_alloc = get_alloc_reserve_per_decode()
cur_kv_lens = [0] * bs cur_kv_lens, nxt_kv_lens, num_needed_tokens = page_aligned_decode_alloc_lens(
nxt_kv_lens = [0] * bs batch.reqs,
num_needed_tokens = 0 reserve=double_alloc,
for i, r in enumerate(batch.reqs): page_size=page_size,
cur = r.kv.kv_allocated_len )
# max(cur, ...) clamps so adaptive downswitch cannot make nxt < cur. for r in batch.reqs:
# kv_committed_len is honest (bonus committed in resolve, not here),
# so it lags batch.seq_lens by ~1 verify in overlap; 2*alloc absorbs.
# Whole-page accounting: the paged allocator hands out full pages, so
# round nxt up to the page boundary or the unaligned tail is allocated
# but never recorded — a stranded-tail leak at page_size > 1.
nxt = max(
cur,
(r.kv_committed_len + double_alloc + page_size - 1)
// page_size
* page_size,
)
cur_kv_lens[i] = cur
nxt_kv_lens[i] = nxt
num_needed_tokens += nxt - cur
r.decode_batch_idx += 1 r.decode_batch_idx += 1
cur_kv_lens_cpu = torch.tensor(cur_kv_lens, dtype=torch.int32, device="cpu") cur_kv_lens_cpu = torch.tensor(cur_kv_lens, dtype=torch.int32, device="cpu")