@modular-prompt/driver 0.16.0 → 0.17.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +93 -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.map +1 -1
- package/dist/driver-registry/config-based-factory.js +6 -0
- 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/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/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 +9 -5
- 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/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 +1 -1
- package/src/mlx-ml/python/server.py +2 -0
- package/src/mlx-ml/python/uv.lock +8 -8
- 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,152 @@
|
|
|
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 images in Phase 1")
|
|
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:
|
|
93
|
+
raise ValueError(
|
|
94
|
+
"PyTorch LIP backend does not support cache trimming in Phase 1"
|
|
95
|
+
)
|
|
96
|
+
|
|
97
|
+
if options is None:
|
|
98
|
+
options = {}
|
|
99
|
+
|
|
100
|
+
final_options = dict(options)
|
|
101
|
+
final_options.pop("trust_remote_code", None)
|
|
102
|
+
|
|
103
|
+
prompt_cache = None
|
|
104
|
+
cache_loaded = None
|
|
105
|
+
cache_read_tokens = 0
|
|
106
|
+
cache_write_tokens = 0
|
|
107
|
+
if cache_path:
|
|
108
|
+
if images:
|
|
109
|
+
cache_loaded = False
|
|
110
|
+
else:
|
|
111
|
+
prompt_cache = backend.load_cache_from_file(
|
|
112
|
+
cache_path,
|
|
113
|
+
images=images,
|
|
114
|
+
max_image_size=max_image_size,
|
|
115
|
+
prompt=prompt,
|
|
116
|
+
)
|
|
117
|
+
cache_loaded = prompt_cache is not None
|
|
118
|
+
|
|
119
|
+
if prompt_cache is not None:
|
|
120
|
+
cache_read_tokens = backend.get_cache_offset(prompt_cache)
|
|
121
|
+
if cache_read_tokens <= 0:
|
|
122
|
+
prompt_cache = None
|
|
123
|
+
cache_loaded = False
|
|
124
|
+
elif isinstance(prompt, (str, list)):
|
|
125
|
+
full_tokens = (
|
|
126
|
+
backend.tokenize_prompt(prompt)
|
|
127
|
+
if isinstance(prompt, str)
|
|
128
|
+
else [int(token_id) for token_id in prompt]
|
|
129
|
+
)
|
|
130
|
+
if cache_read_tokens < len(full_tokens):
|
|
131
|
+
prompt = full_tokens[cache_read_tokens:]
|
|
132
|
+
else:
|
|
133
|
+
# A cache covering the complete prompt cannot be passed
|
|
134
|
+
# with an empty input_ids tensor. The safe fallback is a
|
|
135
|
+
# cold generation for this Phase 1 handler.
|
|
136
|
+
prompt_cache = None
|
|
137
|
+
cache_loaded = False
|
|
138
|
+
cache_read_tokens = 0
|
|
139
|
+
if prompt_cache is not None:
|
|
140
|
+
cache_write_tokens = backend.consume_cache_write_tokens(cache_path)
|
|
141
|
+
|
|
142
|
+
_stream_to_stdout(
|
|
143
|
+
backend,
|
|
144
|
+
prompt,
|
|
145
|
+
final_options,
|
|
146
|
+
images=images,
|
|
147
|
+
primer=primer,
|
|
148
|
+
prompt_cache=prompt_cache,
|
|
149
|
+
cache_loaded=cache_loaded,
|
|
150
|
+
cache_read_tokens=cache_read_tokens,
|
|
151
|
+
cache_write_tokens=cache_write_tokens,
|
|
152
|
+
)
|
|
@@ -0,0 +1,40 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
|
|
5
|
+
from backends.base import ModelBackend
|
|
6
|
+
from utils.template_render import apply_chat_template_prompt
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def handle_render(
|
|
10
|
+
backend: ModelBackend,
|
|
11
|
+
messages: list,
|
|
12
|
+
options: dict | None = None,
|
|
13
|
+
tools: list | None = None,
|
|
14
|
+
reasoning_effort: str | None = None,
|
|
15
|
+
) -> None:
|
|
16
|
+
"""LIP render: apply_chat_template のみ(推論しない)"""
|
|
17
|
+
if options is None:
|
|
18
|
+
options = {}
|
|
19
|
+
|
|
20
|
+
result: dict = {
|
|
21
|
+
"formatted_prompt": None,
|
|
22
|
+
"error": None,
|
|
23
|
+
}
|
|
24
|
+
|
|
25
|
+
try:
|
|
26
|
+
trust_remote_code = options.get("trust_remote_code")
|
|
27
|
+
primer = options.get("primer")
|
|
28
|
+
prompt = apply_chat_template_prompt(
|
|
29
|
+
backend,
|
|
30
|
+
messages,
|
|
31
|
+
primer=primer,
|
|
32
|
+
tools=tools,
|
|
33
|
+
reasoning_effort=reasoning_effort,
|
|
34
|
+
trust_remote_code=trust_remote_code,
|
|
35
|
+
)
|
|
36
|
+
result["formatted_prompt"] = prompt
|
|
37
|
+
except Exception as e:
|
|
38
|
+
result["error"] = str(e)
|
|
39
|
+
|
|
40
|
+
print(json.dumps(result), end="\0", flush=True)
|
|
@@ -0,0 +1,63 @@
|
|
|
1
|
+
import json
|
|
2
|
+
|
|
3
|
+
from backends.base import ModelBackend
|
|
4
|
+
from utils.prompt_builder import generate_merged_prompt, supports_chat_template
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def handle_tokenize(
|
|
8
|
+
backend: ModelBackend,
|
|
9
|
+
capabilities: dict,
|
|
10
|
+
messages: list,
|
|
11
|
+
tools: list | None = None,
|
|
12
|
+
reasoning_effort: str | None = None,
|
|
13
|
+
) -> None:
|
|
14
|
+
"""メッセージをchat template適用後にトークン化して返す"""
|
|
15
|
+
tokenizer = backend.get_tokenizer()
|
|
16
|
+
|
|
17
|
+
result = {
|
|
18
|
+
"token_ids": None,
|
|
19
|
+
"token_count": 0,
|
|
20
|
+
"error": None,
|
|
21
|
+
}
|
|
22
|
+
|
|
23
|
+
try:
|
|
24
|
+
# apply_chat_templateのfallbackパターン (chat.py L165-188 と同じ)
|
|
25
|
+
# add_generation_prompt=False で、アシスタントの開始トークンは含めない
|
|
26
|
+
extra_kwargs = {}
|
|
27
|
+
if tools is not None:
|
|
28
|
+
extra_kwargs["tools"] = tools
|
|
29
|
+
if reasoning_effort is not None:
|
|
30
|
+
extra_kwargs["reasoning_effort"] = reasoning_effort
|
|
31
|
+
|
|
32
|
+
if supports_chat_template(tokenizer):
|
|
33
|
+
# chat.py と同じfallbackチェーン
|
|
34
|
+
prompt = None
|
|
35
|
+
for kwargs in [extra_kwargs, {k: v for k, v in extra_kwargs.items() if k == "tools"}, {}]:
|
|
36
|
+
try:
|
|
37
|
+
prompt = tokenizer.apply_chat_template(
|
|
38
|
+
messages,
|
|
39
|
+
add_generation_prompt=False,
|
|
40
|
+
tokenize=False,
|
|
41
|
+
**kwargs,
|
|
42
|
+
)
|
|
43
|
+
break
|
|
44
|
+
except TypeError:
|
|
45
|
+
continue
|
|
46
|
+
|
|
47
|
+
if prompt is None:
|
|
48
|
+
prompt = str(messages)
|
|
49
|
+
else:
|
|
50
|
+
prompt = generate_merged_prompt(messages, capabilities)
|
|
51
|
+
|
|
52
|
+
# トークン化
|
|
53
|
+
add_special = tokenizer.bos_token is None or not prompt.startswith(
|
|
54
|
+
tokenizer.bos_token or ""
|
|
55
|
+
)
|
|
56
|
+
token_ids = tokenizer.encode(prompt, add_special_tokens=add_special)
|
|
57
|
+
|
|
58
|
+
result["token_ids"] = token_ids
|
|
59
|
+
result["token_count"] = len(token_ids)
|
|
60
|
+
except Exception as e:
|
|
61
|
+
result["error"] = str(e)
|
|
62
|
+
|
|
63
|
+
print(json.dumps(result), end="\0", flush=True)
|
|
@@ -0,0 +1,37 @@
|
|
|
1
|
+
[project]
|
|
2
|
+
name = "pytorch_driver"
|
|
3
|
+
version = "0.1.0"
|
|
4
|
+
description = "PyTorch (Transformers) driver for modular-prompt — CUDA runtime"
|
|
5
|
+
requires-python = ">=3.10,<3.14"
|
|
6
|
+
dependencies = [
|
|
7
|
+
"safetensors==0.8.0",
|
|
8
|
+
"tokenizers==0.22.2",
|
|
9
|
+
"transformers>=5.14.0",
|
|
10
|
+
]
|
|
11
|
+
|
|
12
|
+
[dependency-groups]
|
|
13
|
+
dev = ["pytest>=9.0"]
|
|
14
|
+
|
|
15
|
+
[build-system]
|
|
16
|
+
requires = ["setuptools>=61.0"]
|
|
17
|
+
build-backend = "setuptools.build_meta"
|
|
18
|
+
|
|
19
|
+
[tool.pytest.ini_options]
|
|
20
|
+
testpaths = ["tests"]
|
|
21
|
+
|
|
22
|
+
[tool.setuptools]
|
|
23
|
+
py-modules = ["__main__", "server"]
|
|
24
|
+
|
|
25
|
+
[tool.setuptools.packages.find]
|
|
26
|
+
where = ["."]
|
|
27
|
+
include = ["backends*", "handlers*", "utils*"]
|
|
28
|
+
|
|
29
|
+
# setup-pytorch --variant cuda は CUDA wheel を明示インストールする。
|
|
30
|
+
# デフォルトは cu124。別の index を使う場合は --cuda <version> を指定する。
|
|
31
|
+
[[tool.uv.index]]
|
|
32
|
+
name = "pytorch-cuda"
|
|
33
|
+
url = "https://download.pytorch.org/whl/cu124"
|
|
34
|
+
explicit = true
|
|
35
|
+
|
|
36
|
+
[tool.uv.sources]
|
|
37
|
+
torch = { index = "pytorch-cuda" }
|
|
@@ -0,0 +1,158 @@
|
|
|
1
|
+
"""JSON-RPC風サーバー: stdin/stdoutベースのリクエストディスパッチ"""
|
|
2
|
+
import json
|
|
3
|
+
import sys
|
|
4
|
+
|
|
5
|
+
from backends.base import ModelBackend
|
|
6
|
+
from handlers import handle_cache_prefill, handle_capabilities, handle_completion, handle_format_test, handle_generate, handle_render, handle_tokenize
|
|
7
|
+
from handlers.cancel import request_cancel, reset_cancel
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
MAX_READ_LINES = 10000
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def read():
|
|
14
|
+
lines = []
|
|
15
|
+
while True:
|
|
16
|
+
line = sys.stdin.readline()
|
|
17
|
+
if not line:
|
|
18
|
+
return None
|
|
19
|
+
lines.append(line)
|
|
20
|
+
if len(lines) > MAX_READ_LINES:
|
|
21
|
+
sys.stderr.write(f"Error: read buffer exceeded {MAX_READ_LINES} lines, discarding\n")
|
|
22
|
+
lines.clear()
|
|
23
|
+
continue
|
|
24
|
+
try:
|
|
25
|
+
return json.loads(''.join(lines))
|
|
26
|
+
except json.JSONDecodeError:
|
|
27
|
+
continue
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def _is_valid_generate_prompt(prompt) -> bool:
|
|
31
|
+
"""prompt は非空文字列または非空のトークン ID リスト"""
|
|
32
|
+
if isinstance(prompt, str):
|
|
33
|
+
return bool(prompt)
|
|
34
|
+
if isinstance(prompt, list):
|
|
35
|
+
return len(prompt) > 0
|
|
36
|
+
return False
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
class Server:
|
|
40
|
+
def __init__(self, backend: ModelBackend, capabilities: dict):
|
|
41
|
+
self.backend = backend
|
|
42
|
+
self.capabilities = capabilities
|
|
43
|
+
|
|
44
|
+
def run(self):
|
|
45
|
+
while True:
|
|
46
|
+
req = read()
|
|
47
|
+
if req is None:
|
|
48
|
+
break
|
|
49
|
+
self._dispatch(req)
|
|
50
|
+
|
|
51
|
+
def _error_response(self, message: str) -> None:
|
|
52
|
+
sys.stderr.write(f"Error: {message}\n")
|
|
53
|
+
print(json.dumps({"error": message}), end='\0', flush=True)
|
|
54
|
+
|
|
55
|
+
def _dispatch(self, req: dict):
|
|
56
|
+
method = req.get('method')
|
|
57
|
+
if not method:
|
|
58
|
+
self._error_response("'method' field is required")
|
|
59
|
+
return
|
|
60
|
+
|
|
61
|
+
if method == 'cancel':
|
|
62
|
+
request_cancel()
|
|
63
|
+
return
|
|
64
|
+
|
|
65
|
+
reset_cancel()
|
|
66
|
+
|
|
67
|
+
try:
|
|
68
|
+
if method == 'capabilities':
|
|
69
|
+
handle_capabilities(self.capabilities)
|
|
70
|
+
|
|
71
|
+
elif method == 'format_test':
|
|
72
|
+
messages = req.get('messages')
|
|
73
|
+
if not messages:
|
|
74
|
+
self._error_response("'messages' field is required for format_test method")
|
|
75
|
+
return
|
|
76
|
+
handle_format_test(self.backend, self.capabilities, messages, req.get('options', {}), req.get('tools'))
|
|
77
|
+
|
|
78
|
+
elif method == 'tokenize':
|
|
79
|
+
messages = req.get('messages')
|
|
80
|
+
if messages is None:
|
|
81
|
+
self._error_response("'messages' field is required for tokenize method")
|
|
82
|
+
return
|
|
83
|
+
handle_tokenize(
|
|
84
|
+
self.backend, self.capabilities, messages,
|
|
85
|
+
tools=req.get('tools'),
|
|
86
|
+
reasoning_effort=req.get('reasoning_effort'),
|
|
87
|
+
)
|
|
88
|
+
|
|
89
|
+
elif method == 'cache_prefill':
|
|
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
|
+
)
|
|
109
|
+
|
|
110
|
+
elif method == 'render':
|
|
111
|
+
messages = req.get('messages')
|
|
112
|
+
if not messages:
|
|
113
|
+
self._error_response("'messages' field is required for render method")
|
|
114
|
+
return
|
|
115
|
+
handle_render(
|
|
116
|
+
self.backend,
|
|
117
|
+
messages,
|
|
118
|
+
options=req.get('options', {}),
|
|
119
|
+
tools=req.get('tools'),
|
|
120
|
+
reasoning_effort=req.get('reasoning_effort'),
|
|
121
|
+
)
|
|
122
|
+
|
|
123
|
+
elif method == 'generate':
|
|
124
|
+
prompt = req.get('prompt')
|
|
125
|
+
if not _is_valid_generate_prompt(prompt):
|
|
126
|
+
self._error_response("'prompt' field is required for generate method")
|
|
127
|
+
return
|
|
128
|
+
images = req.get('images', [])
|
|
129
|
+
handle_generate(
|
|
130
|
+
self.backend,
|
|
131
|
+
prompt,
|
|
132
|
+
options=req.get('options', {}),
|
|
133
|
+
images=images if images else None,
|
|
134
|
+
max_image_size=req.get('maxImageSize', 768),
|
|
135
|
+
primer=req.get('primer'),
|
|
136
|
+
cache_path=req.get('cache_path'),
|
|
137
|
+
cache_trim_tokens=req.get('cache_trim_tokens'),
|
|
138
|
+
)
|
|
139
|
+
|
|
140
|
+
elif method == 'completion':
|
|
141
|
+
prompt = req.get('prompt')
|
|
142
|
+
if not prompt:
|
|
143
|
+
self._error_response("'prompt' field is required for completion method")
|
|
144
|
+
return
|
|
145
|
+
images = req.get('images', [])
|
|
146
|
+
handle_completion(
|
|
147
|
+
self.backend,
|
|
148
|
+
prompt,
|
|
149
|
+
options=req.get('options', {}),
|
|
150
|
+
images=images if images else None,
|
|
151
|
+
max_image_size=req.get('maxImageSize', 768),
|
|
152
|
+
)
|
|
153
|
+
|
|
154
|
+
else:
|
|
155
|
+
self._error_response(f"Unknown method '{method}'")
|
|
156
|
+
|
|
157
|
+
except Exception as e:
|
|
158
|
+
self._error_response(f"Error processing request: {e}")
|
|
File without changes
|
|
@@ -0,0 +1,207 @@
|
|
|
1
|
+
import pytest
|
|
2
|
+
|
|
3
|
+
pytest.importorskip("torch")
|
|
4
|
+
|
|
5
|
+
import json
|
|
6
|
+
from types import SimpleNamespace
|
|
7
|
+
|
|
8
|
+
from handlers.cache import handle_cache_prefill
|
|
9
|
+
from handlers.generate import handle_generate
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class _Tokenizer:
|
|
13
|
+
bos_token = "<bos>"
|
|
14
|
+
chat_template = "template"
|
|
15
|
+
|
|
16
|
+
def __init__(self):
|
|
17
|
+
self.calls = []
|
|
18
|
+
|
|
19
|
+
def apply_chat_template(self, messages, **kwargs):
|
|
20
|
+
self.calls.append((messages, kwargs))
|
|
21
|
+
return "prefix"
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class _Cache:
|
|
25
|
+
pass
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class _Backend:
|
|
29
|
+
model_kind = "lm"
|
|
30
|
+
|
|
31
|
+
def __init__(self, cache=None, multi_chunk=False, cache_write_tokens=0):
|
|
32
|
+
self.tokenizer = _Tokenizer()
|
|
33
|
+
self.cache = cache
|
|
34
|
+
self.multi_chunk = multi_chunk
|
|
35
|
+
self.pending_cache_write_tokens = cache_write_tokens
|
|
36
|
+
self.calls = []
|
|
37
|
+
|
|
38
|
+
def get_tokenizer(self):
|
|
39
|
+
return self.tokenizer
|
|
40
|
+
|
|
41
|
+
def cache_prefill(self, cache_path, prompt, **kwargs):
|
|
42
|
+
self.calls.append(("prefill", cache_path, prompt, kwargs))
|
|
43
|
+
self.pending_cache_write_tokens = 2
|
|
44
|
+
return {
|
|
45
|
+
"cache_path": cache_path,
|
|
46
|
+
"token_count": 2,
|
|
47
|
+
"cache_write_tokens": 2,
|
|
48
|
+
}
|
|
49
|
+
|
|
50
|
+
def consume_cache_write_tokens(self, cache_path):
|
|
51
|
+
self.calls.append(("consume-write", cache_path))
|
|
52
|
+
result = self.pending_cache_write_tokens
|
|
53
|
+
self.pending_cache_write_tokens = 0
|
|
54
|
+
return result
|
|
55
|
+
|
|
56
|
+
def load_cache_from_file(self, cache_path, **kwargs):
|
|
57
|
+
self.calls.append(("load", cache_path, kwargs))
|
|
58
|
+
return self.cache
|
|
59
|
+
|
|
60
|
+
def get_cache_offset(self, prompt_cache):
|
|
61
|
+
return 2 if prompt_cache is self.cache else 0
|
|
62
|
+
|
|
63
|
+
def tokenize_prompt(self, prompt):
|
|
64
|
+
assert prompt == "prefix suffix"
|
|
65
|
+
return [1, 2, 3]
|
|
66
|
+
|
|
67
|
+
def stream_generate(self, prompt, options, images=None, prompt_cache=None):
|
|
68
|
+
self.calls.append(("generate", prompt, options, images, prompt_cache))
|
|
69
|
+
cache_read_tokens = (
|
|
70
|
+
2 if self.cache is not None and prompt_cache is self.cache else None
|
|
71
|
+
)
|
|
72
|
+
if self.multi_chunk:
|
|
73
|
+
for index, text in enumerate(("a", "b", "c"), start=1):
|
|
74
|
+
yield SimpleNamespace(
|
|
75
|
+
text=text,
|
|
76
|
+
prompt_tokens=3 if index == 1 else None,
|
|
77
|
+
generation_tokens=index,
|
|
78
|
+
cache_read_tokens=cache_read_tokens if index == 1 else None,
|
|
79
|
+
)
|
|
80
|
+
return
|
|
81
|
+
yield SimpleNamespace(
|
|
82
|
+
text="ok",
|
|
83
|
+
prompt_tokens=3,
|
|
84
|
+
generation_tokens=1,
|
|
85
|
+
cache_read_tokens=cache_read_tokens,
|
|
86
|
+
)
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def _json_response(output):
|
|
90
|
+
return json.loads(output.split("\0", 1)[0])
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def test_cache_prefill_renders_without_generation_prompt(capsys):
|
|
94
|
+
backend = _Backend()
|
|
95
|
+
|
|
96
|
+
handle_cache_prefill(
|
|
97
|
+
backend,
|
|
98
|
+
{"special_tokens": {}},
|
|
99
|
+
"memory://prefix",
|
|
100
|
+
[{"role": "user", "content": "hello"}],
|
|
101
|
+
)
|
|
102
|
+
|
|
103
|
+
assert _json_response(capsys.readouterr().out) == {
|
|
104
|
+
"cache_path": "memory://prefix",
|
|
105
|
+
"token_count": 2,
|
|
106
|
+
"cache_write_tokens": 2,
|
|
107
|
+
}
|
|
108
|
+
assert backend.calls == [
|
|
109
|
+
("prefill", "memory://prefix", "prefix", {
|
|
110
|
+
"base_cache_path": None,
|
|
111
|
+
"trim_to_tokens": None,
|
|
112
|
+
"prefix_offsets": None,
|
|
113
|
+
"prefix_hashes": None,
|
|
114
|
+
"images": None,
|
|
115
|
+
"max_image_size": 768,
|
|
116
|
+
})
|
|
117
|
+
]
|
|
118
|
+
assert backend.tokenizer.calls[0][1]["add_generation_prompt"] is False
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
def test_generate_loads_cache_and_only_generates_suffix(capsys):
|
|
122
|
+
cache = _Cache()
|
|
123
|
+
backend = _Backend(cache, cache_write_tokens=2)
|
|
124
|
+
|
|
125
|
+
handle_generate(
|
|
126
|
+
backend,
|
|
127
|
+
"prefix suffix",
|
|
128
|
+
options={"max_tokens": 1},
|
|
129
|
+
cache_path="memory://prefix",
|
|
130
|
+
)
|
|
131
|
+
|
|
132
|
+
output = capsys.readouterr().out
|
|
133
|
+
assert output.startswith("ok")
|
|
134
|
+
meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
|
|
135
|
+
assert meta == {
|
|
136
|
+
"prompt_tokens": 3,
|
|
137
|
+
"generation_tokens": 1,
|
|
138
|
+
"cache_read_tokens": 2,
|
|
139
|
+
"cache_write_tokens": 2,
|
|
140
|
+
"cache_loaded": True,
|
|
141
|
+
}
|
|
142
|
+
generate_call = next(call for call in backend.calls if call[0] == "generate")
|
|
143
|
+
assert generate_call[1] == [3]
|
|
144
|
+
assert generate_call[4] is cache
|
|
145
|
+
load_call = next(call for call in backend.calls if call[0] == "load")
|
|
146
|
+
assert load_call[2]["prompt"] == "prefix suffix"
|
|
147
|
+
assert load_call[2]["images"] is None
|
|
148
|
+
assert load_call[2]["max_image_size"] == 768
|
|
149
|
+
assert ("consume-write", "memory://prefix") in backend.calls
|
|
150
|
+
|
|
151
|
+
|
|
152
|
+
def test_generate_preserves_usage_meta_across_multiple_chunks(capsys):
|
|
153
|
+
cache = _Cache()
|
|
154
|
+
backend = _Backend(cache, multi_chunk=True, cache_write_tokens=2)
|
|
155
|
+
|
|
156
|
+
handle_generate(
|
|
157
|
+
backend,
|
|
158
|
+
"prefix suffix",
|
|
159
|
+
options={"max_tokens": 3},
|
|
160
|
+
cache_path="memory://prefix",
|
|
161
|
+
)
|
|
162
|
+
|
|
163
|
+
output = capsys.readouterr().out
|
|
164
|
+
assert output.startswith("abc")
|
|
165
|
+
meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
|
|
166
|
+
assert meta == {
|
|
167
|
+
"prompt_tokens": 3,
|
|
168
|
+
"generation_tokens": 3,
|
|
169
|
+
"cache_read_tokens": 2,
|
|
170
|
+
"cache_write_tokens": 2,
|
|
171
|
+
"cache_loaded": True,
|
|
172
|
+
}
|
|
173
|
+
|
|
174
|
+
|
|
175
|
+
def test_generate_uses_cold_path_for_missing_cache(capsys):
|
|
176
|
+
backend = _Backend(cache=None)
|
|
177
|
+
|
|
178
|
+
handle_generate(
|
|
179
|
+
backend,
|
|
180
|
+
"prefix suffix",
|
|
181
|
+
cache_path="memory://missing",
|
|
182
|
+
)
|
|
183
|
+
|
|
184
|
+
output = capsys.readouterr().out
|
|
185
|
+
meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
|
|
186
|
+
assert meta["cache_loaded"] is False
|
|
187
|
+
assert "cache_read_tokens" not in meta
|
|
188
|
+
assert "cache_write_tokens" not in meta
|
|
189
|
+
generate_call = next(call for call in backend.calls if call[0] == "generate")
|
|
190
|
+
assert generate_call[1] == "prefix suffix"
|
|
191
|
+
assert generate_call[4] is None
|
|
192
|
+
|
|
193
|
+
|
|
194
|
+
def test_generate_removes_cached_prefix_from_token_ids(capsys):
|
|
195
|
+
cache = _Cache()
|
|
196
|
+
backend = _Backend(cache)
|
|
197
|
+
|
|
198
|
+
handle_generate(
|
|
199
|
+
backend,
|
|
200
|
+
[1, 2, 3],
|
|
201
|
+
cache_path="memory://prefix",
|
|
202
|
+
)
|
|
203
|
+
|
|
204
|
+
output = capsys.readouterr().out
|
|
205
|
+
assert '"cache_loaded": true' in output
|
|
206
|
+
generate_call = next(call for call in backend.calls if call[0] == "generate")
|
|
207
|
+
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"]
|