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)