Skip to content

Instantly share code, notes, and snippets.

@smqd19
Created May 7, 2026 10:34
Show Gist options
  • Select an option

  • Save smqd19/b22b63f9d000570d9cf628c695a80d45 to your computer and use it in GitHub Desktop.

Select an option

Save smqd19/b22b63f9d000570d9cf628c695a80d45 to your computer and use it in GitHub Desktop.
Embedding Cache with LRU + Disk Persistence — RAG Optimization (Python)
"""
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