@modular-prompt/driver 0.16.0 → 0.17.1
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +98 -10
- package/dist/cache-controller.d.ts +4 -0
- package/dist/cache-controller.d.ts.map +1 -1
- package/dist/driver-registry/config-based-factory.d.ts +3 -0
- package/dist/driver-registry/config-based-factory.d.ts.map +1 -1
- package/dist/driver-registry/config-based-factory.js +8 -1
- package/dist/driver-registry/config-based-factory.js.map +1 -1
- package/dist/driver-registry/factory-helper.d.ts.map +1 -1
- package/dist/driver-registry/factory-helper.js +9 -2
- package/dist/driver-registry/factory-helper.js.map +1 -1
- package/dist/driver-registry/index.d.ts +1 -1
- package/dist/driver-registry/index.d.ts.map +1 -1
- package/dist/driver-registry/types.d.ts +15 -1
- package/dist/driver-registry/types.d.ts.map +1 -1
- package/dist/formatter/converter.d.ts.map +1 -1
- package/dist/formatter/converter.js +31 -2
- package/dist/formatter/converter.js.map +1 -1
- package/dist/google-genai/google-genai-driver.d.ts +1 -0
- package/dist/google-genai/google-genai-driver.d.ts.map +1 -1
- package/dist/google-genai/google-genai-driver.js +36 -27
- package/dist/google-genai/google-genai-driver.js.map +1 -1
- package/dist/index.d.ts +5 -3
- package/dist/index.d.ts.map +1 -1
- package/dist/index.js +5 -3
- package/dist/index.js.map +1 -1
- package/dist/local-inference/adapters.d.ts +6 -0
- package/dist/local-inference/adapters.d.ts.map +1 -1
- package/dist/local-inference/driver.d.ts.map +1 -1
- package/dist/local-inference/driver.js +45 -24
- package/dist/local-inference/driver.js.map +1 -1
- package/dist/local-inference/process-client.d.ts +4 -2
- package/dist/local-inference/process-client.d.ts.map +1 -1
- package/dist/local-inference/process-client.js +24 -8
- package/dist/local-inference/process-client.js.map +1 -1
- package/dist/local-inference/process-communication.d.ts +9 -2
- package/dist/local-inference/process-communication.d.ts.map +1 -1
- package/dist/local-inference/process-communication.js +37 -5
- package/dist/local-inference/process-communication.js.map +1 -1
- package/dist/local-inference/protocol.d.ts +4 -0
- package/dist/local-inference/protocol.d.ts.map +1 -1
- package/dist/local-inference/request-queue.d.ts +1 -1
- package/dist/local-inference/request-queue.d.ts.map +1 -1
- package/dist/local-inference/request-queue.js +26 -7
- package/dist/local-inference/request-queue.js.map +1 -1
- package/dist/local-inference/stream-utils.d.ts +6 -0
- package/dist/local-inference/stream-utils.d.ts.map +1 -1
- package/dist/local-inference/stream-utils.js.map +1 -1
- package/dist/mlx-ml/mlx-cache-controller.d.ts +9 -0
- package/dist/mlx-ml/mlx-cache-controller.d.ts.map +1 -1
- package/dist/mlx-ml/mlx-cache-controller.js +158 -33
- package/dist/mlx-ml/mlx-cache-controller.js.map +1 -1
- package/dist/mlx-ml/mlx-cache-support.d.ts +2 -2
- package/dist/mlx-ml/mlx-cache-support.d.ts.map +1 -1
- package/dist/mlx-ml/mlx-cache-support.js +8 -3
- package/dist/mlx-ml/mlx-cache-support.js.map +1 -1
- package/dist/mlx-ml/mlx-driver.d.ts +0 -1
- package/dist/mlx-ml/mlx-driver.d.ts.map +1 -1
- package/dist/mlx-ml/mlx-driver.js +1 -8
- package/dist/mlx-ml/mlx-driver.js.map +1 -1
- package/dist/mlx-ml/process/index.d.ts +1 -1
- package/dist/mlx-ml/process/index.d.ts.map +1 -1
- package/dist/mlx-ml/process/index.js +2 -2
- package/dist/mlx-ml/process/index.js.map +1 -1
- package/dist/models-config/index.d.ts +1 -1
- package/dist/models-config/index.d.ts.map +1 -1
- package/dist/models-config/index.js +1 -1
- package/dist/models-config/index.js.map +1 -1
- package/dist/models-config/resolve.d.ts +9 -1
- package/dist/models-config/resolve.d.ts.map +1 -1
- package/dist/models-config/resolve.js +94 -2
- package/dist/models-config/resolve.js.map +1 -1
- package/dist/models-config/types.d.ts +3 -1
- package/dist/models-config/types.d.ts.map +1 -1
- package/dist/pytorch/process/index.d.ts +4 -2
- package/dist/pytorch/process/index.d.ts.map +1 -1
- package/dist/pytorch/process/index.js +24 -7
- package/dist/pytorch/process/index.js.map +1 -1
- package/dist/pytorch/pytorch-cache-controller.d.ts +84 -0
- package/dist/pytorch/pytorch-cache-controller.d.ts.map +1 -0
- package/dist/pytorch/pytorch-cache-controller.js +742 -0
- package/dist/pytorch/pytorch-cache-controller.js.map +1 -0
- package/dist/pytorch/pytorch-cache-support.d.ts +23 -0
- package/dist/pytorch/pytorch-cache-support.d.ts.map +1 -0
- package/dist/pytorch/pytorch-cache-support.js +47 -0
- package/dist/pytorch/pytorch-cache-support.js.map +1 -0
- package/dist/pytorch/pytorch-driver.d.ts +8 -1
- package/dist/pytorch/pytorch-driver.d.ts.map +1 -1
- package/dist/pytorch/pytorch-driver.js +40 -0
- package/dist/pytorch/pytorch-driver.js.map +1 -1
- package/dist/runtime/check.d.ts.map +1 -1
- package/dist/runtime/check.js +8 -6
- package/dist/runtime/check.js.map +1 -1
- package/dist/runtime/index.d.ts +2 -2
- package/dist/runtime/index.d.ts.map +1 -1
- package/dist/runtime/index.js +2 -2
- package/dist/runtime/index.js.map +1 -1
- package/dist/runtime/manifest-core.d.mts +1 -0
- package/dist/runtime/manifest-core.mjs +1 -0
- package/dist/runtime/manifest-core.mjs.map +1 -1
- package/dist/runtime/manifest.d.ts +2 -0
- package/dist/runtime/manifest.d.ts.map +1 -1
- package/dist/runtime/manifest.js.map +1 -1
- package/dist/runtime/paths-core.d.mts +15 -1
- package/dist/runtime/paths-core.d.mts.map +1 -1
- package/dist/runtime/paths-core.mjs +50 -5
- package/dist/runtime/paths-core.mjs.map +1 -1
- package/dist/runtime/paths.d.ts +2 -2
- package/dist/runtime/paths.d.ts.map +1 -1
- package/dist/runtime/paths.js +2 -2
- package/dist/runtime/paths.js.map +1 -1
- package/dist/runtime/pytorch-template-core.d.mts +11 -0
- package/dist/runtime/pytorch-template-core.d.mts.map +1 -0
- package/dist/runtime/pytorch-template-core.mjs +54 -0
- package/dist/runtime/pytorch-template-core.mjs.map +1 -0
- package/dist/runtime/setup-commands-core.d.mts +3 -0
- package/dist/runtime/setup-commands-core.d.mts.map +1 -1
- package/dist/runtime/setup-commands-core.mjs +4 -0
- package/dist/runtime/setup-commands-core.mjs.map +1 -1
- package/dist/runtime/setup-commands.d.ts +1 -1
- package/dist/runtime/setup-commands.d.ts.map +1 -1
- package/dist/runtime/setup-commands.js +1 -1
- package/dist/runtime/setup-commands.js.map +1 -1
- package/dist/vertexai/vertexai-driver.d.ts +6 -0
- package/dist/vertexai/vertexai-driver.d.ts.map +1 -1
- package/dist/vertexai/vertexai-driver.js +106 -36
- package/dist/vertexai/vertexai-driver.js.map +1 -1
- package/docs/DRIVER_API.md +455 -0
- package/docs/LOCAL_MODEL_SETUP.md +765 -0
- package/docs/mlx-api-selection.md +301 -0
- package/package.json +10 -6
- package/scripts/runtime-cli.bin.test.ts +142 -0
- package/scripts/runtime-cli.js +305 -35
- package/scripts/runtime-cli.test.ts +163 -0
- package/skills/driver-usage/SKILL.md +29 -0
- package/src/mlx-ml/python/__main__.py +1 -1
- package/src/mlx-ml/python/backends/base.py +88 -18
- package/src/mlx-ml/python/backends/mlx_lm.py +28 -3
- package/src/mlx-ml/python/backends/mlx_vlm.py +679 -2
- package/src/mlx-ml/python/handlers/cache.py +4 -0
- package/src/mlx-ml/python/handlers/generate.py +33 -10
- package/src/mlx-ml/python/handlers/tokenize.py +1 -4
- package/src/mlx-ml/python/pyproject.toml +2 -2
- package/src/mlx-ml/python/server.py +2 -0
- package/src/mlx-ml/python/uv.lock +12 -12
- package/src/pytorch/templates/cpu-minimal/backends/base.py +139 -0
- package/src/pytorch/templates/cpu-minimal/backends/transformers_lm.py +1167 -0
- package/src/pytorch/{python → templates/cpu-minimal}/handlers/__init__.py +1 -0
- package/src/pytorch/templates/cpu-minimal/handlers/cache.py +88 -0
- package/src/pytorch/templates/cpu-minimal/handlers/generate.py +157 -0
- package/src/pytorch/{python → templates/cpu-minimal}/pyproject.toml +2 -2
- package/src/pytorch/{python → templates/cpu-minimal}/server.py +20 -2
- package/src/pytorch/templates/cpu-minimal/tests/test_cache_handler.py +284 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_capabilities.py +14 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_server.py +141 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_transformers_errors.py +89 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_transformers_lm_cache.py +554 -0
- package/src/pytorch/templates/cpu-minimal/utils/__init__.py +0 -0
- package/src/pytorch/{python → templates/cpu-minimal}/utils/token_utils.py +2 -2
- package/src/pytorch/templates/cpu-minimal/utils/transformers_errors.py +54 -0
- package/src/pytorch/{python → templates/cpu-minimal}/uv.lock +149 -109
- package/src/pytorch/templates/cuda/__main__.py +19 -0
- package/src/pytorch/templates/cuda/backends/__init__.py +3 -0
- package/src/pytorch/{python → templates/cuda}/backends/base.py +54 -6
- package/src/pytorch/templates/cuda/backends/transformers_lm.py +379 -0
- package/src/pytorch/templates/cuda/handlers/__init__.py +7 -0
- package/src/pytorch/templates/cuda/handlers/cache.py +93 -0
- package/src/pytorch/templates/cuda/handlers/cancel.py +53 -0
- package/src/pytorch/templates/cuda/handlers/capabilities.py +6 -0
- package/src/pytorch/templates/cuda/handlers/completion.py +15 -0
- package/src/pytorch/templates/cuda/handlers/format_test.py +70 -0
- package/src/pytorch/templates/cuda/handlers/generate.py +152 -0
- package/src/pytorch/templates/cuda/handlers/render.py +40 -0
- package/src/pytorch/templates/cuda/handlers/tokenize.py +63 -0
- package/src/pytorch/templates/cuda/pyproject.toml +37 -0
- package/src/pytorch/templates/cuda/server.py +158 -0
- package/src/pytorch/templates/cuda/tests/__init__.py +0 -0
- package/src/pytorch/templates/cuda/tests/test_cache_handler.py +207 -0
- package/src/pytorch/templates/cuda/tests/test_capabilities.py +14 -0
- package/src/pytorch/templates/cuda/tests/test_server.py +145 -0
- package/src/pytorch/templates/cuda/tests/test_transformers_errors.py +89 -0
- package/src/pytorch/templates/cuda/tests/test_transformers_lm_cache.py +288 -0
- package/src/pytorch/templates/cuda/utils/__init__.py +0 -0
- package/src/pytorch/templates/cuda/utils/chat_template_constraints.py +164 -0
- package/src/pytorch/templates/cuda/utils/prompt_builder.py +54 -0
- package/src/pytorch/templates/cuda/utils/template_render.py +80 -0
- package/src/pytorch/templates/cuda/utils/token_utils.py +376 -0
- package/src/pytorch/templates/cuda/utils/transformers_errors.py +54 -0
- package/src/pytorch/templates/cuda/uv.lock +734 -0
- package/src/pytorch/python/backends/transformers_lm.py +0 -127
- package/src/pytorch/python/handlers/generate.py +0 -68
- /package/src/pytorch/{python → templates/cpu-minimal}/__main__.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/backends/__init__.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/cancel.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/capabilities.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/completion.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/format_test.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/render.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/tokenize.py +0 -0
- /package/src/pytorch/{python/utils → templates/cpu-minimal/tests}/__init__.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/utils/chat_template_constraints.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/utils/prompt_builder.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/utils/template_render.py +0 -0
|
@@ -0,0 +1,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
|
|
File without changes
|
|
@@ -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
|
+
)
|