File size: 4,902 Bytes
d53adc9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
# Exported for HuggingFace trust_remote_code loading.
# This file is intentionally self-contained.

"""HF cache adapter owned by the Sophia HF adapter layer."""

from __future__ import annotations

import weakref

import torch
from transformers.cache_utils import Cache, CacheLayerMixin

from .cache_decode import RuntimeCacheState
from .model_state import LayerCacheSnapshot, RuntimeCacheSnapshot


CacheState = RuntimeCacheState


class SophiaCacheLayer(CacheLayerMixin):
    def __init__(
        self,
        *,
        name: str,
        snapshot: LayerCacheSnapshot,
        seq_length: int,
        max_cache_shape: int,
    ):
        super().__init__()
        self.name = str(name)
        self.snapshot = snapshot.clone()
        self.payload = {
            str(key): value
            for key, value in self.snapshot.to_payload(prefix=self.name).items()
        }
        self._seq_length = int(seq_length)
        self._max_cache_shape = int(max_cache_shape)
        representative = next(
            (
                value
                for _name, value in self.snapshot.tensor_fields()
                if value is not None
            ),
            None,
        )
        if representative is not None:
            self.keys = representative.unsqueeze(1)
            self.values = self.keys
            self.is_initialized = True

    def lazy_initialization(self, key_states: torch.Tensor, value_states: torch.Tensor) -> None:
        del key_states, value_states
        raise NotImplementedError("SophiaCacheLayer is immutable and cannot be initialized lazily")

    def update(
        self,
        key_states: torch.Tensor,
        value_states: torch.Tensor,
        cache_kwargs: dict[str, object] | None = None,
    ) -> tuple[torch.Tensor, torch.Tensor]:
        del key_states, value_states, cache_kwargs
        raise NotImplementedError("SophiaCacheLayer is an exported runtime snapshot and does not support update()")

    def get_mask_sizes(self, cache_position: torch.Tensor) -> tuple[int, int]:
        del cache_position
        return self.get_seq_length(), 0

    def get_seq_length(self) -> int:
        return self._seq_length

    def get_max_cache_shape(self) -> int:
        return self._max_cache_shape

    @property
    def max_batch_size(self) -> int:
        if self.keys is None:
            return 0
        return int(self.keys.size(0))

    @property
    def max_cache_len(self) -> int:
        return int(self._max_cache_shape)

    @property
    def device(self) -> torch.device:
        if self.keys is None:
            return torch.device("cpu")
        return self.keys.device


class SophiaCache(Cache):
    def __init__(
        self,
        *,
        owner: object,
        cache: RuntimeCacheSnapshot,
        cache_pos: int,
        batch_size: int,
    ):
        self._cache = cache.clone()
        self._cache_pos = int(cache_pos)
        self._batch_size = int(batch_size)
        self._owner_ref = weakref.ref(owner)
        super().__init__(layers=self._build_layers())

    def _build_layers(self) -> list[SophiaCacheLayer]:
        owner = self._owner_ref()
        max_cache_shape = (
            -1
            if owner is None
            else int(getattr(owner.config, "max_position_embeddings", 0) or -1)
        )
        return [
            SophiaCacheLayer(
                name=name,
                snapshot=snapshot,
                seq_length=self._cache_pos,
                max_cache_shape=max_cache_shape,
            )
            for name, snapshot in self._cache.named_snapshots()
        ]

    def to_runtime_cache(self) -> RuntimeCacheSnapshot:
        return self._cache.clone()

    def get_seq_length(self, layer_idx: int = 0) -> int:
        del layer_idx
        return int(self._cache_pos)

    def get_max_cache_shape(self, layer_idx: int = 0) -> int:
        del layer_idx
        owner = self._owner_ref()
        if owner is None:
            return -1
        return int(getattr(owner.config, "max_position_embeddings", 0) or -1)

    @property
    def cache_pos(self) -> int:
        return int(self._cache_pos)

    @property
    def batch_size(self) -> int:
        return int(self._batch_size)


def cache_state_from_past_key_values(
    past_key_values: SophiaCache | None,
) -> CacheState | None:
    if past_key_values is None:
        return None
    if isinstance(past_key_values, SophiaCache):
        return CacheState(
            cache=past_key_values.to_runtime_cache(),
            batch_size=int(past_key_values.batch_size),
            cache_pos=int(past_key_values.cache_pos),
        )
    raise TypeError("past_key_values must be a SophiaCache returned by the HF adapter")


__all__ = [
    "CacheState",
    "SophiaCache",
    "SophiaCacheLayer",
    "cache_state_from_past_key_values",
]