@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
@@ -52,33 +52,103 @@ class ModelBackend(ABC):
52
52
  trim_to_tokens: int | None = None,
53
53
  prefix_offsets: list[int] | None = None,
54
54
  prefix_hashes: list[str] | None = None,
55
+ images: list | None = None,
56
+ max_image_size: int = 768,
55
57
  ) -> dict:
56
58
  """Build a KV cache from a prompt prefix."""
57
59
  raise NotImplementedError(
58
60
  f"{type(self).__name__} does not support prompt caching"
59
61
  )
60
62
 
61
- def load_cache_from_file(self, cache_path: str) -> list | None:
62
- """Load a prompt cache from file, or None."""
63
+ def tokenize_prompt(
64
+ self,
65
+ prompt: str,
66
+ images: list | None = None,
67
+ max_image_size: int = 768,
68
+ ) -> list[int]:
69
+ """Tokenize a rendered prompt using this backend's prompt rules.
70
+
71
+ A VLM backend returns a processor from ``get_tokenizer()``. Keeping
72
+ this operation on the backend prevents shared handlers from assuming
73
+ that the returned object is a plain tokenizer. ``images`` and
74
+ ``max_image_size`` let multimodal backends reproduce the same image
75
+ placeholder expansion as their generation path; text-only backends
76
+ ignore them.
77
+ """
78
+ tokenizer = self.get_tokenizer()
79
+ tokenizer = getattr(tokenizer, "tokenizer", tokenizer)
80
+ bos_token = getattr(tokenizer, "bos_token", None)
81
+ add_special = bos_token is None or not prompt.startswith(bos_token or "")
82
+ return list(tokenizer.encode(prompt, add_special_tokens=add_special))
83
+
84
+ def trim_cache(self, prompt_cache: list, tokens: int) -> None:
85
+ """Trim a backend-owned prompt cache in place.
86
+
87
+ Cache implementations are backend-specific. The default handles the
88
+ cache objects exposed by mlx-vlm without importing mlx-lm into the
89
+ shared generate handler.
90
+ """
91
+ if tokens <= 0:
92
+ return
93
+
94
+ for entry in prompt_cache:
95
+ trim = getattr(entry, "trim", None)
96
+ if callable(trim):
97
+ trim(tokens)
98
+ continue
99
+
100
+ children = getattr(entry, "caches", None)
101
+ if children is not None:
102
+ for child in children:
103
+ child_trim = getattr(child, "trim", None)
104
+ if callable(child_trim):
105
+ child_trim(tokens)
106
+
107
+ def load_cache_from_file(
108
+ self,
109
+ cache_path: str,
110
+ images: list | None = None,
111
+ max_image_size: int = 768,
112
+ prompt: str | list[int] | None = None,
113
+ ) -> list | None:
114
+ """Load a prompt cache from file, or None.
115
+
116
+ ``prompt`` is optional so backends that persist prompt identities can
117
+ validate the loaded cache against the current request.
118
+ """
63
119
  return None
64
120
 
65
121
  def get_cache_offset(self, prompt_cache: list) -> int:
66
122
  """Get the number of tokens stored in a loaded prompt cache."""
67
123
  if not prompt_cache:
68
124
  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
125
+
126
+ def offset_of(cache: Any) -> int:
127
+ if hasattr(cache, "offset"):
128
+ off = cache.offset
129
+ return int(off.item() if hasattr(off, "item") else off)
130
+
131
+ size = getattr(cache, "size", None)
132
+ if callable(size):
133
+ try:
134
+ return int(size())
135
+ except Exception:
136
+ pass
137
+
138
+ children = getattr(cache, "caches", None)
139
+ if children is not None:
140
+ return max((offset_of(child) for child in children), default=0)
141
+
142
+ keys = getattr(cache, "keys", None)
143
+ if keys is not None:
144
+ try:
145
+ return int(keys.shape[-2])
146
+ except Exception:
147
+ pass
148
+
149
+ try:
150
+ return int(cache[0].shape[2])
151
+ except Exception:
152
+ return 0
153
+
154
+ return max((offset_of(layer) for layer in prompt_cache), default=0)
@@ -78,12 +78,29 @@ class MlxLmBackend(ModelBackend):
78
78
 
79
79
  # get_cache_offset is inherited from ModelBackend base class
80
80
 
81
- def _tokenize_prompt(self, prompt: str) -> list[int]:
81
+ def _tokenize_prompt(
82
+ self,
83
+ prompt: str,
84
+ images: list | None = None,
85
+ max_image_size: int = 768,
86
+ ) -> list[int]:
82
87
  """Tokenize a prompt string using the same logic as stream_generate."""
83
88
  add_special = self.tokenizer.bos_token is None or not prompt.startswith(
84
89
  self.tokenizer.bos_token
85
90
  )
86
- return self.tokenizer.encode(prompt, add_special_tokens=add_special)
91
+ return list(self.tokenizer.encode(prompt, add_special_tokens=add_special))
92
+
93
+ def tokenize_prompt(
94
+ self,
95
+ prompt: str,
96
+ images: list | None = None,
97
+ max_image_size: int = 768,
98
+ ) -> list[int]:
99
+ return self._tokenize_prompt(prompt, images, max_image_size)
100
+
101
+ def trim_cache(self, prompt_cache: list, tokens: int) -> None:
102
+ if tokens > 0:
103
+ trim_prompt_cache(prompt_cache, tokens)
87
104
 
88
105
  @staticmethod
89
106
  def _write_cache_meta(
@@ -122,6 +139,8 @@ class MlxLmBackend(ModelBackend):
122
139
  trim_to_tokens: int | None = None,
123
140
  prefix_offsets: list[int] | None = None,
124
141
  prefix_hashes: list[str] | None = None,
142
+ images: list | None = None,
143
+ max_image_size: int = 768,
125
144
  ) -> dict:
126
145
  if self.model is None or self.tokenizer is None:
127
146
  raise RuntimeError("Model is not loaded")
@@ -200,7 +219,13 @@ class MlxLmBackend(ModelBackend):
200
219
  sys.stderr.write(f"Cache created: {cache_path} ({token_count} tokens)\n")
201
220
  return {"cache_path": cache_path, "token_count": token_count}
202
221
 
203
- def load_cache_from_file(self, cache_path: str) -> list | None:
222
+ def load_cache_from_file(
223
+ self,
224
+ cache_path: str,
225
+ images: list | None = None,
226
+ max_image_size: int = 768,
227
+ prompt: str | list[int] | None = None,
228
+ ) -> list | None:
204
229
  try:
205
230
  return load_prompt_cache(cache_path)
206
231
  except FileNotFoundError: