@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,379 @@
1
+ from __future__ import annotations
2
+
3
+ import copy
4
+ import os
5
+ import sys
6
+ from dataclasses import dataclass
7
+ from threading import Thread
8
+ from typing import Any, Iterator
9
+
10
+ import torch
11
+ import transformers
12
+ from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer
13
+
14
+ from backends.base import ModelBackend
15
+ from utils.token_utils import is_eod_token
16
+ from utils.transformers_errors import (
17
+ extract_unsupported_model_type,
18
+ unsupported_model_type_error,
19
+ )
20
+
21
+
22
+ class _TokenCountingTextIteratorStreamer(TextIteratorStreamer):
23
+ """TextIteratorStreamer that counts generated token IDs, not text chunks."""
24
+
25
+ def __init__(self, *args, **kwargs):
26
+ super().__init__(*args, **kwargs)
27
+ self.generated_token_count = 0
28
+
29
+ def put(self, value: torch.Tensor) -> None:
30
+ is_prompt = self.skip_prompt and self.next_tokens_are_prompt
31
+ if not is_prompt:
32
+ token_count = value[0].numel() if value.ndim > 1 else value.numel()
33
+ self.generated_token_count += int(token_count)
34
+ super().put(value)
35
+
36
+ def on_finalized_text(self, text: str, stream_end: bool = False) -> None:
37
+ """Queue text together with the token count at its emission point."""
38
+ self.text_queue.put((text, self.generated_token_count), timeout=self.timeout)
39
+ if stream_end:
40
+ self.text_queue.put(self.stop_signal, timeout=self.timeout)
41
+
42
+ def __next__(self) -> tuple[str, int]:
43
+ value = self.text_queue.get(timeout=self.timeout)
44
+ if value == self.stop_signal:
45
+ raise StopIteration()
46
+ return value
47
+
48
+
49
+ @dataclass
50
+ class StreamChunk:
51
+ text: str
52
+ prompt_tokens: int | None = None
53
+ generation_tokens: int | None = None
54
+ finish_reason: str | None = None
55
+ cache_read_tokens: int | None = None
56
+
57
+
58
+ class TransformersLmBackend(ModelBackend):
59
+ """Transformers causal LM backend (text-only, CUDA-first)."""
60
+
61
+ def __init__(self, device: str | None = None) -> None:
62
+ self.model: Any | None = None
63
+ self.tokenizer: Any | None = None
64
+ self._device_name = device or os.environ.get("PYTORCH_DEVICE", "cuda")
65
+ self._device = torch.device(self._device_name)
66
+ if self._device.type == "cuda" and not torch.cuda.is_available():
67
+ raise RuntimeError(
68
+ "CUDA device requested, but CUDA is not available in this PyTorch runtime. "
69
+ "Install a CUDA-enabled torch wheel and verify the NVIDIA driver."
70
+ )
71
+ self._caches: dict[str, Any] = {}
72
+ self._cache_token_counts: dict[str, int] = {}
73
+ self._cache_token_ids: dict[str, tuple[int, ...]] = {}
74
+ self._cache_write_token_counts: dict[str, int] = {}
75
+
76
+ def load(self, model_name: str) -> None:
77
+ self._caches.clear()
78
+ self._cache_token_counts.clear()
79
+ self._cache_token_ids.clear()
80
+ self._cache_write_token_counts.clear()
81
+
82
+ trust_remote_code = os.environ.get("PYTORCH_TRUST_REMOTE_CODE", "").lower() in (
83
+ "1",
84
+ "true",
85
+ "yes",
86
+ )
87
+ try:
88
+ self.tokenizer = AutoTokenizer.from_pretrained(
89
+ model_name,
90
+ trust_remote_code=trust_remote_code,
91
+ )
92
+ if self.tokenizer.pad_token is None and self.tokenizer.eos_token is not None:
93
+ self.tokenizer.pad_token = self.tokenizer.eos_token
94
+
95
+ dtype = torch.float32 if self._device.type == "cpu" else torch.float16
96
+ self.model = AutoModelForCausalLM.from_pretrained(
97
+ model_name,
98
+ trust_remote_code=trust_remote_code,
99
+ dtype=dtype,
100
+ )
101
+ except (KeyError, ValueError) as error:
102
+ model_type = extract_unsupported_model_type(error)
103
+ if model_type is None:
104
+ raise
105
+ raise unsupported_model_type_error(
106
+ model_name,
107
+ model_type,
108
+ getattr(transformers, "__version__", "unknown"),
109
+ ) from error
110
+
111
+ self.model.to(self._device)
112
+ self.model.eval()
113
+
114
+ def get_tokenizer(self) -> Any:
115
+ return self.tokenizer
116
+
117
+ def tokenize_prompt(
118
+ self,
119
+ prompt: str,
120
+ images: list | None = None,
121
+ max_image_size: int = 768,
122
+ ) -> list[int]:
123
+ if images:
124
+ raise ValueError("TransformersLmBackend does not support vision input")
125
+ if self.tokenizer is None:
126
+ raise RuntimeError("Model is not loaded")
127
+
128
+ bos_token = getattr(self.tokenizer, "bos_token", None)
129
+ add_special = bos_token is None or not prompt.startswith(bos_token or "")
130
+ token_ids = self.tokenizer.encode(prompt, add_special_tokens=add_special)
131
+ if hasattr(token_ids, "flatten"):
132
+ token_ids = token_ids.flatten().tolist()
133
+ return [int(token_id) for token_id in token_ids]
134
+
135
+ @staticmethod
136
+ def _clone_cache(prompt_cache: Any) -> Any:
137
+ """Clone model-owned cache state before generation mutates it."""
138
+ if isinstance(prompt_cache, torch.Tensor):
139
+ return prompt_cache.clone()
140
+ if isinstance(prompt_cache, tuple):
141
+ return tuple(TransformersLmBackend._clone_cache(item) for item in prompt_cache)
142
+ if isinstance(prompt_cache, list):
143
+ return [TransformersLmBackend._clone_cache(item) for item in prompt_cache]
144
+ if hasattr(prompt_cache, "layers"):
145
+ try:
146
+ cloned_cache = copy.copy(prompt_cache)
147
+ cloned_cache.layers = [
148
+ TransformersLmBackend._clone_cache(layer)
149
+ for layer in prompt_cache.layers
150
+ ]
151
+ return cloned_cache
152
+ except Exception:
153
+ pass
154
+ if hasattr(prompt_cache, "keys") and hasattr(prompt_cache, "values"):
155
+ try:
156
+ cloned_layer = copy.copy(prompt_cache)
157
+ cloned_layer.keys = TransformersLmBackend._clone_cache(prompt_cache.keys)
158
+ cloned_layer.values = TransformersLmBackend._clone_cache(prompt_cache.values)
159
+ return cloned_layer
160
+ except Exception:
161
+ pass
162
+ if hasattr(prompt_cache, "get_seq_length"):
163
+ try:
164
+ return copy.deepcopy(prompt_cache)
165
+ except Exception:
166
+ # Some third-party Cache implementations cannot be deep-copied.
167
+ # Keep the request usable; those implementations must tolerate
168
+ # in-place generation updates.
169
+ return prompt_cache
170
+ return prompt_cache
171
+
172
+ def get_cache_offset(self, prompt_cache: Any) -> int:
173
+ """Return the token count represented by a Transformers cache."""
174
+ for cache_path, cached in self._caches.items():
175
+ if cached is prompt_cache:
176
+ return self._cache_token_counts.get(cache_path, 0)
177
+
178
+ get_seq_length = getattr(prompt_cache, "get_seq_length", None)
179
+ if callable(get_seq_length):
180
+ try:
181
+ return int(get_seq_length())
182
+ except Exception:
183
+ pass
184
+
185
+ if isinstance(prompt_cache, (list, tuple)):
186
+ for layer in prompt_cache:
187
+ if isinstance(layer, (list, tuple)) and layer:
188
+ key = layer[0]
189
+ else:
190
+ key = layer
191
+ shape = getattr(key, "shape", None)
192
+ if shape is not None and len(shape) >= 2:
193
+ return int(shape[-2])
194
+ return super().get_cache_offset(prompt_cache)
195
+
196
+ def cache_prefill(
197
+ self,
198
+ cache_path: str,
199
+ prompt: str,
200
+ base_cache_path: str | None = None,
201
+ trim_to_tokens: int | None = None,
202
+ prefix_offsets: list[int] | None = None,
203
+ prefix_hashes: list[str] | None = None,
204
+ images: list | None = None,
205
+ max_image_size: int = 768,
206
+ ) -> dict:
207
+ """Prefill and retain a Transformers KV cache in this process.
208
+
209
+ ``cache_path`` is an opaque process-local reference in Phase 1. No
210
+ file is created; persistence and incremental prefill belong to Phase 2.
211
+ """
212
+ if images:
213
+ raise ValueError("TransformersLmBackend does not support vision input")
214
+ if base_cache_path is not None or trim_to_tokens is not None:
215
+ raise ValueError(
216
+ "TransformersLmBackend does not support incremental prefill in Phase 1"
217
+ )
218
+ if prefix_offsets is not None or prefix_hashes is not None:
219
+ raise ValueError(
220
+ "TransformersLmBackend does not support cache prefix metadata in Phase 1"
221
+ )
222
+ if self.model is None or self.tokenizer is None:
223
+ raise RuntimeError("Model is not loaded")
224
+
225
+ token_ids = self.tokenize_prompt(prompt)
226
+ if not token_ids:
227
+ raise ValueError("Cannot prefill an empty prompt")
228
+
229
+ input_ids = torch.tensor([token_ids], dtype=torch.long, device=self._device)
230
+ with torch.no_grad():
231
+ outputs = self.model(input_ids=input_ids, use_cache=True)
232
+
233
+ past_key_values = getattr(outputs, "past_key_values", None)
234
+ if past_key_values is None and isinstance(outputs, (tuple, list)) and len(outputs) > 1:
235
+ past_key_values = outputs[1]
236
+ if past_key_values is None:
237
+ raise RuntimeError("Transformers model did not return past_key_values")
238
+
239
+ self._caches[cache_path] = past_key_values
240
+ self._cache_token_counts[cache_path] = len(token_ids)
241
+ self._cache_token_ids[cache_path] = tuple(token_ids)
242
+ self._cache_write_token_counts[cache_path] = len(token_ids)
243
+ return {
244
+ "cache_path": cache_path,
245
+ "token_count": len(token_ids),
246
+ "cache_write_tokens": len(token_ids),
247
+ }
248
+
249
+ def consume_cache_write_tokens(self, cache_path: str) -> int:
250
+ """Attribute each prefill write to the first generate using the cache."""
251
+ return self._cache_write_token_counts.pop(cache_path, 0)
252
+
253
+ def load_cache_from_file(
254
+ self,
255
+ cache_path: str,
256
+ images: list | None = None,
257
+ max_image_size: int = 768,
258
+ prompt: str | list[int] | None = None,
259
+ ) -> Any | None:
260
+ """Resolve a process-local cache reference; disk loading is Phase 2."""
261
+ if images:
262
+ sys.stderr.write(
263
+ f"PyTorch cache does not support vision input: {cache_path}\n"
264
+ )
265
+ return None
266
+
267
+ prompt_cache = self._caches.get(cache_path)
268
+ if prompt_cache is None:
269
+ sys.stderr.write(f"PyTorch cache not found in this process: {cache_path}\n")
270
+ return None
271
+
272
+ cached_token_ids = self._cache_token_ids.get(cache_path)
273
+ if cached_token_ids is not None and isinstance(prompt, (str, list)):
274
+ current_token_ids = (
275
+ self.tokenize_prompt(prompt)
276
+ if isinstance(prompt, str)
277
+ else [int(token_id) for token_id in prompt]
278
+ )
279
+ if (
280
+ len(current_token_ids) < len(cached_token_ids)
281
+ or tuple(current_token_ids[: len(cached_token_ids)]) != cached_token_ids
282
+ ):
283
+ sys.stderr.write(
284
+ f"PyTorch cache prompt prefix mismatch: {cache_path}\n"
285
+ )
286
+ return None
287
+ return prompt_cache
288
+
289
+ def stream_generate(
290
+ self,
291
+ prompt: str | list[int],
292
+ options: dict,
293
+ images: list | None = None,
294
+ prompt_cache: Any | None = None,
295
+ ) -> Iterator[StreamChunk]:
296
+ if images:
297
+ raise ValueError("TransformersLmBackend does not support vision input")
298
+ if self.model is None or self.tokenizer is None:
299
+ raise RuntimeError("Model is not loaded")
300
+
301
+ final_options = {"max_tokens": 256, **options}
302
+ max_new_tokens = int(final_options.pop("max_tokens", 256))
303
+ temperature = float(final_options.pop("temperature", 1.0))
304
+ top_p = final_options.pop("top_p", None)
305
+ top_k = final_options.pop("top_k", None)
306
+
307
+ if isinstance(prompt, list):
308
+ token_ids = [int(token_id) for token_id in prompt]
309
+ else:
310
+ token_ids = self.tokenize_prompt(prompt)
311
+ if not token_ids:
312
+ raise ValueError("Cannot generate from an empty prompt")
313
+
314
+ input_ids = torch.tensor([token_ids], dtype=torch.long, device=self._device)
315
+ prompt_token_count = len(token_ids)
316
+ cache_read_tokens = 0
317
+
318
+ do_sample = temperature > 0
319
+ gen_kwargs: dict[str, Any] = {
320
+ "input_ids": input_ids,
321
+ "max_new_tokens": max_new_tokens,
322
+ "do_sample": do_sample,
323
+ }
324
+ if do_sample:
325
+ gen_kwargs["temperature"] = temperature
326
+ if top_p is not None:
327
+ gen_kwargs["top_p"] = float(top_p)
328
+ if top_k is not None:
329
+ gen_kwargs["top_k"] = int(top_k)
330
+ if prompt_cache is not None:
331
+ cache_read_tokens = self.get_cache_offset(prompt_cache)
332
+ gen_kwargs["past_key_values"] = self._clone_cache(prompt_cache)
333
+ gen_kwargs["cache_position"] = torch.arange(
334
+ cache_read_tokens,
335
+ cache_read_tokens + prompt_token_count,
336
+ dtype=torch.long,
337
+ device=self._device,
338
+ )
339
+ gen_kwargs["attention_mask"] = torch.ones(
340
+ (1, cache_read_tokens + prompt_token_count),
341
+ dtype=torch.long,
342
+ device=self._device,
343
+ )
344
+
345
+ streamer = _TokenCountingTextIteratorStreamer(
346
+ self.tokenizer,
347
+ skip_special_tokens=True,
348
+ skip_prompt=True,
349
+ )
350
+ gen_kwargs["streamer"] = streamer
351
+
352
+ thread = Thread(target=self.model.generate, kwargs=gen_kwargs)
353
+ thread.start()
354
+
355
+ first_chunk = True
356
+ for text, generation_tokens in streamer:
357
+ chunk = StreamChunk(
358
+ text=text,
359
+ prompt_tokens=(prompt_token_count + cache_read_tokens)
360
+ if first_chunk
361
+ else None,
362
+ generation_tokens=generation_tokens,
363
+ cache_read_tokens=cache_read_tokens if first_chunk else None,
364
+ )
365
+ first_chunk = False
366
+ if is_eod_token(chunk, self.tokenizer):
367
+ chunk.finish_reason = "stop"
368
+ yield chunk
369
+ break
370
+ yield chunk
371
+
372
+ thread.join()
373
+
374
+ def supports_vision(self) -> bool:
375
+ return False
376
+
377
+ @property
378
+ def model_kind(self) -> str:
379
+ return "lm"
@@ -0,0 +1,7 @@
1
+ from handlers.cache import handle_cache_prefill
2
+ from handlers.capabilities import handle_capabilities
3
+ from handlers.completion import handle_completion
4
+ from handlers.format_test import handle_format_test
5
+ from handlers.generate import handle_generate
6
+ from handlers.render import handle_render
7
+ from handlers.tokenize import handle_tokenize
@@ -0,0 +1,93 @@
1
+ from __future__ import annotations
2
+
3
+ import json
4
+
5
+ from backends.base import ModelBackend
6
+ from utils.prompt_builder import generate_merged_prompt, supports_chat_template
7
+
8
+
9
+ def _render_prefill_prompt(
10
+ backend: ModelBackend,
11
+ capabilities: dict,
12
+ messages: list,
13
+ tools: list | None,
14
+ reasoning_effort: str | None,
15
+ ) -> str:
16
+ tokenizer = backend.get_tokenizer()
17
+ extra_kwargs = {}
18
+ if tools is not None:
19
+ extra_kwargs["tools"] = tools
20
+ if reasoning_effort is not None:
21
+ extra_kwargs["reasoning_effort"] = reasoning_effort
22
+
23
+ if not supports_chat_template(tokenizer):
24
+ return generate_merged_prompt(messages, capabilities)
25
+
26
+ try:
27
+ return tokenizer.apply_chat_template(
28
+ messages,
29
+ add_generation_prompt=False,
30
+ tokenize=False,
31
+ **extra_kwargs,
32
+ )
33
+ except TypeError:
34
+ try:
35
+ fallback_kwargs = {"tools": tools} if tools is not None else {}
36
+ return tokenizer.apply_chat_template(
37
+ messages,
38
+ add_generation_prompt=False,
39
+ tokenize=False,
40
+ **fallback_kwargs,
41
+ )
42
+ except TypeError:
43
+ return tokenizer.apply_chat_template(
44
+ messages,
45
+ add_generation_prompt=False,
46
+ tokenize=False,
47
+ )
48
+
49
+
50
+ def handle_cache_prefill(
51
+ backend: ModelBackend,
52
+ capabilities: dict,
53
+ cache_path: str,
54
+ messages: list,
55
+ base_cache_path: str | None = None,
56
+ trim_to_tokens: int | None = None,
57
+ prefix_offsets: list[int] | None = None,
58
+ prefix_hashes: list[str] | None = None,
59
+ tools: list | None = None,
60
+ reasoning_effort: str | None = None,
61
+ images: list | None = None,
62
+ max_image_size: int = 768,
63
+ ) -> None:
64
+ """Build a process-local PyTorch KV cache from chat messages."""
65
+ if images:
66
+ raise ValueError("PyTorch LIP backend does not support vision input")
67
+ if base_cache_path is not None or trim_to_tokens is not None:
68
+ raise ValueError(
69
+ "PyTorch LIP backend does not support incremental prefill in Phase 1"
70
+ )
71
+ if prefix_offsets is not None or prefix_hashes is not None:
72
+ raise ValueError(
73
+ "PyTorch LIP backend does not support cache prefix metadata in Phase 1"
74
+ )
75
+
76
+ prompt = _render_prefill_prompt(
77
+ backend,
78
+ capabilities,
79
+ messages,
80
+ tools,
81
+ reasoning_effort,
82
+ )
83
+ result = backend.cache_prefill(
84
+ cache_path,
85
+ prompt,
86
+ base_cache_path=base_cache_path,
87
+ trim_to_tokens=trim_to_tokens,
88
+ prefix_offsets=prefix_offsets,
89
+ prefix_hashes=prefix_hashes,
90
+ images=images,
91
+ max_image_size=max_image_size,
92
+ )
93
+ print(json.dumps(result), end="\0", flush=True)
@@ -0,0 +1,53 @@
1
+ """Cancel request handling for in-flight streaming generation."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import select
7
+ import sys
8
+
9
+ _cancel_requested = False
10
+
11
+
12
+ def request_cancel() -> None:
13
+ global _cancel_requested
14
+ _cancel_requested = True
15
+
16
+
17
+ def reset_cancel() -> None:
18
+ global _cancel_requested
19
+ _cancel_requested = False
20
+
21
+
22
+ def is_cancel_requested() -> bool:
23
+ return _cancel_requested
24
+
25
+
26
+ def poll_cancel() -> bool:
27
+ """Non-blocking check for a cancel command on stdin during streaming."""
28
+ global _cancel_requested
29
+ if _cancel_requested:
30
+ return True
31
+
32
+ try:
33
+ ready, _, _ = select.select([sys.stdin], [], [], 0)
34
+ except (ValueError, OSError):
35
+ return False
36
+
37
+ if not ready:
38
+ return False
39
+
40
+ line = sys.stdin.readline()
41
+ if not line:
42
+ return False
43
+
44
+ try:
45
+ req = json.loads(line)
46
+ except json.JSONDecodeError:
47
+ return False
48
+
49
+ if req.get("method") == "cancel":
50
+ _cancel_requested = True
51
+ return True
52
+
53
+ return False
@@ -0,0 +1,6 @@
1
+ import json
2
+
3
+
4
+ def handle_capabilities(capabilities: dict) -> None:
5
+ """capabilities API の処理。JSON出力してnull文字で終端"""
6
+ print(json.dumps(capabilities), end="\0", flush=True)
@@ -0,0 +1,15 @@
1
+ from __future__ import annotations
2
+
3
+ from backends.base import ModelBackend
4
+ from handlers.generate import handle_generate
5
+
6
+
7
+ def handle_completion(
8
+ backend: ModelBackend,
9
+ prompt: str | list[int],
10
+ options: dict | None = None,
11
+ images: list | None = None,
12
+ max_image_size: int = 768,
13
+ ) -> None:
14
+ """completion API(後方互換)— generate に委譲"""
15
+ handle_generate(backend, prompt, options, images, max_image_size)
@@ -0,0 +1,70 @@
1
+ import json
2
+
3
+ from backends.base import ModelBackend
4
+ from utils.prompt_builder import generate_merged_prompt, supports_chat_template
5
+
6
+
7
+ def handle_format_test(
8
+ backend: ModelBackend,
9
+ capabilities: dict,
10
+ messages: list,
11
+ options: dict | None = None,
12
+ tools: list | None = None,
13
+ ) -> None:
14
+ """フォーマットテスト API の処理(実際に生成せずフォーマットのみ)"""
15
+ if options is None:
16
+ options = {}
17
+
18
+ tokenizer = backend.get_tokenizer()
19
+ result = {
20
+ "formatted_prompt": None,
21
+ "template_applied": False,
22
+ "model_specific_processing": None,
23
+ "error": None,
24
+ }
25
+
26
+ try:
27
+ if supports_chat_template(tokenizer):
28
+ result["model_specific_processing"] = messages
29
+
30
+ primer = options.get("primer")
31
+ add_generation_prompt = True
32
+ fmt_messages = list(messages)
33
+
34
+ if primer is not None:
35
+ fmt_messages.append({"role": "assistant", "content": primer})
36
+ add_generation_prompt = False
37
+
38
+ try:
39
+ formatted_prompt = tokenizer.apply_chat_template(
40
+ fmt_messages,
41
+ tools=tools,
42
+ add_generation_prompt=add_generation_prompt,
43
+ tokenize=False,
44
+ )
45
+ except TypeError:
46
+ formatted_prompt = tokenizer.apply_chat_template(
47
+ fmt_messages,
48
+ add_generation_prompt=add_generation_prompt,
49
+ tokenize=False,
50
+ )
51
+
52
+ if primer is not None:
53
+ formatted_prompt = (
54
+ primer.join(formatted_prompt.split(primer)[0:-1]) + primer
55
+ )
56
+
57
+ result["formatted_prompt"] = formatted_prompt
58
+ result["template_applied"] = True
59
+ else:
60
+ formatted_prompt = generate_merged_prompt(messages, capabilities)
61
+ primer = options.get("primer")
62
+ if primer is not None:
63
+ formatted_prompt += primer
64
+
65
+ result["formatted_prompt"] = formatted_prompt
66
+ result["template_applied"] = False
67
+ except Exception as e:
68
+ result["error"] = str(e)
69
+
70
+ print(json.dumps(result), end="\0", flush=True)