@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.
Files changed (192) hide show
  1. package/README.md +93 -10
  2. package/dist/cache-controller.d.ts +4 -0
  3. package/dist/cache-controller.d.ts.map +1 -1
  4. package/dist/driver-registry/config-based-factory.d.ts.map +1 -1
  5. package/dist/driver-registry/config-based-factory.js +6 -0
  6. package/dist/driver-registry/config-based-factory.js.map +1 -1
  7. package/dist/driver-registry/factory-helper.d.ts.map +1 -1
  8. package/dist/driver-registry/factory-helper.js +9 -2
  9. package/dist/driver-registry/factory-helper.js.map +1 -1
  10. package/dist/driver-registry/index.d.ts +1 -1
  11. package/dist/driver-registry/index.d.ts.map +1 -1
  12. package/dist/driver-registry/types.d.ts +15 -1
  13. package/dist/driver-registry/types.d.ts.map +1 -1
  14. package/dist/formatter/converter.d.ts.map +1 -1
  15. package/dist/formatter/converter.js +31 -2
  16. package/dist/formatter/converter.js.map +1 -1
  17. package/dist/index.d.ts +5 -3
  18. package/dist/index.d.ts.map +1 -1
  19. package/dist/index.js +5 -3
  20. package/dist/index.js.map +1 -1
  21. package/dist/local-inference/adapters.d.ts +6 -0
  22. package/dist/local-inference/adapters.d.ts.map +1 -1
  23. package/dist/local-inference/driver.d.ts.map +1 -1
  24. package/dist/local-inference/driver.js +45 -24
  25. package/dist/local-inference/driver.js.map +1 -1
  26. package/dist/local-inference/process-client.d.ts +4 -2
  27. package/dist/local-inference/process-client.d.ts.map +1 -1
  28. package/dist/local-inference/process-client.js +24 -8
  29. package/dist/local-inference/process-client.js.map +1 -1
  30. package/dist/local-inference/process-communication.d.ts +9 -2
  31. package/dist/local-inference/process-communication.d.ts.map +1 -1
  32. package/dist/local-inference/process-communication.js +37 -5
  33. package/dist/local-inference/process-communication.js.map +1 -1
  34. package/dist/local-inference/protocol.d.ts +4 -0
  35. package/dist/local-inference/protocol.d.ts.map +1 -1
  36. package/dist/local-inference/request-queue.d.ts +1 -1
  37. package/dist/local-inference/request-queue.d.ts.map +1 -1
  38. package/dist/local-inference/request-queue.js +26 -7
  39. package/dist/local-inference/request-queue.js.map +1 -1
  40. package/dist/local-inference/stream-utils.d.ts +6 -0
  41. package/dist/local-inference/stream-utils.d.ts.map +1 -1
  42. package/dist/local-inference/stream-utils.js.map +1 -1
  43. package/dist/mlx-ml/mlx-cache-controller.d.ts +9 -0
  44. package/dist/mlx-ml/mlx-cache-controller.d.ts.map +1 -1
  45. package/dist/mlx-ml/mlx-cache-controller.js +158 -33
  46. package/dist/mlx-ml/mlx-cache-controller.js.map +1 -1
  47. package/dist/mlx-ml/mlx-cache-support.d.ts +2 -2
  48. package/dist/mlx-ml/mlx-cache-support.d.ts.map +1 -1
  49. package/dist/mlx-ml/mlx-cache-support.js +8 -3
  50. package/dist/mlx-ml/mlx-cache-support.js.map +1 -1
  51. package/dist/mlx-ml/mlx-driver.d.ts +0 -1
  52. package/dist/mlx-ml/mlx-driver.d.ts.map +1 -1
  53. package/dist/mlx-ml/mlx-driver.js +1 -8
  54. package/dist/mlx-ml/mlx-driver.js.map +1 -1
  55. package/dist/mlx-ml/process/index.d.ts +1 -1
  56. package/dist/mlx-ml/process/index.d.ts.map +1 -1
  57. package/dist/mlx-ml/process/index.js +2 -2
  58. package/dist/mlx-ml/process/index.js.map +1 -1
  59. package/dist/models-config/index.d.ts +1 -1
  60. package/dist/models-config/index.d.ts.map +1 -1
  61. package/dist/models-config/index.js +1 -1
  62. package/dist/models-config/index.js.map +1 -1
  63. package/dist/models-config/resolve.d.ts +9 -1
  64. package/dist/models-config/resolve.d.ts.map +1 -1
  65. package/dist/models-config/resolve.js +94 -2
  66. package/dist/models-config/resolve.js.map +1 -1
  67. package/dist/models-config/types.d.ts +3 -1
  68. package/dist/models-config/types.d.ts.map +1 -1
  69. package/dist/pytorch/process/index.d.ts +4 -2
  70. package/dist/pytorch/process/index.d.ts.map +1 -1
  71. package/dist/pytorch/process/index.js +24 -7
  72. package/dist/pytorch/process/index.js.map +1 -1
  73. package/dist/pytorch/pytorch-cache-controller.d.ts +84 -0
  74. package/dist/pytorch/pytorch-cache-controller.d.ts.map +1 -0
  75. package/dist/pytorch/pytorch-cache-controller.js +742 -0
  76. package/dist/pytorch/pytorch-cache-controller.js.map +1 -0
  77. package/dist/pytorch/pytorch-cache-support.d.ts +23 -0
  78. package/dist/pytorch/pytorch-cache-support.d.ts.map +1 -0
  79. package/dist/pytorch/pytorch-cache-support.js +47 -0
  80. package/dist/pytorch/pytorch-cache-support.js.map +1 -0
  81. package/dist/pytorch/pytorch-driver.d.ts +8 -1
  82. package/dist/pytorch/pytorch-driver.d.ts.map +1 -1
  83. package/dist/pytorch/pytorch-driver.js +40 -0
  84. package/dist/pytorch/pytorch-driver.js.map +1 -1
  85. package/dist/runtime/check.d.ts.map +1 -1
  86. package/dist/runtime/check.js +8 -6
  87. package/dist/runtime/check.js.map +1 -1
  88. package/dist/runtime/index.d.ts +2 -2
  89. package/dist/runtime/index.d.ts.map +1 -1
  90. package/dist/runtime/index.js +2 -2
  91. package/dist/runtime/index.js.map +1 -1
  92. package/dist/runtime/manifest-core.d.mts +1 -0
  93. package/dist/runtime/manifest-core.mjs +1 -0
  94. package/dist/runtime/manifest-core.mjs.map +1 -1
  95. package/dist/runtime/manifest.d.ts +2 -0
  96. package/dist/runtime/manifest.d.ts.map +1 -1
  97. package/dist/runtime/manifest.js.map +1 -1
  98. package/dist/runtime/paths-core.d.mts +15 -1
  99. package/dist/runtime/paths-core.d.mts.map +1 -1
  100. package/dist/runtime/paths-core.mjs +50 -5
  101. package/dist/runtime/paths-core.mjs.map +1 -1
  102. package/dist/runtime/paths.d.ts +2 -2
  103. package/dist/runtime/paths.d.ts.map +1 -1
  104. package/dist/runtime/paths.js +2 -2
  105. package/dist/runtime/paths.js.map +1 -1
  106. package/dist/runtime/pytorch-template-core.d.mts +11 -0
  107. package/dist/runtime/pytorch-template-core.d.mts.map +1 -0
  108. package/dist/runtime/pytorch-template-core.mjs +54 -0
  109. package/dist/runtime/pytorch-template-core.mjs.map +1 -0
  110. package/dist/runtime/setup-commands-core.d.mts +3 -0
  111. package/dist/runtime/setup-commands-core.d.mts.map +1 -1
  112. package/dist/runtime/setup-commands-core.mjs +4 -0
  113. package/dist/runtime/setup-commands-core.mjs.map +1 -1
  114. package/dist/runtime/setup-commands.d.ts +1 -1
  115. package/dist/runtime/setup-commands.d.ts.map +1 -1
  116. package/dist/runtime/setup-commands.js +1 -1
  117. package/dist/runtime/setup-commands.js.map +1 -1
  118. package/docs/DRIVER_API.md +455 -0
  119. package/docs/LOCAL_MODEL_SETUP.md +765 -0
  120. package/docs/mlx-api-selection.md +301 -0
  121. package/package.json +9 -5
  122. package/scripts/runtime-cli.bin.test.ts +142 -0
  123. package/scripts/runtime-cli.js +305 -35
  124. package/scripts/runtime-cli.test.ts +163 -0
  125. package/src/mlx-ml/python/__main__.py +1 -1
  126. package/src/mlx-ml/python/backends/base.py +88 -18
  127. package/src/mlx-ml/python/backends/mlx_lm.py +28 -3
  128. package/src/mlx-ml/python/backends/mlx_vlm.py +679 -2
  129. package/src/mlx-ml/python/handlers/cache.py +4 -0
  130. package/src/mlx-ml/python/handlers/generate.py +33 -10
  131. package/src/mlx-ml/python/handlers/tokenize.py +1 -4
  132. package/src/mlx-ml/python/pyproject.toml +1 -1
  133. package/src/mlx-ml/python/server.py +2 -0
  134. package/src/mlx-ml/python/uv.lock +8 -8
  135. package/src/pytorch/templates/cpu-minimal/backends/base.py +139 -0
  136. package/src/pytorch/templates/cpu-minimal/backends/transformers_lm.py +1167 -0
  137. package/src/pytorch/{python → templates/cpu-minimal}/handlers/__init__.py +1 -0
  138. package/src/pytorch/templates/cpu-minimal/handlers/cache.py +88 -0
  139. package/src/pytorch/templates/cpu-minimal/handlers/generate.py +157 -0
  140. package/src/pytorch/{python → templates/cpu-minimal}/pyproject.toml +2 -2
  141. package/src/pytorch/{python → templates/cpu-minimal}/server.py +20 -2
  142. package/src/pytorch/templates/cpu-minimal/tests/test_cache_handler.py +284 -0
  143. package/src/pytorch/templates/cpu-minimal/tests/test_capabilities.py +14 -0
  144. package/src/pytorch/templates/cpu-minimal/tests/test_server.py +141 -0
  145. package/src/pytorch/templates/cpu-minimal/tests/test_transformers_errors.py +89 -0
  146. package/src/pytorch/templates/cpu-minimal/tests/test_transformers_lm_cache.py +554 -0
  147. package/src/pytorch/templates/cpu-minimal/utils/__init__.py +0 -0
  148. package/src/pytorch/{python → templates/cpu-minimal}/utils/token_utils.py +2 -2
  149. package/src/pytorch/templates/cpu-minimal/utils/transformers_errors.py +54 -0
  150. package/src/pytorch/{python → templates/cpu-minimal}/uv.lock +149 -109
  151. package/src/pytorch/templates/cuda/__main__.py +19 -0
  152. package/src/pytorch/templates/cuda/backends/__init__.py +3 -0
  153. package/src/pytorch/{python → templates/cuda}/backends/base.py +54 -6
  154. package/src/pytorch/templates/cuda/backends/transformers_lm.py +379 -0
  155. package/src/pytorch/templates/cuda/handlers/__init__.py +7 -0
  156. package/src/pytorch/templates/cuda/handlers/cache.py +93 -0
  157. package/src/pytorch/templates/cuda/handlers/cancel.py +53 -0
  158. package/src/pytorch/templates/cuda/handlers/capabilities.py +6 -0
  159. package/src/pytorch/templates/cuda/handlers/completion.py +15 -0
  160. package/src/pytorch/templates/cuda/handlers/format_test.py +70 -0
  161. package/src/pytorch/templates/cuda/handlers/generate.py +152 -0
  162. package/src/pytorch/templates/cuda/handlers/render.py +40 -0
  163. package/src/pytorch/templates/cuda/handlers/tokenize.py +63 -0
  164. package/src/pytorch/templates/cuda/pyproject.toml +37 -0
  165. package/src/pytorch/templates/cuda/server.py +158 -0
  166. package/src/pytorch/templates/cuda/tests/__init__.py +0 -0
  167. package/src/pytorch/templates/cuda/tests/test_cache_handler.py +207 -0
  168. package/src/pytorch/templates/cuda/tests/test_capabilities.py +14 -0
  169. package/src/pytorch/templates/cuda/tests/test_server.py +145 -0
  170. package/src/pytorch/templates/cuda/tests/test_transformers_errors.py +89 -0
  171. package/src/pytorch/templates/cuda/tests/test_transformers_lm_cache.py +288 -0
  172. package/src/pytorch/templates/cuda/utils/__init__.py +0 -0
  173. package/src/pytorch/templates/cuda/utils/chat_template_constraints.py +164 -0
  174. package/src/pytorch/templates/cuda/utils/prompt_builder.py +54 -0
  175. package/src/pytorch/templates/cuda/utils/template_render.py +80 -0
  176. package/src/pytorch/templates/cuda/utils/token_utils.py +376 -0
  177. package/src/pytorch/templates/cuda/utils/transformers_errors.py +54 -0
  178. package/src/pytorch/templates/cuda/uv.lock +734 -0
  179. package/src/pytorch/python/backends/transformers_lm.py +0 -127
  180. package/src/pytorch/python/handlers/generate.py +0 -68
  181. /package/src/pytorch/{python → templates/cpu-minimal}/__main__.py +0 -0
  182. /package/src/pytorch/{python → templates/cpu-minimal}/backends/__init__.py +0 -0
  183. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/cancel.py +0 -0
  184. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/capabilities.py +0 -0
  185. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/completion.py +0 -0
  186. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/format_test.py +0 -0
  187. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/render.py +0 -0
  188. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/tokenize.py +0 -0
  189. /package/src/pytorch/{python/utils → templates/cpu-minimal/tests}/__init__.py +0 -0
  190. /package/src/pytorch/{python → templates/cpu-minimal}/utils/chat_template_constraints.py +0 -0
  191. /package/src/pytorch/{python → templates/cpu-minimal}/utils/prompt_builder.py +0 -0
  192. /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"