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