@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.
Files changed (205) hide show
  1. package/README.md +58 -5
  2. package/dist/driver-registry/ai-service.d.ts +23 -1
  3. package/dist/driver-registry/ai-service.d.ts.map +1 -1
  4. package/dist/driver-registry/ai-service.js +44 -10
  5. package/dist/driver-registry/ai-service.js.map +1 -1
  6. package/dist/driver-registry/config-based-factory.d.ts.map +1 -1
  7. package/dist/driver-registry/config-based-factory.js +16 -0
  8. package/dist/driver-registry/config-based-factory.js.map +1 -1
  9. package/dist/driver-registry/factory-helper.d.ts +2 -0
  10. package/dist/driver-registry/factory-helper.d.ts.map +1 -1
  11. package/dist/driver-registry/factory-helper.js +21 -3
  12. package/dist/driver-registry/factory-helper.js.map +1 -1
  13. package/dist/driver-registry/index.d.ts +2 -2
  14. package/dist/driver-registry/index.d.ts.map +1 -1
  15. package/dist/driver-registry/index.js +1 -1
  16. package/dist/driver-registry/index.js.map +1 -1
  17. package/dist/driver-registry/registry.d.ts.map +1 -1
  18. package/dist/driver-registry/registry.js +3 -1
  19. package/dist/driver-registry/registry.js.map +1 -1
  20. package/dist/driver-registry/types.d.ts +8 -2
  21. package/dist/driver-registry/types.d.ts.map +1 -1
  22. package/dist/index.d.ts +10 -2
  23. package/dist/index.d.ts.map +1 -1
  24. package/dist/index.js +9 -1
  25. package/dist/index.js.map +1 -1
  26. package/dist/local-inference/adapters.d.ts +66 -0
  27. package/dist/local-inference/adapters.d.ts.map +1 -0
  28. package/dist/local-inference/adapters.js +2 -0
  29. package/dist/local-inference/adapters.js.map +1 -0
  30. package/dist/local-inference/driver.d.ts +51 -0
  31. package/dist/local-inference/driver.d.ts.map +1 -0
  32. package/dist/local-inference/driver.js +309 -0
  33. package/dist/local-inference/driver.js.map +1 -0
  34. package/dist/local-inference/index.d.ts +22 -0
  35. package/dist/local-inference/index.d.ts.map +1 -0
  36. package/dist/local-inference/index.js +17 -0
  37. package/dist/local-inference/index.js.map +1 -0
  38. package/dist/local-inference/process-client.d.ts +50 -0
  39. package/dist/local-inference/process-client.d.ts.map +1 -0
  40. package/dist/local-inference/process-client.js +92 -0
  41. package/dist/local-inference/process-client.js.map +1 -0
  42. package/dist/local-inference/process-communication.d.ts +41 -0
  43. package/dist/local-inference/process-communication.d.ts.map +1 -0
  44. package/dist/{mlx-ml/process → local-inference}/process-communication.js +25 -61
  45. package/dist/local-inference/process-communication.js.map +1 -0
  46. package/dist/local-inference/process-port.d.ts +12 -0
  47. package/dist/local-inference/process-port.d.ts.map +1 -0
  48. package/dist/local-inference/process-port.js +2 -0
  49. package/dist/local-inference/process-port.js.map +1 -0
  50. package/dist/local-inference/prompt-utils.d.ts +6 -0
  51. package/dist/local-inference/prompt-utils.d.ts.map +1 -0
  52. package/dist/local-inference/prompt-utils.js +17 -0
  53. package/dist/local-inference/prompt-utils.js.map +1 -0
  54. package/dist/local-inference/protocol.d.ts +192 -0
  55. package/dist/local-inference/protocol.d.ts.map +1 -0
  56. package/dist/local-inference/protocol.js +2 -0
  57. package/dist/local-inference/protocol.js.map +1 -0
  58. package/dist/local-inference/queue-types.d.ts +54 -0
  59. package/dist/local-inference/queue-types.d.ts.map +1 -0
  60. package/dist/local-inference/queue-types.js +2 -0
  61. package/dist/local-inference/queue-types.js.map +1 -0
  62. package/dist/local-inference/request-queue.d.ts +36 -0
  63. package/dist/local-inference/request-queue.d.ts.map +1 -0
  64. package/dist/{mlx-ml/process/queue.js → local-inference/request-queue.js} +83 -56
  65. package/dist/local-inference/request-queue.js.map +1 -0
  66. package/dist/local-inference/stream-utils.d.ts +19 -0
  67. package/dist/local-inference/stream-utils.d.ts.map +1 -0
  68. package/dist/local-inference/stream-utils.js +76 -0
  69. package/dist/local-inference/stream-utils.js.map +1 -0
  70. package/dist/mlx-ml/mlx-cache-support.d.ts +23 -0
  71. package/dist/mlx-ml/mlx-cache-support.d.ts.map +1 -0
  72. package/dist/mlx-ml/mlx-cache-support.js +45 -0
  73. package/dist/mlx-ml/mlx-cache-support.js.map +1 -0
  74. package/dist/mlx-ml/mlx-driver.d.ts +20 -59
  75. package/dist/mlx-ml/mlx-driver.d.ts.map +1 -1
  76. package/dist/mlx-ml/mlx-driver.js +86 -460
  77. package/dist/mlx-ml/mlx-driver.js.map +1 -1
  78. package/dist/mlx-ml/mlx-local-inference-adapters.d.ts +3 -0
  79. package/dist/mlx-ml/mlx-local-inference-adapters.d.ts.map +1 -0
  80. package/dist/mlx-ml/mlx-local-inference-adapters.js +19 -0
  81. package/dist/mlx-ml/mlx-local-inference-adapters.js.map +1 -0
  82. package/dist/mlx-ml/mlx-options.d.ts +19 -0
  83. package/dist/mlx-ml/mlx-options.d.ts.map +1 -0
  84. package/dist/mlx-ml/mlx-options.js +30 -0
  85. package/dist/mlx-ml/mlx-options.js.map +1 -0
  86. package/dist/mlx-ml/process/index.d.ts +10 -8
  87. package/dist/mlx-ml/process/index.d.ts.map +1 -1
  88. package/dist/mlx-ml/process/index.js +73 -54
  89. package/dist/mlx-ml/process/index.js.map +1 -1
  90. package/dist/mlx-ml/process/model-specific.d.ts +2 -1
  91. package/dist/mlx-ml/process/model-specific.d.ts.map +1 -1
  92. package/dist/mlx-ml/process/model-specific.js.map +1 -1
  93. package/dist/mlx-ml/process/prompt-builder.d.ts +11 -0
  94. package/dist/mlx-ml/process/prompt-builder.d.ts.map +1 -0
  95. package/dist/mlx-ml/process/prompt-builder.js +51 -0
  96. package/dist/mlx-ml/process/prompt-builder.js.map +1 -0
  97. package/dist/mlx-ml/process/types.d.ts +15 -183
  98. package/dist/mlx-ml/process/types.d.ts.map +1 -1
  99. package/dist/mlx-ml/types.d.ts +2 -45
  100. package/dist/mlx-ml/types.d.ts.map +1 -1
  101. package/dist/models-config/index.d.ts +8 -0
  102. package/dist/models-config/index.d.ts.map +1 -0
  103. package/dist/models-config/index.js +7 -0
  104. package/dist/models-config/index.js.map +1 -0
  105. package/dist/models-config/loader.d.ts +20 -0
  106. package/dist/models-config/loader.d.ts.map +1 -0
  107. package/dist/models-config/loader.js +85 -0
  108. package/dist/models-config/loader.js.map +1 -0
  109. package/dist/models-config/paths.d.ts +7 -0
  110. package/dist/models-config/paths.d.ts.map +1 -0
  111. package/dist/models-config/paths.js +11 -0
  112. package/dist/models-config/paths.js.map +1 -0
  113. package/dist/models-config/resolve.d.ts +57 -0
  114. package/dist/models-config/resolve.d.ts.map +1 -0
  115. package/dist/models-config/resolve.js +187 -0
  116. package/dist/models-config/resolve.js.map +1 -0
  117. package/dist/models-config/types.d.ts +60 -0
  118. package/dist/models-config/types.d.ts.map +1 -0
  119. package/dist/models-config/types.js +5 -0
  120. package/dist/models-config/types.js.map +1 -0
  121. package/dist/pytorch/process/index.d.ts +35 -0
  122. package/dist/pytorch/process/index.d.ts.map +1 -0
  123. package/dist/pytorch/process/index.js +69 -0
  124. package/dist/pytorch/process/index.js.map +1 -0
  125. package/dist/pytorch/pytorch-driver.d.ts +35 -0
  126. package/dist/pytorch/pytorch-driver.d.ts.map +1 -0
  127. package/dist/pytorch/pytorch-driver.js +48 -0
  128. package/dist/pytorch/pytorch-driver.js.map +1 -0
  129. package/dist/pytorch/pytorch-local-inference-adapters.d.ts +3 -0
  130. package/dist/pytorch/pytorch-local-inference-adapters.d.ts.map +1 -0
  131. package/dist/pytorch/pytorch-local-inference-adapters.js +19 -0
  132. package/dist/pytorch/pytorch-local-inference-adapters.js.map +1 -0
  133. package/dist/pytorch/pytorch-options.d.ts +8 -0
  134. package/dist/pytorch/pytorch-options.d.ts.map +1 -0
  135. package/dist/pytorch/pytorch-options.js +21 -0
  136. package/dist/pytorch/pytorch-options.js.map +1 -0
  137. package/dist/query-logger.js +1 -1
  138. package/dist/query-logger.js.map +1 -1
  139. package/dist/runtime/check.d.ts +10 -0
  140. package/dist/runtime/check.d.ts.map +1 -0
  141. package/dist/runtime/check.js +24 -0
  142. package/dist/runtime/check.js.map +1 -0
  143. package/dist/runtime/index.d.ts +4 -0
  144. package/dist/runtime/index.d.ts.map +1 -0
  145. package/dist/runtime/index.js +4 -0
  146. package/dist/runtime/index.js.map +1 -0
  147. package/dist/runtime/manifest-core.d.mts +27 -0
  148. package/dist/runtime/manifest-core.d.mts.map +1 -0
  149. package/dist/runtime/manifest-core.mjs +68 -0
  150. package/dist/runtime/manifest-core.mjs.map +1 -0
  151. package/dist/runtime/manifest.d.ts +18 -0
  152. package/dist/runtime/manifest.d.ts.map +1 -0
  153. package/dist/runtime/manifest.js +9 -0
  154. package/dist/runtime/manifest.js.map +1 -0
  155. package/dist/runtime/paths-core.d.mts +17 -0
  156. package/dist/runtime/paths-core.d.mts.map +1 -0
  157. package/dist/runtime/paths-core.mjs +67 -0
  158. package/dist/runtime/paths-core.mjs.map +1 -0
  159. package/dist/runtime/paths.d.ts +12 -0
  160. package/dist/runtime/paths.d.ts.map +1 -0
  161. package/dist/runtime/paths.js +17 -0
  162. package/dist/runtime/paths.js.map +1 -0
  163. package/dist/types.d.ts +11 -1
  164. package/dist/types.d.ts.map +1 -1
  165. package/dist/types.js.map +1 -1
  166. package/package.json +10 -7
  167. package/scripts/download-model.js +25 -9
  168. package/scripts/runtime-cli.js +315 -0
  169. package/src/mlx-ml/python/__main__.py +43 -4
  170. package/src/mlx-ml/python/handlers/__init__.py +2 -1
  171. package/src/mlx-ml/python/handlers/completion.py +3 -27
  172. package/src/mlx-ml/python/handlers/{chat.py → generate.py} +33 -106
  173. package/src/mlx-ml/python/handlers/render.py +40 -0
  174. package/src/mlx-ml/python/pyproject.toml +3 -2
  175. package/src/mlx-ml/python/server.py +28 -8
  176. package/src/mlx-ml/python/utils/template_render.py +80 -0
  177. package/src/mlx-ml/python/utils/token_utils.py +2 -2
  178. package/src/mlx-ml/python/uv.lock +549 -454
  179. package/src/pytorch/python/__main__.py +19 -0
  180. package/src/pytorch/python/backends/__init__.py +3 -0
  181. package/src/pytorch/python/backends/base.py +84 -0
  182. package/src/pytorch/python/backends/transformers_lm.py +127 -0
  183. package/src/pytorch/python/handlers/__init__.py +6 -0
  184. package/src/pytorch/python/handlers/cancel.py +53 -0
  185. package/src/pytorch/python/handlers/capabilities.py +6 -0
  186. package/src/pytorch/python/handlers/completion.py +15 -0
  187. package/src/pytorch/python/handlers/format_test.py +70 -0
  188. package/src/pytorch/python/handlers/generate.py +68 -0
  189. package/src/pytorch/python/handlers/render.py +40 -0
  190. package/src/pytorch/python/handlers/tokenize.py +63 -0
  191. package/src/pytorch/python/pyproject.toml +36 -0
  192. package/src/pytorch/python/server.py +140 -0
  193. package/src/pytorch/python/utils/__init__.py +0 -0
  194. package/src/pytorch/python/utils/chat_template_constraints.py +164 -0
  195. package/src/pytorch/python/utils/prompt_builder.py +54 -0
  196. package/src/pytorch/python/utils/template_render.py +80 -0
  197. package/src/pytorch/python/utils/token_utils.py +376 -0
  198. package/src/pytorch/python/uv.lock +694 -0
  199. package/dist/mlx-ml/process/process-communication.d.ts +0 -45
  200. package/dist/mlx-ml/process/process-communication.d.ts.map +0 -1
  201. package/dist/mlx-ml/process/process-communication.js.map +0 -1
  202. package/dist/mlx-ml/process/queue.d.ts +0 -35
  203. package/dist/mlx-ml/process/queue.d.ts.map +0 -1
  204. package/dist/mlx-ml/process/queue.js.map +0 -1
  205. 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,3 @@
1
+ from backends.transformers_lm import TransformersLmBackend
2
+
3
+ __all__ = ["TransformersLmBackend"]
@@ -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,6 @@
1
+ import json
2
+
3
+
4
+ def handle_capabilities(capabilities: dict) -> None:
5
+ """capabilities API の処理。JSON出力してnull文字で終端"""
6
+ print(json.dumps(capabilities), end="\0", flush=True)
@@ -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" }