File size: 4,894 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
146
147
148
149
150
151
152
153
154
155
156
157
158
import xxhash
import weakref

import numpy as np

from typing import Callable
from collections import deque
from abc import ABC, abstractmethod
from dataclasses import dataclass, field

from diffulex.config import Config
from diffulex.engine.request import DllmReq
from diffulex.engine.strategy_registry import DiffulexStrategyRegistry


@dataclass
class Page:
    page_id: int
    ref_count: int = 0
    hash: int = -1
    token_ids: list[int] = field(default_factory=list)

    req: DllmReq | None = None

    def update(self, hash: int, token_ids: list[int]):
        self.hash = hash
        self.token_ids = token_ids

    def reset(self):
        self.ref_count = 1
        self.hash = -1
        self.token_ids = []
        self.req = None

    def set_req(self, req: DllmReq):
        self.req = weakref.ref(req)


class KVCacheManagerBase(ABC):
    def __init__(self, config: Config):
        num_pages = config.num_pages
        page_size = config.page_size
        assert num_pages > 0
        self.config = config
        self.page_size = page_size
        self.enable_prefix_caching = bool(config.enable_prefix_caching)
        self.pages: list[Page] = [Page(page_id=i) for i in range(num_pages)]
        self.hash_to_page_id: dict[int, int] = dict()
        self.free_page_ids: deque[int] = deque(range(num_pages))
        self.used_page_ids: set[int] = set()

    @classmethod
    def compute_hash(cls, token_ids: list[int], prefix: int = -1):
        h = xxhash.xxh64()

        if prefix != -1:
            h.update(prefix.to_bytes(8, "little"))

        h.update(np.array(token_ids).tobytes())
        return h.intdigest()

    def _allocate_page(self, page_id: int) -> Page:
        page = self.pages[page_id]
        assert page.ref_count == 0
        page.reset()
        self.free_page_ids.remove(page_id)
        self.used_page_ids.add(page_id)
        return self.pages[page_id]

    def _free_page(self, page_id: int) -> Page:
        assert self.pages[page_id].ref_count == 0
        self.used_page_ids.remove(page_id)
        self.free_page_ids.append(page_id)

    def can_allocate(self, req: DllmReq) -> bool:
        return len(self.free_page_ids) >= req.num_pages

    def allocate(self, req: DllmReq):
        assert not req.page_table
        req.page_cache_missed.clear()
        req.num_cached_tokens = 0
        h = -1
        cache_miss = False

        for i in range(req.num_pages):
            token_ids = req.page(i)
            h = self.compute_hash(token_ids, h) if len(token_ids) == self.page_size else -1
            page_id = self.hash_to_page_id.get(h, -1) if self.enable_prefix_caching else -1

            if page_id == -1 or self.pages[page_id].token_ids != token_ids:
                cache_miss = True

            req.page_cache_missed.append(cache_miss)

            if cache_miss:
                page_id = self.free_page_ids[0]
                page = self._allocate_page(page_id)
            else:
                req.num_cached_tokens += self.page_size
                if page_id in self.used_page_ids:
                    page = self.pages[page_id]
                    page.ref_count += 1
                else:
                    page = self._allocate_page(page_id)

            if h != -1:
                page.update(h, token_ids)
            if self.enable_prefix_caching and h != -1:
                self.hash_to_page_id[h] = page_id

            req.page_table.append(page_id)

    def free(self, req: DllmReq):
        for page_id in reversed(req.page_table):
            page = self.pages[page_id]
            page.ref_count -= 1
            if page.ref_count == 0:
                self._free_page(page_id)

        req.num_cached_tokens = 0
        req.page_cache_missed.clear()
        req.page_table.clear()

    @abstractmethod
    def can_append(self, req: DllmReq) -> bool:
        pass

    @abstractmethod
    def may_append(self, req: DllmReq) -> None:
        pass


KVCacheManagerFactory = Callable[[Config], "KVCacheManagerBase"]


class AutoKVCacheManager(DiffulexStrategyRegistry):
    """Registry-driven factory for page manager implementations."""

    @classmethod
    def from_config(cls, config: Config) -> KVCacheManagerBase:
        cls._ensure_strategies_loaded()
        cls._MODULE_MAPPING: dict[str, KVCacheManagerFactory]
        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 page manager registered for decoding_strategy="
            f"'{config.decoding_strategy}'. Available page managers: {available}."
        )