| """ |
| Frox AI Morph 1.1 — KV Cache |
| |
| Improvements over Morph 1.0: |
| - Paged block allocation (small-scale PagedAttention-style idea): |
| cache grows in fixed-size blocks instead of one giant |
| pre-allocated [1, heads, max_seq_len, head_dim] tensor per layer. |
| This means a short 50-token chat doesn't reserve the same VRAM as |
| a 4096-token one, and growth doesn't require a full reallocation. |
| - Dynamic growth: if a sequence exceeds its current block budget, |
| new blocks are appended rather than raising an error or silently |
| overwriting (the 1.0 behavior silently dropped the oldest tokens, |
| which is wrong for a cache that's supposed to be exact — sliding |
| eviction now only happens if the caller explicitly opts in). |
| - bfloat16 default is preserved, but dtype/device now match the |
| model automatically via `MorphKVCache.for_model(model)`. |
| """ |
| from __future__ import annotations |
|
|
| from typing import Dict, List, Optional, Tuple |
|
|
| import torch |
|
|
|
|
| class MorphKVCache: |
| """ |
| Paged KV Cache for Morph 1.1 inference. |
| |
| Instead of one pre-allocated [1, kv_heads, max_seq_len, head_dim] |
| tensor per layer (Morph 1.0 behavior — wastes VRAM for short |
| sequences and requires knowing max_seq_len up front), Morph 1.1 |
| allocates in fixed-size blocks and appends new blocks on demand. |
| |
| Memory usage (3B model, batch=1, block_size=256): |
| Per block per layer: 2 × 8 kv_heads × 256 × 128 × 2 bytes = 1MB |
| A 300-token conversation needs 2 blocks/layer × 28 layers = 56MB |
| (vs. 224MB pre-allocated for a 2048 max_seq_len in 1.0) |
| """ |
|
|
| def __init__( |
| self, |
| num_layers: int, |
| num_kv_heads: int = 8, |
| head_dim: int = 128, |
| block_size: int = 256, |
| dtype: torch.dtype = torch.bfloat16, |
| device: torch.device = torch.device("cpu"), |
| max_seq_len: Optional[int] = None, |
| evict_on_overflow: bool = False, |
| ): |
| self.num_layers = num_layers |
| self.num_kv_heads = num_kv_heads |
| self.head_dim = head_dim |
| self.block_size = block_size |
| self.dtype = dtype |
| self.device = device |
| self.max_seq_len = max_seq_len |
| self.evict_on_overflow = evict_on_overflow |
|
|
| |
| self._k_blocks: List[List[torch.Tensor]] = [[] for _ in range(num_layers)] |
| self._v_blocks: List[List[torch.Tensor]] = [[] for _ in range(num_layers)] |
| self._seq_len = 0 |
|
|
| @classmethod |
| def for_model( |
| cls, |
| model, |
| block_size: int = 256, |
| max_seq_len: Optional[int] = None, |
| ) -> "MorphKVCache": |
| """Build a cache matching a MorphForCausalLM's config, dtype, device.""" |
| cfg = model.config |
| device = next(model.parameters()).device |
| dtype = next(model.parameters()).dtype |
| return cls( |
| num_layers=cfg.num_hidden_layers, |
| num_kv_heads=cfg.num_key_value_heads, |
| head_dim=cfg.head_dim, |
| block_size=block_size, |
| dtype=dtype, |
| device=device, |
| max_seq_len=max_seq_len, |
| ) |
|
|
| @property |
| def seq_len(self) -> int: |
| return self._seq_len |
|
|
| def _new_block(self) -> torch.Tensor: |
| return torch.zeros( |
| 1, self.num_kv_heads, self.block_size, self.head_dim, |
| dtype=self.dtype, device=self.device, |
| ) |
|
|
| def _ensure_capacity(self, layer_idx: int, needed_len: int): |
| """Append blocks until the layer can hold `needed_len` tokens.""" |
| current_capacity = len(self._k_blocks[layer_idx]) * self.block_size |
| while current_capacity < needed_len: |
| self._k_blocks[layer_idx].append(self._new_block()) |
| self._v_blocks[layer_idx].append(self._new_block()) |
| current_capacity += self.block_size |
|
|
| def update( |
| self, |
| layer_idx: int, |
| k: torch.Tensor, |
| v: torch.Tensor, |
| ) -> Tuple[torch.Tensor, torch.Tensor]: |
| """ |
| Append new K/V for a layer and return the full cached sequence. |
| |
| k, v: [B, num_kv_heads, new_len, head_dim] |
| returns: [B, num_kv_heads, total_len, head_dim] |
| """ |
| new_len = k.shape[2] |
| start = self._seq_len |
| end = start + new_len |
|
|
| if self.max_seq_len is not None and end > self.max_seq_len: |
| if self.evict_on_overflow: |
| self._evict_oldest(layer_idx, end - self.max_seq_len) |
| start = self._seq_len |
| end = start + new_len |
| |
| |
|
|
| self._ensure_capacity(layer_idx, end) |
|
|
| |
| |
| |
| |
| block_idx_start = start // self.block_size |
| offset_in_block = start % self.block_size |
|
|
| remaining = new_len |
| src_offset = 0 |
| b = block_idx_start |
| off = offset_in_block |
|
|
| while remaining > 0: |
| space_in_block = self.block_size - off |
| write_len = min(space_in_block, remaining) |
|
|
| self._k_blocks[layer_idx][b][:, :, off:off + write_len, :] = \ |
| k[:, :, src_offset:src_offset + write_len, :] |
| self._v_blocks[layer_idx][b][:, :, off:off + write_len, :] = \ |
| v[:, :, src_offset:src_offset + write_len, :] |
|
|
| remaining -= write_len |
| src_offset += write_len |
| b += 1 |
| off = 0 |
|
|
| k_full = self._read_range(layer_idx, self._k_blocks, 0, end) |
| v_full = self._read_range(layer_idx, self._v_blocks, 0, end) |
|
|
| return k_full, v_full |
|
|
| def _read_range( |
| self, layer_idx: int, blocks_container: List[List[torch.Tensor]], |
| start: int, end: int, |
| ) -> torch.Tensor: |
| blocks = blocks_container[layer_idx] |
| parts = [] |
| pos = 0 |
| for block in blocks: |
| block_start = pos |
| block_end = pos + self.block_size |
| if block_end > start and block_start < end: |
| lo = max(0, start - block_start) |
| hi = min(self.block_size, end - block_start) |
| parts.append(block[:, :, lo:hi, :]) |
| pos += self.block_size |
| if pos >= end: |
| break |
|
|
| if not parts: |
| |
| |
| |
| |
| return torch.zeros( |
| 1, self.num_kv_heads, 0, self.head_dim, |
| dtype=self.dtype, device=self.device, |
| ) |
| return torch.cat(parts, dim=2) if len(parts) > 1 else parts[0] |
|
|
| def _evict_oldest(self, layer_idx: int, num_to_evict: int): |
| """ |
| Shift the cache left by `num_to_evict` tokens (sliding window). |
| Only called when evict_on_overflow=True — an explicit opt-in. |
| """ |
| k_full = self._read_range(layer_idx, self._k_blocks, 0, self._seq_len) |
| v_full = self._read_range(layer_idx, self._v_blocks, 0, self._seq_len) |
|
|
| k_kept = k_full[:, :, num_to_evict:, :].contiguous() |
| v_kept = v_full[:, :, num_to_evict:, :].contiguous() |
|
|
| |
| self._k_blocks[layer_idx] = [] |
| self._v_blocks[layer_idx] = [] |
| self._seq_len = 0 |
| if k_kept.shape[2] > 0: |
| self._ensure_capacity(layer_idx, k_kept.shape[2]) |
| self._k_blocks[layer_idx][0][:, :, :k_kept.shape[2], :] = k_kept |
| self._v_blocks[layer_idx][0][:, :, :v_kept.shape[2], :] = v_kept |
| self._seq_len = max(0, self._seq_len - num_to_evict) |
|
|
| def step(self, new_tokens: int = 1): |
| """Advance sequence length after generating new_tokens.""" |
| self._seq_len += new_tokens |
|
|
| def reset(self): |
| """Clear cache for a new conversation.""" |
| self._k_blocks = [[] for _ in range(self.num_layers)] |
| self._v_blocks = [[] for _ in range(self.num_layers)] |
| self._seq_len = 0 |
|
|
| def get(self, layer_idx: int) -> Tuple[torch.Tensor, torch.Tensor]: |
| """Get cached KV for a layer as [B, kv_heads, seq_len, head_dim].""" |
| if self._seq_len == 0: |
| empty = torch.zeros( |
| 1, self.num_kv_heads, 0, self.head_dim, |
| dtype=self.dtype, device=self.device, |
| ) |
| return empty, empty |
| k = self._read_range(layer_idx, self._k_blocks, 0, self._seq_len) |
| v = self._read_range(layer_idx, self._v_blocks, 0, self._seq_len) |
| return k, v |
|
|
| def memory_mb(self) -> float: |
| """Current allocated memory (not just used) across all layers.""" |
| total_bytes = 0 |
| for layer_blocks in self._k_blocks: |
| for block in layer_blocks: |
| total_bytes += block.numel() * block.element_size() |
| for layer_blocks in self._v_blocks: |
| for block in layer_blocks: |
| total_bytes += block.numel() * block.element_size() |
| return total_bytes / (1024 ** 2) |
|
|
| def utilization(self) -> float: |
| """Fraction of allocated blocks actually holding real tokens.""" |
| allocated = len(self._k_blocks[0]) * self.block_size if self._k_blocks[0] else 0 |
| if allocated == 0: |
| return 0.0 |
| return round(self._seq_len / allocated, 3) |
|
|
| def to(self, device: torch.device) -> "MorphKVCache": |
| self._k_blocks = [[b.to(device) for b in layer] for layer in self._k_blocks] |
| self._v_blocks = [[b.to(device) for b in layer] for layer in self._v_blocks] |
| self.device = device |
| return self |
|
|
|
|
| class MorphSessionCache: |
| """ |
| Per-session KV cache manager. |
| Maintains separate paged caches for concurrent chat sessions, |
| evicting the least-recently-used session when at capacity. |
| """ |
|
|
| def __init__( |
| self, |
| num_layers: int = 28, |
| max_sessions: int = 16, |
| block_size: int = 256, |
| num_kv_heads: int = 8, |
| head_dim: int = 128, |
| dtype: torch.dtype = torch.bfloat16, |
| device: torch.device = torch.device("cpu"), |
| max_seq_len_per_session: Optional[int] = None, |
| ): |
| self.num_layers = num_layers |
| self.max_sessions = max_sessions |
| self.block_size = block_size |
| self.num_kv_heads = num_kv_heads |
| self.head_dim = head_dim |
| self.dtype = dtype |
| self.device = device |
| self.max_seq_len_per_session = max_seq_len_per_session |
|
|
| self._caches: Dict[str, MorphKVCache] = {} |
| self._lru_order: List[str] = [] |
|
|
| def get_or_create(self, session_id: str) -> MorphKVCache: |
| if session_id in self._caches: |
| self._touch(session_id) |
| return self._caches[session_id] |
|
|
| if len(self._caches) >= self.max_sessions: |
| lru_id = self._lru_order.pop(0) |
| del self._caches[lru_id] |
|
|
| self._caches[session_id] = MorphKVCache( |
| num_layers=self.num_layers, |
| block_size=self.block_size, |
| num_kv_heads=self.num_kv_heads, |
| head_dim=self.head_dim, |
| dtype=self.dtype, |
| device=self.device, |
| max_seq_len=self.max_seq_len_per_session, |
| ) |
| self._lru_order.append(session_id) |
| return self._caches[session_id] |
|
|
| def _touch(self, session_id: str): |
| if session_id in self._lru_order: |
| self._lru_order.remove(session_id) |
| self._lru_order.append(session_id) |
|
|
| def reset_session(self, session_id: str): |
| if session_id in self._caches: |
| self._caches[session_id].reset() |
|
|
| def delete_session(self, session_id: str): |
| if session_id in self._caches: |
| del self._caches[session_id] |
| if session_id in self._lru_order: |
| self._lru_order.remove(session_id) |
|
|
| def total_memory_mb(self) -> float: |
| return sum(c.memory_mb() for c in self._caches.values()) |
|
|
| def stats(self) -> Dict: |
| return { |
| "active_sessions": len(self._caches), |
| "max_sessions": self.max_sessions, |
| "total_memory_mb": round(self.total_memory_mb(), 1), |
| "avg_utilization": round( |
| sum(c.utilization() for c in self._caches.values()) / max(len(self._caches), 1), 3 |
| ), |
| } |
|
|