Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
47 changes: 33 additions & 14 deletions python/infinilm/llm/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -133,6 +133,7 @@ def schedule(self) -> Optional[SchedulerOutput]:
is_prefill = False
current_num_batched_tokens = 0
current_prefill_extra_blocks = 0
running_required_blocks = None

# Process Waiting queue (prefill phase)
while (
Expand Down Expand Up @@ -199,10 +200,18 @@ def schedule(self) -> Optional[SchedulerOutput]:
deferred_requests.append(req)
break

if running_required_blocks is None and (
self.mamba_cache_manager is None
or req.mamba_cache_index is not None
or self.mamba_cache_manager.can_allocate()
):
# Running requests do not advance during this prefill loop.
running_required_blocks = self._get_running_required_blocks()
if not self.can_accept_request(
req,
num_local_computed_tokens,
current_prefill_extra_blocks,
running_required_blocks=running_required_blocks,
):
logger.warning(
"Insufficient KV cache blocks for request %s, deferring.",
Expand Down Expand Up @@ -450,33 +459,43 @@ def complete_requests(self, requests: List[InferenceRequest]):
# Still running, put back in running queue
self.running_queue.sync_q.put(req)

def _get_running_required_blocks(self) -> int:
"""Sum decode reservations while preserving running queue order."""
total_required_blocks = 0
running_queue_size = self.running_queue.sync_q.qsize()
for _ in range(running_queue_size):
req = self.running_queue.sync_q.get()
remaining_tokens = (
req.sampling_params.max_tokens - req.get_num_generated_tokens()
)
num_blocks_needed = (
remaining_tokens + self.block_size - 1
) // self.block_size
total_required_blocks += num_blocks_needed
self.running_queue.sync_q.put(req)
return total_required_blocks

def can_accept_request(
self,
request: InferenceRequest,
num_local_computed_tokens: int,
current_prefill_extra_blocks: int = 0,
*,
running_required_blocks: int | None = None,
) -> bool:
"""Check capacity, optionally reusing this schedule's running reservations."""
if (
self.mamba_cache_manager is not None
and request.mamba_cache_index is None
and not self.mamba_cache_manager.can_allocate()
):
return False

total_required_blocks = 0

# Calculate blocks needed for running requests
running_queue_size = self.running_queue.sync_q.qsize()
for _ in range(running_queue_size):
req = self.running_queue.sync_q.get()
remaining_tokens = (
req.sampling_params.max_tokens - req.get_num_generated_tokens()
)
num_blocks_needed = (
remaining_tokens + self.block_size - 1
) // self.block_size
total_required_blocks += num_blocks_needed
self.running_queue.sync_q.put(req)
total_required_blocks = (
self._get_running_required_blocks()
if running_required_blocks is None
else running_required_blocks
)

# Calculate blocks needed for the new request
total_length = request.get_prompt_length() - num_local_computed_tokens
Expand Down