@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.
- package/README.md +124 -9
- package/dist/cache-controller.d.ts +4 -0
- package/dist/cache-controller.d.ts.map +1 -1
- package/dist/driver-registry/config-based-factory.d.ts.map +1 -1
- package/dist/driver-registry/config-based-factory.js +10 -3
- 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 +10 -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 +18 -1
- package/dist/driver-registry/types.d.ts.map +1 -1
- package/dist/formatter/converter.d.ts.map +1 -1
- package/dist/formatter/converter.js +31 -2
- package/dist/formatter/converter.js.map +1 -1
- package/dist/index.d.ts +5 -3
- package/dist/index.d.ts.map +1 -1
- package/dist/index.js +5 -3
- package/dist/index.js.map +1 -1
- package/dist/local-inference/adapters.d.ts +6 -0
- package/dist/local-inference/adapters.d.ts.map +1 -1
- package/dist/local-inference/driver.d.ts.map +1 -1
- package/dist/local-inference/driver.js +45 -24
- package/dist/local-inference/driver.js.map +1 -1
- package/dist/local-inference/process-client.d.ts +4 -2
- package/dist/local-inference/process-client.d.ts.map +1 -1
- package/dist/local-inference/process-client.js +24 -8
- package/dist/local-inference/process-client.js.map +1 -1
- package/dist/local-inference/process-communication.d.ts +9 -2
- package/dist/local-inference/process-communication.d.ts.map +1 -1
- package/dist/local-inference/process-communication.js +37 -5
- package/dist/local-inference/process-communication.js.map +1 -1
- package/dist/local-inference/protocol.d.ts +4 -0
- package/dist/local-inference/protocol.d.ts.map +1 -1
- package/dist/local-inference/request-queue.d.ts +1 -1
- package/dist/local-inference/request-queue.d.ts.map +1 -1
- package/dist/local-inference/request-queue.js +26 -7
- package/dist/local-inference/request-queue.js.map +1 -1
- package/dist/local-inference/stream-utils.d.ts +6 -0
- package/dist/local-inference/stream-utils.d.ts.map +1 -1
- package/dist/local-inference/stream-utils.js.map +1 -1
- package/dist/mlx-ml/mlx-cache-controller.d.ts +9 -0
- package/dist/mlx-ml/mlx-cache-controller.d.ts.map +1 -1
- package/dist/mlx-ml/mlx-cache-controller.js +158 -32
- 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 +2 -2
- package/dist/models-config/index.d.ts.map +1 -1
- package/dist/models-config/index.js +2 -2
- package/dist/models-config/index.js.map +1 -1
- package/dist/models-config/paths.d.ts +8 -0
- package/dist/models-config/paths.d.ts.map +1 -1
- package/dist/models-config/paths.js +16 -1
- package/dist/models-config/paths.js.map +1 -1
- package/dist/models-config/resolve.d.ts +10 -2
- package/dist/models-config/resolve.d.ts.map +1 -1
- package/dist/models-config/resolve.js +119 -6
- package/dist/models-config/resolve.js.map +1 -1
- package/dist/models-config/types.d.ts +5 -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 +9 -6
- package/dist/runtime/check.js.map +1 -1
- package/dist/runtime/index.d.ts +2 -1
- package/dist/runtime/index.d.ts.map +1 -1
- package/dist/runtime/index.js +2 -1
- 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 +16 -0
- package/dist/runtime/setup-commands-core.d.mts.map +1 -0
- package/dist/runtime/setup-commands-core.mjs +18 -0
- package/dist/runtime/setup-commands-core.mjs.map +1 -0
- package/dist/runtime/setup-commands.d.ts +2 -0
- package/dist/runtime/setup-commands.d.ts.map +1 -0
- package/dist/runtime/setup-commands.js +2 -0
- package/dist/runtime/setup-commands.js.map +1 -0
- 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 +12 -5
- package/scripts/download-model.js +3 -2
- package/scripts/runtime-cli.bin.test.ts +142 -0
- package/scripts/runtime-cli.js +322 -47
- package/scripts/runtime-cli.test.ts +163 -0
- package/src/mlx-ml/python/__main__.py +1 -1
- package/src/mlx-ml/python/backends/base.py +88 -18
- package/src/mlx-ml/python/backends/cache_archive.py +41 -0
- package/src/mlx-ml/python/backends/mlx_lm.py +45 -4
- 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 +9 -3
- package/src/mlx-ml/python/server.py +2 -0
- package/src/mlx-ml/python/uv.lock +193 -433
- 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,88 @@
|
|
|
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 persistent or process-local PyTorch KV cache from chat messages."""
|
|
65
|
+
if images:
|
|
66
|
+
raise ValueError("PyTorch LIP backend does not support vision input")
|
|
67
|
+
|
|
68
|
+
prompt = _render_prefill_prompt(
|
|
69
|
+
backend,
|
|
70
|
+
capabilities,
|
|
71
|
+
messages,
|
|
72
|
+
tools,
|
|
73
|
+
reasoning_effort,
|
|
74
|
+
)
|
|
75
|
+
result = backend.cache_prefill(
|
|
76
|
+
cache_path,
|
|
77
|
+
prompt,
|
|
78
|
+
base_cache_path=base_cache_path,
|
|
79
|
+
trim_to_tokens=trim_to_tokens,
|
|
80
|
+
prefix_offsets=prefix_offsets,
|
|
81
|
+
prefix_hashes=prefix_hashes,
|
|
82
|
+
images=images,
|
|
83
|
+
max_image_size=max_image_size,
|
|
84
|
+
)
|
|
85
|
+
if prefix_offsets is not None and prefix_hashes is not None:
|
|
86
|
+
result["prefix_offsets"] = prefix_offsets
|
|
87
|
+
result["prefix_hashes"] = prefix_hashes
|
|
88
|
+
print(json.dumps(result), end="\0", flush=True)
|
|
@@ -0,0 +1,157 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
|
|
5
|
+
from backends.base import ModelBackend
|
|
6
|
+
from handlers.cancel import poll_cancel
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def _stream_to_stdout(
|
|
10
|
+
backend: ModelBackend,
|
|
11
|
+
prompt: str | list[int],
|
|
12
|
+
options: dict,
|
|
13
|
+
images: list | None = None,
|
|
14
|
+
primer: str | None = None,
|
|
15
|
+
prompt_cache=None,
|
|
16
|
+
cache_loaded: bool | None = None,
|
|
17
|
+
cache_read_tokens: int = 0,
|
|
18
|
+
cache_write_tokens: int = 0,
|
|
19
|
+
) -> None:
|
|
20
|
+
if images:
|
|
21
|
+
raise ValueError("PyTorch LIP backend does not support vision input")
|
|
22
|
+
|
|
23
|
+
if primer is not None:
|
|
24
|
+
print(primer, end="", flush=True)
|
|
25
|
+
|
|
26
|
+
first_prompt_tokens = None
|
|
27
|
+
first_cache_read_tokens = None
|
|
28
|
+
first_cache_write_tokens = None
|
|
29
|
+
reported_generation_tokens = None
|
|
30
|
+
for response in backend.stream_generate(
|
|
31
|
+
prompt,
|
|
32
|
+
options,
|
|
33
|
+
images,
|
|
34
|
+
prompt_cache=prompt_cache,
|
|
35
|
+
):
|
|
36
|
+
if poll_cancel():
|
|
37
|
+
break
|
|
38
|
+
response_prompt_tokens = getattr(response, "prompt_tokens", None)
|
|
39
|
+
if first_prompt_tokens is None and response_prompt_tokens is not None:
|
|
40
|
+
first_prompt_tokens = response_prompt_tokens
|
|
41
|
+
if (
|
|
42
|
+
first_cache_read_tokens is None
|
|
43
|
+
and getattr(response, "cache_read_tokens", None) is not None
|
|
44
|
+
):
|
|
45
|
+
first_cache_read_tokens = response.cache_read_tokens
|
|
46
|
+
if (
|
|
47
|
+
first_cache_write_tokens is None
|
|
48
|
+
and getattr(response, "cache_write_tokens", None) is not None
|
|
49
|
+
):
|
|
50
|
+
first_cache_write_tokens = response.cache_write_tokens
|
|
51
|
+
response_generation_tokens = getattr(response, "generation_tokens", None)
|
|
52
|
+
if response_generation_tokens is not None:
|
|
53
|
+
reported_generation_tokens = max(
|
|
54
|
+
reported_generation_tokens or 0,
|
|
55
|
+
int(response_generation_tokens),
|
|
56
|
+
)
|
|
57
|
+
print(response.text.replace("\0", "").replace("\x1e", ""), end="", flush=True)
|
|
58
|
+
|
|
59
|
+
meta: dict = {}
|
|
60
|
+
if first_prompt_tokens is not None:
|
|
61
|
+
meta["prompt_tokens"] = first_prompt_tokens
|
|
62
|
+
if reported_generation_tokens is not None:
|
|
63
|
+
meta["generation_tokens"] = max(0, reported_generation_tokens)
|
|
64
|
+
if first_cache_read_tokens is not None:
|
|
65
|
+
meta["cache_read_tokens"] = first_cache_read_tokens
|
|
66
|
+
if first_cache_write_tokens is not None:
|
|
67
|
+
meta["cache_write_tokens"] = first_cache_write_tokens
|
|
68
|
+
if cache_read_tokens > 0 and "cache_read_tokens" not in meta:
|
|
69
|
+
meta["cache_read_tokens"] = cache_read_tokens
|
|
70
|
+
if cache_write_tokens > 0 and "cache_write_tokens" not in meta:
|
|
71
|
+
meta["cache_write_tokens"] = cache_write_tokens
|
|
72
|
+
if cache_loaded is not None:
|
|
73
|
+
meta["cache_loaded"] = cache_loaded
|
|
74
|
+
|
|
75
|
+
if meta:
|
|
76
|
+
print(f"\x1e__META__:{json.dumps(meta)}", end="\0", flush=True)
|
|
77
|
+
else:
|
|
78
|
+
print("", end="\0", flush=True)
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def handle_generate(
|
|
82
|
+
backend: ModelBackend,
|
|
83
|
+
prompt: str | list[int],
|
|
84
|
+
options: dict | None = None,
|
|
85
|
+
images: list | None = None,
|
|
86
|
+
max_image_size: int = 768,
|
|
87
|
+
primer: str | None = None,
|
|
88
|
+
cache_path: str | None = None,
|
|
89
|
+
cache_trim_tokens: int | None = None,
|
|
90
|
+
) -> None:
|
|
91
|
+
"""LIP generate: 整形済み prompt のストリーム推論"""
|
|
92
|
+
if cache_trim_tokens is not None and cache_trim_tokens < 0:
|
|
93
|
+
raise ValueError("cache_trim_tokens must be non-negative")
|
|
94
|
+
|
|
95
|
+
if options is None:
|
|
96
|
+
options = {}
|
|
97
|
+
|
|
98
|
+
final_options = dict(options)
|
|
99
|
+
final_options.pop("trust_remote_code", None)
|
|
100
|
+
|
|
101
|
+
prompt_cache = None
|
|
102
|
+
cache_loaded = None
|
|
103
|
+
cache_read_tokens = 0
|
|
104
|
+
cache_write_tokens = 0
|
|
105
|
+
if cache_path:
|
|
106
|
+
if images:
|
|
107
|
+
cache_loaded = False
|
|
108
|
+
else:
|
|
109
|
+
prompt_cache = backend.load_cache_from_file(
|
|
110
|
+
cache_path,
|
|
111
|
+
images=images,
|
|
112
|
+
max_image_size=max_image_size,
|
|
113
|
+
prompt=prompt,
|
|
114
|
+
prefix_token_count=cache_trim_tokens,
|
|
115
|
+
)
|
|
116
|
+
cache_loaded = prompt_cache is not None
|
|
117
|
+
|
|
118
|
+
if prompt_cache is not None:
|
|
119
|
+
cache_read_tokens = backend.get_cache_offset(prompt_cache)
|
|
120
|
+
if cache_trim_tokens is not None and cache_read_tokens > cache_trim_tokens:
|
|
121
|
+
prompt_cache = backend.trim_cache(
|
|
122
|
+
prompt_cache,
|
|
123
|
+
cache_read_tokens - cache_trim_tokens,
|
|
124
|
+
)
|
|
125
|
+
cache_read_tokens = cache_trim_tokens
|
|
126
|
+
if cache_read_tokens <= 0:
|
|
127
|
+
prompt_cache = None
|
|
128
|
+
cache_loaded = False
|
|
129
|
+
elif isinstance(prompt, (str, list)):
|
|
130
|
+
full_tokens = (
|
|
131
|
+
backend.tokenize_prompt(prompt)
|
|
132
|
+
if isinstance(prompt, str)
|
|
133
|
+
else [int(token_id) for token_id in prompt]
|
|
134
|
+
)
|
|
135
|
+
if cache_read_tokens < len(full_tokens):
|
|
136
|
+
prompt = full_tokens[cache_read_tokens:]
|
|
137
|
+
else:
|
|
138
|
+
# A cache covering the complete prompt cannot be passed
|
|
139
|
+
# with an empty input_ids tensor. The safe fallback is a
|
|
140
|
+
# cold generation for this handler.
|
|
141
|
+
prompt_cache = None
|
|
142
|
+
cache_loaded = False
|
|
143
|
+
cache_read_tokens = 0
|
|
144
|
+
if prompt_cache is not None:
|
|
145
|
+
cache_write_tokens = backend.consume_cache_write_tokens(cache_path)
|
|
146
|
+
|
|
147
|
+
_stream_to_stdout(
|
|
148
|
+
backend,
|
|
149
|
+
prompt,
|
|
150
|
+
final_options,
|
|
151
|
+
images=images,
|
|
152
|
+
primer=primer,
|
|
153
|
+
prompt_cache=prompt_cache,
|
|
154
|
+
cache_loaded=cache_loaded,
|
|
155
|
+
cache_read_tokens=cache_read_tokens,
|
|
156
|
+
cache_write_tokens=cache_write_tokens,
|
|
157
|
+
)
|
|
@@ -4,9 +4,9 @@ version = "0.1.0"
|
|
|
4
4
|
description = "PyTorch (Transformers) driver for modular-prompt — cpu-minimal runtime"
|
|
5
5
|
requires-python = ">=3.10,<3.14"
|
|
6
6
|
dependencies = [
|
|
7
|
-
"safetensors==0.
|
|
7
|
+
"safetensors==0.8.0",
|
|
8
8
|
"tokenizers==0.22.2",
|
|
9
|
-
"transformers
|
|
9
|
+
"transformers>=5.14.0",
|
|
10
10
|
]
|
|
11
11
|
|
|
12
12
|
[dependency-groups]
|
|
@@ -3,7 +3,7 @@ import json
|
|
|
3
3
|
import sys
|
|
4
4
|
|
|
5
5
|
from backends.base import ModelBackend
|
|
6
|
-
from handlers import handle_capabilities, handle_completion, handle_format_test, handle_generate, handle_render, handle_tokenize
|
|
6
|
+
from handlers import handle_cache_prefill, handle_capabilities, handle_completion, handle_format_test, handle_generate, handle_render, handle_tokenize
|
|
7
7
|
from handlers.cancel import request_cancel, reset_cancel
|
|
8
8
|
|
|
9
9
|
|
|
@@ -87,7 +87,25 @@ class Server:
|
|
|
87
87
|
)
|
|
88
88
|
|
|
89
89
|
elif method == 'cache_prefill':
|
|
90
|
-
|
|
90
|
+
cache_path = req.get('cache_path')
|
|
91
|
+
messages = req.get('messages')
|
|
92
|
+
if not cache_path or not messages:
|
|
93
|
+
self._error_response("'cache_path' and 'messages' fields are required for cache_prefill")
|
|
94
|
+
return
|
|
95
|
+
handle_cache_prefill(
|
|
96
|
+
self.backend,
|
|
97
|
+
self.capabilities,
|
|
98
|
+
cache_path,
|
|
99
|
+
messages,
|
|
100
|
+
base_cache_path=req.get('base_cache_path'),
|
|
101
|
+
trim_to_tokens=req.get('trim_to_tokens'),
|
|
102
|
+
prefix_offsets=req.get('prefix_offsets'),
|
|
103
|
+
prefix_hashes=req.get('prefix_hashes'),
|
|
104
|
+
tools=req.get('tools'),
|
|
105
|
+
reasoning_effort=req.get('reasoning_effort'),
|
|
106
|
+
images=req.get('images'),
|
|
107
|
+
max_image_size=req.get('maxImageSize', 768),
|
|
108
|
+
)
|
|
91
109
|
|
|
92
110
|
elif method == 'render':
|
|
93
111
|
messages = req.get('messages')
|
|
@@ -0,0 +1,284 @@
|
|
|
1
|
+
import json
|
|
2
|
+
from types import SimpleNamespace
|
|
3
|
+
|
|
4
|
+
from handlers.cache import handle_cache_prefill
|
|
5
|
+
from handlers.generate import handle_generate
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class _Tokenizer:
|
|
9
|
+
bos_token = "<bos>"
|
|
10
|
+
chat_template = "template"
|
|
11
|
+
|
|
12
|
+
def __init__(self):
|
|
13
|
+
self.calls = []
|
|
14
|
+
|
|
15
|
+
def apply_chat_template(self, messages, **kwargs):
|
|
16
|
+
self.calls.append((messages, kwargs))
|
|
17
|
+
return "prefix"
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class _Cache:
|
|
21
|
+
pass
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class _Backend:
|
|
25
|
+
model_kind = "lm"
|
|
26
|
+
|
|
27
|
+
def __init__(self, cache=None, multi_chunk=False, cache_write_tokens=0):
|
|
28
|
+
self.tokenizer = _Tokenizer()
|
|
29
|
+
self.cache = cache
|
|
30
|
+
self.cache_offsets = {cache: 2} if cache is not None else {}
|
|
31
|
+
self.multi_chunk = multi_chunk
|
|
32
|
+
self.pending_cache_write_tokens = cache_write_tokens
|
|
33
|
+
self.calls = []
|
|
34
|
+
|
|
35
|
+
def get_tokenizer(self):
|
|
36
|
+
return self.tokenizer
|
|
37
|
+
|
|
38
|
+
def cache_prefill(self, cache_path, prompt, **kwargs):
|
|
39
|
+
self.calls.append(("prefill", cache_path, prompt, kwargs))
|
|
40
|
+
self.pending_cache_write_tokens = 2
|
|
41
|
+
return {
|
|
42
|
+
"cache_path": cache_path,
|
|
43
|
+
"token_count": 2,
|
|
44
|
+
"cache_write_tokens": 2,
|
|
45
|
+
}
|
|
46
|
+
|
|
47
|
+
def consume_cache_write_tokens(self, cache_path):
|
|
48
|
+
self.calls.append(("consume-write", cache_path))
|
|
49
|
+
result = self.pending_cache_write_tokens
|
|
50
|
+
self.pending_cache_write_tokens = 0
|
|
51
|
+
return result
|
|
52
|
+
|
|
53
|
+
def load_cache_from_file(self, cache_path, **kwargs):
|
|
54
|
+
self.calls.append(("load", cache_path, kwargs))
|
|
55
|
+
return self.cache
|
|
56
|
+
|
|
57
|
+
def get_cache_offset(self, prompt_cache):
|
|
58
|
+
return self.cache_offsets.get(prompt_cache, 0)
|
|
59
|
+
|
|
60
|
+
def trim_cache(self, prompt_cache, tokens):
|
|
61
|
+
self.calls.append(("trim", prompt_cache, tokens))
|
|
62
|
+
trimmed_cache = _Cache()
|
|
63
|
+
self.cache_offsets[trimmed_cache] = max(
|
|
64
|
+
0,
|
|
65
|
+
self.cache_offsets[prompt_cache] - tokens,
|
|
66
|
+
)
|
|
67
|
+
return trimmed_cache
|
|
68
|
+
|
|
69
|
+
def tokenize_prompt(self, prompt):
|
|
70
|
+
assert prompt == "prefix suffix"
|
|
71
|
+
return [1, 2, 3]
|
|
72
|
+
|
|
73
|
+
def stream_generate(self, prompt, options, images=None, prompt_cache=None):
|
|
74
|
+
self.calls.append(("generate", prompt, options, images, prompt_cache))
|
|
75
|
+
cache_read_tokens = self.cache_offsets.get(prompt_cache)
|
|
76
|
+
if self.multi_chunk:
|
|
77
|
+
for index, text in enumerate(("a", "b", "c"), start=1):
|
|
78
|
+
yield SimpleNamespace(
|
|
79
|
+
text=text,
|
|
80
|
+
prompt_tokens=3 if index == 1 else None,
|
|
81
|
+
generation_tokens=index,
|
|
82
|
+
cache_read_tokens=cache_read_tokens if index == 1 else None,
|
|
83
|
+
)
|
|
84
|
+
return
|
|
85
|
+
yield SimpleNamespace(
|
|
86
|
+
text="ok",
|
|
87
|
+
prompt_tokens=3,
|
|
88
|
+
generation_tokens=1,
|
|
89
|
+
cache_read_tokens=cache_read_tokens,
|
|
90
|
+
)
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def _json_response(output):
|
|
94
|
+
return json.loads(output.split("\0", 1)[0])
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
def test_cache_prefill_renders_without_generation_prompt(capsys):
|
|
98
|
+
backend = _Backend()
|
|
99
|
+
|
|
100
|
+
handle_cache_prefill(
|
|
101
|
+
backend,
|
|
102
|
+
{"special_tokens": {}},
|
|
103
|
+
"memory://prefix",
|
|
104
|
+
[{"role": "user", "content": "hello"}],
|
|
105
|
+
)
|
|
106
|
+
|
|
107
|
+
assert _json_response(capsys.readouterr().out) == {
|
|
108
|
+
"cache_path": "memory://prefix",
|
|
109
|
+
"token_count": 2,
|
|
110
|
+
"cache_write_tokens": 2,
|
|
111
|
+
}
|
|
112
|
+
assert backend.calls == [
|
|
113
|
+
("prefill", "memory://prefix", "prefix", {
|
|
114
|
+
"base_cache_path": None,
|
|
115
|
+
"trim_to_tokens": None,
|
|
116
|
+
"prefix_offsets": None,
|
|
117
|
+
"prefix_hashes": None,
|
|
118
|
+
"images": None,
|
|
119
|
+
"max_image_size": 768,
|
|
120
|
+
})
|
|
121
|
+
]
|
|
122
|
+
assert backend.tokenizer.calls[0][1]["add_generation_prompt"] is False
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
def test_cache_prefill_passes_incremental_options_and_propagates_prefix_meta(capsys):
|
|
126
|
+
backend = _Backend()
|
|
127
|
+
|
|
128
|
+
handle_cache_prefill(
|
|
129
|
+
backend,
|
|
130
|
+
{"special_tokens": {}},
|
|
131
|
+
"tmp/extended.pytorch-cache",
|
|
132
|
+
[{"role": "user", "content": "hello"}],
|
|
133
|
+
base_cache_path="tmp/base.pytorch-cache",
|
|
134
|
+
trim_to_tokens=1,
|
|
135
|
+
prefix_offsets=[1, 2],
|
|
136
|
+
prefix_hashes=["hash-prefix", "hash-full"],
|
|
137
|
+
)
|
|
138
|
+
|
|
139
|
+
result = _json_response(capsys.readouterr().out)
|
|
140
|
+
assert result["prefix_offsets"] == [1, 2]
|
|
141
|
+
assert result["prefix_hashes"] == ["hash-prefix", "hash-full"]
|
|
142
|
+
assert backend.calls == [
|
|
143
|
+
("prefill", "tmp/extended.pytorch-cache", "prefix", {
|
|
144
|
+
"base_cache_path": "tmp/base.pytorch-cache",
|
|
145
|
+
"trim_to_tokens": 1,
|
|
146
|
+
"prefix_offsets": [1, 2],
|
|
147
|
+
"prefix_hashes": ["hash-prefix", "hash-full"],
|
|
148
|
+
"images": None,
|
|
149
|
+
"max_image_size": 768,
|
|
150
|
+
})
|
|
151
|
+
]
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def test_generate_loads_cache_and_only_generates_suffix(capsys):
|
|
155
|
+
cache = _Cache()
|
|
156
|
+
backend = _Backend(cache, cache_write_tokens=2)
|
|
157
|
+
|
|
158
|
+
handle_generate(
|
|
159
|
+
backend,
|
|
160
|
+
"prefix suffix",
|
|
161
|
+
options={"max_tokens": 1},
|
|
162
|
+
cache_path="memory://prefix",
|
|
163
|
+
)
|
|
164
|
+
|
|
165
|
+
output = capsys.readouterr().out
|
|
166
|
+
assert output.startswith("ok")
|
|
167
|
+
meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
|
|
168
|
+
assert meta == {
|
|
169
|
+
"prompt_tokens": 3,
|
|
170
|
+
"generation_tokens": 1,
|
|
171
|
+
"cache_read_tokens": 2,
|
|
172
|
+
"cache_write_tokens": 2,
|
|
173
|
+
"cache_loaded": True,
|
|
174
|
+
}
|
|
175
|
+
generate_call = next(call for call in backend.calls if call[0] == "generate")
|
|
176
|
+
assert generate_call[1] == [3]
|
|
177
|
+
assert generate_call[4] is cache
|
|
178
|
+
load_call = next(call for call in backend.calls if call[0] == "load")
|
|
179
|
+
assert load_call[2]["prompt"] == "prefix suffix"
|
|
180
|
+
assert load_call[2]["prefix_token_count"] is None
|
|
181
|
+
assert load_call[2]["images"] is None
|
|
182
|
+
assert load_call[2]["max_image_size"] == 768
|
|
183
|
+
assert ("consume-write", "memory://prefix") in backend.calls
|
|
184
|
+
|
|
185
|
+
|
|
186
|
+
def test_generate_trims_loaded_cache_before_generating(capsys):
|
|
187
|
+
cache = _Cache()
|
|
188
|
+
backend = _Backend(cache, cache_write_tokens=2)
|
|
189
|
+
|
|
190
|
+
handle_generate(
|
|
191
|
+
backend,
|
|
192
|
+
"prefix suffix",
|
|
193
|
+
options={"max_tokens": 1},
|
|
194
|
+
cache_path="memory://prefix",
|
|
195
|
+
cache_trim_tokens=1,
|
|
196
|
+
)
|
|
197
|
+
|
|
198
|
+
output = capsys.readouterr().out
|
|
199
|
+
generate_call = next(call for call in backend.calls if call[0] == "generate")
|
|
200
|
+
assert generate_call[1] == [2, 3]
|
|
201
|
+
assert generate_call[4] is not cache
|
|
202
|
+
assert backend.get_cache_offset(cache) == 2
|
|
203
|
+
assert backend.get_cache_offset(generate_call[4]) == 1
|
|
204
|
+
assert ("trim", cache, 1) in backend.calls
|
|
205
|
+
load_call = next(call for call in backend.calls if call[0] == "load")
|
|
206
|
+
assert load_call[2]["prefix_token_count"] == 1
|
|
207
|
+
meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
|
|
208
|
+
assert meta["cache_read_tokens"] == 1
|
|
209
|
+
|
|
210
|
+
# The original ref remains usable after the trimmed generation.
|
|
211
|
+
handle_generate(
|
|
212
|
+
backend,
|
|
213
|
+
"prefix suffix",
|
|
214
|
+
options={"max_tokens": 1},
|
|
215
|
+
cache_path="memory://prefix",
|
|
216
|
+
)
|
|
217
|
+
second_output = capsys.readouterr().out
|
|
218
|
+
second_generate_call = [
|
|
219
|
+
call for call in backend.calls if call[0] == "generate"
|
|
220
|
+
][-1]
|
|
221
|
+
assert second_generate_call[1] == [3]
|
|
222
|
+
assert second_generate_call[4] is cache
|
|
223
|
+
second_meta = json.loads(
|
|
224
|
+
second_output.split("\x1e__META__:", 1)[1].split("\0", 1)[0]
|
|
225
|
+
)
|
|
226
|
+
assert second_meta["cache_read_tokens"] == 2
|
|
227
|
+
|
|
228
|
+
|
|
229
|
+
def test_generate_preserves_usage_meta_across_multiple_chunks(capsys):
|
|
230
|
+
cache = _Cache()
|
|
231
|
+
backend = _Backend(cache, multi_chunk=True, cache_write_tokens=2)
|
|
232
|
+
|
|
233
|
+
handle_generate(
|
|
234
|
+
backend,
|
|
235
|
+
"prefix suffix",
|
|
236
|
+
options={"max_tokens": 3},
|
|
237
|
+
cache_path="memory://prefix",
|
|
238
|
+
)
|
|
239
|
+
|
|
240
|
+
output = capsys.readouterr().out
|
|
241
|
+
assert output.startswith("abc")
|
|
242
|
+
meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
|
|
243
|
+
assert meta == {
|
|
244
|
+
"prompt_tokens": 3,
|
|
245
|
+
"generation_tokens": 3,
|
|
246
|
+
"cache_read_tokens": 2,
|
|
247
|
+
"cache_write_tokens": 2,
|
|
248
|
+
"cache_loaded": True,
|
|
249
|
+
}
|
|
250
|
+
|
|
251
|
+
|
|
252
|
+
def test_generate_uses_cold_path_for_missing_cache(capsys):
|
|
253
|
+
backend = _Backend(cache=None)
|
|
254
|
+
|
|
255
|
+
handle_generate(
|
|
256
|
+
backend,
|
|
257
|
+
"prefix suffix",
|
|
258
|
+
cache_path="memory://missing",
|
|
259
|
+
)
|
|
260
|
+
|
|
261
|
+
output = capsys.readouterr().out
|
|
262
|
+
meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
|
|
263
|
+
assert meta["cache_loaded"] is False
|
|
264
|
+
assert "cache_read_tokens" not in meta
|
|
265
|
+
assert "cache_write_tokens" not in meta
|
|
266
|
+
generate_call = next(call for call in backend.calls if call[0] == "generate")
|
|
267
|
+
assert generate_call[1] == "prefix suffix"
|
|
268
|
+
assert generate_call[4] is None
|
|
269
|
+
|
|
270
|
+
|
|
271
|
+
def test_generate_removes_cached_prefix_from_token_ids(capsys):
|
|
272
|
+
cache = _Cache()
|
|
273
|
+
backend = _Backend(cache)
|
|
274
|
+
|
|
275
|
+
handle_generate(
|
|
276
|
+
backend,
|
|
277
|
+
[1, 2, 3],
|
|
278
|
+
cache_path="memory://prefix",
|
|
279
|
+
)
|
|
280
|
+
|
|
281
|
+
output = capsys.readouterr().out
|
|
282
|
+
assert '"cache_loaded": true' in output
|
|
283
|
+
generate_call = next(call for call in backend.calls if call[0] == "generate")
|
|
284
|
+
assert generate_call[1] == [3]
|
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
from utils.token_utils import get_capabilities
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
class _Tokenizer:
|
|
5
|
+
apply_chat_template = None
|
|
6
|
+
chat_template = None
|
|
7
|
+
special_tokens_map = {}
|
|
8
|
+
added_tokens_encoder = {}
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def test_capabilities_advertise_cache_prefill():
|
|
12
|
+
capabilities = get_capabilities(_Tokenizer())
|
|
13
|
+
|
|
14
|
+
assert "cache_prefill" in capabilities["methods"]
|