diff --git a/python/infinilm/llm/scheduler.py b/python/infinilm/llm/scheduler.py index c10b55f2f..085837658 100644 --- a/python/infinilm/llm/scheduler.py +++ b/python/infinilm/llm/scheduler.py @@ -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 ( @@ -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.", @@ -450,12 +459,31 @@ 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 @@ -463,20 +491,11 @@ def can_accept_request( ): 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