| from collections import OrderedDict | |
| class LRUCache: | |
| def __init__(self, max_size=50, max_memory_mb=500): | |
| self.cache = OrderedDict() | |
| self.max_size = max_size | |
| self.max_memory_mb = max_memory_mb | |
| self.current_memory_mb = 0 | |
| def _estimate_size_mb(self, tensor): | |
| if hasattr(tensor, 'element_size'): | |
| return tensor.element_size() * tensor.nelement() / (1024 * 1024) | |
| return 0 | |
| def get(self, key): | |
| if key in self.cache: | |
| self.cache.move_to_end(key) | |
| return self.cache[key] | |
| return None | |
| def set(self, key, value): | |
| size_mb = self._estimate_size_mb(value) | |
| if size_mb > self.max_memory_mb: | |
| return | |
| if key in self.cache: | |
| self.current_memory_mb -= self._estimate_size_mb(self.cache[key]) | |
| del self.cache[key] | |
| while (len(self.cache) >= self.max_size or | |
| self.current_memory_mb + size_mb > self.max_memory_mb): | |
| if len(self.cache) == 0: | |
| break | |
| old_key, old_value = self.cache.popitem(last=False) | |
| self.current_memory_mb -= self._estimate_size_mb(old_value) | |
| self.cache[key] = value | |
| self.cache.move_to_end(key) | |
| self.current_memory_mb += size_mb | |
| def clear(self): | |
| self.cache.clear() | |
| self.current_memory_mb = 0 | |