"""Unit-safe KV-cache estimates for decoder-only transformer serving.""" from dataclasses import dataclass GIB = 1024**3 @dataclass(frozen=True) class ModelShape: layers: int kv_heads: int head_dim: int bytes_per_element: float def validate(self) -> None: if min(self.layers, self.kv_heads, self.head_dim) <= 0: raise ValueError("model dimensions must be positive") if self.bytes_per_element <= 0: raise ValueError("bytes_per_element must be positive") def cache_bytes( shape: ModelShape, tokens: int, sequences: int = 1, overhead: float = 1.0, ) -> float: shape.validate() if tokens <= 0 or sequences <= 0: raise ValueError("tokens and sequences must be positive") if overhead < 1: raise ValueError("overhead must be at least 1") return ( sequences * tokens * shape.layers * 2 * shape.kv_heads * shape.head_dim * shape.bytes_per_element * overhead ) def gib(value: float) -> float: return value / GIB def max_sequences( shape: ModelShape, tokens: int, available_gib: float, overhead: float = 1.0, ) -> int: if available_gib <= 0: raise ValueError("available_gib must be positive") per_sequence = cache_bytes(shape, tokens, overhead=overhead) return int((available_gib * GIB) // per_sequence)