@modular-prompt/driver 0.15.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 (198) hide show
  1. package/README.md +124 -9
  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 +10 -3
  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 +10 -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 +18 -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 -32
  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 +2 -2
  60. package/dist/models-config/index.d.ts.map +1 -1
  61. package/dist/models-config/index.js +2 -2
  62. package/dist/models-config/index.js.map +1 -1
  63. package/dist/models-config/paths.d.ts +8 -0
  64. package/dist/models-config/paths.d.ts.map +1 -1
  65. package/dist/models-config/paths.js +16 -1
  66. package/dist/models-config/paths.js.map +1 -1
  67. package/dist/models-config/resolve.d.ts +10 -2
  68. package/dist/models-config/resolve.d.ts.map +1 -1
  69. package/dist/models-config/resolve.js +119 -6
  70. package/dist/models-config/resolve.js.map +1 -1
  71. package/dist/models-config/types.d.ts +5 -1
  72. package/dist/models-config/types.d.ts.map +1 -1
  73. package/dist/pytorch/process/index.d.ts +4 -2
  74. package/dist/pytorch/process/index.d.ts.map +1 -1
  75. package/dist/pytorch/process/index.js +24 -7
  76. package/dist/pytorch/process/index.js.map +1 -1
  77. package/dist/pytorch/pytorch-cache-controller.d.ts +84 -0
  78. package/dist/pytorch/pytorch-cache-controller.d.ts.map +1 -0
  79. package/dist/pytorch/pytorch-cache-controller.js +742 -0
  80. package/dist/pytorch/pytorch-cache-controller.js.map +1 -0
  81. package/dist/pytorch/pytorch-cache-support.d.ts +23 -0
  82. package/dist/pytorch/pytorch-cache-support.d.ts.map +1 -0
  83. package/dist/pytorch/pytorch-cache-support.js +47 -0
  84. package/dist/pytorch/pytorch-cache-support.js.map +1 -0
  85. package/dist/pytorch/pytorch-driver.d.ts +8 -1
  86. package/dist/pytorch/pytorch-driver.d.ts.map +1 -1
  87. package/dist/pytorch/pytorch-driver.js +40 -0
  88. package/dist/pytorch/pytorch-driver.js.map +1 -1
  89. package/dist/runtime/check.d.ts.map +1 -1
  90. package/dist/runtime/check.js +9 -6
  91. package/dist/runtime/check.js.map +1 -1
  92. package/dist/runtime/index.d.ts +2 -1
  93. package/dist/runtime/index.d.ts.map +1 -1
  94. package/dist/runtime/index.js +2 -1
  95. package/dist/runtime/index.js.map +1 -1
  96. package/dist/runtime/manifest-core.d.mts +1 -0
  97. package/dist/runtime/manifest-core.mjs +1 -0
  98. package/dist/runtime/manifest-core.mjs.map +1 -1
  99. package/dist/runtime/manifest.d.ts +2 -0
  100. package/dist/runtime/manifest.d.ts.map +1 -1
  101. package/dist/runtime/manifest.js.map +1 -1
  102. package/dist/runtime/paths-core.d.mts +15 -1
  103. package/dist/runtime/paths-core.d.mts.map +1 -1
  104. package/dist/runtime/paths-core.mjs +50 -5
  105. package/dist/runtime/paths-core.mjs.map +1 -1
  106. package/dist/runtime/paths.d.ts +2 -2
  107. package/dist/runtime/paths.d.ts.map +1 -1
  108. package/dist/runtime/paths.js +2 -2
  109. package/dist/runtime/paths.js.map +1 -1
  110. package/dist/runtime/pytorch-template-core.d.mts +11 -0
  111. package/dist/runtime/pytorch-template-core.d.mts.map +1 -0
  112. package/dist/runtime/pytorch-template-core.mjs +54 -0
  113. package/dist/runtime/pytorch-template-core.mjs.map +1 -0
  114. package/dist/runtime/setup-commands-core.d.mts +16 -0
  115. package/dist/runtime/setup-commands-core.d.mts.map +1 -0
  116. package/dist/runtime/setup-commands-core.mjs +18 -0
  117. package/dist/runtime/setup-commands-core.mjs.map +1 -0
  118. package/dist/runtime/setup-commands.d.ts +2 -0
  119. package/dist/runtime/setup-commands.d.ts.map +1 -0
  120. package/dist/runtime/setup-commands.js +2 -0
  121. package/dist/runtime/setup-commands.js.map +1 -0
  122. package/docs/DRIVER_API.md +455 -0
  123. package/docs/LOCAL_MODEL_SETUP.md +765 -0
  124. package/docs/mlx-api-selection.md +301 -0
  125. package/package.json +12 -5
  126. package/scripts/download-model.js +3 -2
  127. package/scripts/runtime-cli.bin.test.ts +142 -0
  128. package/scripts/runtime-cli.js +322 -47
  129. package/scripts/runtime-cli.test.ts +163 -0
  130. package/src/mlx-ml/python/__main__.py +1 -1
  131. package/src/mlx-ml/python/backends/base.py +88 -18
  132. package/src/mlx-ml/python/backends/cache_archive.py +41 -0
  133. package/src/mlx-ml/python/backends/mlx_lm.py +45 -4
  134. package/src/mlx-ml/python/backends/mlx_vlm.py +679 -2
  135. package/src/mlx-ml/python/handlers/cache.py +4 -0
  136. package/src/mlx-ml/python/handlers/generate.py +33 -10
  137. package/src/mlx-ml/python/handlers/tokenize.py +1 -4
  138. package/src/mlx-ml/python/pyproject.toml +9 -3
  139. package/src/mlx-ml/python/server.py +2 -0
  140. package/src/mlx-ml/python/uv.lock +193 -433
  141. package/src/pytorch/templates/cpu-minimal/backends/base.py +139 -0
  142. package/src/pytorch/templates/cpu-minimal/backends/transformers_lm.py +1167 -0
  143. package/src/pytorch/{python → templates/cpu-minimal}/handlers/__init__.py +1 -0
  144. package/src/pytorch/templates/cpu-minimal/handlers/cache.py +88 -0
  145. package/src/pytorch/templates/cpu-minimal/handlers/generate.py +157 -0
  146. package/src/pytorch/{python → templates/cpu-minimal}/pyproject.toml +2 -2
  147. package/src/pytorch/{python → templates/cpu-minimal}/server.py +20 -2
  148. package/src/pytorch/templates/cpu-minimal/tests/test_cache_handler.py +284 -0
  149. package/src/pytorch/templates/cpu-minimal/tests/test_capabilities.py +14 -0
  150. package/src/pytorch/templates/cpu-minimal/tests/test_server.py +141 -0
  151. package/src/pytorch/templates/cpu-minimal/tests/test_transformers_errors.py +89 -0
  152. package/src/pytorch/templates/cpu-minimal/tests/test_transformers_lm_cache.py +554 -0
  153. package/src/pytorch/templates/cpu-minimal/utils/__init__.py +0 -0
  154. package/src/pytorch/{python → templates/cpu-minimal}/utils/token_utils.py +2 -2
  155. package/src/pytorch/templates/cpu-minimal/utils/transformers_errors.py +54 -0
  156. package/src/pytorch/{python → templates/cpu-minimal}/uv.lock +149 -109
  157. package/src/pytorch/templates/cuda/__main__.py +19 -0
  158. package/src/pytorch/templates/cuda/backends/__init__.py +3 -0
  159. package/src/pytorch/{python → templates/cuda}/backends/base.py +54 -6
  160. package/src/pytorch/templates/cuda/backends/transformers_lm.py +379 -0
  161. package/src/pytorch/templates/cuda/handlers/__init__.py +7 -0
  162. package/src/pytorch/templates/cuda/handlers/cache.py +93 -0
  163. package/src/pytorch/templates/cuda/handlers/cancel.py +53 -0
  164. package/src/pytorch/templates/cuda/handlers/capabilities.py +6 -0
  165. package/src/pytorch/templates/cuda/handlers/completion.py +15 -0
  166. package/src/pytorch/templates/cuda/handlers/format_test.py +70 -0
  167. package/src/pytorch/templates/cuda/handlers/generate.py +152 -0
  168. package/src/pytorch/templates/cuda/handlers/render.py +40 -0
  169. package/src/pytorch/templates/cuda/handlers/tokenize.py +63 -0
  170. package/src/pytorch/templates/cuda/pyproject.toml +37 -0
  171. package/src/pytorch/templates/cuda/server.py +158 -0
  172. package/src/pytorch/templates/cuda/tests/__init__.py +0 -0
  173. package/src/pytorch/templates/cuda/tests/test_cache_handler.py +207 -0
  174. package/src/pytorch/templates/cuda/tests/test_capabilities.py +14 -0
  175. package/src/pytorch/templates/cuda/tests/test_server.py +145 -0
  176. package/src/pytorch/templates/cuda/tests/test_transformers_errors.py +89 -0
  177. package/src/pytorch/templates/cuda/tests/test_transformers_lm_cache.py +288 -0
  178. package/src/pytorch/templates/cuda/utils/__init__.py +0 -0
  179. package/src/pytorch/templates/cuda/utils/chat_template_constraints.py +164 -0
  180. package/src/pytorch/templates/cuda/utils/prompt_builder.py +54 -0
  181. package/src/pytorch/templates/cuda/utils/template_render.py +80 -0
  182. package/src/pytorch/templates/cuda/utils/token_utils.py +376 -0
  183. package/src/pytorch/templates/cuda/utils/transformers_errors.py +54 -0
  184. package/src/pytorch/templates/cuda/uv.lock +734 -0
  185. package/src/pytorch/python/backends/transformers_lm.py +0 -127
  186. package/src/pytorch/python/handlers/generate.py +0 -68
  187. /package/src/pytorch/{python → templates/cpu-minimal}/__main__.py +0 -0
  188. /package/src/pytorch/{python → templates/cpu-minimal}/backends/__init__.py +0 -0
  189. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/cancel.py +0 -0
  190. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/capabilities.py +0 -0
  191. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/completion.py +0 -0
  192. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/format_test.py +0 -0
  193. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/render.py +0 -0
  194. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/tokenize.py +0 -0
  195. /package/src/pytorch/{python/utils → templates/cpu-minimal/tests}/__init__.py +0 -0
  196. /package/src/pytorch/{python → templates/cpu-minimal}/utils/chat_template_constraints.py +0 -0
  197. /package/src/pytorch/{python → templates/cpu-minimal}/utils/prompt_builder.py +0 -0
  198. /package/src/pytorch/{python → templates/cpu-minimal}/utils/template_render.py +0 -0
@@ -0,0 +1,554 @@
1
+ import json
2
+ from types import SimpleNamespace
3
+
4
+ import pytest
5
+
6
+ torch = pytest.importorskip("torch")
7
+
8
+ from backends.transformers_lm import TransformersLmBackend
9
+ from handlers.generate import handle_generate
10
+ from tokenizers import Tokenizer
11
+ from tokenizers.models import WordLevel
12
+ from tokenizers.pre_tokenizers import Whitespace
13
+ from transformers import GPT2Config, GPT2LMHeadModel, PreTrainedTokenizerFast
14
+
15
+
16
+ class _Tokenizer:
17
+ bos_token = "<bos>"
18
+ pad_token = "<pad>"
19
+ eos_token = "<eos>"
20
+
21
+ def __init__(self):
22
+ self.prompts = {
23
+ "prefix": [1, 2],
24
+ "prefix suffix": [1, 2, 3],
25
+ "different suffix": [9, 8, 7],
26
+ "hello alpha": [1, 2],
27
+ "hello beta": [1, 3],
28
+ "hello beta gamma": [1, 3, 4],
29
+ }
30
+ self.token_text = {
31
+ 10: "a",
32
+ 11: "b",
33
+ 12: "c",
34
+ 13: "d",
35
+ }
36
+
37
+ def encode(self, prompt, add_special_tokens):
38
+ del add_special_tokens
39
+ return self.prompts[prompt]
40
+
41
+ def decode(self, token_ids, **kwargs):
42
+ del kwargs
43
+ if isinstance(token_ids, torch.Tensor):
44
+ token_ids = token_ids.flatten().tolist()
45
+ return "".join(self.token_text.get(int(token_id), "x") for token_id in token_ids)
46
+
47
+
48
+ class _Model:
49
+ def __init__(self, past_key_values, generated_texts=None, generated_token_batches=None):
50
+ self.past_key_values = past_key_values
51
+ self.generated_texts = generated_texts or ["generated"]
52
+ self.generated_token_batches = generated_token_batches
53
+ self.forward_calls = []
54
+ self.generate_calls = []
55
+
56
+ def __call__(self, **kwargs):
57
+ self.forward_calls.append(kwargs)
58
+ return SimpleNamespace(past_key_values=self.past_key_values)
59
+
60
+ def _emit_stream(self, streamer, input_ids, token_batches):
61
+ streamer.put(input_ids.cpu())
62
+ for token_batch in token_batches:
63
+ streamer.put(torch.tensor(token_batch, dtype=torch.long))
64
+ streamer.end()
65
+
66
+ def generate(self, **kwargs):
67
+ self.generate_calls.append(kwargs)
68
+ streamer = kwargs["streamer"]
69
+ if self.generated_token_batches is not None:
70
+ self._emit_stream(streamer, kwargs["input_ids"], self.generated_token_batches)
71
+ return
72
+ token_batches = [[10 + index] for index in range(len(self.generated_texts))]
73
+ self._emit_stream(streamer, kwargs["input_ids"], token_batches)
74
+
75
+
76
+ class _IncrementalModel(_Model):
77
+ def __call__(self, **kwargs):
78
+ self.forward_calls.append(kwargs)
79
+ input_ids = kwargs["input_ids"]
80
+ suffix_length = input_ids.shape[-1]
81
+ past_key_values = kwargs.get("past_key_values")
82
+ if past_key_values is None:
83
+ key = torch.zeros((1, 1, suffix_length, 4))
84
+ value = torch.zeros((1, 1, suffix_length, 4))
85
+ else:
86
+ key, value = past_key_values[0]
87
+ key_tail = torch.zeros(
88
+ (key.shape[0], key.shape[1], suffix_length, key.shape[3])
89
+ )
90
+ value_tail = torch.zeros(
91
+ (value.shape[0], value.shape[1], suffix_length, value.shape[3])
92
+ )
93
+ key = torch.cat((key, key_tail), dim=-2)
94
+ value = torch.cat((value, value_tail), dim=-2)
95
+ self.past_key_values = ((key, value),)
96
+ return SimpleNamespace(past_key_values=self.past_key_values)
97
+
98
+
99
+ class _CacheLayer:
100
+ def __init__(self, token_count):
101
+ self.keys = torch.zeros((1, 1, token_count, 4))
102
+ self.values = torch.zeros((1, 1, token_count, 4))
103
+
104
+ def get_seq_length(self):
105
+ return self.keys.shape[-2]
106
+
107
+
108
+ class _Cache:
109
+ def __init__(self, token_count):
110
+ self.layers = [_CacheLayer(token_count)]
111
+
112
+ def get_seq_length(self):
113
+ return self.layers[0].get_seq_length()
114
+
115
+
116
+ def _backend(generated_texts=None, generated_token_batches=None):
117
+ # Legacy tuple-shaped caches remain supported by Transformers and make the
118
+ # test independent of a specific Cache class implementation.
119
+ key = torch.zeros((1, 1, 2, 4))
120
+ value = torch.zeros((1, 1, 2, 4))
121
+ past_key_values = ((key, value),)
122
+
123
+ backend = TransformersLmBackend()
124
+ backend.tokenizer = _Tokenizer()
125
+ backend.model = _Model(past_key_values, generated_texts, generated_token_batches)
126
+ return backend, past_key_values
127
+
128
+
129
+ def _incremental_backend():
130
+ key = torch.zeros((1, 1, 2, 4))
131
+ value = torch.zeros((1, 1, 2, 4))
132
+ model = _IncrementalModel(((key, value),))
133
+ backend = TransformersLmBackend()
134
+ backend.tokenizer = _Tokenizer()
135
+ backend.model = model
136
+ return backend, model
137
+
138
+
139
+ def _tiny_gpt2_backend():
140
+ vocab = {
141
+ "<pad>": 0,
142
+ "<unk>": 1,
143
+ "hello": 2,
144
+ "alpha": 3,
145
+ "beta": 4,
146
+ "gamma": 5,
147
+ }
148
+ tokenizer = Tokenizer(WordLevel(vocab=vocab, unk_token="<unk>"))
149
+ tokenizer.pre_tokenizer = Whitespace()
150
+ fast_tokenizer = PreTrainedTokenizerFast(
151
+ tokenizer_object=tokenizer,
152
+ unk_token="<unk>",
153
+ pad_token="<pad>",
154
+ )
155
+ config = GPT2Config(
156
+ vocab_size=len(vocab),
157
+ n_positions=32,
158
+ n_ctx=32,
159
+ n_embd=16,
160
+ n_layer=1,
161
+ n_head=1,
162
+ pad_token_id=vocab["<pad>"],
163
+ eos_token_id=None,
164
+ )
165
+ model = GPT2LMHeadModel(config)
166
+ # Make greedy output deterministic and non-special so the streamer emits
167
+ # several chunks without downloading a model in CI.
168
+ model.lm_head = torch.nn.Linear(config.n_embd, config.vocab_size, bias=True)
169
+ with torch.no_grad():
170
+ model.lm_head.weight.zero_()
171
+ model.lm_head.bias.fill_(-1)
172
+ model.lm_head.bias[vocab["alpha"]] = 1
173
+ model.eval()
174
+
175
+ backend = TransformersLmBackend()
176
+ backend.tokenizer = fast_tokenizer
177
+ backend.model = model
178
+ return backend
179
+
180
+
181
+ def test_cache_prefill_keeps_past_key_values_in_process():
182
+ backend, past_key_values = _backend()
183
+
184
+ result = backend.cache_prefill("memory://prefix", "prefix")
185
+
186
+ assert result == {
187
+ "cache_path": "memory://prefix",
188
+ "token_count": 2,
189
+ "cache_write_tokens": 2,
190
+ }
191
+ assert backend.load_cache_from_file("memory://prefix") is past_key_values
192
+ assert backend.consume_cache_write_tokens("memory://prefix") == 2
193
+ assert backend.consume_cache_write_tokens("memory://prefix") == 0
194
+ call = backend.model.forward_calls[0]
195
+ assert call["input_ids"].tolist() == [[1, 2]]
196
+ assert call["input_ids"].device == backend._device
197
+ assert call["use_cache"] is True
198
+
199
+
200
+ def test_cache_prefill_persists_cache_and_metadata(tmp_path):
201
+ backend, past_key_values = _backend()
202
+ cache_path = tmp_path / "prefix.pytorch-cache"
203
+
204
+ result = backend.cache_prefill(
205
+ str(cache_path),
206
+ "prefix",
207
+ prefix_offsets=[2],
208
+ prefix_hashes=["hash-prefix"],
209
+ )
210
+
211
+ assert result == {
212
+ "cache_path": str(cache_path),
213
+ "token_count": 2,
214
+ "cache_write_tokens": 2,
215
+ }
216
+ assert cache_path.is_file()
217
+ meta = json.loads(cache_path.with_name(cache_path.name + ".meta.json").read_text())
218
+ assert meta == {
219
+ "layout": "pytorch_kv_v1",
220
+ "token_count": 2,
221
+ "prefix_offsets": [2],
222
+ "prefix_hashes": ["hash-prefix"],
223
+ "model_id": "unknown",
224
+ "dtype": "float32",
225
+ "device": "cpu",
226
+ }
227
+
228
+ restarted_backend, _ = _backend()
229
+ loaded = restarted_backend.load_cache_from_file(
230
+ str(cache_path),
231
+ prompt="prefix suffix",
232
+ )
233
+ assert loaded is not None
234
+ assert loaded is not past_key_values
235
+ assert restarted_backend.get_cache_offset(loaded) == 2
236
+ assert torch.equal(loaded[0][0], past_key_values[0][0])
237
+
238
+
239
+ def test_disk_cache_load_rejects_model_metadata_mismatch(tmp_path):
240
+ backend, _ = _backend()
241
+ cache_path = tmp_path / "prefix.pytorch-cache"
242
+ backend.cache_prefill(str(cache_path), "prefix")
243
+
244
+ restarted_backend, _ = _backend()
245
+ restarted_backend._model_id = "different-model"
246
+
247
+ assert restarted_backend.load_cache_from_file(str(cache_path), prompt="prefix") is None
248
+
249
+
250
+ def test_disk_cache_loads_after_restart_and_generates_with_trim(tmp_path, capsys):
251
+ cache_path = tmp_path / "prefix.pytorch-cache"
252
+ backend = _tiny_gpt2_backend()
253
+ backend.cache_prefill(str(cache_path), "hello alpha")
254
+
255
+ restarted_backend = _tiny_gpt2_backend()
256
+ handle_generate(
257
+ restarted_backend,
258
+ "hello alpha beta",
259
+ options={"max_tokens": 1, "temperature": 0},
260
+ cache_path=str(cache_path),
261
+ cache_trim_tokens=1,
262
+ )
263
+
264
+ output = capsys.readouterr().out
265
+ meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
266
+ assert meta["prompt_tokens"] == 3
267
+ assert meta["generation_tokens"] == 1
268
+ assert meta["cache_read_tokens"] == 1
269
+ assert meta["cache_loaded"] is True
270
+
271
+
272
+ def test_incremental_prefill_loads_base_trims_and_prefills_suffix(tmp_path):
273
+ base_backend, _ = _backend()
274
+ base_path = tmp_path / "base.pytorch-cache"
275
+ base_backend.cache_prefill(str(base_path), "prefix")
276
+
277
+ backend, model = _incremental_backend()
278
+ cache_path = tmp_path / "extended.pytorch-cache"
279
+ result = backend.cache_prefill(
280
+ str(cache_path),
281
+ "prefix suffix",
282
+ base_cache_path=str(base_path),
283
+ trim_to_tokens=2,
284
+ prefix_offsets=[2, 3],
285
+ prefix_hashes=["hash-prefix", "hash-full"],
286
+ )
287
+
288
+ assert result == {
289
+ "cache_path": str(cache_path),
290
+ "token_count": 3,
291
+ "cache_write_tokens": 1,
292
+ }
293
+ call = model.forward_calls[0]
294
+ assert call["input_ids"].tolist() == [[3]]
295
+ assert call["past_key_values"][0][0].shape[-2] == 2
296
+ assert call["cache_position"].tolist() == [2]
297
+ assert call["attention_mask"].tolist() == [[1, 1, 1]]
298
+ assert backend.get_cache_offset(backend._caches[str(cache_path)]) == 3
299
+
300
+ meta = json.loads(cache_path.with_name(cache_path.name + ".meta.json").read_text())
301
+ assert meta["token_count"] == 3
302
+ assert meta["prefix_offsets"] == [2, 3]
303
+ assert meta["prefix_hashes"] == ["hash-prefix", "hash-full"]
304
+
305
+
306
+ def test_incremental_prefill_validates_only_trimmed_prefix(tmp_path, capsys):
307
+ base_backend, _ = _backend()
308
+ base_path = tmp_path / "base.pytorch-cache"
309
+ base_backend.cache_prefill(str(base_path), "hello alpha")
310
+
311
+ backend, model = _incremental_backend()
312
+ cache_path = tmp_path / "diverged.pytorch-cache"
313
+ result = backend.cache_prefill(
314
+ str(cache_path),
315
+ "hello beta",
316
+ base_cache_path=str(base_path),
317
+ trim_to_tokens=1,
318
+ )
319
+
320
+ assert result["token_count"] == 2
321
+ assert result["cache_write_tokens"] == 1
322
+ assert model.forward_calls[0]["input_ids"].tolist() == [[3]]
323
+ assert model.forward_calls[0]["past_key_values"][0][0].shape[-2] == 1
324
+
325
+ handle_generate(
326
+ backend,
327
+ "hello beta gamma",
328
+ options={"max_tokens": 1},
329
+ cache_path=str(cache_path),
330
+ )
331
+ output = capsys.readouterr().out
332
+ meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
333
+ assert meta["cache_loaded"] is True
334
+ assert meta["cache_read_tokens"] == 2
335
+
336
+
337
+ def test_generate_trim_does_not_mutate_reusable_disk_cache(tmp_path, capsys):
338
+ cache_path = tmp_path / "prefix.pytorch-cache"
339
+ backend = _tiny_gpt2_backend()
340
+ backend.cache_prefill(str(cache_path), "hello alpha")
341
+ original_cache = backend.load_cache_from_file(
342
+ str(cache_path),
343
+ prompt="hello alpha beta",
344
+ )
345
+ assert original_cache is not None
346
+ assert backend.get_cache_offset(original_cache) == 2
347
+
348
+ handle_generate(
349
+ backend,
350
+ "hello alpha beta",
351
+ options={"max_tokens": 1, "temperature": 0},
352
+ cache_path=str(cache_path),
353
+ cache_trim_tokens=1,
354
+ )
355
+ first_output = capsys.readouterr().out
356
+ first_meta = json.loads(
357
+ first_output.split("\x1e__META__:", 1)[1].split("\0", 1)[0]
358
+ )
359
+ assert first_meta["cache_loaded"] is True
360
+ assert first_meta["cache_read_tokens"] == 1
361
+ assert backend.get_cache_offset(original_cache) == 2
362
+
363
+ handle_generate(
364
+ backend,
365
+ "hello alpha beta",
366
+ options={"max_tokens": 1, "temperature": 0},
367
+ cache_path=str(cache_path),
368
+ )
369
+ second_output = capsys.readouterr().out
370
+ second_meta = json.loads(
371
+ second_output.split("\x1e__META__:", 1)[1].split("\0", 1)[0]
372
+ )
373
+ assert second_meta["cache_loaded"] is True
374
+ assert second_meta["cache_read_tokens"] == 2
375
+ assert backend.get_cache_offset(original_cache) == 2
376
+
377
+
378
+ def test_generate_trim_validates_only_trimmed_prefix(tmp_path, capsys):
379
+ base_path = tmp_path / "base.pytorch-cache"
380
+ base_backend = _tiny_gpt2_backend()
381
+ base_backend.cache_prefill(str(base_path), "hello alpha")
382
+
383
+ backend = _tiny_gpt2_backend()
384
+ handle_generate(
385
+ backend,
386
+ "hello beta gamma",
387
+ options={"max_tokens": 1, "temperature": 0},
388
+ cache_path=str(base_path),
389
+ cache_trim_tokens=1,
390
+ )
391
+
392
+ output = capsys.readouterr().out
393
+ meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
394
+ assert meta["cache_loaded"] is True
395
+ assert meta["cache_read_tokens"] == 1
396
+
397
+
398
+ def test_trim_cache_supports_legacy_tuple_and_cache_objects():
399
+ backend, past_key_values = _backend()
400
+
401
+ trimmed_tuple = backend.trim_cache(past_key_values, 1)
402
+ assert backend.get_cache_offset(trimmed_tuple) == 1
403
+ assert trimmed_tuple[0][0].shape[-2] == 1
404
+ assert backend.get_cache_offset(past_key_values) == 2
405
+
406
+ cache = _Cache(3)
407
+ trimmed_cache = backend.trim_cache(cache, 1)
408
+ assert trimmed_cache is not cache
409
+ assert backend.get_cache_offset(trimmed_cache) == 2
410
+ assert backend.get_cache_offset(cache) == 3
411
+ assert trimmed_cache.layers[0].keys.shape[-2] == 2
412
+ assert cache.layers[0].keys.shape[-2] == 3
413
+
414
+
415
+ def test_load_cache_validates_prompt_prefix_without_reading_a_file():
416
+ backend, past_key_values = _backend()
417
+ backend.cache_prefill("memory://prefix", "prefix")
418
+
419
+ assert backend.load_cache_from_file(
420
+ "memory://prefix", prompt="prefix suffix"
421
+ ) is past_key_values
422
+ assert backend.load_cache_from_file(
423
+ "memory://prefix", prompt="different suffix"
424
+ ) is None
425
+ assert backend.load_cache_from_file("/tmp/not-created.cache") is None
426
+
427
+
428
+ def test_stream_generate_passes_cached_prefix_and_reports_usage():
429
+ backend, past_key_values = _backend()
430
+ backend.cache_prefill("memory://prefix", "prefix")
431
+
432
+ chunks = list(
433
+ backend.stream_generate(
434
+ [3],
435
+ {"max_tokens": 1, "temperature": 0},
436
+ prompt_cache=past_key_values,
437
+ )
438
+ )
439
+
440
+ assert chunks[-1].text == "a"
441
+ assert chunks[-1].generation_tokens == 1
442
+ call = backend.model.generate_calls[0]
443
+ assert call["input_ids"].tolist() == [[3]]
444
+ assert call["past_key_values"] is not past_key_values
445
+ assert torch.equal(call["past_key_values"][0][0], past_key_values[0][0])
446
+ assert call["cache_position"].tolist() == [2]
447
+ assert call["attention_mask"].tolist() == [[1, 1, 1]]
448
+ first_meta_chunk = next(chunk for chunk in chunks if chunk.prompt_tokens is not None)
449
+ assert first_meta_chunk.prompt_tokens == 3
450
+ assert first_meta_chunk.cache_read_tokens == 2
451
+
452
+
453
+ def test_stream_generate_reports_cumulative_generation_tokens_for_multiple_chunks():
454
+ backend, past_key_values = _backend(["a", "b", "c"])
455
+ backend.cache_prefill("memory://prefix", "prefix")
456
+
457
+ chunks = list(
458
+ backend.stream_generate(
459
+ [3],
460
+ {"max_tokens": 3, "temperature": 0},
461
+ prompt_cache=past_key_values,
462
+ )
463
+ )
464
+
465
+ assert chunks[-1].text == "abc"
466
+ assert chunks[-1].generation_tokens == 3
467
+ first_meta_chunk = next(chunk for chunk in chunks if chunk.prompt_tokens is not None)
468
+ assert first_meta_chunk.prompt_tokens == 3
469
+ assert first_meta_chunk.cache_read_tokens == 2
470
+
471
+
472
+ def test_stream_generate_counts_token_ids_not_empty_text_chunks(capsys):
473
+ backend, _ = _backend(generated_token_batches=[[10], [11], [12], [13]])
474
+
475
+ chunks = list(
476
+ backend.stream_generate(
477
+ "prefix suffix",
478
+ {"max_tokens": 4, "temperature": 0},
479
+ )
480
+ )
481
+
482
+ # TextIteratorStreamer emits an empty chunk while buffering each token and
483
+ # one final text chunk, so chunk count is five for four generated tokens.
484
+ assert len(chunks) == 5
485
+ assert chunks[0].text == ""
486
+ assert chunks[-1].text == "abcd"
487
+ assert chunks[-1].generation_tokens == 4
488
+
489
+ handle_generate(
490
+ backend,
491
+ "prefix suffix",
492
+ options={"max_tokens": 4, "temperature": 0},
493
+ )
494
+ output = capsys.readouterr().out
495
+ meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
496
+ assert meta["prompt_tokens"] == 3
497
+ assert meta["generation_tokens"] == 4
498
+
499
+
500
+ def test_generate_handler_reports_backend_prefill_write_usage(capsys):
501
+ backend, _ = _backend()
502
+ backend.cache_prefill("memory://prefix", "prefix")
503
+
504
+ handle_generate(
505
+ backend,
506
+ "prefix suffix",
507
+ options={"max_tokens": 1, "temperature": 0},
508
+ cache_path="memory://prefix",
509
+ )
510
+
511
+ output = capsys.readouterr().out
512
+ assert output.startswith("a")
513
+ meta = output.split("\x1e__META__:", 1)[1].split("\0", 1)[0]
514
+ assert json.loads(meta) == {
515
+ "prompt_tokens": 3,
516
+ "generation_tokens": 1,
517
+ "cache_read_tokens": 2,
518
+ "cache_write_tokens": 2,
519
+ "cache_loaded": True,
520
+ }
521
+
522
+
523
+ def test_generate_handler_preserves_usage_for_tiny_gpt2_multiple_chunks(capsys):
524
+ backend = _tiny_gpt2_backend()
525
+ options = {"max_tokens": 4, "temperature": 0}
526
+
527
+ expected_chunks = list(backend.stream_generate("hello", options))
528
+ # The real TextIteratorStreamer buffers the generated word, so the first
529
+ # text chunk is empty and the final chunk flushes all four generated
530
+ # tokens. Token usage must not be derived from this five-chunk sequence.
531
+ assert len(expected_chunks) == 5
532
+ assert expected_chunks[0].text == ""
533
+ expected_generation_tokens = expected_chunks[-1].generation_tokens
534
+ assert expected_generation_tokens == 4
535
+
536
+ handle_generate(backend, "hello", options=options)
537
+
538
+ output = capsys.readouterr().out
539
+ meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
540
+ assert meta["prompt_tokens"] == len(backend.tokenize_prompt("hello"))
541
+ assert meta["generation_tokens"] == expected_generation_tokens
542
+
543
+
544
+ def test_cache_prefill_uses_cold_path_when_base_cache_is_missing():
545
+ backend, _ = _backend()
546
+
547
+ result = backend.cache_prefill(
548
+ "memory://prefix",
549
+ "prefix",
550
+ base_cache_path="memory://base",
551
+ )
552
+
553
+ assert result["token_count"] == 2
554
+ assert result["cache_write_tokens"] == 2
@@ -342,7 +342,7 @@ def get_capabilities(tokenizer):
342
342
  dict: capabilities情報
343
343
  """
344
344
  # 基本メソッド
345
- methods = ["capabilities", "completion", "generate", "format_test"]
345
+ methods = ["capabilities", "completion", "generate", "format_test", "cache_prefill"]
346
346
 
347
347
  # apply_chat_templateがある場合はchatメソッドを追加
348
348
  if hasattr(tokenizer, 'apply_chat_template'):
@@ -373,4 +373,4 @@ def get_capabilities(tokenizer):
373
373
  if chat_restrictions:
374
374
  capabilities["chat_restrictions"] = chat_restrictions
375
375
 
376
- return capabilities
376
+ return capabilities
@@ -0,0 +1,54 @@
1
+ """Helpful errors for model architectures not supported by the runtime."""
2
+
3
+ import re
4
+
5
+
6
+ MIN_TRANSFORMERS_VERSION = "5.14.0"
7
+
8
+ _MODEL_TYPE_MESSAGE = re.compile(
9
+ r"model type [`'\"](?P<model_type>[A-Za-z0-9_.-]+)[`'\"]",
10
+ re.IGNORECASE,
11
+ )
12
+ _MODEL_TYPE_KEY = re.compile(r"[A-Za-z][A-Za-z0-9_.-]*")
13
+
14
+
15
+ def extract_unsupported_model_type(error: BaseException) -> str | None:
16
+ """Extract a model type from Transformers' unknown-architecture errors.
17
+
18
+ Transformers normally raises ``ValueError`` with a message containing the
19
+ model type. Older releases can leak the registry ``KeyError`` instead,
20
+ so handle that form as well while leaving unrelated errors untouched.
21
+ """
22
+
23
+ match = _MODEL_TYPE_MESSAGE.search(str(error))
24
+ if match:
25
+ return match.group("model_type")
26
+
27
+ if isinstance(error, KeyError) and len(error.args) == 1:
28
+ candidate = error.args[0]
29
+ if (
30
+ isinstance(candidate, str)
31
+ and candidate != "model_type"
32
+ and _MODEL_TYPE_KEY.fullmatch(candidate)
33
+ ):
34
+ return candidate
35
+
36
+ return None
37
+
38
+
39
+ def unsupported_model_type_error(
40
+ model_name: str,
41
+ model_type: str,
42
+ transformers_version: str,
43
+ ) -> RuntimeError:
44
+ """Build the actionable load error shown to PyTorch runtime users."""
45
+
46
+ return RuntimeError(
47
+ f"Cannot load model '{model_name}': Transformers does not recognize "
48
+ f"model_type '{model_type}'. The PyTorch runtime is using transformers "
49
+ f"{transformers_version}, but this model requires transformers>="
50
+ f"{MIN_TRANSFORMERS_VERSION}. Run `setup-pytorch` again after updating "
51
+ "@modular-prompt/driver. If the runtime already has a Python project, "
52
+ "update its transformers requirement and run "
53
+ "`modular-prompt-runtime sync pytorch`."
54
+ )