Created
May 7, 2026 10:34
-
-
Save smqd19/b22b63f9d000570d9cf628c695a80d45 to your computer and use it in GitHub Desktop.
Embedding Cache with LRU + Disk Persistence — RAG Optimization (Python)
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| """ | |
| Embedding Cache with LRU + Disk Persistence | |
| Author: Sheikh Muhammad Qasim | ML Architect | |
| Production embedding cache that reduced our API costs by ~70%. | |
| Supports LRU eviction, disk persistence, and TTL expiry. | |
| """ | |
| import hashlib | |
| import json | |
| import time | |
| from collections import OrderedDict | |
| from pathlib import Path | |
| from dataclasses import dataclass | |
| @dataclass | |
| class CacheStats: | |
| hits: int = 0 | |
| misses: int = 0 | |
| evictions: int = 0 | |
| @property | |
| def hit_rate(self) -> float: | |
| total = self.hits + self.misses | |
| return self.hits / total if total > 0 else 0.0 | |
| class EmbeddingCache: | |
| def __init__(self, max_size: int = 10000, ttl_seconds: int = 86400, | |
| persist_path: str | None = None): | |
| self.max_size = max_size | |
| self.ttl_seconds = ttl_seconds | |
| self.persist_path = Path(persist_path) if persist_path else None | |
| self._cache: OrderedDict[str, tuple[list[float], float]] = OrderedDict() | |
| self.stats = CacheStats() | |
| if self.persist_path and self.persist_path.exists(): | |
| self._load() | |
| @staticmethod | |
| def _key(text: str, model: str = "default") -> str: | |
| return hashlib.sha256(f"{model}:{text}".encode()).hexdigest() | |
| def get(self, text: str, model: str = "default") -> list[float] | None: | |
| key = self._key(text, model) | |
| if key not in self._cache: | |
| self.stats.misses += 1 | |
| return None | |
| embedding, timestamp = self._cache[key] | |
| if time.time() - timestamp > self.ttl_seconds: | |
| del self._cache[key] | |
| self.stats.misses += 1 | |
| return None | |
| self._cache.move_to_end(key) | |
| self.stats.hits += 1 | |
| return embedding | |
| def put(self, text: str, embedding: list[float], model: str = "default"): | |
| key = self._key(text, model) | |
| if key in self._cache: | |
| self._cache.move_to_end(key) | |
| self._cache[key] = (embedding, time.time()) | |
| while len(self._cache) > self.max_size: | |
| self._cache.popitem(last=False) | |
| self.stats.evictions += 1 | |
| def save(self): | |
| if not self.persist_path: | |
| return | |
| data = {k: {"emb": v[0], "ts": v[1]} for k, v in self._cache.items()} | |
| self.persist_path.write_text(json.dumps(data), encoding="utf-8") | |
| def _load(self): | |
| data = json.loads(self.persist_path.read_text(encoding="utf-8")) | |
| now = time.time() | |
| for k, v in data.items(): | |
| if now - v["ts"] < self.ttl_seconds: | |
| self._cache[k] = (v["emb"], v["ts"]) | |
| if __name__ == "__main__": | |
| cache = EmbeddingCache(max_size=100, ttl_seconds=3600) | |
| cache.put("hello world", [0.1, 0.2, 0.3]) | |
| cache.put("machine learning", [0.4, 0.5, 0.6]) | |
| print("Hit:", cache.get("hello world")) | |
| print("Miss:", cache.get("unknown")) | |
| print(f"Stats: {cache.stats.hits} hits, {cache.stats.misses} misses, " | |
| f"rate={cache.stats.hit_rate:.0%}") |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment