@modular-prompt/driver 0.16.0 → 0.17.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +93 -10
- package/dist/cache-controller.d.ts +4 -0
- package/dist/cache-controller.d.ts.map +1 -1
- package/dist/driver-registry/config-based-factory.d.ts.map +1 -1
- package/dist/driver-registry/config-based-factory.js +6 -0
- package/dist/driver-registry/config-based-factory.js.map +1 -1
- package/dist/driver-registry/factory-helper.d.ts.map +1 -1
- package/dist/driver-registry/factory-helper.js +9 -2
- package/dist/driver-registry/factory-helper.js.map +1 -1
- package/dist/driver-registry/index.d.ts +1 -1
- package/dist/driver-registry/index.d.ts.map +1 -1
- package/dist/driver-registry/types.d.ts +15 -1
- package/dist/driver-registry/types.d.ts.map +1 -1
- package/dist/formatter/converter.d.ts.map +1 -1
- package/dist/formatter/converter.js +31 -2
- package/dist/formatter/converter.js.map +1 -1
- package/dist/index.d.ts +5 -3
- package/dist/index.d.ts.map +1 -1
- package/dist/index.js +5 -3
- package/dist/index.js.map +1 -1
- package/dist/local-inference/adapters.d.ts +6 -0
- package/dist/local-inference/adapters.d.ts.map +1 -1
- package/dist/local-inference/driver.d.ts.map +1 -1
- package/dist/local-inference/driver.js +45 -24
- package/dist/local-inference/driver.js.map +1 -1
- package/dist/local-inference/process-client.d.ts +4 -2
- package/dist/local-inference/process-client.d.ts.map +1 -1
- package/dist/local-inference/process-client.js +24 -8
- package/dist/local-inference/process-client.js.map +1 -1
- package/dist/local-inference/process-communication.d.ts +9 -2
- package/dist/local-inference/process-communication.d.ts.map +1 -1
- package/dist/local-inference/process-communication.js +37 -5
- package/dist/local-inference/process-communication.js.map +1 -1
- package/dist/local-inference/protocol.d.ts +4 -0
- package/dist/local-inference/protocol.d.ts.map +1 -1
- package/dist/local-inference/request-queue.d.ts +1 -1
- package/dist/local-inference/request-queue.d.ts.map +1 -1
- package/dist/local-inference/request-queue.js +26 -7
- package/dist/local-inference/request-queue.js.map +1 -1
- package/dist/local-inference/stream-utils.d.ts +6 -0
- package/dist/local-inference/stream-utils.d.ts.map +1 -1
- package/dist/local-inference/stream-utils.js.map +1 -1
- package/dist/mlx-ml/mlx-cache-controller.d.ts +9 -0
- package/dist/mlx-ml/mlx-cache-controller.d.ts.map +1 -1
- package/dist/mlx-ml/mlx-cache-controller.js +158 -33
- package/dist/mlx-ml/mlx-cache-controller.js.map +1 -1
- package/dist/mlx-ml/mlx-cache-support.d.ts +2 -2
- package/dist/mlx-ml/mlx-cache-support.d.ts.map +1 -1
- package/dist/mlx-ml/mlx-cache-support.js +8 -3
- package/dist/mlx-ml/mlx-cache-support.js.map +1 -1
- package/dist/mlx-ml/mlx-driver.d.ts +0 -1
- package/dist/mlx-ml/mlx-driver.d.ts.map +1 -1
- package/dist/mlx-ml/mlx-driver.js +1 -8
- package/dist/mlx-ml/mlx-driver.js.map +1 -1
- package/dist/mlx-ml/process/index.d.ts +1 -1
- package/dist/mlx-ml/process/index.d.ts.map +1 -1
- package/dist/mlx-ml/process/index.js +2 -2
- package/dist/mlx-ml/process/index.js.map +1 -1
- package/dist/models-config/index.d.ts +1 -1
- package/dist/models-config/index.d.ts.map +1 -1
- package/dist/models-config/index.js +1 -1
- package/dist/models-config/index.js.map +1 -1
- package/dist/models-config/resolve.d.ts +9 -1
- package/dist/models-config/resolve.d.ts.map +1 -1
- package/dist/models-config/resolve.js +94 -2
- package/dist/models-config/resolve.js.map +1 -1
- package/dist/models-config/types.d.ts +3 -1
- package/dist/models-config/types.d.ts.map +1 -1
- package/dist/pytorch/process/index.d.ts +4 -2
- package/dist/pytorch/process/index.d.ts.map +1 -1
- package/dist/pytorch/process/index.js +24 -7
- package/dist/pytorch/process/index.js.map +1 -1
- package/dist/pytorch/pytorch-cache-controller.d.ts +84 -0
- package/dist/pytorch/pytorch-cache-controller.d.ts.map +1 -0
- package/dist/pytorch/pytorch-cache-controller.js +742 -0
- package/dist/pytorch/pytorch-cache-controller.js.map +1 -0
- package/dist/pytorch/pytorch-cache-support.d.ts +23 -0
- package/dist/pytorch/pytorch-cache-support.d.ts.map +1 -0
- package/dist/pytorch/pytorch-cache-support.js +47 -0
- package/dist/pytorch/pytorch-cache-support.js.map +1 -0
- package/dist/pytorch/pytorch-driver.d.ts +8 -1
- package/dist/pytorch/pytorch-driver.d.ts.map +1 -1
- package/dist/pytorch/pytorch-driver.js +40 -0
- package/dist/pytorch/pytorch-driver.js.map +1 -1
- package/dist/runtime/check.d.ts.map +1 -1
- package/dist/runtime/check.js +8 -6
- package/dist/runtime/check.js.map +1 -1
- package/dist/runtime/index.d.ts +2 -2
- package/dist/runtime/index.d.ts.map +1 -1
- package/dist/runtime/index.js +2 -2
- package/dist/runtime/index.js.map +1 -1
- package/dist/runtime/manifest-core.d.mts +1 -0
- package/dist/runtime/manifest-core.mjs +1 -0
- package/dist/runtime/manifest-core.mjs.map +1 -1
- package/dist/runtime/manifest.d.ts +2 -0
- package/dist/runtime/manifest.d.ts.map +1 -1
- package/dist/runtime/manifest.js.map +1 -1
- package/dist/runtime/paths-core.d.mts +15 -1
- package/dist/runtime/paths-core.d.mts.map +1 -1
- package/dist/runtime/paths-core.mjs +50 -5
- package/dist/runtime/paths-core.mjs.map +1 -1
- package/dist/runtime/paths.d.ts +2 -2
- package/dist/runtime/paths.d.ts.map +1 -1
- package/dist/runtime/paths.js +2 -2
- package/dist/runtime/paths.js.map +1 -1
- package/dist/runtime/pytorch-template-core.d.mts +11 -0
- package/dist/runtime/pytorch-template-core.d.mts.map +1 -0
- package/dist/runtime/pytorch-template-core.mjs +54 -0
- package/dist/runtime/pytorch-template-core.mjs.map +1 -0
- package/dist/runtime/setup-commands-core.d.mts +3 -0
- package/dist/runtime/setup-commands-core.d.mts.map +1 -1
- package/dist/runtime/setup-commands-core.mjs +4 -0
- package/dist/runtime/setup-commands-core.mjs.map +1 -1
- package/dist/runtime/setup-commands.d.ts +1 -1
- package/dist/runtime/setup-commands.d.ts.map +1 -1
- package/dist/runtime/setup-commands.js +1 -1
- package/dist/runtime/setup-commands.js.map +1 -1
- package/docs/DRIVER_API.md +455 -0
- package/docs/LOCAL_MODEL_SETUP.md +765 -0
- package/docs/mlx-api-selection.md +301 -0
- package/package.json +9 -5
- package/scripts/runtime-cli.bin.test.ts +142 -0
- package/scripts/runtime-cli.js +305 -35
- package/scripts/runtime-cli.test.ts +163 -0
- package/src/mlx-ml/python/__main__.py +1 -1
- package/src/mlx-ml/python/backends/base.py +88 -18
- package/src/mlx-ml/python/backends/mlx_lm.py +28 -3
- package/src/mlx-ml/python/backends/mlx_vlm.py +679 -2
- package/src/mlx-ml/python/handlers/cache.py +4 -0
- package/src/mlx-ml/python/handlers/generate.py +33 -10
- package/src/mlx-ml/python/handlers/tokenize.py +1 -4
- package/src/mlx-ml/python/pyproject.toml +1 -1
- package/src/mlx-ml/python/server.py +2 -0
- package/src/mlx-ml/python/uv.lock +8 -8
- package/src/pytorch/templates/cpu-minimal/backends/base.py +139 -0
- package/src/pytorch/templates/cpu-minimal/backends/transformers_lm.py +1167 -0
- package/src/pytorch/{python → templates/cpu-minimal}/handlers/__init__.py +1 -0
- package/src/pytorch/templates/cpu-minimal/handlers/cache.py +88 -0
- package/src/pytorch/templates/cpu-minimal/handlers/generate.py +157 -0
- package/src/pytorch/{python → templates/cpu-minimal}/pyproject.toml +2 -2
- package/src/pytorch/{python → templates/cpu-minimal}/server.py +20 -2
- package/src/pytorch/templates/cpu-minimal/tests/test_cache_handler.py +284 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_capabilities.py +14 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_server.py +141 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_transformers_errors.py +89 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_transformers_lm_cache.py +554 -0
- package/src/pytorch/templates/cpu-minimal/utils/__init__.py +0 -0
- package/src/pytorch/{python → templates/cpu-minimal}/utils/token_utils.py +2 -2
- package/src/pytorch/templates/cpu-minimal/utils/transformers_errors.py +54 -0
- package/src/pytorch/{python → templates/cpu-minimal}/uv.lock +149 -109
- package/src/pytorch/templates/cuda/__main__.py +19 -0
- package/src/pytorch/templates/cuda/backends/__init__.py +3 -0
- package/src/pytorch/{python → templates/cuda}/backends/base.py +54 -6
- package/src/pytorch/templates/cuda/backends/transformers_lm.py +379 -0
- package/src/pytorch/templates/cuda/handlers/__init__.py +7 -0
- package/src/pytorch/templates/cuda/handlers/cache.py +93 -0
- package/src/pytorch/templates/cuda/handlers/cancel.py +53 -0
- package/src/pytorch/templates/cuda/handlers/capabilities.py +6 -0
- package/src/pytorch/templates/cuda/handlers/completion.py +15 -0
- package/src/pytorch/templates/cuda/handlers/format_test.py +70 -0
- package/src/pytorch/templates/cuda/handlers/generate.py +152 -0
- package/src/pytorch/templates/cuda/handlers/render.py +40 -0
- package/src/pytorch/templates/cuda/handlers/tokenize.py +63 -0
- package/src/pytorch/templates/cuda/pyproject.toml +37 -0
- package/src/pytorch/templates/cuda/server.py +158 -0
- package/src/pytorch/templates/cuda/tests/__init__.py +0 -0
- package/src/pytorch/templates/cuda/tests/test_cache_handler.py +207 -0
- package/src/pytorch/templates/cuda/tests/test_capabilities.py +14 -0
- package/src/pytorch/templates/cuda/tests/test_server.py +145 -0
- package/src/pytorch/templates/cuda/tests/test_transformers_errors.py +89 -0
- package/src/pytorch/templates/cuda/tests/test_transformers_lm_cache.py +288 -0
- package/src/pytorch/templates/cuda/utils/__init__.py +0 -0
- package/src/pytorch/templates/cuda/utils/chat_template_constraints.py +164 -0
- package/src/pytorch/templates/cuda/utils/prompt_builder.py +54 -0
- package/src/pytorch/templates/cuda/utils/template_render.py +80 -0
- package/src/pytorch/templates/cuda/utils/token_utils.py +376 -0
- package/src/pytorch/templates/cuda/utils/transformers_errors.py +54 -0
- package/src/pytorch/templates/cuda/uv.lock +734 -0
- package/src/pytorch/python/backends/transformers_lm.py +0 -127
- package/src/pytorch/python/handlers/generate.py +0 -68
- /package/src/pytorch/{python → templates/cpu-minimal}/__main__.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/backends/__init__.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/cancel.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/capabilities.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/completion.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/format_test.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/render.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/tokenize.py +0 -0
- /package/src/pytorch/{python/utils → templates/cpu-minimal/tests}/__init__.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/utils/chat_template_constraints.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/utils/prompt_builder.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/utils/template_render.py +0 -0
|
@@ -0,0 +1,1167 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import copy
|
|
4
|
+
import inspect
|
|
5
|
+
import json
|
|
6
|
+
import os
|
|
7
|
+
import sys
|
|
8
|
+
from dataclasses import dataclass
|
|
9
|
+
from tempfile import mkstemp
|
|
10
|
+
from threading import Thread
|
|
11
|
+
from typing import Any, Iterator
|
|
12
|
+
|
|
13
|
+
import torch
|
|
14
|
+
import transformers
|
|
15
|
+
from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer
|
|
16
|
+
|
|
17
|
+
from backends.base import ModelBackend
|
|
18
|
+
from utils.token_utils import is_eod_token
|
|
19
|
+
from utils.transformers_errors import (
|
|
20
|
+
extract_unsupported_model_type,
|
|
21
|
+
unsupported_model_type_error,
|
|
22
|
+
)
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class _TokenCountingTextIteratorStreamer(TextIteratorStreamer):
|
|
26
|
+
"""TextIteratorStreamer that counts generated token IDs, not text chunks."""
|
|
27
|
+
|
|
28
|
+
def __init__(self, *args, **kwargs):
|
|
29
|
+
super().__init__(*args, **kwargs)
|
|
30
|
+
self.generated_token_count = 0
|
|
31
|
+
|
|
32
|
+
def put(self, value: torch.Tensor) -> None:
|
|
33
|
+
is_prompt = self.skip_prompt and self.next_tokens_are_prompt
|
|
34
|
+
if not is_prompt:
|
|
35
|
+
token_count = value[0].numel() if value.ndim > 1 else value.numel()
|
|
36
|
+
self.generated_token_count += int(token_count)
|
|
37
|
+
super().put(value)
|
|
38
|
+
|
|
39
|
+
def on_finalized_text(self, text: str, stream_end: bool = False) -> None:
|
|
40
|
+
"""Queue text together with the token count at its emission point."""
|
|
41
|
+
self.text_queue.put((text, self.generated_token_count), timeout=self.timeout)
|
|
42
|
+
if stream_end:
|
|
43
|
+
self.text_queue.put(self.stop_signal, timeout=self.timeout)
|
|
44
|
+
|
|
45
|
+
def __next__(self) -> tuple[str, int]:
|
|
46
|
+
value = self.text_queue.get(timeout=self.timeout)
|
|
47
|
+
if value == self.stop_signal:
|
|
48
|
+
raise StopIteration()
|
|
49
|
+
return value
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
@dataclass
|
|
53
|
+
class StreamChunk:
|
|
54
|
+
text: str
|
|
55
|
+
prompt_tokens: int | None = None
|
|
56
|
+
generation_tokens: int | None = None
|
|
57
|
+
finish_reason: str | None = None
|
|
58
|
+
cache_read_tokens: int | None = None
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
class TransformersLmBackend(ModelBackend):
|
|
62
|
+
"""Transformers causal LM backend (text-only, CPU-first)."""
|
|
63
|
+
|
|
64
|
+
CACHE_LAYOUT = "pytorch_kv_v1"
|
|
65
|
+
CACHE_META_SUFFIX = ".meta.json"
|
|
66
|
+
|
|
67
|
+
def __init__(self, device: str | None = None) -> None:
|
|
68
|
+
self.model: Any | None = None
|
|
69
|
+
self.tokenizer: Any | None = None
|
|
70
|
+
self._device_name = device or os.environ.get("PYTORCH_DEVICE", "cpu")
|
|
71
|
+
self._device = torch.device(self._device_name)
|
|
72
|
+
if self._device.type == "cuda" and not torch.cuda.is_available():
|
|
73
|
+
raise RuntimeError(
|
|
74
|
+
"CUDA device requested, but CUDA is not available in this PyTorch runtime. "
|
|
75
|
+
"Install a CUDA-enabled torch wheel and verify the NVIDIA driver."
|
|
76
|
+
)
|
|
77
|
+
self._model_id: str | None = None
|
|
78
|
+
self._model_dtype: str | None = None
|
|
79
|
+
self._caches: dict[str, Any] = {}
|
|
80
|
+
self._cache_token_counts: dict[str, int] = {}
|
|
81
|
+
self._cache_token_ids: dict[str, tuple[int, ...]] = {}
|
|
82
|
+
self._cache_write_token_counts: dict[str, int] = {}
|
|
83
|
+
self._cache_meta: dict[str, dict[str, Any]] = {}
|
|
84
|
+
|
|
85
|
+
def load(self, model_name: str) -> None:
|
|
86
|
+
self._caches.clear()
|
|
87
|
+
self._cache_token_counts.clear()
|
|
88
|
+
self._cache_token_ids.clear()
|
|
89
|
+
self._cache_write_token_counts.clear()
|
|
90
|
+
self._cache_meta.clear()
|
|
91
|
+
self._model_id = None
|
|
92
|
+
self._model_dtype = None
|
|
93
|
+
|
|
94
|
+
trust_remote_code = os.environ.get("PYTORCH_TRUST_REMOTE_CODE", "").lower() in (
|
|
95
|
+
"1",
|
|
96
|
+
"true",
|
|
97
|
+
"yes",
|
|
98
|
+
)
|
|
99
|
+
try:
|
|
100
|
+
self.tokenizer = AutoTokenizer.from_pretrained(
|
|
101
|
+
model_name,
|
|
102
|
+
trust_remote_code=trust_remote_code,
|
|
103
|
+
)
|
|
104
|
+
if self.tokenizer.pad_token is None and self.tokenizer.eos_token is not None:
|
|
105
|
+
self.tokenizer.pad_token = self.tokenizer.eos_token
|
|
106
|
+
|
|
107
|
+
dtype = torch.float32 if self._device.type == "cpu" else torch.float16
|
|
108
|
+
self.model = AutoModelForCausalLM.from_pretrained(
|
|
109
|
+
model_name,
|
|
110
|
+
trust_remote_code=trust_remote_code,
|
|
111
|
+
dtype=dtype,
|
|
112
|
+
)
|
|
113
|
+
except (KeyError, ValueError) as error:
|
|
114
|
+
model_type = extract_unsupported_model_type(error)
|
|
115
|
+
if model_type is None:
|
|
116
|
+
raise
|
|
117
|
+
raise unsupported_model_type_error(
|
|
118
|
+
model_name,
|
|
119
|
+
model_type,
|
|
120
|
+
getattr(transformers, "__version__", "unknown"),
|
|
121
|
+
) from error
|
|
122
|
+
|
|
123
|
+
self.model.to(self._device)
|
|
124
|
+
self.model.eval()
|
|
125
|
+
self._model_id = model_name
|
|
126
|
+
self._model_dtype = self._dtype_name(dtype)
|
|
127
|
+
|
|
128
|
+
def get_tokenizer(self) -> Any:
|
|
129
|
+
return self.tokenizer
|
|
130
|
+
|
|
131
|
+
def tokenize_prompt(
|
|
132
|
+
self,
|
|
133
|
+
prompt: str,
|
|
134
|
+
images: list | None = None,
|
|
135
|
+
max_image_size: int = 768,
|
|
136
|
+
) -> list[int]:
|
|
137
|
+
if images:
|
|
138
|
+
raise ValueError("TransformersLmBackend does not support vision input")
|
|
139
|
+
if self.tokenizer is None:
|
|
140
|
+
raise RuntimeError("Model is not loaded")
|
|
141
|
+
|
|
142
|
+
bos_token = getattr(self.tokenizer, "bos_token", None)
|
|
143
|
+
add_special = bos_token is None or not prompt.startswith(bos_token or "")
|
|
144
|
+
token_ids = self.tokenizer.encode(prompt, add_special_tokens=add_special)
|
|
145
|
+
if hasattr(token_ids, "flatten"):
|
|
146
|
+
token_ids = token_ids.flatten().tolist()
|
|
147
|
+
return [int(token_id) for token_id in token_ids]
|
|
148
|
+
|
|
149
|
+
@staticmethod
|
|
150
|
+
def _clone_cache(prompt_cache: Any) -> Any:
|
|
151
|
+
"""Clone model-owned cache state before generation mutates it."""
|
|
152
|
+
if isinstance(prompt_cache, torch.Tensor):
|
|
153
|
+
return prompt_cache.clone()
|
|
154
|
+
if isinstance(prompt_cache, tuple):
|
|
155
|
+
return tuple(TransformersLmBackend._clone_cache(item) for item in prompt_cache)
|
|
156
|
+
if isinstance(prompt_cache, list):
|
|
157
|
+
return [TransformersLmBackend._clone_cache(item) for item in prompt_cache]
|
|
158
|
+
if isinstance(prompt_cache, dict):
|
|
159
|
+
return {
|
|
160
|
+
key: TransformersLmBackend._clone_cache(value)
|
|
161
|
+
for key, value in prompt_cache.items()
|
|
162
|
+
}
|
|
163
|
+
if hasattr(prompt_cache, "layers"):
|
|
164
|
+
try:
|
|
165
|
+
cloned_cache = copy.copy(prompt_cache)
|
|
166
|
+
cloned_cache.layers = [
|
|
167
|
+
TransformersLmBackend._clone_cache(layer)
|
|
168
|
+
for layer in prompt_cache.layers
|
|
169
|
+
]
|
|
170
|
+
return cloned_cache
|
|
171
|
+
except Exception:
|
|
172
|
+
pass
|
|
173
|
+
if (
|
|
174
|
+
hasattr(prompt_cache, "keys")
|
|
175
|
+
and hasattr(prompt_cache, "values")
|
|
176
|
+
) or any(
|
|
177
|
+
hasattr(prompt_cache, attribute)
|
|
178
|
+
for attribute in ("conv_states", "recurrent_states")
|
|
179
|
+
):
|
|
180
|
+
try:
|
|
181
|
+
cloned_layer = copy.copy(prompt_cache)
|
|
182
|
+
for attribute in (
|
|
183
|
+
"keys",
|
|
184
|
+
"values",
|
|
185
|
+
"indexer_keys",
|
|
186
|
+
"indexer_cumulative_length",
|
|
187
|
+
"cumulative_length",
|
|
188
|
+
"cumulative_length_int",
|
|
189
|
+
"conv_states",
|
|
190
|
+
"recurrent_states",
|
|
191
|
+
):
|
|
192
|
+
if hasattr(prompt_cache, attribute):
|
|
193
|
+
setattr(
|
|
194
|
+
cloned_layer,
|
|
195
|
+
attribute,
|
|
196
|
+
TransformersLmBackend._clone_cache(
|
|
197
|
+
getattr(prompt_cache, attribute)
|
|
198
|
+
),
|
|
199
|
+
)
|
|
200
|
+
return cloned_layer
|
|
201
|
+
except Exception:
|
|
202
|
+
pass
|
|
203
|
+
if hasattr(prompt_cache, "get_seq_length"):
|
|
204
|
+
try:
|
|
205
|
+
return copy.deepcopy(prompt_cache)
|
|
206
|
+
except Exception:
|
|
207
|
+
# Some third-party Cache implementations cannot be deep-copied.
|
|
208
|
+
# Keep the request usable; those implementations must tolerate
|
|
209
|
+
# in-place generation updates.
|
|
210
|
+
return prompt_cache
|
|
211
|
+
return prompt_cache
|
|
212
|
+
|
|
213
|
+
@staticmethod
|
|
214
|
+
def _dtype_name(dtype: Any) -> str:
|
|
215
|
+
"""Return a stable, human-readable dtype name for cache metadata."""
|
|
216
|
+
value = str(dtype)
|
|
217
|
+
return value.removeprefix("torch.")
|
|
218
|
+
|
|
219
|
+
def _current_model_id(self) -> str:
|
|
220
|
+
if self._model_id:
|
|
221
|
+
return self._model_id
|
|
222
|
+
|
|
223
|
+
for candidate in (
|
|
224
|
+
getattr(self.model, "name_or_path", None),
|
|
225
|
+
getattr(getattr(self.model, "config", None), "_name_or_path", None),
|
|
226
|
+
getattr(getattr(self.model, "config", None), "name_or_path", None),
|
|
227
|
+
):
|
|
228
|
+
if candidate:
|
|
229
|
+
return str(candidate)
|
|
230
|
+
return "unknown"
|
|
231
|
+
|
|
232
|
+
def _current_dtype(self) -> str:
|
|
233
|
+
if self._model_dtype:
|
|
234
|
+
return self._model_dtype
|
|
235
|
+
|
|
236
|
+
if self.model is not None:
|
|
237
|
+
try:
|
|
238
|
+
parameter = next(self.model.parameters())
|
|
239
|
+
return self._dtype_name(parameter.dtype)
|
|
240
|
+
except (AttributeError, RuntimeError, StopIteration, TypeError):
|
|
241
|
+
pass
|
|
242
|
+
|
|
243
|
+
return "float32" if self._device.type == "cpu" else "float16"
|
|
244
|
+
|
|
245
|
+
def _supports_model_kwarg(self, name: str) -> bool:
|
|
246
|
+
"""Check whether Transformers exposes a kwarg on the model forward."""
|
|
247
|
+
if self.model is None:
|
|
248
|
+
return False
|
|
249
|
+
try:
|
|
250
|
+
parameters = inspect.signature(self.model.forward).parameters
|
|
251
|
+
except (AttributeError, TypeError, ValueError):
|
|
252
|
+
# Test doubles and custom remote-code models may not expose a
|
|
253
|
+
# useful signature. Preserve the existing kwargs in that case.
|
|
254
|
+
return True
|
|
255
|
+
return name in parameters
|
|
256
|
+
|
|
257
|
+
def _cache_meta_for(
|
|
258
|
+
self,
|
|
259
|
+
token_count: int,
|
|
260
|
+
prefix_offsets: list[int] | None = None,
|
|
261
|
+
prefix_hashes: list[str] | None = None,
|
|
262
|
+
) -> dict[str, Any]:
|
|
263
|
+
if (prefix_offsets is None) != (prefix_hashes is None):
|
|
264
|
+
raise ValueError("prefix_offsets and prefix_hashes must be provided together")
|
|
265
|
+
if prefix_offsets is not None and len(prefix_offsets) != len(prefix_hashes or []):
|
|
266
|
+
raise ValueError("prefix_offsets and prefix_hashes must have the same length")
|
|
267
|
+
|
|
268
|
+
return {
|
|
269
|
+
"layout": self.CACHE_LAYOUT,
|
|
270
|
+
"token_count": int(token_count),
|
|
271
|
+
"prefix_offsets": list(prefix_offsets or []),
|
|
272
|
+
"prefix_hashes": list(prefix_hashes or []),
|
|
273
|
+
"model_id": self._current_model_id(),
|
|
274
|
+
"dtype": self._current_dtype(),
|
|
275
|
+
"device": str(self._device),
|
|
276
|
+
}
|
|
277
|
+
|
|
278
|
+
@classmethod
|
|
279
|
+
def _is_memory_cache_path(cls, cache_path: str) -> bool:
|
|
280
|
+
return cache_path.startswith("memory://")
|
|
281
|
+
|
|
282
|
+
@classmethod
|
|
283
|
+
def _meta_path(cls, cache_path: str) -> str:
|
|
284
|
+
return cache_path + cls.CACHE_META_SUFFIX
|
|
285
|
+
|
|
286
|
+
@staticmethod
|
|
287
|
+
def _atomic_torch_save(payload: dict[str, Any], cache_path: str) -> None:
|
|
288
|
+
directory = os.path.dirname(os.path.abspath(cache_path))
|
|
289
|
+
os.makedirs(directory, exist_ok=True)
|
|
290
|
+
fd, temporary_path = mkstemp(
|
|
291
|
+
prefix=f".{os.path.basename(cache_path)}.",
|
|
292
|
+
suffix=".tmp",
|
|
293
|
+
dir=directory,
|
|
294
|
+
)
|
|
295
|
+
try:
|
|
296
|
+
with os.fdopen(fd, "wb") as output:
|
|
297
|
+
torch.save(payload, output)
|
|
298
|
+
os.replace(temporary_path, cache_path)
|
|
299
|
+
except Exception:
|
|
300
|
+
try:
|
|
301
|
+
os.unlink(temporary_path)
|
|
302
|
+
except FileNotFoundError:
|
|
303
|
+
pass
|
|
304
|
+
raise
|
|
305
|
+
|
|
306
|
+
@staticmethod
|
|
307
|
+
def _atomic_json_save(meta: dict[str, Any], meta_path: str) -> None:
|
|
308
|
+
directory = os.path.dirname(os.path.abspath(meta_path))
|
|
309
|
+
os.makedirs(directory, exist_ok=True)
|
|
310
|
+
fd, temporary_path = mkstemp(
|
|
311
|
+
prefix=f".{os.path.basename(meta_path)}.",
|
|
312
|
+
suffix=".tmp",
|
|
313
|
+
dir=directory,
|
|
314
|
+
)
|
|
315
|
+
try:
|
|
316
|
+
with os.fdopen(fd, "w", encoding="utf-8") as output:
|
|
317
|
+
json.dump(meta, output)
|
|
318
|
+
os.replace(temporary_path, meta_path)
|
|
319
|
+
except Exception:
|
|
320
|
+
try:
|
|
321
|
+
os.unlink(temporary_path)
|
|
322
|
+
except FileNotFoundError:
|
|
323
|
+
pass
|
|
324
|
+
raise
|
|
325
|
+
|
|
326
|
+
@staticmethod
|
|
327
|
+
def _serialize_cache_value(value: Any) -> Any:
|
|
328
|
+
if isinstance(value, torch.Tensor):
|
|
329
|
+
return {"kind": "tensor", "value": value.detach().cpu()}
|
|
330
|
+
if isinstance(value, torch.dtype):
|
|
331
|
+
return {
|
|
332
|
+
"kind": "dtype",
|
|
333
|
+
"value": str(value).removeprefix("torch."),
|
|
334
|
+
}
|
|
335
|
+
if isinstance(value, torch.device):
|
|
336
|
+
return {"kind": "device", "value": str(value)}
|
|
337
|
+
if value is None or isinstance(value, (bool, int, float, str)):
|
|
338
|
+
return {"kind": "value", "value": value}
|
|
339
|
+
if isinstance(value, dict):
|
|
340
|
+
return {
|
|
341
|
+
"kind": "mapping",
|
|
342
|
+
"items": [
|
|
343
|
+
[
|
|
344
|
+
TransformersLmBackend._serialize_cache_value(key),
|
|
345
|
+
TransformersLmBackend._serialize_cache_value(item),
|
|
346
|
+
]
|
|
347
|
+
for key, item in value.items()
|
|
348
|
+
],
|
|
349
|
+
}
|
|
350
|
+
if isinstance(value, tuple):
|
|
351
|
+
return {
|
|
352
|
+
"kind": "sequence",
|
|
353
|
+
"sequence_type": "tuple",
|
|
354
|
+
"items": [TransformersLmBackend._serialize_cache_value(item) for item in value],
|
|
355
|
+
}
|
|
356
|
+
if isinstance(value, list):
|
|
357
|
+
return {
|
|
358
|
+
"kind": "sequence",
|
|
359
|
+
"sequence_type": "list",
|
|
360
|
+
"items": [TransformersLmBackend._serialize_cache_value(item) for item in value],
|
|
361
|
+
}
|
|
362
|
+
raise TypeError(f"Unsupported PyTorch cache value: {type(value).__name__}")
|
|
363
|
+
|
|
364
|
+
@classmethod
|
|
365
|
+
def _serialize_cache_layer(cls, layer: Any) -> dict[str, Any]:
|
|
366
|
+
keys = getattr(layer, "keys", None)
|
|
367
|
+
values = getattr(layer, "values", None)
|
|
368
|
+
attributes = getattr(layer, "__dict__", None)
|
|
369
|
+
if not isinstance(attributes, dict):
|
|
370
|
+
attributes = {}
|
|
371
|
+
serialized = {
|
|
372
|
+
"kind": "layer",
|
|
373
|
+
"class_name": type(layer).__name__,
|
|
374
|
+
"attributes": {
|
|
375
|
+
name: cls._serialize_cache_value(value)
|
|
376
|
+
for name, value in attributes.items()
|
|
377
|
+
},
|
|
378
|
+
"token_count": cls._cache_offset_from_value(layer),
|
|
379
|
+
}
|
|
380
|
+
if not attributes and not (
|
|
381
|
+
isinstance(keys, torch.Tensor) and isinstance(values, torch.Tensor)
|
|
382
|
+
):
|
|
383
|
+
raise TypeError(f"Unsupported PyTorch cache layer: {type(layer).__name__}")
|
|
384
|
+
return serialized
|
|
385
|
+
|
|
386
|
+
@classmethod
|
|
387
|
+
def _serialize_cache(cls, prompt_cache: Any) -> Any:
|
|
388
|
+
if isinstance(prompt_cache, torch.Tensor):
|
|
389
|
+
return cls._serialize_cache_value(prompt_cache)
|
|
390
|
+
if isinstance(prompt_cache, (tuple, list)):
|
|
391
|
+
return cls._serialize_cache_value(prompt_cache)
|
|
392
|
+
if hasattr(prompt_cache, "layers"):
|
|
393
|
+
return {
|
|
394
|
+
"kind": "cache",
|
|
395
|
+
"cache_type": type(prompt_cache).__name__,
|
|
396
|
+
"layers": [
|
|
397
|
+
cls._serialize_cache_layer(layer)
|
|
398
|
+
for layer in prompt_cache.layers
|
|
399
|
+
],
|
|
400
|
+
}
|
|
401
|
+
if hasattr(prompt_cache, "keys") and hasattr(prompt_cache, "values"):
|
|
402
|
+
return cls._serialize_cache_layer(prompt_cache)
|
|
403
|
+
raise TypeError(f"Unsupported PyTorch cache: {type(prompt_cache).__name__}")
|
|
404
|
+
|
|
405
|
+
@classmethod
|
|
406
|
+
def _deserialize_cache_value(cls, value: Any) -> Any:
|
|
407
|
+
if not isinstance(value, dict):
|
|
408
|
+
raise ValueError("Invalid PyTorch cache payload")
|
|
409
|
+
|
|
410
|
+
kind = value.get("kind")
|
|
411
|
+
if kind == "tensor":
|
|
412
|
+
tensor = value.get("value")
|
|
413
|
+
if not isinstance(tensor, torch.Tensor):
|
|
414
|
+
raise ValueError("Invalid tensor in PyTorch cache payload")
|
|
415
|
+
return tensor
|
|
416
|
+
if kind == "dtype":
|
|
417
|
+
try:
|
|
418
|
+
return getattr(torch, value["value"])
|
|
419
|
+
except (KeyError, AttributeError):
|
|
420
|
+
raise ValueError("Invalid dtype in PyTorch cache payload")
|
|
421
|
+
if kind == "device":
|
|
422
|
+
try:
|
|
423
|
+
return torch.device(value["value"])
|
|
424
|
+
except (KeyError, RuntimeError, TypeError):
|
|
425
|
+
raise ValueError("Invalid device in PyTorch cache payload")
|
|
426
|
+
if kind == "value":
|
|
427
|
+
return value.get("value")
|
|
428
|
+
if kind == "mapping":
|
|
429
|
+
items = value.get("items", [])
|
|
430
|
+
if not isinstance(items, list):
|
|
431
|
+
raise ValueError("Invalid mapping in PyTorch cache payload")
|
|
432
|
+
return {
|
|
433
|
+
cls._deserialize_cache_value(key): cls._deserialize_cache_value(item)
|
|
434
|
+
for key, item in items
|
|
435
|
+
}
|
|
436
|
+
if kind == "sequence":
|
|
437
|
+
items = [cls._deserialize_cache_value(item) for item in value.get("items", [])]
|
|
438
|
+
return tuple(items) if value.get("sequence_type") == "tuple" else items
|
|
439
|
+
if kind == "layer":
|
|
440
|
+
if "attributes" in value:
|
|
441
|
+
return cls._deserialize_cache_layer(value)
|
|
442
|
+
return (
|
|
443
|
+
cls._deserialize_cache_value(value["keys"]),
|
|
444
|
+
cls._deserialize_cache_value(value["values"]),
|
|
445
|
+
)
|
|
446
|
+
raise ValueError(f"Unknown PyTorch cache payload kind: {kind!r}")
|
|
447
|
+
|
|
448
|
+
@classmethod
|
|
449
|
+
def _deserialize_cache_layer(cls, value: Any) -> Any:
|
|
450
|
+
if not isinstance(value, dict) or value.get("kind") != "layer":
|
|
451
|
+
raise ValueError("Invalid PyTorch cache layer payload")
|
|
452
|
+
if "attributes" in value:
|
|
453
|
+
attributes = value["attributes"]
|
|
454
|
+
if not isinstance(attributes, dict):
|
|
455
|
+
raise ValueError("Invalid PyTorch cache layer attributes")
|
|
456
|
+
return {
|
|
457
|
+
"class_name": value.get("class_name"),
|
|
458
|
+
"token_count": value.get("token_count"),
|
|
459
|
+
"attributes": {
|
|
460
|
+
name: cls._deserialize_cache_value(item)
|
|
461
|
+
for name, item in attributes.items()
|
|
462
|
+
},
|
|
463
|
+
}
|
|
464
|
+
return (
|
|
465
|
+
cls._deserialize_cache_value(value["keys"]),
|
|
466
|
+
cls._deserialize_cache_value(value["values"]),
|
|
467
|
+
)
|
|
468
|
+
|
|
469
|
+
def _deserialize_cache(self, value: Any) -> Any:
|
|
470
|
+
if not isinstance(value, dict):
|
|
471
|
+
raise ValueError("Invalid PyTorch cache payload")
|
|
472
|
+
|
|
473
|
+
if value.get("kind") != "cache":
|
|
474
|
+
return self._deserialize_cache_value(value)
|
|
475
|
+
|
|
476
|
+
raw_layers = value.get("layers")
|
|
477
|
+
if not isinstance(raw_layers, list):
|
|
478
|
+
raise ValueError("Invalid PyTorch cache layers")
|
|
479
|
+
layers = [self._deserialize_cache_layer(layer) for layer in raw_layers]
|
|
480
|
+
legacy_layers = []
|
|
481
|
+
for layer in layers:
|
|
482
|
+
if isinstance(layer, tuple):
|
|
483
|
+
legacy_layers.append(layer)
|
|
484
|
+
continue
|
|
485
|
+
if not isinstance(layer, dict):
|
|
486
|
+
break
|
|
487
|
+
attributes = layer["attributes"]
|
|
488
|
+
keys = attributes.get("keys")
|
|
489
|
+
values = attributes.get("values")
|
|
490
|
+
if (
|
|
491
|
+
not isinstance(keys, torch.Tensor)
|
|
492
|
+
or not isinstance(values, torch.Tensor)
|
|
493
|
+
or "conv_states" in attributes
|
|
494
|
+
or "recurrent_states" in attributes
|
|
495
|
+
or "indexer_keys" in attributes
|
|
496
|
+
or "cumulative_length" in attributes
|
|
497
|
+
or "max_cache_len" in attributes
|
|
498
|
+
):
|
|
499
|
+
break
|
|
500
|
+
layer_token_count = layer.get("token_count")
|
|
501
|
+
if layer_token_count is not None:
|
|
502
|
+
try:
|
|
503
|
+
layer_token_count = int(layer_token_count)
|
|
504
|
+
except (TypeError, ValueError):
|
|
505
|
+
break
|
|
506
|
+
if layer_token_count < 0:
|
|
507
|
+
break
|
|
508
|
+
if keys.ndim >= 2 and layer_token_count < keys.shape[-2]:
|
|
509
|
+
keys = keys[..., :layer_token_count, :]
|
|
510
|
+
values = values[..., :layer_token_count, :]
|
|
511
|
+
legacy_layers.append((keys, values))
|
|
512
|
+
if len(legacy_layers) == len(layers):
|
|
513
|
+
try:
|
|
514
|
+
from transformers.cache_utils import DynamicCache
|
|
515
|
+
|
|
516
|
+
return DynamicCache(legacy_layers)
|
|
517
|
+
except Exception:
|
|
518
|
+
return tuple(legacy_layers)
|
|
519
|
+
|
|
520
|
+
config = getattr(self.model, "config", None)
|
|
521
|
+
if config is None:
|
|
522
|
+
raise ValueError(
|
|
523
|
+
"A model config is required to restore a non-KV Transformers cache"
|
|
524
|
+
)
|
|
525
|
+
try:
|
|
526
|
+
from transformers.cache_utils import DynamicCache
|
|
527
|
+
|
|
528
|
+
prompt_cache = DynamicCache(config=config)
|
|
529
|
+
except Exception:
|
|
530
|
+
raise ValueError("Unable to initialize the Transformers cache")
|
|
531
|
+
|
|
532
|
+
if len(prompt_cache.layers) != len(layers):
|
|
533
|
+
raise ValueError("Transformers cache layer count does not match model")
|
|
534
|
+
for target_layer, source_layer in zip(prompt_cache.layers, layers):
|
|
535
|
+
if not isinstance(source_layer, dict):
|
|
536
|
+
raise ValueError("Invalid Transformers cache layer")
|
|
537
|
+
class_name = source_layer.get("class_name")
|
|
538
|
+
if class_name and type(target_layer).__name__ != class_name:
|
|
539
|
+
raise ValueError("Transformers cache layer type does not match model")
|
|
540
|
+
for name, item in source_layer["attributes"].items():
|
|
541
|
+
setattr(target_layer, name, item)
|
|
542
|
+
return prompt_cache
|
|
543
|
+
|
|
544
|
+
@staticmethod
|
|
545
|
+
def _cache_offset_from_value(prompt_cache: Any) -> int:
|
|
546
|
+
get_seq_length = getattr(prompt_cache, "get_seq_length", None)
|
|
547
|
+
if callable(get_seq_length):
|
|
548
|
+
try:
|
|
549
|
+
value = get_seq_length()
|
|
550
|
+
return int(value.item() if hasattr(value, "item") else value)
|
|
551
|
+
except Exception:
|
|
552
|
+
pass
|
|
553
|
+
|
|
554
|
+
if isinstance(prompt_cache, torch.Tensor):
|
|
555
|
+
shape = getattr(prompt_cache, "shape", None)
|
|
556
|
+
if shape is not None and len(shape) >= 2:
|
|
557
|
+
return int(shape[-2])
|
|
558
|
+
|
|
559
|
+
if isinstance(prompt_cache, (list, tuple)):
|
|
560
|
+
offsets = [
|
|
561
|
+
TransformersLmBackend._cache_offset_from_value(item)
|
|
562
|
+
for item in prompt_cache
|
|
563
|
+
]
|
|
564
|
+
return max(offsets, default=0)
|
|
565
|
+
|
|
566
|
+
keys = getattr(prompt_cache, "keys", None)
|
|
567
|
+
shape = getattr(keys, "shape", None)
|
|
568
|
+
if shape is not None and len(shape) >= 2:
|
|
569
|
+
return int(shape[-2])
|
|
570
|
+
|
|
571
|
+
layers = getattr(prompt_cache, "layers", None)
|
|
572
|
+
if layers is not None:
|
|
573
|
+
offsets = [TransformersLmBackend._cache_offset_from_value(item) for item in layers]
|
|
574
|
+
return max(offsets, default=0)
|
|
575
|
+
return 0
|
|
576
|
+
|
|
577
|
+
def _write_cache_meta(
|
|
578
|
+
self,
|
|
579
|
+
cache_path: str,
|
|
580
|
+
token_count: int,
|
|
581
|
+
prefix_offsets: list[int] | None = None,
|
|
582
|
+
prefix_hashes: list[str] | None = None,
|
|
583
|
+
) -> dict[str, Any]:
|
|
584
|
+
meta = self._cache_meta_for(token_count, prefix_offsets, prefix_hashes)
|
|
585
|
+
self._atomic_json_save(meta, self._meta_path(cache_path))
|
|
586
|
+
self._cache_meta[cache_path] = meta
|
|
587
|
+
return meta
|
|
588
|
+
|
|
589
|
+
@classmethod
|
|
590
|
+
def _read_cache_meta(cls, cache_path: str) -> dict[str, Any] | None:
|
|
591
|
+
try:
|
|
592
|
+
with open(cls._meta_path(cache_path), encoding="utf-8") as input_file:
|
|
593
|
+
meta = json.load(input_file)
|
|
594
|
+
except (FileNotFoundError, json.JSONDecodeError, OSError):
|
|
595
|
+
return None
|
|
596
|
+
|
|
597
|
+
if not isinstance(meta, dict) or meta.get("layout") != cls.CACHE_LAYOUT:
|
|
598
|
+
return None
|
|
599
|
+
try:
|
|
600
|
+
token_count = int(meta["token_count"])
|
|
601
|
+
except (KeyError, TypeError, ValueError):
|
|
602
|
+
return None
|
|
603
|
+
if token_count < 0:
|
|
604
|
+
return None
|
|
605
|
+
return meta
|
|
606
|
+
|
|
607
|
+
def _cache_meta_matches_current(self, meta: dict[str, Any]) -> bool:
|
|
608
|
+
expected = {
|
|
609
|
+
"layout": self.CACHE_LAYOUT,
|
|
610
|
+
"model_id": self._current_model_id(),
|
|
611
|
+
"dtype": self._current_dtype(),
|
|
612
|
+
"device": str(self._device),
|
|
613
|
+
}
|
|
614
|
+
return all(meta.get(key) == value for key, value in expected.items())
|
|
615
|
+
|
|
616
|
+
@staticmethod
|
|
617
|
+
def _token_ids_from_payload(value: Any) -> tuple[int, ...]:
|
|
618
|
+
if isinstance(value, torch.Tensor):
|
|
619
|
+
value = value.flatten().tolist()
|
|
620
|
+
if not isinstance(value, (list, tuple)):
|
|
621
|
+
raise ValueError("Invalid token IDs in PyTorch cache payload")
|
|
622
|
+
return tuple(int(token_id) for token_id in value)
|
|
623
|
+
|
|
624
|
+
def _prompt_matches_cache(
|
|
625
|
+
self,
|
|
626
|
+
prompt: str | list[int] | None,
|
|
627
|
+
cached_token_ids: tuple[int, ...],
|
|
628
|
+
prefix_token_count: int | None = None,
|
|
629
|
+
) -> bool:
|
|
630
|
+
if prompt is None:
|
|
631
|
+
return True
|
|
632
|
+
current_token_ids = (
|
|
633
|
+
self.tokenize_prompt(prompt)
|
|
634
|
+
if isinstance(prompt, str)
|
|
635
|
+
else [int(token_id) for token_id in prompt]
|
|
636
|
+
)
|
|
637
|
+
if prefix_token_count is None:
|
|
638
|
+
compare_count = len(cached_token_ids)
|
|
639
|
+
else:
|
|
640
|
+
try:
|
|
641
|
+
prefix_token_count = int(prefix_token_count)
|
|
642
|
+
except (TypeError, ValueError):
|
|
643
|
+
return False
|
|
644
|
+
if prefix_token_count < 0 or len(current_token_ids) < prefix_token_count:
|
|
645
|
+
return False
|
|
646
|
+
compare_count = min(prefix_token_count, len(cached_token_ids))
|
|
647
|
+
return (
|
|
648
|
+
len(current_token_ids) >= compare_count
|
|
649
|
+
and tuple(current_token_ids[:compare_count])
|
|
650
|
+
== cached_token_ids[:compare_count]
|
|
651
|
+
)
|
|
652
|
+
|
|
653
|
+
def _load_disk_cache(
|
|
654
|
+
self,
|
|
655
|
+
cache_path: str,
|
|
656
|
+
prompt: str | list[int] | None = None,
|
|
657
|
+
prefix_token_count: int | None = None,
|
|
658
|
+
) -> Any | None:
|
|
659
|
+
meta = self._read_cache_meta(cache_path)
|
|
660
|
+
if meta is None:
|
|
661
|
+
sys.stderr.write(f"PyTorch cache metadata not found or invalid: {cache_path}\n")
|
|
662
|
+
return None
|
|
663
|
+
if not self._cache_meta_matches_current(meta):
|
|
664
|
+
sys.stderr.write(f"PyTorch cache metadata mismatch: {cache_path}\n")
|
|
665
|
+
return None
|
|
666
|
+
|
|
667
|
+
try:
|
|
668
|
+
try:
|
|
669
|
+
payload = torch.load(
|
|
670
|
+
cache_path,
|
|
671
|
+
map_location=self._device,
|
|
672
|
+
weights_only=True,
|
|
673
|
+
)
|
|
674
|
+
except TypeError:
|
|
675
|
+
payload = torch.load(cache_path, map_location=self._device)
|
|
676
|
+
if not isinstance(payload, dict) or payload.get("layout") != self.CACHE_LAYOUT:
|
|
677
|
+
raise ValueError("unsupported cache layout")
|
|
678
|
+
cached_token_ids = self._token_ids_from_payload(payload["token_ids"])
|
|
679
|
+
token_count = int(meta["token_count"])
|
|
680
|
+
if len(cached_token_ids) != token_count:
|
|
681
|
+
raise ValueError("cache token count does not match metadata")
|
|
682
|
+
prompt_cache = self._deserialize_cache(payload["cache"])
|
|
683
|
+
cache_offset = self._cache_offset_from_value(prompt_cache)
|
|
684
|
+
if cache_offset not in (0, token_count):
|
|
685
|
+
raise ValueError("cache offset does not match metadata")
|
|
686
|
+
if not self._prompt_matches_cache(
|
|
687
|
+
prompt,
|
|
688
|
+
cached_token_ids,
|
|
689
|
+
prefix_token_count=prefix_token_count,
|
|
690
|
+
):
|
|
691
|
+
sys.stderr.write(f"PyTorch cache prompt prefix mismatch: {cache_path}\n")
|
|
692
|
+
return None
|
|
693
|
+
except Exception as error:
|
|
694
|
+
sys.stderr.write(f"Failed to load PyTorch cache {cache_path}: {error}\n")
|
|
695
|
+
return None
|
|
696
|
+
|
|
697
|
+
self._caches[cache_path] = prompt_cache
|
|
698
|
+
self._cache_token_counts[cache_path] = token_count
|
|
699
|
+
self._cache_token_ids[cache_path] = cached_token_ids
|
|
700
|
+
self._cache_meta[cache_path] = meta
|
|
701
|
+
return prompt_cache
|
|
702
|
+
|
|
703
|
+
def get_cache_offset(self, prompt_cache: Any) -> int:
|
|
704
|
+
"""Return the token count represented by a Transformers cache."""
|
|
705
|
+
offset = self._cache_offset_from_value(prompt_cache)
|
|
706
|
+
if offset > 0:
|
|
707
|
+
return offset
|
|
708
|
+
|
|
709
|
+
for cache_path, cached in self._caches.items():
|
|
710
|
+
if cached is prompt_cache:
|
|
711
|
+
return self._cache_token_counts.get(cache_path, 0)
|
|
712
|
+
|
|
713
|
+
return super().get_cache_offset(prompt_cache)
|
|
714
|
+
|
|
715
|
+
@staticmethod
|
|
716
|
+
def _trim_tensor(tensor: torch.Tensor, target_tokens: int) -> torch.Tensor:
|
|
717
|
+
if tensor.ndim < 2:
|
|
718
|
+
return tensor
|
|
719
|
+
return tensor[..., :target_tokens, :]
|
|
720
|
+
|
|
721
|
+
@staticmethod
|
|
722
|
+
def _set_cache_layer_length(layer: Any, token_count: int) -> None:
|
|
723
|
+
for attribute in (
|
|
724
|
+
"cumulative_length",
|
|
725
|
+
"cumulative_length_int",
|
|
726
|
+
"indexer_cumulative_length",
|
|
727
|
+
):
|
|
728
|
+
if not hasattr(layer, attribute):
|
|
729
|
+
continue
|
|
730
|
+
value = getattr(layer, attribute)
|
|
731
|
+
try:
|
|
732
|
+
if isinstance(value, torch.Tensor):
|
|
733
|
+
value.fill_(token_count)
|
|
734
|
+
else:
|
|
735
|
+
setattr(layer, attribute, token_count)
|
|
736
|
+
except (AttributeError, RuntimeError, TypeError):
|
|
737
|
+
pass
|
|
738
|
+
|
|
739
|
+
def _trim_cache_layer(self, layer: Any, target_tokens: int, tokens: int) -> bool:
|
|
740
|
+
keys = getattr(layer, "keys", None)
|
|
741
|
+
values = getattr(layer, "values", None)
|
|
742
|
+
has_auxiliary_state = hasattr(layer, "conv_states") or hasattr(
|
|
743
|
+
layer, "recurrent_states"
|
|
744
|
+
)
|
|
745
|
+
if has_auxiliary_state:
|
|
746
|
+
trim = getattr(layer, "trim", None)
|
|
747
|
+
if callable(trim):
|
|
748
|
+
trim(tokens)
|
|
749
|
+
return True
|
|
750
|
+
crop = getattr(layer, "crop", None)
|
|
751
|
+
if callable(crop):
|
|
752
|
+
crop(-tokens)
|
|
753
|
+
return True
|
|
754
|
+
|
|
755
|
+
# Sliding-window layers may retain only the working window while
|
|
756
|
+
# tracking a larger logical sequence. Their crop implementation
|
|
757
|
+
# knows how to select the corresponding suffix and update that
|
|
758
|
+
# logical length; slicing keys directly would keep the wrong window.
|
|
759
|
+
crop = getattr(layer, "crop", None)
|
|
760
|
+
if callable(crop) and getattr(layer, "is_sliding", False):
|
|
761
|
+
crop(-tokens)
|
|
762
|
+
return True
|
|
763
|
+
|
|
764
|
+
if isinstance(keys, torch.Tensor) and isinstance(values, torch.Tensor):
|
|
765
|
+
current_tokens = self._cache_offset_from_value(layer)
|
|
766
|
+
physical_tokens = keys.shape[-2] if keys.ndim >= 2 else current_tokens
|
|
767
|
+
if target_tokens < current_tokens and current_tokens > physical_tokens:
|
|
768
|
+
raise ValueError(
|
|
769
|
+
"Cannot trim a sliding-window cache after its prefix was discarded"
|
|
770
|
+
)
|
|
771
|
+
|
|
772
|
+
get_max_length = getattr(layer, "get_max_length", None)
|
|
773
|
+
max_length = None
|
|
774
|
+
if callable(get_max_length):
|
|
775
|
+
try:
|
|
776
|
+
max_length = int(get_max_length())
|
|
777
|
+
except (TypeError, ValueError, RuntimeError):
|
|
778
|
+
pass
|
|
779
|
+
is_static = max_length is not None and max_length > 0 and physical_tokens >= max_length
|
|
780
|
+
if is_static:
|
|
781
|
+
# StaticCache owns a fixed-capacity tensor. Keep the capacity and
|
|
782
|
+
# move only the logical length; generation will overwrite the tail.
|
|
783
|
+
self._set_cache_layer_length(layer, target_tokens)
|
|
784
|
+
else:
|
|
785
|
+
layer.keys = self._trim_tensor(keys, target_tokens)
|
|
786
|
+
layer.values = self._trim_tensor(values, target_tokens)
|
|
787
|
+
self._set_cache_layer_length(layer, target_tokens)
|
|
788
|
+
|
|
789
|
+
indexer_keys = getattr(layer, "indexer_keys", None)
|
|
790
|
+
if (
|
|
791
|
+
not is_static
|
|
792
|
+
and isinstance(indexer_keys, torch.Tensor)
|
|
793
|
+
and indexer_keys.ndim >= 2
|
|
794
|
+
):
|
|
795
|
+
layer.indexer_keys = indexer_keys[:, :target_tokens, ...]
|
|
796
|
+
return True
|
|
797
|
+
|
|
798
|
+
trim = getattr(layer, "trim", None)
|
|
799
|
+
if callable(trim):
|
|
800
|
+
trim(tokens)
|
|
801
|
+
return True
|
|
802
|
+
crop = getattr(layer, "crop", None)
|
|
803
|
+
if callable(crop):
|
|
804
|
+
# Transformers 5.x uses a negative value for the number of tokens
|
|
805
|
+
# to remove; the positive absolute-length form is deprecated.
|
|
806
|
+
crop(-tokens)
|
|
807
|
+
return True
|
|
808
|
+
return False
|
|
809
|
+
|
|
810
|
+
def _trim_cache_sequence(self, prompt_cache: Any, target_tokens: int) -> Any:
|
|
811
|
+
if isinstance(prompt_cache, torch.Tensor):
|
|
812
|
+
return self._trim_tensor(prompt_cache, target_tokens)
|
|
813
|
+
if isinstance(prompt_cache, tuple):
|
|
814
|
+
return tuple(
|
|
815
|
+
self._trim_cache_sequence(item, target_tokens)
|
|
816
|
+
for item in prompt_cache
|
|
817
|
+
)
|
|
818
|
+
if isinstance(prompt_cache, list):
|
|
819
|
+
return [
|
|
820
|
+
self._trim_cache_sequence(item, target_tokens)
|
|
821
|
+
for item in prompt_cache
|
|
822
|
+
]
|
|
823
|
+
return prompt_cache
|
|
824
|
+
|
|
825
|
+
def trim_cache(self, prompt_cache: Any, tokens: int) -> Any:
|
|
826
|
+
"""Remove trailing tokens from a Transformers KV cache.
|
|
827
|
+
|
|
828
|
+
Legacy tuple caches and Transformers ``Cache`` instances are returned
|
|
829
|
+
as trimmed copies. The input cache is never modified, so callers can
|
|
830
|
+
safely reuse a registered cache reference after a trim.
|
|
831
|
+
"""
|
|
832
|
+
tokens = int(tokens)
|
|
833
|
+
if tokens < 0:
|
|
834
|
+
raise ValueError("tokens to trim must be non-negative")
|
|
835
|
+
if tokens == 0:
|
|
836
|
+
return prompt_cache
|
|
837
|
+
|
|
838
|
+
current_tokens = self.get_cache_offset(prompt_cache)
|
|
839
|
+
target_tokens = max(0, current_tokens - tokens)
|
|
840
|
+
if target_tokens == current_tokens:
|
|
841
|
+
return prompt_cache
|
|
842
|
+
|
|
843
|
+
if isinstance(prompt_cache, (torch.Tensor, tuple, list)):
|
|
844
|
+
return self._trim_cache_sequence(prompt_cache, target_tokens)
|
|
845
|
+
|
|
846
|
+
trimmed_cache = self._clone_cache(prompt_cache)
|
|
847
|
+
if trimmed_cache is prompt_cache:
|
|
848
|
+
raise ValueError(
|
|
849
|
+
"Unable to clone Transformers cache for non-destructive trim"
|
|
850
|
+
)
|
|
851
|
+
prompt_cache = trimmed_cache
|
|
852
|
+
layers = getattr(prompt_cache, "layers", None)
|
|
853
|
+
if layers is not None:
|
|
854
|
+
cache_crop = getattr(prompt_cache, "crop", None)
|
|
855
|
+
requires_cache_crop = any(
|
|
856
|
+
(
|
|
857
|
+
not callable(getattr(layer, "trim", None))
|
|
858
|
+
and not callable(getattr(layer, "crop", None))
|
|
859
|
+
and (
|
|
860
|
+
not (
|
|
861
|
+
isinstance(getattr(layer, "keys", None), torch.Tensor)
|
|
862
|
+
and isinstance(getattr(layer, "values", None), torch.Tensor)
|
|
863
|
+
)
|
|
864
|
+
or hasattr(layer, "conv_states")
|
|
865
|
+
or hasattr(layer, "recurrent_states")
|
|
866
|
+
)
|
|
867
|
+
)
|
|
868
|
+
for layer in layers
|
|
869
|
+
)
|
|
870
|
+
if requires_cache_crop:
|
|
871
|
+
if not callable(cache_crop):
|
|
872
|
+
raise ValueError(
|
|
873
|
+
"Unsupported Transformers cache layer for trimming"
|
|
874
|
+
)
|
|
875
|
+
# Transformers 5.x interprets a negative value as the number
|
|
876
|
+
# of tokens to remove.
|
|
877
|
+
cache_crop(-tokens)
|
|
878
|
+
return prompt_cache
|
|
879
|
+
|
|
880
|
+
for layer in layers:
|
|
881
|
+
layer_tokens = self._cache_offset_from_value(layer)
|
|
882
|
+
if layer_tokens <= 0:
|
|
883
|
+
continue
|
|
884
|
+
self._trim_cache_layer(
|
|
885
|
+
layer,
|
|
886
|
+
max(0, layer_tokens - tokens),
|
|
887
|
+
tokens,
|
|
888
|
+
)
|
|
889
|
+
return prompt_cache
|
|
890
|
+
|
|
891
|
+
crop = getattr(prompt_cache, "crop", None)
|
|
892
|
+
if callable(crop):
|
|
893
|
+
crop(-tokens)
|
|
894
|
+
return prompt_cache
|
|
895
|
+
raise ValueError(f"Unsupported Transformers cache: {type(prompt_cache).__name__}")
|
|
896
|
+
|
|
897
|
+
def cache_prefill(
|
|
898
|
+
self,
|
|
899
|
+
cache_path: str,
|
|
900
|
+
prompt: str,
|
|
901
|
+
base_cache_path: str | None = None,
|
|
902
|
+
trim_to_tokens: int | None = None,
|
|
903
|
+
prefix_offsets: list[int] | None = None,
|
|
904
|
+
prefix_hashes: list[str] | None = None,
|
|
905
|
+
images: list | None = None,
|
|
906
|
+
max_image_size: int = 768,
|
|
907
|
+
) -> dict:
|
|
908
|
+
"""Prefill and persist a Transformers KV cache.
|
|
909
|
+
|
|
910
|
+
``memory://`` refs retain the Phase 1 process-local behavior. Other
|
|
911
|
+
refs use the backend-owned ``pytorch_kv_v1`` disk layout.
|
|
912
|
+
"""
|
|
913
|
+
if images:
|
|
914
|
+
raise ValueError("TransformersLmBackend does not support vision input")
|
|
915
|
+
if self.model is None or self.tokenizer is None:
|
|
916
|
+
raise RuntimeError("Model is not loaded")
|
|
917
|
+
if trim_to_tokens is not None and trim_to_tokens < 0:
|
|
918
|
+
raise ValueError("trim_to_tokens must be non-negative")
|
|
919
|
+
if base_cache_path is None and trim_to_tokens is not None:
|
|
920
|
+
raise ValueError("trim_to_tokens requires base_cache_path")
|
|
921
|
+
|
|
922
|
+
token_ids = self.tokenize_prompt(prompt)
|
|
923
|
+
if not token_ids:
|
|
924
|
+
raise ValueError("Cannot prefill an empty prompt")
|
|
925
|
+
|
|
926
|
+
prompt_cache = None
|
|
927
|
+
cache_offset = 0
|
|
928
|
+
cache_write_tokens = len(token_ids)
|
|
929
|
+
if base_cache_path is not None:
|
|
930
|
+
base_cache = self.load_cache_from_file(
|
|
931
|
+
base_cache_path,
|
|
932
|
+
prompt=token_ids,
|
|
933
|
+
prefix_token_count=trim_to_tokens,
|
|
934
|
+
)
|
|
935
|
+
if base_cache is not None:
|
|
936
|
+
cache_offset = self.get_cache_offset(base_cache)
|
|
937
|
+
if trim_to_tokens is not None and cache_offset > trim_to_tokens:
|
|
938
|
+
prompt_cache = self.trim_cache(
|
|
939
|
+
base_cache,
|
|
940
|
+
cache_offset - trim_to_tokens,
|
|
941
|
+
)
|
|
942
|
+
cache_offset = trim_to_tokens
|
|
943
|
+
else:
|
|
944
|
+
prompt_cache = self._clone_cache(base_cache)
|
|
945
|
+
cloned_offset = self.get_cache_offset(prompt_cache)
|
|
946
|
+
if cloned_offset > 0:
|
|
947
|
+
cache_offset = cloned_offset
|
|
948
|
+
|
|
949
|
+
if cache_offset <= 0:
|
|
950
|
+
prompt_cache = None
|
|
951
|
+
cache_offset = 0
|
|
952
|
+
elif cache_offset >= len(token_ids):
|
|
953
|
+
cache_write_tokens = 0
|
|
954
|
+
|
|
955
|
+
if prompt_cache is None:
|
|
956
|
+
input_token_ids = token_ids
|
|
957
|
+
model_kwargs: dict[str, Any] = {"use_cache": True}
|
|
958
|
+
cache_offset = 0
|
|
959
|
+
elif cache_offset >= len(token_ids):
|
|
960
|
+
input_token_ids = []
|
|
961
|
+
model_kwargs = {}
|
|
962
|
+
else:
|
|
963
|
+
input_token_ids = token_ids[cache_offset:]
|
|
964
|
+
cache_write_tokens = len(input_token_ids)
|
|
965
|
+
model_kwargs = {
|
|
966
|
+
"use_cache": True,
|
|
967
|
+
"past_key_values": prompt_cache,
|
|
968
|
+
"attention_mask": torch.ones(
|
|
969
|
+
(1, len(token_ids)),
|
|
970
|
+
dtype=torch.long,
|
|
971
|
+
device=self._device,
|
|
972
|
+
),
|
|
973
|
+
}
|
|
974
|
+
if self._supports_model_kwarg("cache_position"):
|
|
975
|
+
model_kwargs["cache_position"] = torch.arange(
|
|
976
|
+
cache_offset,
|
|
977
|
+
cache_offset + len(input_token_ids),
|
|
978
|
+
dtype=torch.long,
|
|
979
|
+
device=self._device,
|
|
980
|
+
)
|
|
981
|
+
|
|
982
|
+
if input_token_ids:
|
|
983
|
+
input_ids = torch.tensor(
|
|
984
|
+
[input_token_ids],
|
|
985
|
+
dtype=torch.long,
|
|
986
|
+
device=self._device,
|
|
987
|
+
)
|
|
988
|
+
with torch.no_grad():
|
|
989
|
+
outputs = self.model(input_ids=input_ids, **model_kwargs)
|
|
990
|
+
|
|
991
|
+
past_key_values = getattr(outputs, "past_key_values", None)
|
|
992
|
+
if past_key_values is None and isinstance(outputs, (tuple, list)) and len(outputs) > 1:
|
|
993
|
+
past_key_values = outputs[1]
|
|
994
|
+
if past_key_values is None:
|
|
995
|
+
raise RuntimeError("Transformers model did not return past_key_values")
|
|
996
|
+
prompt_cache = past_key_values
|
|
997
|
+
|
|
998
|
+
if prompt_cache is None:
|
|
999
|
+
raise RuntimeError("Transformers model did not return past_key_values")
|
|
1000
|
+
|
|
1001
|
+
token_count = len(token_ids)
|
|
1002
|
+
cache_path = os.fspath(cache_path)
|
|
1003
|
+
self._caches[cache_path] = prompt_cache
|
|
1004
|
+
self._cache_token_counts[cache_path] = token_count
|
|
1005
|
+
self._cache_token_ids[cache_path] = tuple(token_ids)
|
|
1006
|
+
self._cache_write_token_counts[cache_path] = cache_write_tokens
|
|
1007
|
+
|
|
1008
|
+
if not self._is_memory_cache_path(cache_path):
|
|
1009
|
+
payload = {
|
|
1010
|
+
"layout": self.CACHE_LAYOUT,
|
|
1011
|
+
"cache": self._serialize_cache(prompt_cache),
|
|
1012
|
+
"token_ids": torch.tensor(token_ids, dtype=torch.long),
|
|
1013
|
+
}
|
|
1014
|
+
self._atomic_torch_save(payload, cache_path)
|
|
1015
|
+
self._write_cache_meta(
|
|
1016
|
+
cache_path,
|
|
1017
|
+
token_count,
|
|
1018
|
+
prefix_offsets,
|
|
1019
|
+
prefix_hashes,
|
|
1020
|
+
)
|
|
1021
|
+
else:
|
|
1022
|
+
self._cache_meta[cache_path] = self._cache_meta_for(
|
|
1023
|
+
token_count,
|
|
1024
|
+
prefix_offsets,
|
|
1025
|
+
prefix_hashes,
|
|
1026
|
+
)
|
|
1027
|
+
|
|
1028
|
+
return {
|
|
1029
|
+
"cache_path": cache_path,
|
|
1030
|
+
"token_count": token_count,
|
|
1031
|
+
"cache_write_tokens": cache_write_tokens,
|
|
1032
|
+
}
|
|
1033
|
+
|
|
1034
|
+
def consume_cache_write_tokens(self, cache_path: str) -> int:
|
|
1035
|
+
"""Attribute each prefill write to the first generate using the cache."""
|
|
1036
|
+
return self._cache_write_token_counts.pop(cache_path, 0)
|
|
1037
|
+
|
|
1038
|
+
def load_cache_from_file(
|
|
1039
|
+
self,
|
|
1040
|
+
cache_path: str,
|
|
1041
|
+
images: list | None = None,
|
|
1042
|
+
max_image_size: int = 768,
|
|
1043
|
+
prompt: str | list[int] | None = None,
|
|
1044
|
+
prefix_token_count: int | None = None,
|
|
1045
|
+
) -> Any | None:
|
|
1046
|
+
"""Load a cache, optionally validating only ``prefix_token_count`` tokens."""
|
|
1047
|
+
if images:
|
|
1048
|
+
sys.stderr.write(
|
|
1049
|
+
f"PyTorch cache does not support vision input: {cache_path}\n"
|
|
1050
|
+
)
|
|
1051
|
+
return None
|
|
1052
|
+
|
|
1053
|
+
cache_path = os.fspath(cache_path)
|
|
1054
|
+
|
|
1055
|
+
prompt_cache = self._caches.get(cache_path)
|
|
1056
|
+
if prompt_cache is None and not self._is_memory_cache_path(cache_path):
|
|
1057
|
+
return self._load_disk_cache(
|
|
1058
|
+
cache_path,
|
|
1059
|
+
prompt=prompt,
|
|
1060
|
+
prefix_token_count=prefix_token_count,
|
|
1061
|
+
)
|
|
1062
|
+
if prompt_cache is None:
|
|
1063
|
+
sys.stderr.write(f"PyTorch cache not found in this process: {cache_path}\n")
|
|
1064
|
+
return None
|
|
1065
|
+
|
|
1066
|
+
cached_token_ids = self._cache_token_ids.get(cache_path)
|
|
1067
|
+
if cached_token_ids is not None and not self._prompt_matches_cache(
|
|
1068
|
+
prompt,
|
|
1069
|
+
cached_token_ids,
|
|
1070
|
+
prefix_token_count=prefix_token_count,
|
|
1071
|
+
):
|
|
1072
|
+
sys.stderr.write(f"PyTorch cache prompt prefix mismatch: {cache_path}\n")
|
|
1073
|
+
return None
|
|
1074
|
+
return prompt_cache
|
|
1075
|
+
|
|
1076
|
+
def stream_generate(
|
|
1077
|
+
self,
|
|
1078
|
+
prompt: str | list[int],
|
|
1079
|
+
options: dict,
|
|
1080
|
+
images: list | None = None,
|
|
1081
|
+
prompt_cache: Any | None = None,
|
|
1082
|
+
) -> Iterator[StreamChunk]:
|
|
1083
|
+
if images:
|
|
1084
|
+
raise ValueError("TransformersLmBackend does not support vision input")
|
|
1085
|
+
if self.model is None or self.tokenizer is None:
|
|
1086
|
+
raise RuntimeError("Model is not loaded")
|
|
1087
|
+
|
|
1088
|
+
final_options = {"max_tokens": 256, **options}
|
|
1089
|
+
max_new_tokens = int(final_options.pop("max_tokens", 256))
|
|
1090
|
+
temperature = float(final_options.pop("temperature", 1.0))
|
|
1091
|
+
top_p = final_options.pop("top_p", None)
|
|
1092
|
+
top_k = final_options.pop("top_k", None)
|
|
1093
|
+
|
|
1094
|
+
if isinstance(prompt, list):
|
|
1095
|
+
token_ids = [int(token_id) for token_id in prompt]
|
|
1096
|
+
else:
|
|
1097
|
+
token_ids = self.tokenize_prompt(prompt)
|
|
1098
|
+
if not token_ids:
|
|
1099
|
+
raise ValueError("Cannot generate from an empty prompt")
|
|
1100
|
+
|
|
1101
|
+
input_ids = torch.tensor([token_ids], dtype=torch.long, device=self._device)
|
|
1102
|
+
prompt_token_count = len(token_ids)
|
|
1103
|
+
cache_read_tokens = 0
|
|
1104
|
+
|
|
1105
|
+
do_sample = temperature > 0
|
|
1106
|
+
gen_kwargs: dict[str, Any] = {
|
|
1107
|
+
"input_ids": input_ids,
|
|
1108
|
+
"max_new_tokens": max_new_tokens,
|
|
1109
|
+
"do_sample": do_sample,
|
|
1110
|
+
}
|
|
1111
|
+
if do_sample:
|
|
1112
|
+
gen_kwargs["temperature"] = temperature
|
|
1113
|
+
if top_p is not None:
|
|
1114
|
+
gen_kwargs["top_p"] = float(top_p)
|
|
1115
|
+
if top_k is not None:
|
|
1116
|
+
gen_kwargs["top_k"] = int(top_k)
|
|
1117
|
+
if prompt_cache is not None:
|
|
1118
|
+
cache_read_tokens = self.get_cache_offset(prompt_cache)
|
|
1119
|
+
gen_kwargs["past_key_values"] = self._clone_cache(prompt_cache)
|
|
1120
|
+
gen_kwargs["attention_mask"] = torch.ones(
|
|
1121
|
+
(1, cache_read_tokens + prompt_token_count),
|
|
1122
|
+
dtype=torch.long,
|
|
1123
|
+
device=self._device,
|
|
1124
|
+
)
|
|
1125
|
+
if self._supports_model_kwarg("cache_position"):
|
|
1126
|
+
gen_kwargs["cache_position"] = torch.arange(
|
|
1127
|
+
cache_read_tokens,
|
|
1128
|
+
cache_read_tokens + prompt_token_count,
|
|
1129
|
+
dtype=torch.long,
|
|
1130
|
+
device=self._device,
|
|
1131
|
+
)
|
|
1132
|
+
|
|
1133
|
+
streamer = _TokenCountingTextIteratorStreamer(
|
|
1134
|
+
self.tokenizer,
|
|
1135
|
+
skip_special_tokens=True,
|
|
1136
|
+
skip_prompt=True,
|
|
1137
|
+
)
|
|
1138
|
+
gen_kwargs["streamer"] = streamer
|
|
1139
|
+
|
|
1140
|
+
thread = Thread(target=self.model.generate, kwargs=gen_kwargs)
|
|
1141
|
+
thread.start()
|
|
1142
|
+
|
|
1143
|
+
first_chunk = True
|
|
1144
|
+
for text, generation_tokens in streamer:
|
|
1145
|
+
chunk = StreamChunk(
|
|
1146
|
+
text=text,
|
|
1147
|
+
prompt_tokens=(prompt_token_count + cache_read_tokens)
|
|
1148
|
+
if first_chunk
|
|
1149
|
+
else None,
|
|
1150
|
+
generation_tokens=generation_tokens,
|
|
1151
|
+
cache_read_tokens=cache_read_tokens if first_chunk else None,
|
|
1152
|
+
)
|
|
1153
|
+
first_chunk = False
|
|
1154
|
+
if is_eod_token(chunk, self.tokenizer):
|
|
1155
|
+
chunk.finish_reason = "stop"
|
|
1156
|
+
yield chunk
|
|
1157
|
+
break
|
|
1158
|
+
yield chunk
|
|
1159
|
+
|
|
1160
|
+
thread.join()
|
|
1161
|
+
|
|
1162
|
+
def supports_vision(self) -> bool:
|
|
1163
|
+
return False
|
|
1164
|
+
|
|
1165
|
+
@property
|
|
1166
|
+
def model_kind(self) -> str:
|
|
1167
|
+
return "lm"
|