File size: 5,157 Bytes
d91766b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 | from __future__ import annotations
from abc import ABC, abstractmethod
from collections import deque, defaultdict
from typing import Callable
from diffulex.config import Config
from diffulex.engine.kv_cache_manager import AutoKVCacheManager
from diffulex.engine.request import DllmReq
from diffulex.engine.status import DllmReqStatus
from diffulex.engine.strategy_registry import DiffulexStrategyRegistry
class SchedulerBase(ABC):
def __init__(self, config: Config):
self.config = config
self.max_num_reqs = config.max_num_reqs
self.max_num_batched_tokens = config.max_num_batched_tokens
self.eos = config.eos
self.kv_cache_manager = AutoKVCacheManager.from_config(config)
self.waiting_reqs: deque[DllmReq] = deque()
self.running_reqs: deque[DllmReq] = deque()
def is_finished(self) -> bool:
return not self.waiting_reqs and not self.running_reqs
def abort_request(self, req_id: int) -> bool:
for req in list(self.waiting_reqs):
if req.req_id == req_id:
self.waiting_reqs.remove(req)
req.status = DllmReqStatus.FINISHED
setattr(req, "completion_reason", "aborted")
return True
for req in list(self.running_reqs):
if req.req_id == req_id:
self.running_reqs.remove(req)
setattr(req, "completion_reason", "aborted")
req.status = DllmReqStatus.FINISHED
self.kv_cache_manager.free(req)
return True
return False
@abstractmethod
def add(self, req: DllmReq) -> None:
pass
@abstractmethod
def schedule(self) -> tuple[list[DllmReq], bool]:
pass
@abstractmethod
def preempt(self, req: DllmReq) -> None:
pass
@abstractmethod
def postprocess(self, reqs: list[DllmReq], sampler_output):
pass
class DataParallelScheduler:
def __init__(self, config: Config, scheduler_factory: Callable[[Config], SchedulerBase]):
self.config = config
self.dp_size = config.data_parallel_size
self.schedulers = [scheduler_factory(config) for _ in range(self.dp_size)]
def _owner_for_req(self, req: DllmReq) -> int:
owner_assigned = bool(getattr(req, "_dp_owner_assigned", False))
owner = getattr(req, "dp_rank", None)
if not owner_assigned or owner is None or not (0 <= owner < self.dp_size):
owner = req.req_id % self.dp_size
assign_fn = getattr(req, "assign_dp_rank", None)
if callable(assign_fn):
assign_fn(owner)
else:
req.dp_rank = owner
return owner
def add(self, req: DllmReq) -> None:
owner = self._owner_for_req(req)
self.schedulers[owner].add(req)
def schedule(self) -> tuple[list[DllmReq], bool]:
scheduled: list[DllmReq] = []
saw_prefill = False
for scheduler in self.schedulers:
if scheduler.is_finished():
continue
local_reqs, is_prefill = scheduler.schedule()
scheduled.extend(local_reqs)
saw_prefill = saw_prefill or is_prefill
return scheduled, saw_prefill
def postprocess(self, reqs: list[DllmReq], sampler_output) -> None:
reqs_by_owner: dict[int, list[DllmReq]] = defaultdict(list)
for req in reqs:
reqs_by_owner[self._owner_for_req(req)].append(req)
for owner, local_reqs in reqs_by_owner.items():
if local_reqs:
self.schedulers[owner].postprocess(local_reqs, sampler_output)
def is_finished(self) -> bool:
return all(scheduler.is_finished() for scheduler in self.schedulers)
def abort_request(self, req_id: int) -> bool:
for scheduler in self.schedulers:
if scheduler.abort_request(req_id):
return True
return False
SchedulerFactory = Callable[[Config], SchedulerBase]
class AutoScheduler(DiffulexStrategyRegistry):
"""Registry-driven factory for scheduler implementations."""
@classmethod
def _from_single_config(cls, config: Config) -> SchedulerBase:
cls._ensure_strategies_loaded()
cls._MODULE_MAPPING: dict[str, SchedulerFactory]
candidates: list[str] = []
if config.decoding_strategy:
candidates.append(config.decoding_strategy)
candidates.append(cls._DEFAULT_KEY)
for key in candidates:
factory = cls._MODULE_MAPPING.get(key)
if factory is not None:
return factory(config)
available = ", ".join(cls.available_modules()) or "<none>"
raise ValueError(
"No scheduler registered for decoding_strategy="
f"'{config.decoding_strategy}'. Available schedulers: {available}."
)
@classmethod
def from_config(cls, config: Config) -> SchedulerBase | DataParallelScheduler:
if config.data_parallel_size > 1:
return DataParallelScheduler(config, scheduler_factory=cls._from_single_config)
return cls._from_single_config(config)
|