Spaces:
Running
Running
| 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 | |