Avoid implicit field-based side channel in Scheduler planning (#29408)

This commit is contained in:
fzyzcjy
2026-07-10 08:55:51 +08:00
committed by GitHub
parent 32c8973ce8
commit 1e75ba236e
13 changed files with 275 additions and 134 deletions
+34 -20
View File
@@ -60,7 +60,11 @@ from sglang.srt.disaggregation.utils import (
setup_state_kv_args,
)
from sglang.srt.environ import envs
from sglang.srt.managers.schedule_batch import FINISH_ABORT, ScheduleBatch
from sglang.srt.managers.schedule_batch import (
FINISH_ABORT,
NextBatchPlan,
ScheduleBatch,
)
from sglang.srt.managers.schedule_policy import match_prefix_for_req
from sglang.srt.managers.utils import GenerationBatchResult
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
@@ -1900,7 +1904,11 @@ class SchedulerDisaggregationDecodeMixin:
self.process_decode_queue()
# Get the next batch to run
batch = self.get_next_disagg_decode_batch_to_run()
plan = self.get_next_disagg_decode_batch_to_run(
running_batch=self.running_batch
)
self.running_batch = plan.running_batch
batch = plan.batch_to_run
self.cur_batch_for_debug = batch
# Launch the current batch
@@ -1934,10 +1942,16 @@ class SchedulerDisaggregationDecodeMixin:
self._apply_war_barrier()
# Get the next batch to run
batch = self.get_next_disagg_decode_batch_to_run()
plan = self.get_next_disagg_decode_batch_to_run(
running_batch=self.running_batch
)
self.running_batch = plan.running_batch
batch = plan.batch_to_run
self.cur_batch_for_debug = batch
# overlap + spec + grammar is unsupported (would desync DP ranks).
disable_overlap_for_batch = self.is_disable_overlap_for_batch(batch)
disable_overlap_for_batch = self.is_disable_overlap_for_batch(
batch, last_batch=self.last_batch
)
if disable_overlap_for_batch and self.last_batch:
pop_and_process()
@@ -1976,11 +1990,11 @@ class SchedulerDisaggregationDecodeMixin:
@scheduler_nvtx_method("scheduler.get_next_batch_to_run")
def get_next_disagg_decode_batch_to_run(
self: Scheduler,
) -> Optional[ScheduleBatch]:
self: Scheduler, running_batch: ScheduleBatch
) -> NextBatchPlan:
"""Process prebuilt batch and schedule the next decode batch."""
# Process pending prebuilt batch: output processing + filter + merge
new_prebuilt_batch = self.get_new_prebuilt_batch()
new_prebuilt_batch = self.get_new_prebuilt_batch(running_batch)
if new_prebuilt_batch:
assert self.chunked_req is None
self.batch_result_processor.process_batch_result_prebuilt(
@@ -1988,28 +2002,28 @@ class SchedulerDisaggregationDecodeMixin:
)
new_prebuilt_batch.filter_batch()
if not new_prebuilt_batch.is_empty():
if self.running_batch.is_empty():
self.running_batch = new_prebuilt_batch
if running_batch.is_empty():
running_batch = new_prebuilt_batch
if self.enable_hisparse:
self.running_batch.hisparse_coordinator = (
self.hisparse_coordinator
)
running_batch.hisparse_coordinator = self.hisparse_coordinator
else:
self.running_batch.merge_batch(new_prebuilt_batch)
running_batch.merge_batch(new_prebuilt_batch)
# Schedule decode batch
if self.running_batch.is_empty():
if running_batch.is_empty():
ret = None
else:
self.running_batch = self.update_running_batch(self.running_batch)
ret = self.running_batch if not self.running_batch.is_empty() else None
running_batch = self.update_running_batch(running_batch)
ret = running_batch if not running_batch.is_empty() else None
ret = self.dp_attn_adapter.maybe_prepare_mlp_sync_batch(ret)
if ret:
set_schedule_time_batch(ret)
return ret
return NextBatchPlan(batch_to_run=ret, running_batch=running_batch)
def get_new_prebuilt_batch(self: Scheduler) -> Optional[ScheduleBatch]:
def get_new_prebuilt_batch(
self: Scheduler, running_batch: ScheduleBatch
) -> Optional[ScheduleBatch]:
"""Create a schedulebatch for fake completed prefill"""
if self.grammar_manager.has_waiting_grammars():
ready_grammar_requests = self.grammar_manager.get_ready_grammar_requests()
@@ -2020,9 +2034,9 @@ class SchedulerDisaggregationDecodeMixin:
return None
if self.enable_priority_scheduling:
self.policy.calc_priority(self.waiting_queue, self.running_batch)
self.policy.calc_priority(self.waiting_queue, running_batch)
curr_batch_size = self.running_batch.batch_size()
curr_batch_size = running_batch.batch_size()
batch_size = min(self.req_to_token_pool.size, self.max_running_requests)
+33 -18
View File
@@ -52,6 +52,7 @@ from sglang.srt.environ import envs
from sglang.srt.managers.schedule_batch import (
FINISH_ABORT,
FINISH_LENGTH,
NextBatchPlan,
Req,
ScheduleBatch,
)
@@ -468,24 +469,28 @@ class SchedulerDisaggregationPrefillMixin:
@scheduler_nvtx_method("scheduler.get_next_batch_to_run")
def get_next_disagg_prefill_batch_to_run(
self: Scheduler,
) -> Optional[ScheduleBatch]:
running_batch: ScheduleBatch,
last_batch: Optional[ScheduleBatch],
) -> NextBatchPlan:
self.process_pending_chunked_abort()
# HACK (byronhsu): reset the batch_is_full flag because we never enter update_running_batch which resets it
# Otherwise, it hangs under high concurrency
self.running_batch.batch_is_full = False
running_batch.batch_is_full = False
self.process_prefill_chunk()
self.process_prefill_chunk(last_batch=last_batch, running_batch=running_batch)
self.resolve_waiting_queue_bootstrap()
batch = self.get_new_batch_prefill()
prefill_plan = self.get_new_batch_prefill(running_batch)
batch = prefill_plan.batch_to_run
running_batch = prefill_plan.running_batch
batch = self.dp_attn_adapter.maybe_prepare_mlp_sync_batch(batch)
if batch:
set_schedule_time_batch(batch)
return batch
return NextBatchPlan(batch_to_run=batch, running_batch=running_batch)
@torch.no_grad()
def event_loop_normal_disagg_prefill(self: Scheduler) -> None:
@@ -501,7 +506,11 @@ class SchedulerDisaggregationPrefillMixin:
)
# Get the next batch to run
batch = self.get_next_disagg_prefill_batch_to_run()
plan = self.get_next_disagg_prefill_batch_to_run(
running_batch=self.running_batch, last_batch=self.last_batch
)
self.running_batch = plan.running_batch
batch = plan.batch_to_run
self.cur_batch_for_debug = batch
# Launch the current batch
@@ -535,7 +544,11 @@ class SchedulerDisaggregationPrefillMixin:
self._apply_war_barrier()
# Get the next batch to run
batch = self.get_next_disagg_prefill_batch_to_run()
plan = self.get_next_disagg_prefill_batch_to_run(
running_batch=self.running_batch, last_batch=self.last_batch
)
self.running_batch = plan.running_batch
batch = plan.batch_to_run
self.cur_batch_for_debug = batch
# Launch the current batch
@@ -940,7 +953,11 @@ class SchedulerDisaggregationPrefillMixin:
req, polls[0], defer_release=self.enable_overlap
)
def process_prefill_chunk(self: Scheduler) -> None:
def process_prefill_chunk(
self: Scheduler,
last_batch: Optional[ScheduleBatch],
running_batch: ScheduleBatch,
) -> None:
chunked_req_to_exclude = set()
if self.chunked_req:
chunked_req_to_exclude.add(self.chunked_req)
@@ -958,20 +975,18 @@ class SchedulerDisaggregationPrefillMixin:
self.send_kv_chunk(self.chunked_req)
if self.chunked_req is not None:
self.running_batch.batch_is_full = False
running_batch.batch_is_full = False
if self.last_batch and self.last_batch.forward_mode.is_extend():
if self.last_batch.chunked_req:
if last_batch and last_batch.forward_mode.is_extend():
if last_batch.chunked_req:
# In the context pipeline parallelism, after the last chunk, the current microbatch still track outdated chunked_req.
# We need to discard it.
chunked_req_to_exclude.add(self.last_batch.chunked_req)
chunked_req_to_exclude.add(last_batch.chunked_req)
last_bs = self.last_batch.batch_size()
self.last_batch.filter_batch(
chunked_req_to_exclude=list(chunked_req_to_exclude)
)
if self.last_batch.batch_size() < last_bs:
self.running_batch.batch_is_full = False
last_bs = last_batch.batch_size()
last_batch.filter_batch(chunked_req_to_exclude=list(chunked_req_to_exclude))
if last_batch.batch_size() < last_bs:
running_batch.batch_is_full = False
def maybe_send_cached_prefix_chunk(self: Scheduler, req: Req) -> None:
# Only bootstrap-finalized requests; staging excluded.
+4 -1
View File
@@ -27,8 +27,11 @@ class SchedulerDllmMixin:
)
self.dllm_manager = DllmManager(dllm_config=self.dllm_config)
def get_new_batch_dllm(self: Scheduler) -> Optional[ScheduleBatch]:
def get_new_batch_dllm(
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
@@ -246,7 +246,11 @@ class SchedulerMlxOverlapMixin:
self._finalize_mlx_pending_job(pending_next)
self.result_queue.popleft()
pending_next = None
next_batch = self.get_next_batch_to_run()
plan = self.get_next_batch_to_run(
running_batch=self.running_batch, last_batch=self.last_batch
)
self.running_batch = plan.running_batch
next_batch = plan.batch_to_run
self.cur_batch_for_debug = next_batch
if next_batch:
pending_curr = _launch_fresh(next_batch)
@@ -58,6 +58,7 @@ from typing import (
Union,
)
import msgspec
import numpy as np
import torch
@@ -3058,3 +3059,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
f"ScheduleBatch(forward_mode={self.forward_mode.name if self.forward_mode else 'None'}, "
f"#req={(len(self.reqs))})"
)
class NextBatchPlan(msgspec.Struct):
batch_to_run: Optional[ScheduleBatch]
running_batch: ScheduleBatch
+97 -82
View File
@@ -164,6 +164,7 @@ from sglang.srt.managers.prefill_delayer import (
from sglang.srt.managers.schedule_batch import (
FINISH_ABORT,
MultimodalInputs,
NextBatchPlan,
Req,
ScheduleBatch,
)
@@ -1455,16 +1456,16 @@ class Scheduler(
]
)
def _abort_on_running_timeout(self):
def _abort_on_running_timeout(self, running_batch: ScheduleBatch):
# NOTE: this should be called before a batch is launched.
timeout_s = envs.SGLANG_REQ_RUNNING_TIMEOUT.get()
if timeout_s <= 0:
return
if self.running_batch.is_empty():
if running_batch.is_empty():
return
deadline = time.perf_counter() - timeout_s
for req in self.running_batch.reqs:
for req in running_batch.reqs:
if not req.finished() and 0 < req.time_stats.forward_entry_time < deadline:
req.to_finish = FINISH_ABORT(
"Request running timeout reached.", HTTPStatus.SERVICE_UNAVAILABLE
@@ -1541,7 +1542,11 @@ class Scheduler(
continue
# Get the next batch to run
batch = self.get_next_batch_to_run()
plan = self.get_next_batch_to_run(
running_batch=self.running_batch, last_batch=self.last_batch
)
self.running_batch = plan.running_batch
batch = plan.batch_to_run
self.cur_batch_for_debug = batch
# Launch the current batch
@@ -1582,9 +1587,15 @@ class Scheduler(
self._apply_war_barrier()
# Get the next batch to run
batch = self.get_next_batch_to_run()
plan = self.get_next_batch_to_run(
running_batch=self.running_batch, last_batch=self.last_batch
)
self.running_batch = plan.running_batch
batch = plan.batch_to_run
self.cur_batch_for_debug = batch
disable_overlap_for_batch = self.is_disable_overlap_for_batch(batch)
disable_overlap_for_batch = self.is_disable_overlap_for_batch(
batch, last_batch=self.last_batch
)
# If we do not need to overlap the current batch with the last batch,
# we can process the last batch immediately.
@@ -1625,7 +1636,9 @@ class Scheduler(
if envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.get():
self.invariant_checker.self_check_during_busy()
def is_disable_overlap_for_batch(self, batch: ScheduleBatch) -> bool:
def is_disable_overlap_for_batch(
self, batch: ScheduleBatch, last_batch: Optional[ScheduleBatch]
) -> bool:
# For two consecutive prefill batches, we disable overlap to improve the TTFT of the first batch.
# This might slightly hurt the throughput, so we use an environment variable to control it.
# In DP attention mode, use the globally synchronized is_extend_in_batch
@@ -1637,7 +1650,7 @@ class Scheduler(
is_extend = lambda b: b and b.forward_mode.is_extend()
batch_is_extend = is_extend(batch)
last_batch_is_extend = is_extend(self.last_batch)
last_batch_is_extend = is_extend(last_batch)
disable_overlap_for_batch = (
envs.SGLANG_DISABLE_CONSECUTIVE_PREFILL_OVERLAP.get()
@@ -2593,13 +2606,15 @@ class Scheduler(
return batch
@scheduler_nvtx_method("scheduler.get_next_batch_to_run")
def get_next_batch_to_run(self) -> Optional[ScheduleBatch]:
def get_next_batch_to_run(
self, running_batch: ScheduleBatch, last_batch: Optional[ScheduleBatch]
) -> NextBatchPlan:
self.process_pending_chunked_abort()
if self.enable_fpm:
self._fpm_batch_t0 = time.monotonic()
self._abort_on_waiting_timeout()
self._abort_on_running_timeout()
self._abort_on_running_timeout(running_batch)
if self.dllm_config is not None:
self.dllm_manager.filter_finished_reqs()
@@ -2629,57 +2644,57 @@ class Scheduler(
ready_reqs = self.hisparse_coordinator.collect_ready_reqs()
if len(ready_reqs) > 0:
new_batch = self._build_hisparse_decode_batch(ready_reqs)
if self.running_batch.is_empty():
self.running_batch = new_batch
if running_batch.is_empty():
running_batch = new_batch
else:
self.running_batch.merge_batch(new_batch)
self.running_batch.hisparse_coordinator = self.hisparse_coordinator
running_batch.merge_batch(new_batch)
running_batch.hisparse_coordinator = self.hisparse_coordinator
# Reset batch_is_full so the scheduler can schedule more prefills.
self.running_batch.batch_is_full = False
running_batch.batch_is_full = False
if (
not self.enable_hisparse
and self.last_batch
and self.last_batch.forward_mode.is_extend()
and last_batch
and last_batch.forward_mode.is_extend()
):
if self.last_batch.chunked_req is not None:
if last_batch.chunked_req is not None:
# In the context pipeline parallelism, after the last chunk, the current microbatch still track outdated chunked_req.
# We need to discard it.
chunked_req_to_exclude.add(self.last_batch.chunked_req)
chunked_req_to_exclude.add(last_batch.chunked_req)
if self.dllm_config is not None and self.last_batch.reqs:
chunked_req_to_exclude.update(self.last_batch.reqs)
if self.dllm_config is not None and last_batch.reqs:
chunked_req_to_exclude.update(last_batch.reqs)
# Filter batch
last_bs = self.last_batch.batch_size()
self.last_batch.filter_batch(
chunked_req_to_exclude=list(chunked_req_to_exclude)
)
if self.last_batch.batch_size() < last_bs:
self.running_batch.batch_is_full = False
last_bs = last_batch.batch_size()
last_batch.filter_batch(chunked_req_to_exclude=list(chunked_req_to_exclude))
if last_batch.batch_size() < last_bs:
running_batch.batch_is_full = False
# Merge the new batch into the running batch.
if not self.last_batch.is_empty():
if self.running_batch.is_empty():
self.running_batch = self.last_batch
if not last_batch.is_empty():
if running_batch.is_empty():
running_batch = last_batch
else:
# Merge running_batch with prefill batch
self.running_batch.merge_batch(self.last_batch)
running_batch.merge_batch(last_batch)
# For prefill-only batch, filter out finished requests since they
# won't go through the decode step. This keeps running_batch accurate
# for load reporting (num_running_reqs via /v1/loads).
# Runs outside the last_batch block so stale requests are cleaned
# even when no new batches arrive (e.g. traffic stops).
if self.running_batch.is_prefill_only:
self.running_batch.filter_batch()
if self.running_batch.is_empty():
self.running_batch.batch_is_full = False
if running_batch.is_prefill_only:
running_batch.filter_batch()
if running_batch.is_empty():
running_batch.batch_is_full = False
if self.dllm_config is not None:
new_batch = self.get_new_batch_dllm()
new_batch = self.get_new_batch_dllm(running_batch)
else:
new_batch = self.get_new_batch_prefill()
prefill_plan = self.get_new_batch_prefill(running_batch)
new_batch = prefill_plan.batch_to_run
running_batch = prefill_plan.running_batch
need_mlp_sync = self.require_mlp_sync
if (
@@ -2699,12 +2714,9 @@ class Scheduler(
ret = new_batch
else:
# Run decode (skip for prefill-only batches)
if (
not self.running_batch.is_empty()
and not self.running_batch.is_prefill_only
):
self.running_batch = self.update_running_batch(self.running_batch)
ret = self.running_batch if not self.running_batch.is_empty() else None
if not running_batch.is_empty() and not running_batch.is_prefill_only:
running_batch = self.update_running_batch(running_batch)
ret = running_batch if not running_batch.is_empty() else None
else:
ret = None
@@ -2721,14 +2733,14 @@ class Scheduler(
if self.enable_fpm:
ret.fpm_start_time = self._fpm_batch_t0
return ret
return NextBatchPlan(batch_to_run=ret, running_batch=running_batch)
def get_num_allocatable_reqs(self, running_bs):
res = get_server_args().pp_max_micro_batch_size - running_bs
res = min(res, self.req_to_token_pool.available_size())
return res
def get_new_batch_prefill(self) -> Optional[ScheduleBatch]:
def get_new_batch_prefill(self, running_batch: ScheduleBatch) -> NextBatchPlan:
prefill_delayer_single_pass = None
if self.prefill_delayer:
# Get max usage across all pools for prefill delay decision
@@ -2739,18 +2751,21 @@ class Scheduler(
self.prefill_delayer, token_usage=max_pool_usage
)
ret = self._get_new_batch_prefill_raw(
prefill_delayer_single_pass=prefill_delayer_single_pass
ret, running_batch = self._get_new_batch_prefill_raw(
prefill_delayer_single_pass=prefill_delayer_single_pass,
running_batch=running_batch,
)
if self.prefill_delayer:
prefill_delayer_single_pass.finalize(actual_prefill=ret is not None)
return ret
return NextBatchPlan(batch_to_run=ret, running_batch=running_batch)
def _get_new_batch_prefill_raw(
self, prefill_delayer_single_pass: Optional[PrefillDelayerSinglePassExecutor]
) -> Optional[ScheduleBatch]:
self,
prefill_delayer_single_pass: Optional[PrefillDelayerSinglePassExecutor],
running_batch: ScheduleBatch,
) -> Tuple[Optional[ScheduleBatch], ScheduleBatch]:
# Check if the grammar is ready in the grammar queue
if self.grammar_manager.has_waiting_grammars():
ready_grammar_requests = self.grammar_manager.get_ready_grammar_requests()
@@ -2762,14 +2777,14 @@ class Scheduler(
if self.enable_priority_preemption or self.is_hybrid_swa:
# Reset batch_is_full to try preemption with a prefill adder.
self.running_batch.batch_is_full = False
running_batch.batch_is_full = False
if (
self.running_batch.batch_is_full or len(self.waiting_queue) == 0
running_batch.batch_is_full or len(self.waiting_queue) == 0
) and self.chunked_req is None:
return None
return None, running_batch
running_bs = len(self.running_batch.reqs)
running_bs = len(running_batch.reqs)
# Skipped during a chunked prefill: that pass must proceed regardless.
if (
self.min_free_slots_delayer is not None
@@ -2779,7 +2794,7 @@ class Scheduler(
num_allocatable_reqs=self.get_num_allocatable_reqs(running_bs),
)
):
return None
return None, running_batch
# Ignore the check if self.chunked_req is not None.
# In the non-PP case, when self.chunked_req is not None, num_allocatable_reqs should always be greater than 0,
@@ -2791,17 +2806,17 @@ class Scheduler(
and self.chunked_req is None
and not self.enable_priority_preemption
):
self.running_batch.batch_is_full = True
return None
running_batch.batch_is_full = True
return None, running_batch
# Get priority queue
self.policy.calc_priority(self.waiting_queue, self.running_batch)
self.policy.calc_priority(self.waiting_queue, running_batch)
if TEST_RETRACT and running_bs > TEST_RETRACT_NO_PREFILL_BS:
# If we are testing retraction and the running batch size exceeds
# TEST_RETRACT_NO_PREFILL_BS, we skip the prefill to keep the requests
# in the waiting queue.
return None
return None, running_batch
# Determine chunked_prefill_size for this batch
chunked_prefill_size = self.chunked_prefill_size
@@ -2816,7 +2831,7 @@ class Scheduler(
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,
chunked_prefill_size,
@@ -2836,7 +2851,7 @@ class Scheduler(
if self.enable_lora:
running_loras = {
req.lora_id for req in self.running_batch.reqs if not req.finished()
req.lora_id for req in running_batch.reqs if not req.finished()
}
# Account for LoRAs that are already loaded in the adder, such as chunked requests
running_loras.update(req.lora_id for req in adder.can_run_list)
@@ -2844,7 +2859,7 @@ class Scheduler(
if self.lora_drainer:
self.lora_drainer.update_draining_state(
self.waiting_queue,
self.running_batch.reqs,
running_batch.reqs,
)
mamba_allocator = getattr(self.req_to_token_pool, "mamba_allocator", None)
@@ -2855,16 +2870,16 @@ class Scheduler(
if self.enable_lora and not self._can_schedule_lora_req(req, running_loras):
continue
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
if self.disaggregation_mode == DisaggregationMode.PREFILL:
# In prefill mode, prealloc queue and transfer queue can also take memory,
# so we need to check if the available size for the actual available size.
if len(adder.can_run_list) >= self.req_to_token_pool.available_size():
self.running_batch.batch_is_full = True
running_batch.batch_is_full = True
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)
@@ -2895,11 +2910,11 @@ class Scheduler(
if res == AddReqResult.NO_TOKEN:
if self.enable_hierarchical_cache:
# Set batch_is_full after making sure there are requests that can be served
self.running_batch.batch_is_full = len(
adder.can_run_list
) > 0 or (not self.running_batch.is_empty())
running_batch.batch_is_full = len(adder.can_run_list) > 0 or (
not running_batch.is_empty()
)
else:
self.running_batch.batch_is_full = True
running_batch.batch_is_full = True
# revert matched mamba idx to avoid memory leak, if req is not added.
# Only free if the slot was freshly allocated in this batch (not
# pre-existing from a session). Session-held slots have their own
@@ -2925,7 +2940,7 @@ class Scheduler(
# Update waiting queue
can_run_list: List[Req] = adder.can_run_list
if len(can_run_list) == 0:
return None
return None, running_batch
can_run_set = set(can_run_list)
self.waiting_queue = [x for x in self.waiting_queue if x not in can_run_set]
@@ -2975,7 +2990,7 @@ class Scheduler(
# Record prefill stats for logging after forward.
new_batch.prefill_stats = PrefillStats.from_adder(
adder,
self.running_batch.reqs,
running_batch.reqs,
self.enable_priority_scheduling,
num_pending_tokens=self.load_inquirer._get_num_pending_tokens(
chunk_deduct=(
@@ -2989,24 +3004,24 @@ class Scheduler(
# Mixed-style chunked prefill
if (
self.is_mixed_chunk
and not self.running_batch.is_empty()
and not (new_batch.return_logprob or self.running_batch.return_logprob)
and not running_batch.is_empty()
and not (new_batch.return_logprob or running_batch.return_logprob)
# mix_with_running cats input_ids but not input_embeds — shapes would mismatch
and new_batch.input_embeds is None
):
# TODO (lianmin): support return_logprob + mixed chunked prefill
self.running_batch.filter_batch()
if not self.running_batch.is_empty():
self.running_batch.prepare_for_decode()
new_batch.mix_with_running(self.running_batch)
new_batch.decoding_reqs = self.running_batch.reqs
self.running_batch = ScheduleBatch(
reqs=[], batch_is_full=self.running_batch.batch_is_full
running_batch.filter_batch()
if not running_batch.is_empty():
running_batch.prepare_for_decode()
new_batch.mix_with_running(running_batch)
new_batch.decoding_reqs = running_batch.reqs
running_batch = ScheduleBatch(
reqs=[], batch_is_full=running_batch.batch_is_full
)
else:
new_batch.decoding_reqs = None
return new_batch
return new_batch, running_batch
def _can_schedule_lora_req(
self, req: Req, running_loras: set[Optional[str]]
@@ -108,7 +108,11 @@ class SchedulerPPMixin:
async_send=True,
)
with torch.profiler.record_function("get_next_batch_to_run"):
self.mbs[mb_id] = self.get_next_batch_to_run()
plan = self.get_next_batch_to_run(
running_batch=self.running_batch, last_batch=self.last_batch
)
self.running_batch = plan.running_batch
self.mbs[mb_id] = plan.batch_to_run
self.running_mbs[mb_id] = self.running_batch
cur_batch: Optional[ScheduleBatch] = self.mbs[mb_id]
self.cur_batch_for_debug = cur_batch
@@ -248,8 +252,12 @@ class SchedulerPPMixin:
self._pp_commit_comm_work(send_transfer_work)
tmbs[mb_id] = transferred_rids
self.process_prefill_chunk()
batch = self.get_new_batch_prefill()
self.process_prefill_chunk(
last_batch=self.last_batch, running_batch=self.running_batch
)
prefill_plan = self.get_new_batch_prefill(self.running_batch)
batch = prefill_plan.batch_to_run
self.running_batch = prefill_plan.running_batch
batch = self.dp_attn_adapter.maybe_prepare_mlp_sync_batch(batch)
self.mbs[mb_id] = batch
self.running_mbs[mb_id] = self.running_batch
@@ -402,7 +410,11 @@ class SchedulerPPMixin:
self._pp_commit_comm_work(send_transfer_work)
# get batch to run and proxy tensors if needed
batch = self.get_next_disagg_decode_batch_to_run()
plan = self.get_next_disagg_decode_batch_to_run(
running_batch=self.running_batch
)
self.running_batch = plan.running_batch
batch = plan.batch_to_run
self.mbs[mb_id] = batch
self.running_mbs[mb_id] = self.running_batch
@@ -83,7 +83,9 @@ class SchedulerMultiplexMixin:
return False
# add new request
batch = self.get_new_batch_prefill()
prefill_plan = self.get_new_batch_prefill(self.running_batch)
batch = prefill_plan.batch_to_run
self.running_batch = prefill_plan.running_batch
if batch and not batch.is_empty():
batch.forward_mode = (
ForwardMode.SPLIT_PREFILL