suryatmodulus
/

File size: 3,840 Bytes
96a4100
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
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)