@modular-prompt/driver 0.16.0 → 0.17.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 (192) hide show
  1. package/README.md +93 -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.map +1 -1
  5. package/dist/driver-registry/config-based-factory.js +6 -0
  6. package/dist/driver-registry/config-based-factory.js.map +1 -1
  7. package/dist/driver-registry/factory-helper.d.ts.map +1 -1
  8. package/dist/driver-registry/factory-helper.js +9 -2
  9. package/dist/driver-registry/factory-helper.js.map +1 -1
  10. package/dist/driver-registry/index.d.ts +1 -1
  11. package/dist/driver-registry/index.d.ts.map +1 -1
  12. package/dist/driver-registry/types.d.ts +15 -1
  13. package/dist/driver-registry/types.d.ts.map +1 -1
  14. package/dist/formatter/converter.d.ts.map +1 -1
  15. package/dist/formatter/converter.js +31 -2
  16. package/dist/formatter/converter.js.map +1 -1
  17. package/dist/index.d.ts +5 -3
  18. package/dist/index.d.ts.map +1 -1
  19. package/dist/index.js +5 -3
  20. package/dist/index.js.map +1 -1
  21. package/dist/local-inference/adapters.d.ts +6 -0
  22. package/dist/local-inference/adapters.d.ts.map +1 -1
  23. package/dist/local-inference/driver.d.ts.map +1 -1
  24. package/dist/local-inference/driver.js +45 -24
  25. package/dist/local-inference/driver.js.map +1 -1
  26. package/dist/local-inference/process-client.d.ts +4 -2
  27. package/dist/local-inference/process-client.d.ts.map +1 -1
  28. package/dist/local-inference/process-client.js +24 -8
  29. package/dist/local-inference/process-client.js.map +1 -1
  30. package/dist/local-inference/process-communication.d.ts +9 -2
  31. package/dist/local-inference/process-communication.d.ts.map +1 -1
  32. package/dist/local-inference/process-communication.js +37 -5
  33. package/dist/local-inference/process-communication.js.map +1 -1
  34. package/dist/local-inference/protocol.d.ts +4 -0
  35. package/dist/local-inference/protocol.d.ts.map +1 -1
  36. package/dist/local-inference/request-queue.d.ts +1 -1
  37. package/dist/local-inference/request-queue.d.ts.map +1 -1
  38. package/dist/local-inference/request-queue.js +26 -7
  39. package/dist/local-inference/request-queue.js.map +1 -1
  40. package/dist/local-inference/stream-utils.d.ts +6 -0
  41. package/dist/local-inference/stream-utils.d.ts.map +1 -1
  42. package/dist/local-inference/stream-utils.js.map +1 -1
  43. package/dist/mlx-ml/mlx-cache-controller.d.ts +9 -0
  44. package/dist/mlx-ml/mlx-cache-controller.d.ts.map +1 -1
  45. package/dist/mlx-ml/mlx-cache-controller.js +158 -33
  46. package/dist/mlx-ml/mlx-cache-controller.js.map +1 -1
  47. package/dist/mlx-ml/mlx-cache-support.d.ts +2 -2
  48. package/dist/mlx-ml/mlx-cache-support.d.ts.map +1 -1
  49. package/dist/mlx-ml/mlx-cache-support.js +8 -3
  50. package/dist/mlx-ml/mlx-cache-support.js.map +1 -1
  51. package/dist/mlx-ml/mlx-driver.d.ts +0 -1
  52. package/dist/mlx-ml/mlx-driver.d.ts.map +1 -1
  53. package/dist/mlx-ml/mlx-driver.js +1 -8
  54. package/dist/mlx-ml/mlx-driver.js.map +1 -1
  55. package/dist/mlx-ml/process/index.d.ts +1 -1
  56. package/dist/mlx-ml/process/index.d.ts.map +1 -1
  57. package/dist/mlx-ml/process/index.js +2 -2
  58. package/dist/mlx-ml/process/index.js.map +1 -1
  59. package/dist/models-config/index.d.ts +1 -1
  60. package/dist/models-config/index.d.ts.map +1 -1
  61. package/dist/models-config/index.js +1 -1
  62. package/dist/models-config/index.js.map +1 -1
  63. package/dist/models-config/resolve.d.ts +9 -1
  64. package/dist/models-config/resolve.d.ts.map +1 -1
  65. package/dist/models-config/resolve.js +94 -2
  66. package/dist/models-config/resolve.js.map +1 -1
  67. package/dist/models-config/types.d.ts +3 -1
  68. package/dist/models-config/types.d.ts.map +1 -1
  69. package/dist/pytorch/process/index.d.ts +4 -2
  70. package/dist/pytorch/process/index.d.ts.map +1 -1
  71. package/dist/pytorch/process/index.js +24 -7
  72. package/dist/pytorch/process/index.js.map +1 -1
  73. package/dist/pytorch/pytorch-cache-controller.d.ts +84 -0
  74. package/dist/pytorch/pytorch-cache-controller.d.ts.map +1 -0
  75. package/dist/pytorch/pytorch-cache-controller.js +742 -0
  76. package/dist/pytorch/pytorch-cache-controller.js.map +1 -0
  77. package/dist/pytorch/pytorch-cache-support.d.ts +23 -0
  78. package/dist/pytorch/pytorch-cache-support.d.ts.map +1 -0
  79. package/dist/pytorch/pytorch-cache-support.js +47 -0
  80. package/dist/pytorch/pytorch-cache-support.js.map +1 -0
  81. package/dist/pytorch/pytorch-driver.d.ts +8 -1
  82. package/dist/pytorch/pytorch-driver.d.ts.map +1 -1
  83. package/dist/pytorch/pytorch-driver.js +40 -0
  84. package/dist/pytorch/pytorch-driver.js.map +1 -1
  85. package/dist/runtime/check.d.ts.map +1 -1
  86. package/dist/runtime/check.js +8 -6
  87. package/dist/runtime/check.js.map +1 -1
  88. package/dist/runtime/index.d.ts +2 -2
  89. package/dist/runtime/index.d.ts.map +1 -1
  90. package/dist/runtime/index.js +2 -2
  91. package/dist/runtime/index.js.map +1 -1
  92. package/dist/runtime/manifest-core.d.mts +1 -0
  93. package/dist/runtime/manifest-core.mjs +1 -0
  94. package/dist/runtime/manifest-core.mjs.map +1 -1
  95. package/dist/runtime/manifest.d.ts +2 -0
  96. package/dist/runtime/manifest.d.ts.map +1 -1
  97. package/dist/runtime/manifest.js.map +1 -1
  98. package/dist/runtime/paths-core.d.mts +15 -1
  99. package/dist/runtime/paths-core.d.mts.map +1 -1
  100. package/dist/runtime/paths-core.mjs +50 -5
  101. package/dist/runtime/paths-core.mjs.map +1 -1
  102. package/dist/runtime/paths.d.ts +2 -2
  103. package/dist/runtime/paths.d.ts.map +1 -1
  104. package/dist/runtime/paths.js +2 -2
  105. package/dist/runtime/paths.js.map +1 -1
  106. package/dist/runtime/pytorch-template-core.d.mts +11 -0
  107. package/dist/runtime/pytorch-template-core.d.mts.map +1 -0
  108. package/dist/runtime/pytorch-template-core.mjs +54 -0
  109. package/dist/runtime/pytorch-template-core.mjs.map +1 -0
  110. package/dist/runtime/setup-commands-core.d.mts +3 -0
  111. package/dist/runtime/setup-commands-core.d.mts.map +1 -1
  112. package/dist/runtime/setup-commands-core.mjs +4 -0
  113. package/dist/runtime/setup-commands-core.mjs.map +1 -1
  114. package/dist/runtime/setup-commands.d.ts +1 -1
  115. package/dist/runtime/setup-commands.d.ts.map +1 -1
  116. package/dist/runtime/setup-commands.js +1 -1
  117. package/dist/runtime/setup-commands.js.map +1 -1
  118. package/docs/DRIVER_API.md +455 -0
  119. package/docs/LOCAL_MODEL_SETUP.md +765 -0
  120. package/docs/mlx-api-selection.md +301 -0
  121. package/package.json +9 -5
  122. package/scripts/runtime-cli.bin.test.ts +142 -0
  123. package/scripts/runtime-cli.js +305 -35
  124. package/scripts/runtime-cli.test.ts +163 -0
  125. package/src/mlx-ml/python/__main__.py +1 -1
  126. package/src/mlx-ml/python/backends/base.py +88 -18
  127. package/src/mlx-ml/python/backends/mlx_lm.py +28 -3
  128. package/src/mlx-ml/python/backends/mlx_vlm.py +679 -2
  129. package/src/mlx-ml/python/handlers/cache.py +4 -0
  130. package/src/mlx-ml/python/handlers/generate.py +33 -10
  131. package/src/mlx-ml/python/handlers/tokenize.py +1 -4
  132. package/src/mlx-ml/python/pyproject.toml +1 -1
  133. package/src/mlx-ml/python/server.py +2 -0
  134. package/src/mlx-ml/python/uv.lock +8 -8
  135. package/src/pytorch/templates/cpu-minimal/backends/base.py +139 -0
  136. package/src/pytorch/templates/cpu-minimal/backends/transformers_lm.py +1167 -0
  137. package/src/pytorch/{python → templates/cpu-minimal}/handlers/__init__.py +1 -0
  138. package/src/pytorch/templates/cpu-minimal/handlers/cache.py +88 -0
  139. package/src/pytorch/templates/cpu-minimal/handlers/generate.py +157 -0
  140. package/src/pytorch/{python → templates/cpu-minimal}/pyproject.toml +2 -2
  141. package/src/pytorch/{python → templates/cpu-minimal}/server.py +20 -2
  142. package/src/pytorch/templates/cpu-minimal/tests/test_cache_handler.py +284 -0
  143. package/src/pytorch/templates/cpu-minimal/tests/test_capabilities.py +14 -0
  144. package/src/pytorch/templates/cpu-minimal/tests/test_server.py +141 -0
  145. package/src/pytorch/templates/cpu-minimal/tests/test_transformers_errors.py +89 -0
  146. package/src/pytorch/templates/cpu-minimal/tests/test_transformers_lm_cache.py +554 -0
  147. package/src/pytorch/templates/cpu-minimal/utils/__init__.py +0 -0
  148. package/src/pytorch/{python → templates/cpu-minimal}/utils/token_utils.py +2 -2
  149. package/src/pytorch/templates/cpu-minimal/utils/transformers_errors.py +54 -0
  150. package/src/pytorch/{python → templates/cpu-minimal}/uv.lock +149 -109
  151. package/src/pytorch/templates/cuda/__main__.py +19 -0
  152. package/src/pytorch/templates/cuda/backends/__init__.py +3 -0
  153. package/src/pytorch/{python → templates/cuda}/backends/base.py +54 -6
  154. package/src/pytorch/templates/cuda/backends/transformers_lm.py +379 -0
  155. package/src/pytorch/templates/cuda/handlers/__init__.py +7 -0
  156. package/src/pytorch/templates/cuda/handlers/cache.py +93 -0
  157. package/src/pytorch/templates/cuda/handlers/cancel.py +53 -0
  158. package/src/pytorch/templates/cuda/handlers/capabilities.py +6 -0
  159. package/src/pytorch/templates/cuda/handlers/completion.py +15 -0
  160. package/src/pytorch/templates/cuda/handlers/format_test.py +70 -0
  161. package/src/pytorch/templates/cuda/handlers/generate.py +152 -0
  162. package/src/pytorch/templates/cuda/handlers/render.py +40 -0
  163. package/src/pytorch/templates/cuda/handlers/tokenize.py +63 -0
  164. package/src/pytorch/templates/cuda/pyproject.toml +37 -0
  165. package/src/pytorch/templates/cuda/server.py +158 -0
  166. package/src/pytorch/templates/cuda/tests/__init__.py +0 -0
  167. package/src/pytorch/templates/cuda/tests/test_cache_handler.py +207 -0
  168. package/src/pytorch/templates/cuda/tests/test_capabilities.py +14 -0
  169. package/src/pytorch/templates/cuda/tests/test_server.py +145 -0
  170. package/src/pytorch/templates/cuda/tests/test_transformers_errors.py +89 -0
  171. package/src/pytorch/templates/cuda/tests/test_transformers_lm_cache.py +288 -0
  172. package/src/pytorch/templates/cuda/utils/__init__.py +0 -0
  173. package/src/pytorch/templates/cuda/utils/chat_template_constraints.py +164 -0
  174. package/src/pytorch/templates/cuda/utils/prompt_builder.py +54 -0
  175. package/src/pytorch/templates/cuda/utils/template_render.py +80 -0
  176. package/src/pytorch/templates/cuda/utils/token_utils.py +376 -0
  177. package/src/pytorch/templates/cuda/utils/transformers_errors.py +54 -0
  178. package/src/pytorch/templates/cuda/uv.lock +734 -0
  179. package/src/pytorch/python/backends/transformers_lm.py +0 -127
  180. package/src/pytorch/python/handlers/generate.py +0 -68
  181. /package/src/pytorch/{python → templates/cpu-minimal}/__main__.py +0 -0
  182. /package/src/pytorch/{python → templates/cpu-minimal}/backends/__init__.py +0 -0
  183. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/cancel.py +0 -0
  184. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/capabilities.py +0 -0
  185. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/completion.py +0 -0
  186. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/format_test.py +0 -0
  187. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/render.py +0 -0
  188. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/tokenize.py +0 -0
  189. /package/src/pytorch/{python/utils → templates/cpu-minimal/tests}/__init__.py +0 -0
  190. /package/src/pytorch/{python → templates/cpu-minimal}/utils/chat_template_constraints.py +0 -0
  191. /package/src/pytorch/{python → templates/cpu-minimal}/utils/prompt_builder.py +0 -0
  192. /package/src/pytorch/{python → templates/cpu-minimal}/utils/template_render.py +0 -0
@@ -1,15 +1,253 @@
1
1
  from __future__ import annotations
2
2
 
3
+ from copy import deepcopy
4
+ import hashlib
5
+ import json
6
+ import os
7
+ from pathlib import Path
3
8
  import sys
4
9
  from typing import Any, Iterator
5
10
 
6
11
  from mlx_vlm import load as mlx_vlm_load
7
12
  from mlx_vlm import stream_generate as mlx_vlm_stream_generate
8
13
 
14
+ try:
15
+ from mlx_vlm.utils import prepare_inputs as mlx_vlm_prepare_inputs
16
+ from mlx_vlm.utils import should_add_special_tokens as mlx_vlm_should_add_special_tokens
17
+ except ImportError: # pragma: no cover - only older mlx-vlm installations
18
+ mlx_vlm_prepare_inputs = None
19
+ mlx_vlm_should_add_special_tokens = None
20
+
21
+ try:
22
+ from mlx_vlm import VisionFeatureCache
23
+ except ImportError: # pragma: no cover - only older mlx-vlm installations
24
+ VisionFeatureCache = None
25
+
9
26
  from backends.base import ModelBackend
10
27
  from utils.vlm_utils import load_and_resize_images
11
28
 
12
29
 
30
+ VLM_EXACT_CACHE_LAYOUT = "exact_cache_v1"
31
+ VLM_IMAGE_CACHE_LAYOUT = "vision_cache_v1"
32
+ VLM_VISION_FEATURE_CACHE_VERSION = "mlx-vlm-0.7.0"
33
+
34
+
35
+ def _vlm_cache_hash(cache_path: str) -> int:
36
+ """Derive a stable APC exact-cache key from the logical cache path.
37
+
38
+ The TypeScript controller owns the logical path. The key is persisted in
39
+ the sidecar because mlx-vlm's DiskBlockStore names the actual safetensors
40
+ file from this value and the returned file path is different from the
41
+ logical path.
42
+ """
43
+ digest = hashlib.sha256(cache_path.encode("utf-8")).digest()
44
+ return int.from_bytes(digest[:8], "little", signed=True)
45
+
46
+
47
+ def _new_vlm_disk_store(cache_path: str, *, logical_path: bool) -> Any:
48
+ """Open the mlx-vlm 0.7.0 DiskBlockStore for a VLM cache.
49
+
50
+ A logical cache path is used as the store namespace. Once the APC writer
51
+ has produced ``exact_*.safetensors``, the returned path is inside that
52
+ namespace and can be used to reconstruct the same store after restart.
53
+ """
54
+ from mlx_vlm.apc import DiskBlockStore
55
+
56
+ path = Path(cache_path)
57
+ if logical_path:
58
+ root = path.parent
59
+ namespace = path.name
60
+ else:
61
+ namespace_dir = path.parent
62
+ root = namespace_dir.parent
63
+ namespace = namespace_dir.name
64
+ return DiskBlockStore(root, namespace=namespace, num_workers=1)
65
+
66
+
67
+ def _vlm_exact_cache_path(store: Any, cache_hash: int) -> Path:
68
+ """Return the exact snapshot path used by DiskBlockStore 0.7.0.
69
+
70
+ ``_exact_id_for`` is intentionally private in mlx-vlm. Its layout is
71
+ stable in the pinned 0.7.0 API: SHA-256 of the unsigned little-endian
72
+ 64-bit cache hash, truncated to 32 hex characters.
73
+ """
74
+ unsigned_hash = int(cache_hash & ((1 << 64) - 1)).to_bytes(8, "little")
75
+ exact_id = hashlib.sha256(unsigned_hash).hexdigest()[:32]
76
+ return store.dir / f"{store.EXACT_PREFIX}{exact_id}{store.SUFFIX}"
77
+
78
+
79
+ def _vlm_exact_cache_filename(cache_hash: int) -> str:
80
+ """Return the pinned 0.7.0 exact snapshot filename for ``cache_hash``."""
81
+ unsigned_hash = int(cache_hash & ((1 << 64) - 1)).to_bytes(8, "little")
82
+ exact_id = hashlib.sha256(unsigned_hash).hexdigest()[:32]
83
+ return f"exact_{exact_id}.safetensors"
84
+
85
+
86
+ def _read_vlm_cache_meta(cache_path: str) -> dict[str, Any] | None:
87
+ try:
88
+ with open(cache_path + ".meta.json") as f:
89
+ meta = json.load(f)
90
+ if not isinstance(meta, dict) or meta.get("layout") not in {
91
+ VLM_EXACT_CACHE_LAYOUT,
92
+ VLM_IMAGE_CACHE_LAYOUT,
93
+ }:
94
+ return None
95
+ if meta.get("cache_hash") is None or meta.get("token_count") is None:
96
+ return None
97
+ cache_hash = int(meta["cache_hash"])
98
+ if Path(cache_path).name != _vlm_exact_cache_filename(cache_hash):
99
+ return None
100
+ return {
101
+ **meta,
102
+ "cache_hash": cache_hash,
103
+ "token_count": int(meta["token_count"]),
104
+ }
105
+ except (FileNotFoundError, json.JSONDecodeError, ValueError, TypeError):
106
+ return None
107
+
108
+
109
+ def _write_vlm_cache_meta(
110
+ cache_path: str,
111
+ token_count: int,
112
+ cache_hash: int,
113
+ prefix_offsets: list[int] | None = None,
114
+ prefix_hashes: list[str] | None = None,
115
+ *,
116
+ layout: str = VLM_EXACT_CACHE_LAYOUT,
117
+ extra_hash: int = 0,
118
+ images: list[str] | None = None,
119
+ max_image_size: int = 768,
120
+ ) -> None:
121
+ meta: dict[str, Any] = {
122
+ "backend": "mlx-vlm",
123
+ "layout": layout,
124
+ "cache_hash": int(cache_hash),
125
+ "token_count": int(token_count),
126
+ }
127
+ if layout == VLM_IMAGE_CACHE_LAYOUT:
128
+ meta.update({
129
+ "extra_hash": int(extra_hash),
130
+ "image_hash": f"{extra_hash & ((1 << 64) - 1):016x}",
131
+ "image_count": len(images or []),
132
+ "image_refs": list(images or []),
133
+ "max_image_size": int(max_image_size),
134
+ "vision_feature_cache_version": VLM_VISION_FEATURE_CACHE_VERSION,
135
+ })
136
+ if prefix_offsets is not None and prefix_hashes is not None:
137
+ meta["prefix_offsets"] = prefix_offsets
138
+ meta["prefix_hashes"] = prefix_hashes
139
+ with open(cache_path + ".meta.json", "w") as f:
140
+ json.dump(meta, f)
141
+
142
+
143
+ def _image_extra_hash(images: list[Any]) -> int:
144
+ """Hash the normalized image payload used by the VLM cache.
145
+
146
+ mlx-vlm's APC uses an image payload hash as the ``extra_hash`` component
147
+ of an exact cache key. The backend does not expose its internal pixel
148
+ tensor before dispatch, so this fallback hashes the exact PIL payload
149
+ produced by ``load_and_resize_images`` (mode, shape, and bytes). It is
150
+ deterministic across processes and keeps same-token/different-image
151
+ snapshots disjoint.
152
+ """
153
+ digest = hashlib.sha256()
154
+ digest.update(b"modular-prompt:vision-cache-v1\0")
155
+ digest.update(len(images).to_bytes(4, "little", signed=False))
156
+ for image in images:
157
+ mode = str(getattr(image, "mode", "")).encode("utf-8")
158
+ width, height = getattr(image, "size", (0, 0))
159
+ tobytes = getattr(image, "tobytes", None)
160
+ raw = tobytes() if callable(tobytes) else str(image).encode("utf-8")
161
+ digest.update(len(mode).to_bytes(4, "little", signed=False))
162
+ digest.update(mode)
163
+ digest.update(int(width).to_bytes(8, "little", signed=False))
164
+ digest.update(int(height).to_bytes(8, "little", signed=False))
165
+ digest.update(len(raw).to_bytes(8, "little", signed=False))
166
+ digest.update(raw)
167
+ return int.from_bytes(digest.digest()[:8], "little", signed=True)
168
+
169
+
170
+ class _CollisionResistantVisionFeatureCache:
171
+ """Wrap mlx-vlm's process-local cache with a complete image identity.
172
+
173
+ mlx-vlm 0.7.0 hashes only ``PIL.Image.tobytes()`` for PIL inputs. That
174
+ allows images with the same byte payload but different mode or dimensions
175
+ to share an entry. The dispatch API only requires ``get``/``put`` (plus
176
+ the usual cache housekeeping methods), so pass a digest string to the
177
+ pinned cache instead. The digest includes the image position in a list,
178
+ mode, dimensions, and bytes before it reaches upstream's key builder.
179
+ """
180
+
181
+ _KEY_DOMAIN = b"modular-prompt:vision-feature-cache-v2\0"
182
+
183
+ def __init__(self, cache: Any) -> None:
184
+ self._cache = cache
185
+
186
+ @classmethod
187
+ def _source_key(cls, image_source: Any) -> str:
188
+ if isinstance(image_source, (list, tuple)):
189
+ digest = hashlib.sha256()
190
+ digest.update(cls._KEY_DOMAIN)
191
+ digest.update(b"list\0")
192
+ digest.update(len(image_source).to_bytes(4, "little", signed=False))
193
+ for index, item in enumerate(image_source):
194
+ item_key = cls._source_key(item).encode("utf-8")
195
+ digest.update(index.to_bytes(4, "little", signed=False))
196
+ digest.update(len(item_key).to_bytes(4, "little", signed=False))
197
+ digest.update(item_key)
198
+ return f"modular-prompt:vision:list:{digest.hexdigest()}"
199
+
200
+ if isinstance(image_source, str):
201
+ digest = hashlib.sha256()
202
+ digest.update(cls._KEY_DOMAIN)
203
+ digest.update(b"source\0")
204
+ raw_source = image_source.encode("utf-8")
205
+ digest.update(len(raw_source).to_bytes(8, "little", signed=False))
206
+ digest.update(raw_source)
207
+ return f"modular-prompt:vision:source:{digest.hexdigest()}"
208
+
209
+ tobytes = getattr(image_source, "tobytes", None)
210
+ if callable(tobytes):
211
+ raw = bytes(tobytes())
212
+ mode = str(getattr(image_source, "mode", "")).encode("utf-8")
213
+ size = getattr(image_source, "size", (0, 0))
214
+ width, height = size if isinstance(size, (list, tuple)) else (0, 0)
215
+ shape = getattr(image_source, "shape", ())
216
+ dtype = str(getattr(image_source, "dtype", "")).encode("utf-8")
217
+ digest = hashlib.sha256()
218
+ digest.update(cls._KEY_DOMAIN)
219
+ digest.update(b"image\0")
220
+ digest.update(len(mode).to_bytes(4, "little", signed=False))
221
+ digest.update(mode)
222
+ digest.update(int(width).to_bytes(8, "little", signed=False))
223
+ digest.update(int(height).to_bytes(8, "little", signed=False))
224
+ shape_bytes = repr(tuple(shape)).encode("utf-8")
225
+ digest.update(len(shape_bytes).to_bytes(4, "little", signed=False))
226
+ digest.update(shape_bytes)
227
+ digest.update(len(dtype).to_bytes(4, "little", signed=False))
228
+ digest.update(dtype)
229
+ digest.update(len(raw).to_bytes(8, "little", signed=False))
230
+ digest.update(raw)
231
+ return f"modular-prompt:vision:image:{digest.hexdigest()}"
232
+
233
+ return f"modular-prompt:vision:object:{type(image_source).__name__}:{id(image_source)}"
234
+
235
+ def get(self, image_source: Any) -> Any:
236
+ return self._cache.get(self._source_key(image_source))
237
+
238
+ def put(self, image_source: Any, features: Any) -> None:
239
+ self._cache.put(self._source_key(image_source), features)
240
+
241
+ def clear(self) -> None:
242
+ self._cache.clear()
243
+
244
+ def __len__(self) -> int:
245
+ return len(self._cache)
246
+
247
+ def __contains__(self, image_source: Any) -> bool:
248
+ return self._cache.__contains__(self._source_key(image_source))
249
+
250
+
13
251
  class MlxVlmBackend(ModelBackend):
14
252
  """`mlx_vlm` backend for vision-language models."""
15
253
 
@@ -19,8 +257,27 @@ class MlxVlmBackend(ModelBackend):
19
257
  self.drafter: Any | None = None
20
258
  self.drafter_kind: str | None = None
21
259
  self.draft_block_size: int | None = None
260
+ # ``mlx-vlm-memory://`` remains a compatibility fallback for callers
261
+ # that provide an explicit Phase 1 ref. The controller's normal path
262
+ # is a DiskBlockStore-backed VLM exact snapshot.
263
+ self._prompt_caches: dict[str, list[Any]] = {}
264
+ self._prompt_cache_meta: dict[str, dict[str, Any]] = {}
265
+ # Projected image features are intentionally process-local. The
266
+ # persisted VLM cache stores prompt/KV state plus a sidecar identity;
267
+ # opaque feature tensors are never mixed with the text-only store.
268
+ self._vision_caches: dict[int, Any] = {}
269
+
270
+ def _get_vision_cache(self, max_image_size: int) -> Any | None:
271
+ if VisionFeatureCache is None:
272
+ return None
273
+ cache = self._vision_caches.get(int(max_image_size))
274
+ if cache is None:
275
+ cache = _CollisionResistantVisionFeatureCache(VisionFeatureCache())
276
+ self._vision_caches[int(max_image_size)] = cache
277
+ return cache
22
278
 
23
279
  def load(self, model_name: str) -> None:
280
+ self._vision_caches.clear()
24
281
  self.model, self.processor = mlx_vlm_load(model_name)
25
282
 
26
283
  def load_drafter(self, drafter_model: str) -> None:
@@ -34,6 +291,157 @@ class MlxVlmBackend(ModelBackend):
34
291
  def get_tokenizer(self) -> Any:
35
292
  return self.processor
36
293
 
294
+ def _text_tokenizer(self) -> Any:
295
+ if self.processor is None:
296
+ raise RuntimeError("Model is not loaded")
297
+ return getattr(self.processor, "tokenizer", self.processor)
298
+
299
+ def _tokenize_text_prompt(self, prompt: str) -> list[int]:
300
+ tokenizer = self._text_tokenizer()
301
+ add_special = getattr(tokenizer, "bos_token", None) is None or not prompt.startswith(
302
+ getattr(tokenizer, "bos_token", None) or ""
303
+ )
304
+
305
+ # Match mlx-vlm's model-specific chat-template marker handling when
306
+ # available, while keeping the backend testable with plain processors.
307
+ try:
308
+ from mlx_vlm.utils import should_add_special_tokens
309
+
310
+ model_type = getattr(getattr(self.model, "config", None), "model_type", "")
311
+ add_special = should_add_special_tokens(model_type, self.processor)
312
+ except (ImportError, AttributeError, TypeError):
313
+ pass
314
+
315
+ return list(tokenizer.encode(prompt, add_special_tokens=add_special))
316
+
317
+ def _tokenize_image_prompt(
318
+ self,
319
+ prompt: str,
320
+ images: list[str],
321
+ max_image_size: int,
322
+ ) -> list[int]:
323
+ """Match mlx-vlm's image-aware input preparation for cache offsets.
324
+
325
+ Dynamic-resolution processors expand one image marker into a model-
326
+ dependent number of image tokens. Calling the same 0.7.0
327
+ ``prepare_inputs`` helper as ``stream_generate`` keeps the persisted
328
+ token count and the generation suffix boundary aligned.
329
+ """
330
+ if mlx_vlm_prepare_inputs is None or mlx_vlm_should_add_special_tokens is None:
331
+ # Older mlx-vlm versions did not expose the public helper. The
332
+ # pinned Phase 3 dependency does, but retain the text fallback for
333
+ # compatible installations that cannot build an image cache.
334
+ return self._tokenize_text_prompt(prompt)
335
+
336
+ processed_images = load_and_resize_images(images, max_image_size)
337
+ inputs = self._prepare_image_inputs(prompt, processed_images)
338
+ if inputs is None:
339
+ # Older mlx-vlm installations did not expose the public helper.
340
+ # The pinned Phase 3 dependency does, but retain the text fallback
341
+ # for compatible installations that cannot expand image tokens.
342
+ return self._tokenize_text_prompt(prompt)
343
+ return self._input_ids_from_prepared_inputs(inputs)
344
+
345
+ def _prepare_image_inputs(
346
+ self,
347
+ prompt: str,
348
+ processed_images: list[Any],
349
+ ) -> dict[str, Any] | None:
350
+ if mlx_vlm_prepare_inputs is None or mlx_vlm_should_add_special_tokens is None:
351
+ return None
352
+ model_type = getattr(getattr(self.model, "config", None), "model_type", "")
353
+ return mlx_vlm_prepare_inputs(
354
+ self.processor,
355
+ images=processed_images,
356
+ prompts=prompt,
357
+ image_token_index=getattr(
358
+ getattr(self.model, "config", None), "image_token_index", None
359
+ ),
360
+ add_special_tokens=mlx_vlm_should_add_special_tokens(
361
+ model_type, self.processor
362
+ ),
363
+ )
364
+
365
+ @staticmethod
366
+ def _input_ids_from_prepared_inputs(inputs: dict[str, Any]) -> list[int]:
367
+ input_ids = inputs.get("input_ids")
368
+ if input_ids is None:
369
+ raise ValueError("mlx-vlm image preparation returned no input_ids")
370
+ if hasattr(input_ids, "flatten"):
371
+ input_ids = input_ids.flatten()
372
+ values = input_ids.tolist() if hasattr(input_ids, "tolist") else input_ids
373
+ while values and isinstance(values[0], (list, tuple)):
374
+ values = values[0]
375
+ return [int(value) for value in values]
376
+
377
+ def _matches_cached_prefix(
378
+ self,
379
+ prompt: str | list[int] | None,
380
+ token_ids: list[int] | tuple[int, ...],
381
+ images: list[str] | None,
382
+ max_image_size: int,
383
+ ) -> bool:
384
+ """Check that the request starts with the snapshot's token sequence."""
385
+ if prompt is None or not isinstance(prompt, str):
386
+ # A list prompt is already a suffix selected by the shared
387
+ # generate handler, so there is no complete request to compare.
388
+ return True
389
+ try:
390
+ current_tokens = self.tokenize_prompt(
391
+ prompt,
392
+ images=images,
393
+ max_image_size=max_image_size,
394
+ )
395
+ cached = [int(token) for token in token_ids]
396
+ return len(current_tokens) >= len(cached) and current_tokens[: len(cached)] == cached
397
+ except Exception as e:
398
+ sys.stderr.write(f"Failed to validate VLM cache token prefix: {e}\n")
399
+ return False
400
+
401
+ def tokenize_prompt(
402
+ self,
403
+ prompt: str,
404
+ images: list[str] | None = None,
405
+ max_image_size: int = 768,
406
+ ) -> list[int]:
407
+ if images:
408
+ return self._tokenize_image_prompt(prompt, images, max_image_size)
409
+ return self._tokenize_text_prompt(prompt)
410
+
411
+ def trim_cache(self, prompt_cache: list, tokens: int) -> None:
412
+ super().trim_cache(prompt_cache, tokens)
413
+
414
+ @staticmethod
415
+ def _clone_prompt_cache(prompt_cache: list[Any]) -> list[Any]:
416
+ """Detach a cache before generation mutates it.
417
+
418
+ mlx-vlm's APC adapters provide a model-cache-aware clone for the cache
419
+ classes shipped in 0.7.0. The deepcopy fallback keeps this backend
420
+ usable with compatible custom cache objects.
421
+ """
422
+ try:
423
+ from mlx_vlm.apc_adapters import clone_cache_entry
424
+ import mlx.core as mx
425
+
426
+ eval_targets: list[Any] = []
427
+ cloned: list[Any] = []
428
+ for entry in prompt_cache:
429
+ copy_entry = clone_cache_entry(
430
+ entry,
431
+ min_capacity_tokens=None,
432
+ eval_targets=eval_targets,
433
+ )
434
+ if copy_entry is None:
435
+ raise RuntimeError(
436
+ f"unsupported cache entry: {type(entry).__name__}"
437
+ )
438
+ cloned.append(copy_entry)
439
+ if eval_targets:
440
+ mx.eval(eval_targets)
441
+ return cloned
442
+ except Exception:
443
+ return deepcopy(prompt_cache)
444
+
37
445
  def stream_generate(
38
446
  self, prompt: str | list[int], options: dict, images: list | None = None,
39
447
  prompt_cache: list | None = None,
@@ -48,8 +456,9 @@ class MlxVlmBackend(ModelBackend):
48
456
  top_k = final_options.pop("top_k", 0)
49
457
 
50
458
  processed_images = None
459
+ max_image_size = 768
51
460
  if images:
52
- max_image_size = final_options.pop("max_image_size", 768)
461
+ max_image_size = int(final_options.pop("max_image_size", 768))
53
462
  processed_images = load_and_resize_images(images, max_image_size)
54
463
 
55
464
  draft_kwargs = {}
@@ -59,6 +468,23 @@ class MlxVlmBackend(ModelBackend):
59
468
  if self.draft_block_size is not None:
60
469
  draft_kwargs["draft_block_size"] = self.draft_block_size
61
470
 
471
+ if prompt_cache is not None:
472
+ draft_kwargs["prompt_cache"] = prompt_cache
473
+ if processed_images is not None:
474
+ vision_cache = self._get_vision_cache(max_image_size)
475
+ if vision_cache is not None:
476
+ # mlx-vlm 0.7.0 resolves this cache before model dispatch and
477
+ # supplies cached_image_features to supported VLM models.
478
+ draft_kwargs["vision_cache"] = vision_cache
479
+ if isinstance(prompt, list):
480
+ # mlx-lm accepts token IDs as its prompt argument, while
481
+ # mlx-vlm's public prompt argument is text. Its dispatch path
482
+ # does accept pre-tokenized ``input_ids``; use that path for the
483
+ # suffix left after a cached prefix.
484
+ import mlx.core as mx
485
+
486
+ draft_kwargs["input_ids"] = mx.array([prompt])
487
+
62
488
  # Generate and collect tokens
63
489
  token_count = 0
64
490
  for result in mlx_vlm_stream_generate(
@@ -81,7 +507,6 @@ class MlxVlmBackend(ModelBackend):
81
507
  if accept_lens:
82
508
  avg_accepted = sum(accept_lens) / len(accept_lens)
83
509
  # Show stats unless MLX_NO_STATS environment variable is set
84
- import os
85
510
  if not os.getenv('MLX_NO_STATS'):
86
511
  sys.stderr.write(f"\n[Speculative Decoding Stats]\n")
87
512
  sys.stderr.write(f" Rounds: {len(accept_lens)}\n")
@@ -91,6 +516,258 @@ class MlxVlmBackend(ModelBackend):
91
516
  # Clear for next generation
92
517
  self.drafter.accept_lens = []
93
518
 
519
+ def cache_prefill(
520
+ self,
521
+ cache_path: str,
522
+ prompt: str,
523
+ base_cache_path: str | None = None,
524
+ trim_to_tokens: int | None = None,
525
+ prefix_offsets: list[int] | None = None,
526
+ prefix_hashes: list[str] | None = None,
527
+ images: list[str] | None = None,
528
+ max_image_size: int = 768,
529
+ ) -> dict:
530
+ """Prefill a VLM cache in the current backend process.
531
+
532
+ mlx-vlm owns a cache module separate from mlx-lm. Its 0.7.0
533
+ ``DiskBlockStore.save_exact_cache`` API stores the whole prompt cache
534
+ as an ``exact_cache_v1`` snapshot. Image-bearing snapshots use a
535
+ separate ``vision_cache_v1`` sidecar and namespace, and carry the
536
+ image payload in ``extra_hash``. This is deliberately not the
537
+ ``mlx-lm`` ``.safetensors.zip`` format.
538
+ """
539
+ if self.model is None or self.processor is None:
540
+ raise RuntimeError("Model is not loaded")
541
+ if not isinstance(prompt, str) or not prompt:
542
+ raise ValueError("VLM cache_prefill requires a non-empty text prompt")
543
+
544
+ is_memory_ref = cache_path.startswith("mlx-vlm-memory://")
545
+ if base_cache_path is not None or trim_to_tokens is not None:
546
+ sys.stderr.write(
547
+ "VLM cache_prefill ignores base/trim arguments; "
548
+ "VLM incremental prefill is not implemented in Phase 3.\n"
549
+ )
550
+
551
+ processed_images = load_and_resize_images(images, max_image_size) if images else None
552
+ prepared_image_inputs = (
553
+ self._prepare_image_inputs(prompt, processed_images)
554
+ if processed_images is not None
555
+ else None
556
+ )
557
+ extra_hash = _image_extra_hash(processed_images) if processed_images is not None else 0
558
+ cache_layout = VLM_IMAGE_CACHE_LAYOUT if processed_images is not None else VLM_EXACT_CACHE_LAYOUT
559
+
560
+ from mlx_vlm.models.cache import make_prompt_cache
561
+
562
+ language_model = getattr(self.model, "language_model", None)
563
+ if language_model is None:
564
+ raise RuntimeError("VLM model does not expose language_model")
565
+
566
+ prompt_cache = make_prompt_cache(language_model)
567
+ full_tokens = (
568
+ self._input_ids_from_prepared_inputs(prepared_image_inputs)
569
+ if prepared_image_inputs is not None
570
+ else self.tokenize_prompt(
571
+ prompt,
572
+ images=images,
573
+ max_image_size=max_image_size,
574
+ )
575
+ )
576
+ token_count = len(full_tokens)
577
+ prefill_options: dict[str, Any] = {"max_tokens": 0}
578
+ if images:
579
+ prefill_options["max_image_size"] = max_image_size
580
+ for _ in self.stream_generate(
581
+ prompt,
582
+ prefill_options,
583
+ images=images,
584
+ prompt_cache=prompt_cache,
585
+ ):
586
+ break
587
+
588
+ if is_memory_ref:
589
+ self._prompt_caches[cache_path] = prompt_cache
590
+ self._prompt_cache_meta[cache_path] = {
591
+ "layout": cache_layout,
592
+ "cache_hash": _vlm_cache_hash(cache_path),
593
+ "extra_hash": extra_hash,
594
+ "token_count": token_count,
595
+ "token_ids": full_tokens,
596
+ "image_count": len(images or []),
597
+ "max_image_size": max_image_size,
598
+ }
599
+ if os.getenv("MLX_DEBUG"):
600
+ sys.stderr.write(
601
+ f"VLM cache created in memory: {cache_path} ({token_count} tokens)\n"
602
+ )
603
+ return {"cache_path": cache_path, "token_count": token_count}
604
+
605
+ cache_hash = _vlm_cache_hash(cache_path)
606
+ store = None
607
+ actual_path: Path
608
+ try:
609
+ store = _new_vlm_disk_store(cache_path, logical_path=True)
610
+ actual_path = _vlm_exact_cache_path(store, cache_hash)
611
+ # DiskBlockStore writes asynchronously. Clone and evaluate on the
612
+ # producer thread before handing the snapshot to its writer, as
613
+ # required by mlx-vlm's APC implementation.
614
+ detached_cache = self._clone_prompt_cache(prompt_cache)
615
+ store.save_exact_cache(cache_hash, full_tokens, extra_hash, detached_cache)
616
+ store.close()
617
+ store = None
618
+ if not actual_path.is_file():
619
+ raise FileNotFoundError(
620
+ f"mlx-vlm exact cache writer did not create {actual_path}"
621
+ )
622
+ _write_vlm_cache_meta(
623
+ str(actual_path),
624
+ token_count,
625
+ cache_hash,
626
+ prefix_offsets,
627
+ prefix_hashes,
628
+ layout=cache_layout,
629
+ extra_hash=extra_hash,
630
+ images=images,
631
+ max_image_size=max_image_size,
632
+ )
633
+ finally:
634
+ if store is not None:
635
+ store.close()
636
+
637
+ if os.getenv("MLX_DEBUG"):
638
+ sys.stderr.write(
639
+ f"VLM cache created on disk: {actual_path} ({token_count} tokens)\n"
640
+ )
641
+ return {"cache_path": str(actual_path), "token_count": token_count}
642
+
643
+ def load_cache_from_file(
644
+ self,
645
+ cache_path: str,
646
+ images: list[str] | None = None,
647
+ max_image_size: int = 768,
648
+ prompt: str | list[int] | None = None,
649
+ ) -> list[Any] | None:
650
+ if cache_path.startswith("mlx-vlm-memory://"):
651
+ prompt_cache = self._prompt_caches.get(cache_path)
652
+ if prompt_cache is None:
653
+ sys.stderr.write(
654
+ f"VLM cache ref not found in this process: {cache_path}\n"
655
+ )
656
+ return None
657
+ meta = self._prompt_cache_meta.get(cache_path)
658
+ if images:
659
+ if not meta or meta.get("layout") != VLM_IMAGE_CACHE_LAYOUT:
660
+ sys.stderr.write(
661
+ f"VLM memory cache has no image metadata: {cache_path}\n"
662
+ )
663
+ return None
664
+ try:
665
+ expected_hash = _image_extra_hash(
666
+ load_and_resize_images(images, max_image_size)
667
+ )
668
+ if (
669
+ int(meta.get("extra_hash")) != expected_hash
670
+ or int(meta.get("image_count")) != len(images)
671
+ or int(meta.get("max_image_size")) != int(max_image_size)
672
+ ):
673
+ sys.stderr.write(
674
+ f"VLM memory cache image metadata mismatch: {cache_path}\n"
675
+ )
676
+ return None
677
+ except Exception as e:
678
+ sys.stderr.write(f"Failed to validate VLM memory cache: {e}\n")
679
+ return None
680
+ elif meta and meta.get("layout") == VLM_IMAGE_CACHE_LAYOUT:
681
+ sys.stderr.write(
682
+ f"VLM image cache requires image inputs: {cache_path}\n"
683
+ )
684
+ return None
685
+ cached_token_ids = meta.get("token_ids") if meta else None
686
+ if isinstance(cached_token_ids, list) and not self._matches_cached_prefix(
687
+ prompt,
688
+ cached_token_ids,
689
+ images,
690
+ max_image_size,
691
+ ):
692
+ sys.stderr.write(
693
+ f"VLM memory cache token prefix mismatch: {cache_path}\n"
694
+ )
695
+ return None
696
+ try:
697
+ return self._clone_prompt_cache(prompt_cache)
698
+ except Exception as e:
699
+ sys.stderr.write(f"Failed to clone VLM cache: {e}\n")
700
+ return None
701
+
702
+ meta = _read_vlm_cache_meta(cache_path)
703
+ if meta is None:
704
+ sys.stderr.write(f"VLM cache metadata not found or invalid: {cache_path}\n")
705
+ return None
706
+
707
+ is_image_cache = meta.get("layout") == VLM_IMAGE_CACHE_LAYOUT
708
+ if bool(images) != is_image_cache:
709
+ sys.stderr.write(
710
+ f"VLM cache image/text layout mismatch: {cache_path}\n"
711
+ )
712
+ return None
713
+
714
+ expected_extra_hash = 0
715
+ if is_image_cache:
716
+ try:
717
+ processed_images = load_and_resize_images(images or [], max_image_size)
718
+ expected_extra_hash = _image_extra_hash(processed_images)
719
+ if (
720
+ int(meta.get("extra_hash")) != expected_extra_hash
721
+ or int(meta.get("image_count")) != len(images or [])
722
+ or int(meta.get("max_image_size")) != int(max_image_size)
723
+ or meta.get("vision_feature_cache_version")
724
+ != VLM_VISION_FEATURE_CACHE_VERSION
725
+ ):
726
+ sys.stderr.write(
727
+ f"VLM image cache metadata mismatch: {cache_path}\n"
728
+ )
729
+ return None
730
+ except Exception as e:
731
+ sys.stderr.write(f"Failed to validate VLM image cache: {e}\n")
732
+ return None
733
+
734
+ store = None
735
+ try:
736
+ store = _new_vlm_disk_store(cache_path, logical_path=False)
737
+ expected_path = _vlm_exact_cache_path(store, meta["cache_hash"])
738
+ requested_path = Path(cache_path).resolve()
739
+ if requested_path != expected_path.resolve():
740
+ sys.stderr.write(
741
+ f"VLM exact cache hash does not match snapshot path: {cache_path}\n"
742
+ )
743
+ return None
744
+ if not expected_path.is_file():
745
+ sys.stderr.write(f"VLM exact cache snapshot missing: {cache_path}\n")
746
+ return None
747
+ loaded = store.load_exact_cache(meta["cache_hash"])
748
+ if loaded is None:
749
+ sys.stderr.write(f"VLM exact cache not found: {cache_path}\n")
750
+ return None
751
+ token_ids, extra_hash, prompt_cache = loaded
752
+ if extra_hash != expected_extra_hash or len(token_ids) != meta["token_count"]:
753
+ sys.stderr.write(f"VLM exact cache metadata mismatch: {cache_path}\n")
754
+ return None
755
+ if not self._matches_cached_prefix(
756
+ prompt,
757
+ token_ids,
758
+ images,
759
+ max_image_size,
760
+ ):
761
+ sys.stderr.write(f"VLM exact cache token prefix mismatch: {cache_path}\n")
762
+ return None
763
+ return prompt_cache
764
+ except Exception as e:
765
+ sys.stderr.write(f"Failed to load VLM cache: {e}\n")
766
+ return None
767
+ finally:
768
+ if store is not None:
769
+ store.close()
770
+
94
771
  def supports_vision(self) -> bool:
95
772
  return True
96
773