File size: 2,145 Bytes
0c6c82c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
44745f2
 
 
0c6c82c
 
44745f2
 
 
0c6c82c
44745f2
0c6c82c
44745f2
0c6c82c
 
 
 
 
 
 
 
 
44745f2
0c6c82c
 
 
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
from __future__ import annotations

import math

from .latency import AnalyticalLatencyModel
from .models import Request, SimulationConfig


class KVCacheModel:
    def __init__(self, latency_model: AnalyticalLatencyModel, cfg: SimulationConfig):
        self.latency_model = latency_model
        self.cfg = cfg
        self.model_weight_gb = latency_model.model_weight_gb
        total_vram = latency_model.accelerator.vram_gb
        remaining = max(0.0, total_vram - self.model_weight_gb - 1.2)
        self.capacity_gb = remaining * cfg.kv_memory_fraction
        self.shared_prefix_gb = 0.0
        if cfg.prefix_cache_enabled and cfg.shared_prefix_tokens > 0 and cfg.prefix_reuse_fraction > 0:
            self.shared_prefix_gb = cfg.shared_prefix_tokens * latency_model.kv_bytes_per_token() / 1e9

    def _allocated_tokens(self, req: Request, include_output_reservation: bool = False) -> int:
        # Cached prefix state is represented once by shared_prefix_gb. Per-request
        # allocation therefore contains only the uncached suffix + generated state.
        live_tokens = req.uncached_prompt_tokens + req.generated_tokens
        if self.cfg.scheduler == "static_fcfs" or include_output_reservation:
            return req.uncached_prompt_tokens + req.output_tokens
        block = max(self.cfg.kv_block_tokens, 1)
        return int(math.ceil(max(live_tokens, 0) / block) * block)

    def request_gb(self, req: Request, include_output_reservation: bool = False) -> float:
        tokens = self._allocated_tokens(req, include_output_reservation)
        return tokens * self.latency_model.kv_bytes_per_token() / 1e9

    def used_gb(self, active: list[Request], prefill_pending: list[Request] | None = None) -> float:
        requests = list(active)
        if prefill_pending:
            requests.extend(prefill_pending)
        return self.shared_prefix_gb + sum(self.request_gb(r) for r in requests)

    def can_admit(self, req: Request, active: list[Request], prefill_pending: list[Request] | None = None) -> bool:
        return self.used_gb(active, prefill_pending) + self.request_gb(req) <= self.capacity_gb + 1e-12