Introduce global alloc_len_per_decode & clean check decode memory (#15115)

This commit is contained in:
Liangsheng Yin
2026-01-26 10:26:20 -08:00
committed by GitHub
parent 8643fb2f52
commit 85d077f44d
10 changed files with 76 additions and 101 deletions

View File

@@ -11,6 +11,7 @@ from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.overlap_utils import FutureIndices
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.model_executor.forward_batch_info import PPProxyTensors
from sglang.srt.server_args import ServerArgs
if TYPE_CHECKING:
from sglang.srt.managers.scheduler import GenerationBatchResult
@@ -176,3 +177,30 @@ def get_logprob_from_pp_outputs(
]
return logits_output, extend_input_len_per_req, extend_logprob_start_len_per_req
def get_alloc_len_per_decode(server_args: Optional[ServerArgs] = None) -> int:
if server_args is None:
from sglang.srt.server_args import get_global_server_args
server_args = get_global_server_args()
if server_args.speculative_algorithm is None:
return 1
# Spec v1:
# 1) alloc topk * num_steps when draft decoding and then restore the allocation
# 2) alloc num_draft_tokens when verifying the drafts
# Sepc v2: allocate max(topk * num_steps, num_draft_tokens)
spec_steps = server_args.speculative_num_steps or 1
spec_topk = server_args.speculative_eagle_topk or 1
spec_tokens = server_args.speculative_num_draft_tokens
page_size = server_args.page_size
if page_size == 1 or spec_topk == 1:
return max(spec_steps * spec_topk, spec_tokens)
else:
raise NotImplementedError(
"get_alloc_len_per_decode not implemented for page_size > 1 and spec_topk > 1"
)