@modular-prompt/driver 0.14.0 → 0.15.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 +58 -5
- package/dist/driver-registry/ai-service.d.ts +23 -1
- package/dist/driver-registry/ai-service.d.ts.map +1 -1
- package/dist/driver-registry/ai-service.js +44 -10
- package/dist/driver-registry/ai-service.js.map +1 -1
- package/dist/driver-registry/config-based-factory.d.ts.map +1 -1
- package/dist/driver-registry/config-based-factory.js +16 -0
- package/dist/driver-registry/config-based-factory.js.map +1 -1
- package/dist/driver-registry/factory-helper.d.ts +2 -0
- package/dist/driver-registry/factory-helper.d.ts.map +1 -1
- package/dist/driver-registry/factory-helper.js +21 -3
- package/dist/driver-registry/factory-helper.js.map +1 -1
- package/dist/driver-registry/index.d.ts +2 -2
- package/dist/driver-registry/index.d.ts.map +1 -1
- package/dist/driver-registry/index.js +1 -1
- package/dist/driver-registry/index.js.map +1 -1
- package/dist/driver-registry/registry.d.ts.map +1 -1
- package/dist/driver-registry/registry.js +3 -1
- package/dist/driver-registry/registry.js.map +1 -1
- package/dist/driver-registry/types.d.ts +8 -2
- package/dist/driver-registry/types.d.ts.map +1 -1
- package/dist/index.d.ts +10 -2
- package/dist/index.d.ts.map +1 -1
- package/dist/index.js +9 -1
- package/dist/index.js.map +1 -1
- package/dist/local-inference/adapters.d.ts +66 -0
- package/dist/local-inference/adapters.d.ts.map +1 -0
- package/dist/local-inference/adapters.js +2 -0
- package/dist/local-inference/adapters.js.map +1 -0
- package/dist/local-inference/driver.d.ts +51 -0
- package/dist/local-inference/driver.d.ts.map +1 -0
- package/dist/local-inference/driver.js +309 -0
- package/dist/local-inference/driver.js.map +1 -0
- package/dist/local-inference/index.d.ts +22 -0
- package/dist/local-inference/index.d.ts.map +1 -0
- package/dist/local-inference/index.js +17 -0
- package/dist/local-inference/index.js.map +1 -0
- package/dist/local-inference/process-client.d.ts +50 -0
- package/dist/local-inference/process-client.d.ts.map +1 -0
- package/dist/local-inference/process-client.js +92 -0
- package/dist/local-inference/process-client.js.map +1 -0
- package/dist/local-inference/process-communication.d.ts +41 -0
- package/dist/local-inference/process-communication.d.ts.map +1 -0
- package/dist/{mlx-ml/process → local-inference}/process-communication.js +25 -61
- package/dist/local-inference/process-communication.js.map +1 -0
- package/dist/local-inference/process-port.d.ts +12 -0
- package/dist/local-inference/process-port.d.ts.map +1 -0
- package/dist/local-inference/process-port.js +2 -0
- package/dist/local-inference/process-port.js.map +1 -0
- package/dist/local-inference/prompt-utils.d.ts +6 -0
- package/dist/local-inference/prompt-utils.d.ts.map +1 -0
- package/dist/local-inference/prompt-utils.js +17 -0
- package/dist/local-inference/prompt-utils.js.map +1 -0
- package/dist/local-inference/protocol.d.ts +192 -0
- package/dist/local-inference/protocol.d.ts.map +1 -0
- package/dist/local-inference/protocol.js +2 -0
- package/dist/local-inference/protocol.js.map +1 -0
- package/dist/local-inference/queue-types.d.ts +54 -0
- package/dist/local-inference/queue-types.d.ts.map +1 -0
- package/dist/local-inference/queue-types.js +2 -0
- package/dist/local-inference/queue-types.js.map +1 -0
- package/dist/local-inference/request-queue.d.ts +36 -0
- package/dist/local-inference/request-queue.d.ts.map +1 -0
- package/dist/{mlx-ml/process/queue.js → local-inference/request-queue.js} +83 -56
- package/dist/local-inference/request-queue.js.map +1 -0
- package/dist/local-inference/stream-utils.d.ts +19 -0
- package/dist/local-inference/stream-utils.d.ts.map +1 -0
- package/dist/local-inference/stream-utils.js +76 -0
- package/dist/local-inference/stream-utils.js.map +1 -0
- package/dist/mlx-ml/mlx-cache-support.d.ts +23 -0
- package/dist/mlx-ml/mlx-cache-support.d.ts.map +1 -0
- package/dist/mlx-ml/mlx-cache-support.js +45 -0
- package/dist/mlx-ml/mlx-cache-support.js.map +1 -0
- package/dist/mlx-ml/mlx-driver.d.ts +20 -59
- package/dist/mlx-ml/mlx-driver.d.ts.map +1 -1
- package/dist/mlx-ml/mlx-driver.js +86 -460
- package/dist/mlx-ml/mlx-driver.js.map +1 -1
- package/dist/mlx-ml/mlx-local-inference-adapters.d.ts +3 -0
- package/dist/mlx-ml/mlx-local-inference-adapters.d.ts.map +1 -0
- package/dist/mlx-ml/mlx-local-inference-adapters.js +19 -0
- package/dist/mlx-ml/mlx-local-inference-adapters.js.map +1 -0
- package/dist/mlx-ml/mlx-options.d.ts +19 -0
- package/dist/mlx-ml/mlx-options.d.ts.map +1 -0
- package/dist/mlx-ml/mlx-options.js +30 -0
- package/dist/mlx-ml/mlx-options.js.map +1 -0
- package/dist/mlx-ml/process/index.d.ts +10 -8
- package/dist/mlx-ml/process/index.d.ts.map +1 -1
- package/dist/mlx-ml/process/index.js +73 -54
- package/dist/mlx-ml/process/index.js.map +1 -1
- package/dist/mlx-ml/process/model-specific.d.ts +2 -1
- package/dist/mlx-ml/process/model-specific.d.ts.map +1 -1
- package/dist/mlx-ml/process/model-specific.js.map +1 -1
- package/dist/mlx-ml/process/prompt-builder.d.ts +11 -0
- package/dist/mlx-ml/process/prompt-builder.d.ts.map +1 -0
- package/dist/mlx-ml/process/prompt-builder.js +51 -0
- package/dist/mlx-ml/process/prompt-builder.js.map +1 -0
- package/dist/mlx-ml/process/types.d.ts +15 -183
- package/dist/mlx-ml/process/types.d.ts.map +1 -1
- package/dist/mlx-ml/types.d.ts +2 -45
- package/dist/mlx-ml/types.d.ts.map +1 -1
- package/dist/models-config/index.d.ts +8 -0
- package/dist/models-config/index.d.ts.map +1 -0
- package/dist/models-config/index.js +7 -0
- package/dist/models-config/index.js.map +1 -0
- package/dist/models-config/loader.d.ts +20 -0
- package/dist/models-config/loader.d.ts.map +1 -0
- package/dist/models-config/loader.js +85 -0
- package/dist/models-config/loader.js.map +1 -0
- package/dist/models-config/paths.d.ts +7 -0
- package/dist/models-config/paths.d.ts.map +1 -0
- package/dist/models-config/paths.js +11 -0
- package/dist/models-config/paths.js.map +1 -0
- package/dist/models-config/resolve.d.ts +57 -0
- package/dist/models-config/resolve.d.ts.map +1 -0
- package/dist/models-config/resolve.js +187 -0
- package/dist/models-config/resolve.js.map +1 -0
- package/dist/models-config/types.d.ts +60 -0
- package/dist/models-config/types.d.ts.map +1 -0
- package/dist/models-config/types.js +5 -0
- package/dist/models-config/types.js.map +1 -0
- package/dist/pytorch/process/index.d.ts +35 -0
- package/dist/pytorch/process/index.d.ts.map +1 -0
- package/dist/pytorch/process/index.js +69 -0
- package/dist/pytorch/process/index.js.map +1 -0
- package/dist/pytorch/pytorch-driver.d.ts +35 -0
- package/dist/pytorch/pytorch-driver.d.ts.map +1 -0
- package/dist/pytorch/pytorch-driver.js +48 -0
- package/dist/pytorch/pytorch-driver.js.map +1 -0
- package/dist/pytorch/pytorch-local-inference-adapters.d.ts +3 -0
- package/dist/pytorch/pytorch-local-inference-adapters.d.ts.map +1 -0
- package/dist/pytorch/pytorch-local-inference-adapters.js +19 -0
- package/dist/pytorch/pytorch-local-inference-adapters.js.map +1 -0
- package/dist/pytorch/pytorch-options.d.ts +8 -0
- package/dist/pytorch/pytorch-options.d.ts.map +1 -0
- package/dist/pytorch/pytorch-options.js +21 -0
- package/dist/pytorch/pytorch-options.js.map +1 -0
- package/dist/query-logger.js +1 -1
- package/dist/query-logger.js.map +1 -1
- package/dist/runtime/check.d.ts +10 -0
- package/dist/runtime/check.d.ts.map +1 -0
- package/dist/runtime/check.js +24 -0
- package/dist/runtime/check.js.map +1 -0
- package/dist/runtime/index.d.ts +4 -0
- package/dist/runtime/index.d.ts.map +1 -0
- package/dist/runtime/index.js +4 -0
- package/dist/runtime/index.js.map +1 -0
- package/dist/runtime/manifest-core.d.mts +27 -0
- package/dist/runtime/manifest-core.d.mts.map +1 -0
- package/dist/runtime/manifest-core.mjs +68 -0
- package/dist/runtime/manifest-core.mjs.map +1 -0
- package/dist/runtime/manifest.d.ts +18 -0
- package/dist/runtime/manifest.d.ts.map +1 -0
- package/dist/runtime/manifest.js +9 -0
- package/dist/runtime/manifest.js.map +1 -0
- package/dist/runtime/paths-core.d.mts +17 -0
- package/dist/runtime/paths-core.d.mts.map +1 -0
- package/dist/runtime/paths-core.mjs +67 -0
- package/dist/runtime/paths-core.mjs.map +1 -0
- package/dist/runtime/paths.d.ts +12 -0
- package/dist/runtime/paths.d.ts.map +1 -0
- package/dist/runtime/paths.js +17 -0
- package/dist/runtime/paths.js.map +1 -0
- package/dist/types.d.ts +11 -1
- package/dist/types.d.ts.map +1 -1
- package/dist/types.js.map +1 -1
- package/package.json +10 -7
- package/scripts/download-model.js +25 -9
- package/scripts/runtime-cli.js +315 -0
- package/src/mlx-ml/python/__main__.py +43 -4
- package/src/mlx-ml/python/handlers/__init__.py +2 -1
- package/src/mlx-ml/python/handlers/completion.py +3 -27
- package/src/mlx-ml/python/handlers/{chat.py → generate.py} +33 -106
- package/src/mlx-ml/python/handlers/render.py +40 -0
- package/src/mlx-ml/python/pyproject.toml +3 -2
- package/src/mlx-ml/python/server.py +28 -8
- package/src/mlx-ml/python/utils/template_render.py +80 -0
- package/src/mlx-ml/python/utils/token_utils.py +2 -2
- package/src/mlx-ml/python/uv.lock +549 -454
- package/src/pytorch/python/__main__.py +19 -0
- package/src/pytorch/python/backends/__init__.py +3 -0
- package/src/pytorch/python/backends/base.py +84 -0
- package/src/pytorch/python/backends/transformers_lm.py +127 -0
- package/src/pytorch/python/handlers/__init__.py +6 -0
- package/src/pytorch/python/handlers/cancel.py +53 -0
- package/src/pytorch/python/handlers/capabilities.py +6 -0
- package/src/pytorch/python/handlers/completion.py +15 -0
- package/src/pytorch/python/handlers/format_test.py +70 -0
- package/src/pytorch/python/handlers/generate.py +68 -0
- package/src/pytorch/python/handlers/render.py +40 -0
- package/src/pytorch/python/handlers/tokenize.py +63 -0
- package/src/pytorch/python/pyproject.toml +36 -0
- package/src/pytorch/python/server.py +140 -0
- package/src/pytorch/python/utils/__init__.py +0 -0
- package/src/pytorch/python/utils/chat_template_constraints.py +164 -0
- package/src/pytorch/python/utils/prompt_builder.py +54 -0
- package/src/pytorch/python/utils/template_render.py +80 -0
- package/src/pytorch/python/utils/token_utils.py +376 -0
- package/src/pytorch/python/uv.lock +694 -0
- package/dist/mlx-ml/process/process-communication.d.ts +0 -45
- package/dist/mlx-ml/process/process-communication.d.ts.map +0 -1
- package/dist/mlx-ml/process/process-communication.js.map +0 -1
- package/dist/mlx-ml/process/queue.d.ts +0 -35
- package/dist/mlx-ml/process/queue.d.ts.map +0 -1
- package/dist/mlx-ml/process/queue.js.map +0 -1
- package/scripts/setup-mlx.js +0 -53
|
@@ -0,0 +1,19 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import sys
|
|
3
|
+
|
|
4
|
+
from backends import TransformersLmBackend
|
|
5
|
+
from server import Server
|
|
6
|
+
from utils.token_utils import get_capabilities
|
|
7
|
+
|
|
8
|
+
model_name = sys.argv[1] if len(sys.argv) > 1 else "gpt2"
|
|
9
|
+
device = os.environ.get("PYTORCH_DEVICE", "cpu")
|
|
10
|
+
|
|
11
|
+
if __name__ == "__main__":
|
|
12
|
+
backend = TransformersLmBackend(device=device)
|
|
13
|
+
backend.load(model_name)
|
|
14
|
+
|
|
15
|
+
capabilities = get_capabilities(backend.get_tokenizer())
|
|
16
|
+
capabilities["model_kind"] = "lm"
|
|
17
|
+
|
|
18
|
+
server = Server(backend, capabilities)
|
|
19
|
+
server.run()
|
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
from abc import ABC, abstractmethod
|
|
2
|
+
from typing import Any, Iterator
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
class ModelBackend(ABC):
|
|
6
|
+
"""Abstract base class for model backends."""
|
|
7
|
+
|
|
8
|
+
@abstractmethod
|
|
9
|
+
def load(self, model_name: str) -> None:
|
|
10
|
+
"""Load the target model."""
|
|
11
|
+
raise NotImplementedError
|
|
12
|
+
|
|
13
|
+
@abstractmethod
|
|
14
|
+
def get_tokenizer(self) -> Any:
|
|
15
|
+
"""Return the tokenizer or processor."""
|
|
16
|
+
raise NotImplementedError
|
|
17
|
+
|
|
18
|
+
@abstractmethod
|
|
19
|
+
def stream_generate(
|
|
20
|
+
self, prompt: str | list[int], options: dict, images: list | None = None,
|
|
21
|
+
prompt_cache: list | None = None,
|
|
22
|
+
) -> Iterator[Any]:
|
|
23
|
+
"""Stream generation results."""
|
|
24
|
+
raise NotImplementedError
|
|
25
|
+
|
|
26
|
+
@abstractmethod
|
|
27
|
+
def supports_vision(self) -> bool:
|
|
28
|
+
"""Return whether image input is supported."""
|
|
29
|
+
raise NotImplementedError
|
|
30
|
+
|
|
31
|
+
@property
|
|
32
|
+
@abstractmethod
|
|
33
|
+
def model_kind(self) -> str:
|
|
34
|
+
"""Return "lm" or "vlm"."""
|
|
35
|
+
raise NotImplementedError
|
|
36
|
+
|
|
37
|
+
def load_drafter(self, drafter_model: str) -> None:
|
|
38
|
+
"""Load a drafter model for speculative decoding."""
|
|
39
|
+
raise NotImplementedError(
|
|
40
|
+
f"{type(self).__name__} does not support drafter models"
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
def has_drafter(self) -> bool:
|
|
44
|
+
"""Return whether a drafter model is loaded."""
|
|
45
|
+
return False
|
|
46
|
+
|
|
47
|
+
def cache_prefill(
|
|
48
|
+
self,
|
|
49
|
+
cache_path: str,
|
|
50
|
+
prompt: str,
|
|
51
|
+
base_cache_path: str | None = None,
|
|
52
|
+
trim_to_tokens: int | None = None,
|
|
53
|
+
prefix_offsets: list[int] | None = None,
|
|
54
|
+
prefix_hashes: list[str] | None = None,
|
|
55
|
+
) -> dict:
|
|
56
|
+
"""Build a KV cache from a prompt prefix."""
|
|
57
|
+
raise NotImplementedError(
|
|
58
|
+
f"{type(self).__name__} does not support prompt caching"
|
|
59
|
+
)
|
|
60
|
+
|
|
61
|
+
def load_cache_from_file(self, cache_path: str) -> list | None:
|
|
62
|
+
"""Load a prompt cache from file, or None."""
|
|
63
|
+
return None
|
|
64
|
+
|
|
65
|
+
def get_cache_offset(self, prompt_cache: list) -> int:
|
|
66
|
+
"""Get the number of tokens stored in a loaded prompt cache."""
|
|
67
|
+
if not prompt_cache:
|
|
68
|
+
return 0
|
|
69
|
+
layer0 = prompt_cache[0]
|
|
70
|
+
if hasattr(layer0, 'offset'):
|
|
71
|
+
off = layer0.offset
|
|
72
|
+
return int(off.item() if hasattr(off, 'item') else off)
|
|
73
|
+
if hasattr(layer0, 'caches'):
|
|
74
|
+
for c in layer0.caches:
|
|
75
|
+
if hasattr(c, 'offset'):
|
|
76
|
+
off = c.offset
|
|
77
|
+
return int(off.item() if hasattr(off, 'item') else off)
|
|
78
|
+
try:
|
|
79
|
+
return int(layer0[0].shape[2])
|
|
80
|
+
except Exception:
|
|
81
|
+
pass
|
|
82
|
+
if hasattr(layer0, 'keys') and layer0.keys is not None:
|
|
83
|
+
return int(layer0.keys.shape[2])
|
|
84
|
+
return 0
|
|
@@ -0,0 +1,127 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import os
|
|
4
|
+
from dataclasses import dataclass
|
|
5
|
+
from threading import Thread
|
|
6
|
+
from typing import Any, Iterator
|
|
7
|
+
|
|
8
|
+
import torch
|
|
9
|
+
from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer
|
|
10
|
+
|
|
11
|
+
from backends.base import ModelBackend
|
|
12
|
+
from utils.token_utils import is_eod_token
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
@dataclass
|
|
16
|
+
class StreamChunk:
|
|
17
|
+
text: str
|
|
18
|
+
prompt_tokens: int | None = None
|
|
19
|
+
generation_tokens: int | None = None
|
|
20
|
+
finish_reason: str | None = None
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class TransformersLmBackend(ModelBackend):
|
|
24
|
+
"""Transformers causal LM backend (text-only, CPU-first)."""
|
|
25
|
+
|
|
26
|
+
def __init__(self, device: str | None = None) -> None:
|
|
27
|
+
self.model: Any | None = None
|
|
28
|
+
self.tokenizer: Any | None = None
|
|
29
|
+
self._device_name = device or os.environ.get("PYTORCH_DEVICE", "cpu")
|
|
30
|
+
self._device = torch.device(self._device_name)
|
|
31
|
+
|
|
32
|
+
def load(self, model_name: str) -> None:
|
|
33
|
+
trust_remote_code = os.environ.get("PYTORCH_TRUST_REMOTE_CODE", "").lower() in (
|
|
34
|
+
"1",
|
|
35
|
+
"true",
|
|
36
|
+
"yes",
|
|
37
|
+
)
|
|
38
|
+
self.tokenizer = AutoTokenizer.from_pretrained(
|
|
39
|
+
model_name,
|
|
40
|
+
trust_remote_code=trust_remote_code,
|
|
41
|
+
)
|
|
42
|
+
if self.tokenizer.pad_token is None and self.tokenizer.eos_token is not None:
|
|
43
|
+
self.tokenizer.pad_token = self.tokenizer.eos_token
|
|
44
|
+
|
|
45
|
+
dtype = torch.float32 if self._device.type == "cpu" else torch.float16
|
|
46
|
+
self.model = AutoModelForCausalLM.from_pretrained(
|
|
47
|
+
model_name,
|
|
48
|
+
trust_remote_code=trust_remote_code,
|
|
49
|
+
torch_dtype=dtype,
|
|
50
|
+
)
|
|
51
|
+
self.model.to(self._device)
|
|
52
|
+
self.model.eval()
|
|
53
|
+
|
|
54
|
+
def get_tokenizer(self) -> Any:
|
|
55
|
+
return self.tokenizer
|
|
56
|
+
|
|
57
|
+
def stream_generate(
|
|
58
|
+
self,
|
|
59
|
+
prompt: str | list[int],
|
|
60
|
+
options: dict,
|
|
61
|
+
images: list | None = None,
|
|
62
|
+
prompt_cache: list | None = None,
|
|
63
|
+
) -> Iterator[StreamChunk]:
|
|
64
|
+
if images:
|
|
65
|
+
raise ValueError("TransformersLmBackend does not support vision input")
|
|
66
|
+
if self.model is None or self.tokenizer is None:
|
|
67
|
+
raise RuntimeError("Model is not loaded")
|
|
68
|
+
|
|
69
|
+
final_options = {"max_tokens": 256, **options}
|
|
70
|
+
max_new_tokens = int(final_options.pop("max_tokens", 256))
|
|
71
|
+
temperature = float(final_options.pop("temperature", 1.0))
|
|
72
|
+
top_p = final_options.pop("top_p", None)
|
|
73
|
+
top_k = final_options.pop("top_k", None)
|
|
74
|
+
|
|
75
|
+
if isinstance(prompt, list):
|
|
76
|
+
input_ids = torch.tensor([prompt], device=self._device)
|
|
77
|
+
prompt_token_count = len(prompt)
|
|
78
|
+
else:
|
|
79
|
+
encoded = self.tokenizer(prompt, return_tensors="pt")
|
|
80
|
+
input_ids = encoded["input_ids"].to(self._device)
|
|
81
|
+
prompt_token_count = int(input_ids.shape[-1])
|
|
82
|
+
|
|
83
|
+
do_sample = temperature > 0
|
|
84
|
+
gen_kwargs: dict[str, Any] = {
|
|
85
|
+
"input_ids": input_ids,
|
|
86
|
+
"max_new_tokens": max_new_tokens,
|
|
87
|
+
"do_sample": do_sample,
|
|
88
|
+
}
|
|
89
|
+
if do_sample:
|
|
90
|
+
gen_kwargs["temperature"] = temperature
|
|
91
|
+
if top_p is not None:
|
|
92
|
+
gen_kwargs["top_p"] = float(top_p)
|
|
93
|
+
if top_k is not None:
|
|
94
|
+
gen_kwargs["top_k"] = int(top_k)
|
|
95
|
+
|
|
96
|
+
streamer = TextIteratorStreamer(
|
|
97
|
+
self.tokenizer,
|
|
98
|
+
skip_special_tokens=True,
|
|
99
|
+
skip_prompt=True,
|
|
100
|
+
)
|
|
101
|
+
gen_kwargs["streamer"] = streamer
|
|
102
|
+
|
|
103
|
+
thread = Thread(target=self.model.generate, kwargs=gen_kwargs)
|
|
104
|
+
thread.start()
|
|
105
|
+
|
|
106
|
+
generation_tokens = 0
|
|
107
|
+
for text in streamer:
|
|
108
|
+
generation_tokens += 1
|
|
109
|
+
chunk = StreamChunk(
|
|
110
|
+
text=text,
|
|
111
|
+
prompt_tokens=prompt_token_count if generation_tokens == 1 else None,
|
|
112
|
+
generation_tokens=generation_tokens if generation_tokens == 1 else None,
|
|
113
|
+
)
|
|
114
|
+
if is_eod_token(chunk, self.tokenizer):
|
|
115
|
+
chunk.finish_reason = "stop"
|
|
116
|
+
yield chunk
|
|
117
|
+
break
|
|
118
|
+
yield chunk
|
|
119
|
+
|
|
120
|
+
thread.join()
|
|
121
|
+
|
|
122
|
+
def supports_vision(self) -> bool:
|
|
123
|
+
return False
|
|
124
|
+
|
|
125
|
+
@property
|
|
126
|
+
def model_kind(self) -> str:
|
|
127
|
+
return "lm"
|
|
@@ -0,0 +1,6 @@
|
|
|
1
|
+
from handlers.capabilities import handle_capabilities
|
|
2
|
+
from handlers.completion import handle_completion
|
|
3
|
+
from handlers.format_test import handle_format_test
|
|
4
|
+
from handlers.generate import handle_generate
|
|
5
|
+
from handlers.render import handle_render
|
|
6
|
+
from handlers.tokenize import handle_tokenize
|
|
@@ -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)
|
|
@@ -0,0 +1,68 @@
|
|
|
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
|
+
) -> None:
|
|
16
|
+
if images:
|
|
17
|
+
raise ValueError("PyTorch LIP backend does not support images in Phase 6")
|
|
18
|
+
|
|
19
|
+
if primer is not None:
|
|
20
|
+
print(primer, end="", flush=True)
|
|
21
|
+
|
|
22
|
+
last_response = None
|
|
23
|
+
for response in backend.stream_generate(prompt, options, images):
|
|
24
|
+
if poll_cancel():
|
|
25
|
+
break
|
|
26
|
+
print(response.text.replace("\0", "").replace("\x1e", ""), end="", flush=True)
|
|
27
|
+
last_response = response
|
|
28
|
+
|
|
29
|
+
meta: dict = {}
|
|
30
|
+
if last_response is not None:
|
|
31
|
+
if last_response.prompt_tokens is not None:
|
|
32
|
+
meta["prompt_tokens"] = last_response.prompt_tokens
|
|
33
|
+
if last_response.generation_tokens is not None:
|
|
34
|
+
meta["generation_tokens"] = last_response.generation_tokens
|
|
35
|
+
|
|
36
|
+
if meta:
|
|
37
|
+
print(f"\x1e__META__:{json.dumps(meta)}", end="\0", flush=True)
|
|
38
|
+
else:
|
|
39
|
+
print("", end="\0", flush=True)
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def handle_generate(
|
|
43
|
+
backend: ModelBackend,
|
|
44
|
+
prompt: str | list[int],
|
|
45
|
+
options: dict | None = None,
|
|
46
|
+
images: list | None = None,
|
|
47
|
+
max_image_size: int = 768,
|
|
48
|
+
primer: str | None = None,
|
|
49
|
+
cache_path: str | None = None,
|
|
50
|
+
cache_trim_tokens: int | None = None,
|
|
51
|
+
) -> None:
|
|
52
|
+
"""LIP generate: 整形済み prompt のストリーム推論(KV キャッシュ非対応)"""
|
|
53
|
+
if cache_path or cache_trim_tokens is not None:
|
|
54
|
+
raise ValueError("PyTorch LIP backend does not support prompt caching")
|
|
55
|
+
|
|
56
|
+
if options is None:
|
|
57
|
+
options = {}
|
|
58
|
+
|
|
59
|
+
final_options = dict(options)
|
|
60
|
+
final_options.pop("trust_remote_code", None)
|
|
61
|
+
|
|
62
|
+
_stream_to_stdout(
|
|
63
|
+
backend,
|
|
64
|
+
prompt,
|
|
65
|
+
final_options,
|
|
66
|
+
images=images,
|
|
67
|
+
primer=primer,
|
|
68
|
+
)
|
|
@@ -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,36 @@
|
|
|
1
|
+
[project]
|
|
2
|
+
name = "pytorch_driver"
|
|
3
|
+
version = "0.1.0"
|
|
4
|
+
description = "PyTorch (Transformers) driver for modular-prompt — cpu-minimal runtime"
|
|
5
|
+
requires-python = ">=3.10,<3.14"
|
|
6
|
+
dependencies = [
|
|
7
|
+
"safetensors==0.7.0",
|
|
8
|
+
"tokenizers==0.22.2",
|
|
9
|
+
"transformers==4.57.6",
|
|
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 は CPU wheel を明示インストールする。手動で CUDA 等に差し替える場合はドキュメント参照。
|
|
30
|
+
[[tool.uv.index]]
|
|
31
|
+
name = "pytorch-cpu"
|
|
32
|
+
url = "https://download.pytorch.org/whl/cpu"
|
|
33
|
+
explicit = true
|
|
34
|
+
|
|
35
|
+
[tool.uv.sources]
|
|
36
|
+
torch = { index = "pytorch-cpu" }
|