suryatmodulus
/

GPC-1 / gpc1_server /vision_cache.py
suryatmodulus's picture harshatheg's picture
Duplicate from harshatheg/GPC-1
96a4100
Raw
History Blame Contribute Delete
3.84 kB
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)