import time from typing import List import numpy as np from sarathi.config import ( CacheConfig, ModelConfig, ParallelConfig, SarathiSchedulerConfig, ) from sarathi.core.block_space_manager.sarathi_block_space_manager import ( SarathiBlockSpaceManager, ) from sarathi.core.datatypes.scheduler_output import SchedulerOutputs from sarathi.core.datatypes.sequence import Sequence, SequenceScheduleMetadata from sarathi.core.scheduler.base_scheduler import BaseScheduler from sarathi.logger import init_logger logger = init_logger(__name__) class SarathiScheduler(BaseScheduler): def __init__( self, model_config: ModelConfig, scheduler_config: SarathiSchedulerConfig, cache_config: CacheConfig, parallel_config: ParallelConfig, ) -> None: super().__init__(model_config, scheduler_config, cache_config, parallel_config) self.chunk_size = self.scheduler_config.chunk_size self.enable_dynamic_chunking_schedule = ( self.scheduler_config.enable_dynamic_chunking_schedule ) # next four params apply only when using dynamic schedule self.low_chunk_size = self.scheduler_config.low_chunk_size self.high_chunk_size = self.scheduler_config.high_chunk_size self.chunk_schedule_max_tokens = self.scheduler_config.chunk_schedule_max_tokens self.chunk_schedule_stages = self.scheduler_config.chunk_schedule_stages if self.enable_dynamic_chunking_schedule: assert self.chunk_schedule_stages > 0 assert self.chunk_schedule_max_tokens > 0 assert self.low_chunk_size % 32 == 0 assert self.high_chunk_size % 32 == 0 self._chunk_sizes = self._compute_chunk_size_schedule() self._tokens_per_stage = int( np.ceil(self.chunk_schedule_max_tokens / self.chunk_schedule_stages) ) def _compute_chunk_size_schedule(self): # create num_steps equally spaced chunk sizes between low_chunk_size and high_chunk_size chunk_sizes = np.linspace( self.low_chunk_size, self.high_chunk_size, self.chunk_schedule_stages, dtype=np.int32, )[::-1] # align each chunk size to the nearest multiple of 32 or self.low_chunk_size round_of_chunk_sizes = min(32, self.low_chunk_size) chunk_sizes = ( np.round(chunk_sizes / round_of_chunk_sizes) * round_of_chunk_sizes ) chunk_sizes = chunk_sizes.astype(np.int64).tolist() return chunk_sizes def get_block_space_manager_class(self): return SarathiBlockSpaceManager def _get_seq_next_num_prefill_tokens( self, seq: Sequence, num_batched_tokens: int ) -> int: assert not seq.is_finished() if self.enable_dynamic_chunking_schedule: request_stage_idx = int( np.ceil( seq.get_num_prompt_tokens_stage_processed() // self._tokens_per_stage ) ) assert request_stage_idx < len(self._chunk_sizes) chunk_size = self._chunk_sizes[request_stage_idx] else: chunk_size = self.chunk_size next_num_tokens = min( seq.get_prompt_len() - seq.get_num_prompt_tokens_stage_processed(), chunk_size - num_batched_tokens, ) return next_num_tokens def _schedule(self) -> SchedulerOutputs: # Fix the current time. now = time.monotonic() running: List[Sequence] = [] ignored_seq_ids: List[str] = [] preempted_seq_ids: List[str] = [] scheduled_seq_metadata_list: List[SequenceScheduleMetadata] = [] num_batched_tokens: int = 0 ###################################################################### # Phase 1: Add existing running sequence groups to the batch. # There are two cases: # 1. The sequence group has incomplete prefill. The routine # remains identical to the one in sarathi scheduler for such sequences. # 2. The sequence group has completed prefill. In this case, we need to # check for memory availability for the next chunk of decode tokens, and preempt # some sequence groups if necessary. Note that, the preempted sequence groups # might belong to either of the two categories. ###################################################################### # NOTE(woosuk): Preemption happens only when there is no available slot # to keep all the sequence groups in the RUNNING state. # In this case, the policy is responsible for deciding which sequence # groups to preempt. self.running = self.policy.sort_by_priority(now, self.running) # in first pass process all the requests with prefill completed # this allows us to accurately account for the number of decode tokens running_prefills: List[Sequence] = [] while self.running: seq = self.running.pop(0) if not seq.is_paused(): running.append(seq) continue if not seq.prompt_stage_processing_finished: running_prefills.append(seq) continue while not self.block_manager.can_append_slot(): if self.running: # Preempt the lowest-priority sequence groups. victim_seq = self.running.pop(-1) self._preempt(victim_seq) preempted_seq_ids.append(victim_seq.seq_id) else: # No other sequence groups can be preempted. # Preempt the current sequence group. self._preempt(seq) preempted_seq_ids.append(seq.seq_id) break else: # Append new slots to the sequence group. self._append_slot(seq) running.append(seq) num_batched_tokens += 1 scheduled_seq_metadata_list.append( SequenceScheduleMetadata.from_sequence(seq) ) # now add the requests with prefill incomplete # the memory for all these prefills has already been allocated # so we should be able to run all of them for seq in running_prefills: assert not seq.prompt_stage_processing_finished next_num_prefill_tokens = self._get_seq_next_num_prefill_tokens( seq, num_batched_tokens ) # as long as the request could fit in the batch previously # it should be able to fit in the batch now # so in non-pipeline case this condition should always be false # however, in pipeline case, the grouping of requests can change # between different microbatches, so this is not guaranteed to be always true if next_num_prefill_tokens == 0: running.append(seq) continue num_batched_tokens += next_num_prefill_tokens scheduled_seq_metadata_list.append( SequenceScheduleMetadata.from_sequence( seq, prompt_chunk_len=next_num_prefill_tokens ) ) running.append(seq) ###################################################################### # Phase 2: Add waiting (new) sequence groups to the batch. # This routine is nearly-identical to the one in sarathi scheduler ###################################################################### # Optimization: We do not sort the waiting queue since the preempted # sequence groups are added to the front and the new sequence groups # are added to the back. while self.waiting: seq = self.waiting[0] # This is required to handle benchmarking where we set request arrival time ahead of time if seq.arrival_time > now: break if not self._check_request_prompt_length(seq): ignored_seq_ids.append(seq.seq_id) continue # If the sequence group cannot be allocated, stop. if not self.block_manager.can_allocate(seq): # this is different from vllm scheduler # even if we cannot allocate this sequence group # there might be other sequence groups that can be allocated break # The total number of sequences in the RUNNING state should not # exceed the maximum number of sequences. if len(running) >= self.scheduler_config.max_num_seqs: break # check if we can fit the prefill in the batch next_num_prefill_tokens = self._get_seq_next_num_prefill_tokens( seq, num_batched_tokens ) if next_num_prefill_tokens == 0: break seq = self.waiting.pop(0) self._allocate(seq) num_batched_tokens += next_num_prefill_tokens scheduled_seq_metadata_list.append( SequenceScheduleMetadata.from_sequence( seq, prompt_chunk_len=next_num_prefill_tokens ) ) running.append(seq) # make sure that prefills are at the start of the batch, so that we don't violate assumptions # made in the original vllm codebase self.running = running return SchedulerOutputs( id=self._iteration_id, ignored_seq_ids=ignored_seq_ids, preempted_seq_ids=preempted_seq_ids, scheduled_seq_metadata_list=scheduled_seq_metadata_list, )