File size: 4,601 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
# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle.
# Exported for HuggingFace trust_remote_code loading.
# This file is intentionally self-contained.

from __future__ import annotations

from collections.abc import Mapping
from dataclasses import dataclass, field
from threading import RLock

import torch


def _clone(value: torch.Tensor | None) -> torch.Tensor | None:
    return None if value is None else value.detach().clone()


_CACHE_SUFFIXES = (
    ("_recurrent", "recurrent"),
    ("_conv", "conv"),
    ("_latent", "latent"),
)


@dataclass(frozen=True)
class LayerCacheSnapshot:
    recurrent: torch.Tensor | None = None
    conv: torch.Tensor | None = None
    latent: torch.Tensor | None = None

    def clone(self) -> LayerCacheSnapshot:
        return LayerCacheSnapshot(
            recurrent=_clone(self.recurrent),
            conv=_clone(self.conv),
            latent=_clone(self.latent),
        )

    def tensor_fields(self) -> tuple[tuple[str, torch.Tensor | None], ...]:
        return (
            ("recurrent", self.recurrent),
            ("conv", self.conv),
            ("latent", self.latent),
        )

    def is_empty(self) -> bool:
        return all(value is None for _name, value in self.tensor_fields())

    def batch_size(self) -> int:
        return max(
            (int(value.size(0)) for _name, value in self.tensor_fields() if value is not None),
            default=0,
        )

    def to_payload(self, *, prefix: str) -> dict[str, torch.Tensor]:
        return {
            f"{prefix}_{name}": value.detach().clone()
            for name, value in self.tensor_fields()
            if value is not None
        }

    @classmethod
    def from_field_map(
        cls,
        field_map: Mapping[str, object] | None,
    ) -> LayerCacheSnapshot | None:
        if field_map is None:
            return None
        values: dict[str, torch.Tensor | None] = {}
        for name in ("recurrent", "conv", "latent"):
            value = field_map.get(name)
            if value is not None and not torch.is_tensor(value):
                raise TypeError(f"layer cache field {name!r} must be a tensor or None")
            values[name] = _clone(value)
        snapshot = cls(**values)
        return None if snapshot.is_empty() else snapshot


@dataclass(frozen=True)
class RuntimeCacheSnapshot:
    layers: tuple[LayerCacheSnapshot | None, ...] = ()

    def clone(self) -> RuntimeCacheSnapshot:
        return RuntimeCacheSnapshot(
            layers=tuple(None if item is None else item.clone() for item in self.layers)
        )

    def is_empty(self) -> bool:
        return all(item is None or item.is_empty() for item in self.layers)

    def batch_size(self) -> int:
        return max((item.batch_size() for item in self.layers if item is not None), default=0)

    def named_snapshots(self) -> tuple[tuple[str, LayerCacheSnapshot], ...]:
        return tuple(
            (f"layer_{index}", item)
            for index, item in enumerate(self.layers)
            if item is not None and not item.is_empty()
        )

    def to_payload(self) -> dict[str, torch.Tensor]:
        payload: dict[str, torch.Tensor] = {}
        for prefix, snapshot in self.named_snapshots():
            payload.update(snapshot.to_payload(prefix=prefix))
        return payload

    @classmethod
    def from_payload(cls, payload: Mapping[str, object]) -> RuntimeCacheSnapshot:
        fields_by_layer: dict[int, dict[str, object]] = {}
        for raw_key, value in payload.items():
            key = str(raw_key)
            for suffix, field_name in _CACHE_SUFFIXES:
                if key.endswith(suffix):
                    prefix = key[: -len(suffix)]
                    if not prefix.startswith("layer_"):
                        break
                    index = int(prefix.removeprefix("layer_"))
                    fields_by_layer.setdefault(index, {})[field_name] = value
                    break
            else:
                raise ValueError(f"unrecognized runtime cache payload key: {key}")
        if not fields_by_layer:
            return cls()
        layers: list[LayerCacheSnapshot | None] = [None] * (max(fields_by_layer) + 1)
        for index, field_map in fields_by_layer.items():
            layers[index] = LayerCacheSnapshot.from_field_map(field_map)
        return cls(layers=tuple(layers))


@dataclass
class TransformerRuntimeState:
    lock: RLock = field(default_factory=RLock)
    loss_chunk_size: int = 0


__all__ = ["LayerCacheSnapshot", "RuntimeCacheSnapshot", "TransformerRuntimeState"]