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