from __future__ import annotations import hashlib import threading from typing import Any class VisionFeatureCache: """One-image, memory-bounded cache for a frozen vision encoder.""" def __init__(self, visual: Any, torch: Any, max_bytes: int = 64 * 1024 * 1024): self.torch = torch self.max_bytes = max_bytes self._key = None self._value = None self._lock = threading.RLock() original = visual.forward def forward(*args, **kwargs): if visual.training or torch.is_grad_enabled() or self.max_bytes <= 0: self.clear() return original(*args, **kwargs) with self._lock: key = self._fingerprint((args, kwargs)) if key is not None and key == self._key: return self._clone(self._value) output = original(*args, **kwargs) self._key = self._value = None size = self._size(output) if key is not None and size is not None and size <= self.max_bytes: self._value = self._clone(output) self._key = key return output visual.forward = forward def clear(self): with self._lock: self._key = self._value = None def _fingerprint(self, value): torch = self.torch digest = hashlib.sha256() def update(item): if isinstance(item, torch.Tensor): if item.layout != torch.strided: raise TypeError digest.update(repr((str(item.dtype), str(item.device), tuple(item.shape), tuple(item.stride()))).encode()) digest.update(item.detach().contiguous().cpu().view(torch.uint8).numpy().tobytes()) elif isinstance(item, (tuple, list)): digest.update(type(item).__name__.encode()) for child in item: update(child) digest.update(b'\0') elif isinstance(item, dict): digest.update(b'dict') for key in sorted(item): update(key) update(item[key]) elif item is None or type(item) in (bool, int, float, str): digest.update(repr((type(item).__name__, item)).encode()) else: raise TypeError try: update(value) update((torch.backends.cuda.matmul.allow_tf32, torch.backends.cudnn.allow_tf32, torch.backends.cudnn.enabled, torch.backends.cudnn.benchmark, torch.backends.cudnn.deterministic, torch.get_float32_matmul_precision())) except (TypeError, ValueError): return None return digest.digest() def _clone(self, value): if isinstance(value, self.torch.Tensor): return value.detach().clone() if isinstance(value, dict): cloned = {k: self._clone(v) for k, v in value.items()} return cloned if type(value) is dict else type(value)(**cloned) if isinstance(value, tuple): return tuple(self._clone(v) for v in value) if isinstance(value, list): return [self._clone(v) for v in value] return value def _size(self, value): if isinstance(value, self.torch.Tensor): return value.numel() * value.element_size() if isinstance(value, dict): sizes = [self._size(v) for v in value.values()] elif isinstance(value, (tuple, list)): sizes = [self._size(v) for v in value] elif value is None or type(value) in (bool, int, float, str): return 0 else: return None return None if any(v is None for v in sizes) else sum(sizes)