BiGRU_T_version / src /bigru_t /tools /tool_cache.py
PowerMachine's picture
V6: Kohonen 4D SOM + EWC per-neuron + MTP+entropy + Xeon V6
79e8e52 verified
Raw
History Blame Contribute Delete
2.95 kB
"""tool_cache.py — V6: tool cache for repeated invocations.
Cache LRU para resultados de ferramentas. Reduz computação redundante
quando a mesma entrada é processada múltiplas vezes.
"""
from __future__ import annotations
import hashlib
import time
from typing import Optional, Dict, Any, Callable
from collections import OrderedDict
import torch
class ToolCache:
"""V6: cache LRU para ferramentas.
Usage:
cache = ToolCache(max_size=1024, ttl_seconds=3600)
result = cache.get_or_compute("tool_name", input_tensor, compute_fn)
"""
def __init__(
self,
max_size: int = 1024,
ttl_seconds: float = 3600.0,
hash_fn: Optional[Callable] = None,
):
self.max_size = max_size
self.ttl_seconds = ttl_seconds
self.hash_fn = hash_fn or self._default_hash
self._cache: OrderedDict[str, tuple] = OrderedDict() # key -> (value, timestamp)
self._stats = {"hits": 0, "misses": 0, "evictions": 0}
@staticmethod
def _default_hash(key: Any) -> str:
if isinstance(key, torch.Tensor):
key = key.detach().cpu().numpy().tobytes()
elif isinstance(key, (list, tuple)):
key = str(key)
return hashlib.sha256(str(key).encode("utf-8")).hexdigest()
def _make_key(self, tool_name: str, input_key: Any) -> str:
return f"{tool_name}:{self.hash_fn(input_key)}"
def get_or_compute(
self,
tool_name: str,
input_key: Any,
compute_fn: Callable,
) -> Any:
key = self._make_key(tool_name, input_key)
now = time.time()
# Check cache
if key in self._cache:
value, ts = self._cache[key]
if now - ts < self.ttl_seconds:
self._cache.move_to_end(key)
self._stats["hits"] += 1
return value
else:
del self._cache[key]
# Compute
value = compute_fn()
self._cache[key] = (value, now)
self._stats["misses"] += 1
# Evict LRU
while len(self._cache) > self.max_size:
self._cache.popitem(last=False)
self._stats["evictions"] += 1
return value
def invalidate(self, tool_name: str) -> None:
"""Remove todas as entradas de uma ferramenta."""
keys_to_remove = [k for k in self._cache if k.startswith(f"{tool_name}:")]
for k in keys_to_remove:
del self._cache[k]
def clear(self) -> None:
self._cache.clear()
self._stats = {"hits": 0, "misses": 0, "evictions": 0}
def get_stats(self) -> Dict[str, Any]:
total = self._stats["hits"] + self._stats["misses"]
hit_rate = self._stats["hits"] / max(1, total)
return {
**self._stats,
"hit_rate": hit_rate,
"size": len(self._cache),
"max_size": self.max_size,
}
__all__ = ["ToolCache"]