[DLLM] Implement initial dynamic batching for diffusion LLM (#14883)

This commit is contained in:
Zehuan Li
2026-01-17 16:48:15 +08:00
committed by GitHub
parent 737a1183d6
commit d2c863878c
8 changed files with 380 additions and 87 deletions
+94 -22
View File
@@ -32,8 +32,9 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Set, Union
import torch
from sglang.srt.dllm.config import DllmConfig
from sglang.srt.layers.attention.nsa.utils import is_nsa_prefill_cp_in_seq_split
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
from sglang.srt.managers.schedule_batch import DllmStagingReqs, Req, ScheduleBatch
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
@@ -373,6 +374,7 @@ class PrefillAdder:
priority_scheduling_preemption_threshold: int = 0,
prefill_max_requests: Optional[int] = None,
prefill_delayer_single_pass: Optional[PrefillDelayerSinglePassExecutor] = None,
dllm_config: Optional[DllmConfig] = None,
):
self.page_size = page_size
self.tree_cache = tree_cache
@@ -381,6 +383,11 @@ class PrefillAdder:
self.new_token_ratio = new_token_ratio
self.rem_input_tokens = rem_input_tokens - mixed_with_decode_tokens
self.rem_chunk_tokens = rem_chunk_tokens
self.dllm_config = dllm_config
if self.dllm_config is not None:
self._init_dllm_meta(dllm_config)
if self.rem_chunk_tokens is not None:
self.rem_chunk_tokens -= mixed_with_decode_tokens
self.rem_total_token_offset = mixed_with_decode_tokens
@@ -414,6 +421,13 @@ class PrefillAdder:
self.prefill_max_requests = prefill_max_requests
self.prefill_delayer_single_pass = prefill_delayer_single_pass
def _init_dllm_meta(self, dllm_config: DllmConfig):
self.dllm_block_size = dllm_config.block_size
max_running_reqs = dllm_config.max_running_requests
self.rem_dllm_tokens = max_running_reqs * self.dllm_block_size
self.dllm_staging_reqs = DllmStagingReqs(dllm_config=dllm_config)
def _get_running_request_total_token_offset(self, req: Req) -> int:
return (
min(
@@ -473,11 +487,16 @@ class PrefillAdder:
if self.rem_total_tokens <= 0 or self.cur_rem_tokens <= 0:
return AddReqResult.NO_TOKEN
if self.rem_input_tokens <= 0 or (
self.rem_chunk_tokens is not None and self.rem_chunk_tokens <= 0
):
if self.rem_input_tokens <= 0:
return AddReqResult.OTHER
if self.dllm_config is not None:
if self.rem_dllm_tokens <= 0:
return AddReqResult.OTHER
else:
if self.rem_chunk_tokens is not None and self.rem_chunk_tokens <= 0:
return AddReqResult.OTHER
return AddReqResult.CONTINUE
def _update_prefill_budget(
@@ -489,18 +508,61 @@ class PrefillAdder:
self.rem_total_token_offset += extend_input_len + max_new_tokens
self.cur_rem_token_offset += extend_input_len
self.rem_input_tokens -= extend_input_len
if self.rem_chunk_tokens is not None:
if self.dllm_config is not None:
self.rem_dllm_tokens -= extend_input_len
elif self.rem_chunk_tokens is not None:
self.rem_chunk_tokens -= extend_input_len
self.log_hit_tokens += prefix_len
self.log_input_tokens += extend_input_len
def add_chunked_req(self, req: Req):
_rem_tokens = min(self.rem_chunk_tokens, int(self.rem_total_tokens))
# The chunked_req must be added to the list; otherwise, it will cause a memory leak.
# Therefore, in certain cases where _rem_tokens <= 0, it should be replaced with rem_chunk_tokens.
def _get_dllm_remain_tokens(self) -> int:
_rem_tokens = min(
self.rem_dllm_tokens,
self.dllm_block_size,
int(self.rem_total_tokens),
)
if _rem_tokens <= 0:
_rem_tokens = self.rem_chunk_tokens
_rem_tokens = self.rem_dllm_tokens
return _rem_tokens
def _add_dllm_req(self, req: Req, prefix_len: int):
# FIXME: consider the case when rem_dllm_tokens < dllm_block_size,
# the diffusion unmask process may have some problems
# Make sure at least one page is available
trunc_len = (
min(self.rem_dllm_tokens, self.dllm_block_size)
// self.page_size
* self.page_size
)
req.extend_input_len = trunc_len
req.fill_ids = req.fill_ids[: prefix_len + trunc_len]
self.can_run_list.append(req)
self.dllm_staging_reqs.add_reqs(req)
self._update_prefill_budget(prefix_len, trunc_len, 0)
def _req_inc_lock_ref(self, req: Req):
if self.is_hybrid_swa:
swa_uuid_for_lock = self.tree_cache.inc_lock_ref(req.last_node)
req.swa_uuid_for_lock = swa_uuid_for_lock
else:
self.tree_cache.inc_lock_ref(req.last_node)
def add_chunked_req(self, req: Req):
if self.dllm_config is not None:
_rem_tokens = self._get_dllm_remain_tokens()
else:
_rem_tokens = min(self.rem_chunk_tokens, int(self.rem_total_tokens))
# The chunked_req must be added to the list; otherwise, it will cause a memory leak.
# Therefore, in certain cases where _rem_tokens <= 0, it should be replaced with rem_chunk_tokens.
if _rem_tokens <= 0:
_rem_tokens = self.rem_chunk_tokens
truncated = req.extend_input_len > _rem_tokens
req.set_extend_input_len(min(req.extend_input_len, _rem_tokens))
req.fill_ids = req.fill_ids[: len(req.prefix_indices) + req.extend_input_len]
@@ -588,7 +650,12 @@ class PrefillAdder:
return AddReqResult.NO_TOKEN
tokens_freed += tokens_occupied
if (
if self.dllm_config is not None:
if self.rem_dllm_tokens <= 0:
return AddReqResult.OTHER
self._add_dllm_req(req, 0)
elif (
self.rem_chunk_tokens is None # chunked prefill is disabled
or req.extend_input_len <= self.rem_chunk_tokens # it is the last chunk
):
@@ -671,14 +738,21 @@ class PrefillAdder:
):
return AddReqResult.OTHER
if self.rem_chunk_tokens is None or input_tokens <= self.rem_chunk_tokens:
if self.dllm_config is not None:
if self.rem_dllm_tokens <= 0:
return AddReqResult.OTHER
assert (
truncation_align_size is None
), "truncation_align_size is not supported for dllm prefill"
self._add_dllm_req(req, prefix_len)
self._req_inc_lock_ref(req)
elif self.rem_chunk_tokens is None or input_tokens <= self.rem_chunk_tokens:
# Non-chunked prefill
self.can_run_list.append(req)
if self.is_hybrid_swa:
swa_uuid_for_lock = self.tree_cache.inc_lock_ref(req.last_node)
req.swa_uuid_for_lock = swa_uuid_for_lock
else:
self.tree_cache.inc_lock_ref(req.last_node)
self._req_inc_lock_ref(req)
self._update_prefill_budget(
prefix_len,
input_tokens,
@@ -690,6 +764,7 @@ class PrefillAdder:
else:
# Make sure at least one page is available
trunc_len = self.rem_chunk_tokens // self.page_size * self.page_size
if trunc_len <= 0:
return AddReqResult.OTHER
@@ -710,11 +785,8 @@ class PrefillAdder:
self.can_run_list.append(req)
self.new_chunked_req = req
if self.is_hybrid_swa:
swa_uuid_for_lock = self.tree_cache.inc_lock_ref(req.last_node)
req.swa_uuid_for_lock = swa_uuid_for_lock
else:
self.tree_cache.inc_lock_ref(req.last_node)
self._req_inc_lock_ref(req)
self._update_prefill_budget(prefix_len, trunc_len, 0)
return self.budget_state()