File size: 4,031 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
"""Req base class and registry."""

from __future__ import annotations

from copy import copy
from itertools import count
from typing import Callable

from diffulex.config import Config
from diffulex.sampling_params import SamplingParams
from diffulex.engine.strategy_registry import DiffulexStrategyRegistry
from diffulex.engine.status import DllmReqStatus
from diffulex.mixin.request_state import ReqStateMixin


class DllmReq(ReqStateMixin):
    """Minimal base class that tracks prompt tokens and cache bookkeeping."""

    page_size = 32
    counter = count()

    def __init__(self, token_ids: list[int], sampling_params: SamplingParams = SamplingParams()):
        self.req_id = next(DllmReq.counter)
        self.status = DllmReqStatus.WAITING
        self.dp_rank = 0
        self._dp_owner_assigned = False
        self.token_ids = copy(token_ids)
        self.last_token = token_ids[-1]
        self.num_prompt_tokens = len(token_ids)
        self.num_cached_tokens = 0
        self.page_table: list[int] = []
        self.page_cache_missed: list[bool] = []
        self.temperature = sampling_params.temperature
        self.max_tokens = sampling_params.max_tokens
        self.max_nfe = sampling_params.max_nfe
        self.max_repetition_run = sampling_params.max_repetition_run
        self.ignore_eos = sampling_params.ignore_eos
        self.new_tokens = 0
        self.nfe = 0
        self.meet_eos = False
        self.is_multi_block = False
        self._execution_prepared = False

    def __len__(self) -> int:
        return self.num_tokens

    def __getitem__(self, key) -> int:
        return self.token_ids[key]

    @property
    def num_tokens(self) -> int:
        return len(self.token_ids)

    @property
    def is_finished(self) -> bool:
        return self.status == DllmReqStatus.FINISHED

    @property
    def prompt_token_ids(self) -> list[int]:
        return self.token_ids[: self.num_prompt_tokens]

    @property
    def num_pages(self) -> int:
        if self.is_multi_block:
            # return (self.running_len + self.page_size - 1) // self.page_size
            return (self.to_cache_len + self.page_size - 1) // self.page_size
        else:
            return (self.num_tokens + self.page_size - 1) // self.page_size

    @property
    def last_page_num_tokens(self) -> int:
        return self.num_tokens - (self.num_pages - 1) * self.page_size

    def reset_new_tokens(self):
        self.new_tokens = 0

    @property
    def is_execution_prepared(self) -> bool:
        return bool(self._execution_prepared)

    def mark_execution_prepared(self) -> None:
        self._execution_prepared = True

    def clear_execution_prepared(self) -> None:
        self._execution_prepared = False

    def assign_dp_rank(self, dp_rank: int) -> None:
        self.dp_rank = dp_rank
        self._dp_owner_assigned = True

    def page(self, index: int) -> list[int]:
        assert 0 <= index < self.num_pages
        return self.token_ids[index * self.page_size : (index + 1) * self.page_size]


ReqFactory = Callable[[list[int], SamplingParams, Config], DllmReq]


class AutoReq(DiffulexStrategyRegistry):
    """Registry-driven factory for req implementations."""

    @classmethod
    def create(
        cls,
        config: Config,
        token_ids: list[int],
        sampling_params: SamplingParams = SamplingParams(),
    ) -> DllmReq:
        cls._MODULE_MAPPING: dict[str, ReqFactory]
        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(token_ids, sampling_params, config)

        available = ", ".join(cls.available_modules()) or "<none>"
        raise ValueError(
            "No req registered for decoding_strategy="
            f"'{config.decoding_strategy}'. Available reqs: {available}."
        )