Avoid implicit running_batch access in dllm and pdmux scheduling (#30676)

This commit is contained in:
fzyzcjy
2026-07-15 14:33:00 +08:00
committed by GitHub
parent 21c62b9830
commit 01343d2759
2 changed files with 72 additions and 50 deletions
+41 -28
View File
@@ -31,26 +31,25 @@ class SchedulerDllmMixin:
self: Scheduler, running_batch: ScheduleBatch
) -> Optional[ScheduleBatch]:
"""Generate a new batch for DLLM (Diffusion LLM) scheduling."""
self.running_batch = running_batch
if self.enable_priority_preemption:
self.running_batch.batch_is_full = False
running_batch.batch_is_full = False
# Early exit if batch is full or no requests available
if self._should_skip_prefill():
if self._should_skip_prefill(running_batch=running_batch):
return None
running_bs = len(self.running_batch.reqs)
running_bs = len(running_batch.reqs)
self.policy.calc_priority(self.waiting_queue)
# Create prefill adder with resource constraints
adder = self._create_dllm_prefill_adder(running_bs)
adder = self._create_dllm_prefill_adder(running_bs, running_batch=running_batch)
# Initialize DLLM manager and transfer requests
self.dllm_manager.init_next_round()
self._fetch_waiting_reqs()
# Process batches
forward_mode = self._process_dllm_batches(adder)
forward_mode = self._process_dllm_batches(adder, running_batch=running_batch)
can_run_list = adder.can_run_list
if not can_run_list:
@@ -58,10 +57,12 @@ class SchedulerDllmMixin:
# Record metrics and update state
set_time_batch(can_run_list, "set_forward_entry_time")
self._update_state_for_batch(can_run_list, adder, running_bs)
self._update_state_for_batch(can_run_list, adder)
# Create and prepare batch
new_batch = self._create_dllm_batch(can_run_list, forward_mode)
new_batch = self._create_dllm_batch(
can_run_list, forward_mode, adder=adder, running_batch=running_batch
)
return new_batch
def process_batch_result_dllm(
@@ -177,31 +178,33 @@ class SchedulerDllmMixin:
self.dllm_manager.add_waiting_reqs(requests_to_add)
self.waiting_queue = self.waiting_queue[num_requests_to_add:]
def _should_skip_prefill(self: Scheduler) -> bool:
def _should_skip_prefill(self: Scheduler, running_batch: ScheduleBatch) -> bool:
"""Check if DLLM prefill should be skipped."""
if (
self.running_batch.batch_is_full or not self.waiting_queue
running_batch.batch_is_full or not self.waiting_queue
) and self.dllm_manager.is_empty():
return True
running_bs = len(self.running_batch.reqs)
running_bs = len(running_batch.reqs)
if (
self.get_num_allocatable_reqs(running_bs) <= 0
and self.dllm_manager.is_empty()
and not self.enable_priority_preemption
):
self.running_batch.batch_is_full = True
running_batch.batch_is_full = True
return True
return False
def _create_dllm_prefill_adder(self: Scheduler, running_bs: int) -> PrefillAdder:
def _create_dllm_prefill_adder(
self: Scheduler, running_bs: int, running_batch: ScheduleBatch
) -> PrefillAdder:
"""Create a prefill adder configured for DLLM scheduling."""
return PrefillAdder(
self.page_size,
self.tree_cache,
self.token_to_kv_pool_allocator,
self.running_batch,
running_batch,
self.new_token_ratio_tracker.current,
self.max_prefill_tokens,
self.chunked_prefill_size,
@@ -211,7 +214,9 @@ class SchedulerDllmMixin:
dllm_config=self.dllm_config,
)
def _process_dllm_batches(self: Scheduler, adder: PrefillAdder) -> ForwardMode:
def _process_dllm_batches(
self: Scheduler, adder: PrefillAdder, running_batch: ScheduleBatch
) -> ForwardMode:
"""Process prefill or decode batches for DLLM."""
forward_mode = ForwardMode.DLLM_EXTEND
@@ -223,6 +228,7 @@ class SchedulerDllmMixin:
prefill_reqs,
DllmReqPhase.STAGING_PREFILL,
DllmReqPhase.INCOMING_PREFILL,
running_batch=running_batch,
)
else:
# Fall back to decode batch
@@ -232,6 +238,7 @@ class SchedulerDllmMixin:
decode_reqs,
DllmReqPhase.STAGING_DECODE,
DllmReqPhase.INCOMING_DECODE,
running_batch=running_batch,
)
return forward_mode
@@ -242,6 +249,7 @@ class SchedulerDllmMixin:
batch: List[Req],
staging_phase: DllmReqPhase,
incoming_phase: DllmReqPhase,
running_batch: ScheduleBatch,
) -> None:
"""Process a batch, separating staging and incoming requests."""
staging_reqs = [req for req in batch if req.dllm_phase == staging_phase]
@@ -252,10 +260,12 @@ class SchedulerDllmMixin:
incoming_reqs = [req for req in batch if req.dllm_phase == incoming_phase]
if incoming_reqs:
self.process_dllm_incoming_reqs(adder, incoming_reqs)
self.process_dllm_incoming_reqs(
adder, incoming_reqs, running_batch=running_batch
)
def _update_state_for_batch(
self: Scheduler, can_run_list: List[Req], adder: PrefillAdder, running_bs: int
self: Scheduler, can_run_list: List[Req], adder: PrefillAdder
) -> None:
"""Update state for the batch."""
@@ -267,12 +277,12 @@ class SchedulerDllmMixin:
self.dllm_manager.add_staging_reqs(can_run_list)
self.dllm_manager.increment_inflight_middle_chunks()
self.adder = adder
self.can_run_list = can_run_list
self.running_bs = len(self.running_batch.reqs)
def _create_dllm_batch(
self: Scheduler, can_run_list: List[Req], forward_mode: ForwardMode
self: Scheduler,
can_run_list: List[Req],
forward_mode: ForwardMode,
adder: PrefillAdder,
running_batch: ScheduleBatch,
) -> ScheduleBatch:
"""Create and prepare a new DLLM batch."""
new_batch = ScheduleBatch.init_new(
@@ -295,24 +305,27 @@ class SchedulerDllmMixin:
)
new_batch.prefill_stats = PrefillStats.from_adder(
self.adder, self.running_batch.reqs, self.enable_priority_scheduling
adder, running_batch.reqs, self.enable_priority_scheduling
)
return new_batch
def process_dllm_incoming_reqs(
self: Scheduler, adder: PrefillAdder, reqs: List[Req]
self: Scheduler,
adder: PrefillAdder,
reqs: List[Req],
running_batch: ScheduleBatch,
) -> AddReqResult:
"""Process incoming DLLM requests with resource allocation and preemption."""
res = AddReqResult.CONTINUE
for req in reqs:
# Check if batch is full
running_bs = len(self.running_batch.reqs)
running_bs = len(running_batch.reqs)
if len(adder.can_run_list) >= self.get_num_allocatable_reqs(running_bs):
self.running_batch.batch_is_full = True
running_batch.batch_is_full = True
# Try preemption if batch is full
if self.running_batch.batch_is_full:
if running_batch.batch_is_full:
if (
not self.enable_priority_preemption
or not adder.preempt_to_schedule(req, self.server_args)
@@ -329,7 +342,7 @@ class SchedulerDllmMixin:
if res != AddReqResult.CONTINUE:
if res == AddReqResult.NO_TOKEN:
self.running_batch.batch_is_full = True
running_batch.batch_is_full = True
break
return res
@@ -47,10 +47,10 @@ class SchedulerMultiplexMixin:
# TODO(jason-fxz): This is a temporary demo
def adjust_stream_groups(
self: Scheduler,
self: Scheduler, running_batch: ScheduleBatch
) -> tuple[int, tuple[ExternalStream, ExternalStream]]:
if not self.running_batch.is_empty() and self.split_prefill_batch:
decode_bs = self.running_batch.batch_size()
if not running_batch.is_empty() and self.split_prefill_batch:
decode_bs = running_batch.batch_size()
manual_divisions = self.pdmux_config.manual_divisions
if manual_divisions:
for i in range(len(manual_divisions)):
@@ -68,7 +68,7 @@ class SchedulerMultiplexMixin:
),
)
set_current_stream_idx(stream_idx)
elif not self.running_batch.is_empty():
elif not running_batch.is_empty():
set_current_stream_idx(self.real_sm_group_num - 1)
else:
set_current_stream_idx(0)
@@ -78,21 +78,23 @@ class SchedulerMultiplexMixin:
self.tp_worker.model_runner.update_decode_attn_backend(stream_idx)
return stream_idx, self.stream_groups[stream_idx]
def update_split_prefill_batch(self: Scheduler, sm_count: int) -> bool:
def update_split_prefill_batch(
self: Scheduler, sm_count: int, running_batch: ScheduleBatch
) -> tuple[bool, ScheduleBatch]:
if self.split_prefill_batch:
return False
return False, running_batch
# add new request
prefill_plan = self.get_new_batch_prefill(self.running_batch)
prefill_plan = self.get_new_batch_prefill(running_batch)
batch = prefill_plan.batch_to_run
self.running_batch = prefill_plan.running_batch
running_batch = prefill_plan.running_batch
if batch and not batch.is_empty():
batch.forward_mode = (
ForwardMode.SPLIT_PREFILL
) # Set forward mode for split prefill
self.split_prefill_batch = batch
return True
return False
return True, running_batch
return False, running_batch
@torch.inference_mode()
def event_loop_pdmux(self: Scheduler):
@@ -114,28 +116,34 @@ class SchedulerMultiplexMixin:
set_pdmux_status(False)
recv_reqs = self.request_receiver.recv_requests()
self.process_input_requests(recv_reqs)
running_batch = self.running_batch
with torch.cuda.stream(prefill_stream):
set_pdmux_status(True)
sm_count = self.sm_counts[stream_idx][0]
if not wait_prefill_kernel_done:
adjust_stream_group = (
self.update_split_prefill_batch(sm_count) or adjust_stream_group
created, running_batch = self.update_split_prefill_batch(
sm_count, running_batch=running_batch
)
self.running_batch = running_batch
adjust_stream_group = created or adjust_stream_group
with torch.cuda.stream(decode_stream):
set_pdmux_status(False)
self.running_batch = self.update_running_batch(self.running_batch)
running_batch = self.update_running_batch(running_batch)
self.running_batch = running_batch
adjust_stream_group = adjust_stream_group or (
stream_idx > 0 and self.running_batch.is_empty()
stream_idx > 0 and running_batch.is_empty()
)
if self.running_batch.is_empty() and self.split_prefill_batch is None:
if running_batch.is_empty() and self.split_prefill_batch is None:
self.on_idle()
if adjust_stream_group:
prefill_stream.synchronize()
decode_stream.synchronize()
stream_idx, stream_group = self.adjust_stream_groups()
stream_idx, stream_group = self.adjust_stream_groups(
running_batch=running_batch
)
prefill_stream = stream_group[0]
decode_stream = stream_group[1]
adjust_stream_group = False
@@ -146,8 +154,8 @@ class SchedulerMultiplexMixin:
with torch.cuda.stream(decode_stream):
set_pdmux_status(False)
# process decode batch
if self.running_batch and not self.running_batch.is_empty():
decode_result = self.run_batch(self.running_batch)
if running_batch and not running_batch.is_empty():
decode_result = self.run_batch(running_batch)
decode_done = True
else:
decode_done = False
@@ -192,7 +200,7 @@ class SchedulerMultiplexMixin:
set_pdmux_status(False)
decode_stream.synchronize()
if decode_done:
self.process_batch_result(self.running_batch, decode_result)
self.process_batch_result(running_batch, decode_result)
with torch.cuda.stream(prefill_stream):
set_pdmux_status(True)
@@ -210,10 +218,11 @@ class SchedulerMultiplexMixin:
self.process_batch_result(
self.split_prefill_batch, prefill_result
)
if self.running_batch and not self.running_batch.is_empty():
self.running_batch.merge_batch(self.split_prefill_batch)
if running_batch and not running_batch.is_empty():
running_batch.merge_batch(self.split_prefill_batch)
else:
self.running_batch = self.split_prefill_batch
running_batch = self.split_prefill_batch
self.running_batch = running_batch
self.split_prefill_batch = None
wait_prefill_kernel_done = False