File size: 5,001 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
# 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 dataclasses import dataclass
from typing import Protocol

import torch

from .model_state import RuntimeCacheSnapshot


@dataclass
class RuntimeCacheState:
    cache: RuntimeCacheSnapshot
    batch_size: int
    cache_pos: int


class RuntimeCacheDecodeModel(Protocol):
    def replay_with_cache(
        self,
        input_ids: torch.Tensor,
        *,
        start_pos: int = 0,
        return_all_logits: bool = True,
    ) -> tuple[torch.Tensor, torch.Tensor | None]: ...

    def cache_dump(
        self,
        device: str = "cpu",
        *,
        cache_pos: int | None = None,
        batch_size: int | None = None,
    ) -> RuntimeCacheSnapshot: ...

    def cache_load(self, cache_snapshot: RuntimeCacheSnapshot) -> None: ...

    def reset_runtime_cache(self) -> None: ...

def clone_runtime_cache_state(cache_state: RuntimeCacheState) -> RuntimeCacheState:
    return RuntimeCacheState(
        cache=cache_state.cache.clone(),
        batch_size=int(cache_state.batch_size),
        cache_pos=int(cache_state.cache_pos),
    )


def cache_state_from_runtime_cache(
    cache_input: RuntimeCacheState | None,
) -> RuntimeCacheState | None:
    if cache_input is None:
        return None
    if isinstance(cache_input, RuntimeCacheState):
        return clone_runtime_cache_state(cache_input)
    raise TypeError("cache must be a Sophia RuntimeCacheState returned by SophiaDecoder")


def prefill_runtime_cache(
    model: RuntimeCacheDecodeModel,
    input_ids: torch.Tensor,
    *,
    start_pos: int,
    logits_to_keep: int | None = None,
) -> torch.Tensor:
    if int(input_ids.size(1)) <= 0:
        raise ValueError("input_ids must contain at least one token")
    return_all_logits = int(logits_to_keep or 0) != 1
    logits, _ = model.replay_with_cache(
        input_ids,
        start_pos=int(start_pos),
        return_all_logits=bool(return_all_logits),
    )
    if not bool(return_all_logits):
        logits = logits.unsqueeze(1)
    if logits_to_keep is not None and int(logits_to_keep) > 0:
        return logits[:, -int(logits_to_keep) :, :]
    return logits


def forward_cached_decode(
    *,
    model: RuntimeCacheDecodeModel,
    input_ids: torch.Tensor,
    cache_state: RuntimeCacheState | None,
    start_pos: int | None,
    logits_to_keep: int | None,
) -> tuple[torch.Tensor, RuntimeCacheState]:
    if logits_to_keep is not None and int(logits_to_keep) < 0:
        raise ValueError(f"logits_to_keep must be >= 0 when provided, got {logits_to_keep}")
    cache_start = 0 if start_pos is None else int(start_pos)
    if cache_state is None:
        if cache_start != 0:
            raise ValueError("cache_state is required when start_pos > 0 for cached decode")
        model.reset_runtime_cache()
        logits = prefill_runtime_cache(
            model,
            input_ids,
            start_pos=cache_start,
            logits_to_keep=logits_to_keep,
        )
        next_cache_pos = int(cache_start) + int(input_ids.size(1))
        return logits, RuntimeCacheState(
            cache=model.cache_dump(
                device="cpu",
                cache_pos=int(next_cache_pos),
                batch_size=int(input_ids.size(0)),
            ),
            cache_pos=int(next_cache_pos),
            batch_size=int(input_ids.size(0)),
        )

    cache_state = clone_runtime_cache_state(cache_state)
    cache_pos = int(cache_state.cache_pos)
    if cache_pos < 0:
        raise ValueError(f"cache cache_pos must be >= 0, got {cache_pos}")
    if start_pos is not None and int(start_pos) != cache_pos:
        raise ValueError(
            f"start_pos ({start_pos}) must match cache cache_pos ({cache_pos})"
        )
    batch_size = int(cache_state.batch_size or int(input_ids.size(0)))
    if batch_size <= 0:
        raise ValueError(f"cache batch_size must be > 0, got {batch_size}")
    if batch_size != int(input_ids.size(0)):
        raise ValueError(
            "cache batch_size does not match input_ids: "
            f"{batch_size} != {int(input_ids.size(0))}"
        )
    model.cache_load(cache_state.cache)
    logits = prefill_runtime_cache(
        model,
        input_ids,
        start_pos=cache_pos,
        logits_to_keep=logits_to_keep,
    )
    next_cache_pos = int(cache_pos) + int(input_ids.size(1))
    return logits, RuntimeCacheState(
        cache=model.cache_dump(
            device="cpu",
            cache_pos=int(next_cache_pos),
            batch_size=int(input_ids.size(0)),
        ),
        cache_pos=int(next_cache_pos),
        batch_size=int(input_ids.size(0)),
    )


__all__ = [
    "RuntimeCacheDecodeModel",
    "RuntimeCacheState",
    "cache_state_from_runtime_cache",
    "clone_runtime_cache_state",
    "forward_cached_decode",
    "prefill_runtime_cache",
]