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