@modular-prompt/driver 0.15.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.
- package/README.md +124 -9
- package/dist/cache-controller.d.ts +4 -0
- package/dist/cache-controller.d.ts.map +1 -1
- package/dist/driver-registry/config-based-factory.d.ts.map +1 -1
- package/dist/driver-registry/config-based-factory.js +10 -3
- package/dist/driver-registry/config-based-factory.js.map +1 -1
- package/dist/driver-registry/factory-helper.d.ts.map +1 -1
- package/dist/driver-registry/factory-helper.js +10 -2
- package/dist/driver-registry/factory-helper.js.map +1 -1
- package/dist/driver-registry/index.d.ts +1 -1
- package/dist/driver-registry/index.d.ts.map +1 -1
- package/dist/driver-registry/types.d.ts +18 -1
- package/dist/driver-registry/types.d.ts.map +1 -1
- package/dist/formatter/converter.d.ts.map +1 -1
- package/dist/formatter/converter.js +31 -2
- package/dist/formatter/converter.js.map +1 -1
- package/dist/index.d.ts +5 -3
- package/dist/index.d.ts.map +1 -1
- package/dist/index.js +5 -3
- package/dist/index.js.map +1 -1
- package/dist/local-inference/adapters.d.ts +6 -0
- package/dist/local-inference/adapters.d.ts.map +1 -1
- package/dist/local-inference/driver.d.ts.map +1 -1
- package/dist/local-inference/driver.js +45 -24
- package/dist/local-inference/driver.js.map +1 -1
- package/dist/local-inference/process-client.d.ts +4 -2
- package/dist/local-inference/process-client.d.ts.map +1 -1
- package/dist/local-inference/process-client.js +24 -8
- package/dist/local-inference/process-client.js.map +1 -1
- package/dist/local-inference/process-communication.d.ts +9 -2
- package/dist/local-inference/process-communication.d.ts.map +1 -1
- package/dist/local-inference/process-communication.js +37 -5
- package/dist/local-inference/process-communication.js.map +1 -1
- package/dist/local-inference/protocol.d.ts +4 -0
- package/dist/local-inference/protocol.d.ts.map +1 -1
- package/dist/local-inference/request-queue.d.ts +1 -1
- package/dist/local-inference/request-queue.d.ts.map +1 -1
- package/dist/local-inference/request-queue.js +26 -7
- package/dist/local-inference/request-queue.js.map +1 -1
- package/dist/local-inference/stream-utils.d.ts +6 -0
- package/dist/local-inference/stream-utils.d.ts.map +1 -1
- package/dist/local-inference/stream-utils.js.map +1 -1
- package/dist/mlx-ml/mlx-cache-controller.d.ts +9 -0
- package/dist/mlx-ml/mlx-cache-controller.d.ts.map +1 -1
- package/dist/mlx-ml/mlx-cache-controller.js +158 -32
- package/dist/mlx-ml/mlx-cache-controller.js.map +1 -1
- package/dist/mlx-ml/mlx-cache-support.d.ts +2 -2
- package/dist/mlx-ml/mlx-cache-support.d.ts.map +1 -1
- package/dist/mlx-ml/mlx-cache-support.js +8 -3
- package/dist/mlx-ml/mlx-cache-support.js.map +1 -1
- package/dist/mlx-ml/mlx-driver.d.ts +0 -1
- package/dist/mlx-ml/mlx-driver.d.ts.map +1 -1
- package/dist/mlx-ml/mlx-driver.js +1 -8
- package/dist/mlx-ml/mlx-driver.js.map +1 -1
- package/dist/mlx-ml/process/index.d.ts +1 -1
- package/dist/mlx-ml/process/index.d.ts.map +1 -1
- package/dist/mlx-ml/process/index.js +2 -2
- package/dist/mlx-ml/process/index.js.map +1 -1
- package/dist/models-config/index.d.ts +2 -2
- package/dist/models-config/index.d.ts.map +1 -1
- package/dist/models-config/index.js +2 -2
- package/dist/models-config/index.js.map +1 -1
- package/dist/models-config/paths.d.ts +8 -0
- package/dist/models-config/paths.d.ts.map +1 -1
- package/dist/models-config/paths.js +16 -1
- package/dist/models-config/paths.js.map +1 -1
- package/dist/models-config/resolve.d.ts +10 -2
- package/dist/models-config/resolve.d.ts.map +1 -1
- package/dist/models-config/resolve.js +119 -6
- package/dist/models-config/resolve.js.map +1 -1
- package/dist/models-config/types.d.ts +5 -1
- package/dist/models-config/types.d.ts.map +1 -1
- package/dist/pytorch/process/index.d.ts +4 -2
- package/dist/pytorch/process/index.d.ts.map +1 -1
- package/dist/pytorch/process/index.js +24 -7
- package/dist/pytorch/process/index.js.map +1 -1
- package/dist/pytorch/pytorch-cache-controller.d.ts +84 -0
- package/dist/pytorch/pytorch-cache-controller.d.ts.map +1 -0
- package/dist/pytorch/pytorch-cache-controller.js +742 -0
- package/dist/pytorch/pytorch-cache-controller.js.map +1 -0
- package/dist/pytorch/pytorch-cache-support.d.ts +23 -0
- package/dist/pytorch/pytorch-cache-support.d.ts.map +1 -0
- package/dist/pytorch/pytorch-cache-support.js +47 -0
- package/dist/pytorch/pytorch-cache-support.js.map +1 -0
- package/dist/pytorch/pytorch-driver.d.ts +8 -1
- package/dist/pytorch/pytorch-driver.d.ts.map +1 -1
- package/dist/pytorch/pytorch-driver.js +40 -0
- package/dist/pytorch/pytorch-driver.js.map +1 -1
- package/dist/runtime/check.d.ts.map +1 -1
- package/dist/runtime/check.js +9 -6
- package/dist/runtime/check.js.map +1 -1
- package/dist/runtime/index.d.ts +2 -1
- package/dist/runtime/index.d.ts.map +1 -1
- package/dist/runtime/index.js +2 -1
- package/dist/runtime/index.js.map +1 -1
- package/dist/runtime/manifest-core.d.mts +1 -0
- package/dist/runtime/manifest-core.mjs +1 -0
- package/dist/runtime/manifest-core.mjs.map +1 -1
- package/dist/runtime/manifest.d.ts +2 -0
- package/dist/runtime/manifest.d.ts.map +1 -1
- package/dist/runtime/manifest.js.map +1 -1
- package/dist/runtime/paths-core.d.mts +15 -1
- package/dist/runtime/paths-core.d.mts.map +1 -1
- package/dist/runtime/paths-core.mjs +50 -5
- package/dist/runtime/paths-core.mjs.map +1 -1
- package/dist/runtime/paths.d.ts +2 -2
- package/dist/runtime/paths.d.ts.map +1 -1
- package/dist/runtime/paths.js +2 -2
- package/dist/runtime/paths.js.map +1 -1
- package/dist/runtime/pytorch-template-core.d.mts +11 -0
- package/dist/runtime/pytorch-template-core.d.mts.map +1 -0
- package/dist/runtime/pytorch-template-core.mjs +54 -0
- package/dist/runtime/pytorch-template-core.mjs.map +1 -0
- package/dist/runtime/setup-commands-core.d.mts +16 -0
- package/dist/runtime/setup-commands-core.d.mts.map +1 -0
- package/dist/runtime/setup-commands-core.mjs +18 -0
- package/dist/runtime/setup-commands-core.mjs.map +1 -0
- package/dist/runtime/setup-commands.d.ts +2 -0
- package/dist/runtime/setup-commands.d.ts.map +1 -0
- package/dist/runtime/setup-commands.js +2 -0
- package/dist/runtime/setup-commands.js.map +1 -0
- package/docs/DRIVER_API.md +455 -0
- package/docs/LOCAL_MODEL_SETUP.md +765 -0
- package/docs/mlx-api-selection.md +301 -0
- package/package.json +12 -5
- package/scripts/download-model.js +3 -2
- package/scripts/runtime-cli.bin.test.ts +142 -0
- package/scripts/runtime-cli.js +322 -47
- package/scripts/runtime-cli.test.ts +163 -0
- package/src/mlx-ml/python/__main__.py +1 -1
- package/src/mlx-ml/python/backends/base.py +88 -18
- package/src/mlx-ml/python/backends/cache_archive.py +41 -0
- package/src/mlx-ml/python/backends/mlx_lm.py +45 -4
- package/src/mlx-ml/python/backends/mlx_vlm.py +679 -2
- package/src/mlx-ml/python/handlers/cache.py +4 -0
- package/src/mlx-ml/python/handlers/generate.py +33 -10
- package/src/mlx-ml/python/handlers/tokenize.py +1 -4
- package/src/mlx-ml/python/pyproject.toml +9 -3
- package/src/mlx-ml/python/server.py +2 -0
- package/src/mlx-ml/python/uv.lock +193 -433
- package/src/pytorch/templates/cpu-minimal/backends/base.py +139 -0
- package/src/pytorch/templates/cpu-minimal/backends/transformers_lm.py +1167 -0
- package/src/pytorch/{python → templates/cpu-minimal}/handlers/__init__.py +1 -0
- package/src/pytorch/templates/cpu-minimal/handlers/cache.py +88 -0
- package/src/pytorch/templates/cpu-minimal/handlers/generate.py +157 -0
- package/src/pytorch/{python → templates/cpu-minimal}/pyproject.toml +2 -2
- package/src/pytorch/{python → templates/cpu-minimal}/server.py +20 -2
- package/src/pytorch/templates/cpu-minimal/tests/test_cache_handler.py +284 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_capabilities.py +14 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_server.py +141 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_transformers_errors.py +89 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_transformers_lm_cache.py +554 -0
- package/src/pytorch/templates/cpu-minimal/utils/__init__.py +0 -0
- package/src/pytorch/{python → templates/cpu-minimal}/utils/token_utils.py +2 -2
- package/src/pytorch/templates/cpu-minimal/utils/transformers_errors.py +54 -0
- package/src/pytorch/{python → templates/cpu-minimal}/uv.lock +149 -109
- package/src/pytorch/templates/cuda/__main__.py +19 -0
- package/src/pytorch/templates/cuda/backends/__init__.py +3 -0
- package/src/pytorch/{python → templates/cuda}/backends/base.py +54 -6
- package/src/pytorch/templates/cuda/backends/transformers_lm.py +379 -0
- package/src/pytorch/templates/cuda/handlers/__init__.py +7 -0
- package/src/pytorch/templates/cuda/handlers/cache.py +93 -0
- package/src/pytorch/templates/cuda/handlers/cancel.py +53 -0
- package/src/pytorch/templates/cuda/handlers/capabilities.py +6 -0
- package/src/pytorch/templates/cuda/handlers/completion.py +15 -0
- package/src/pytorch/templates/cuda/handlers/format_test.py +70 -0
- package/src/pytorch/templates/cuda/handlers/generate.py +152 -0
- package/src/pytorch/templates/cuda/handlers/render.py +40 -0
- package/src/pytorch/templates/cuda/handlers/tokenize.py +63 -0
- package/src/pytorch/templates/cuda/pyproject.toml +37 -0
- package/src/pytorch/templates/cuda/server.py +158 -0
- package/src/pytorch/templates/cuda/tests/__init__.py +0 -0
- package/src/pytorch/templates/cuda/tests/test_cache_handler.py +207 -0
- package/src/pytorch/templates/cuda/tests/test_capabilities.py +14 -0
- package/src/pytorch/templates/cuda/tests/test_server.py +145 -0
- package/src/pytorch/templates/cuda/tests/test_transformers_errors.py +89 -0
- package/src/pytorch/templates/cuda/tests/test_transformers_lm_cache.py +288 -0
- package/src/pytorch/templates/cuda/utils/__init__.py +0 -0
- package/src/pytorch/templates/cuda/utils/chat_template_constraints.py +164 -0
- package/src/pytorch/templates/cuda/utils/prompt_builder.py +54 -0
- package/src/pytorch/templates/cuda/utils/template_render.py +80 -0
- package/src/pytorch/templates/cuda/utils/token_utils.py +376 -0
- package/src/pytorch/templates/cuda/utils/transformers_errors.py +54 -0
- package/src/pytorch/templates/cuda/uv.lock +734 -0
- package/src/pytorch/python/backends/transformers_lm.py +0 -127
- package/src/pytorch/python/handlers/generate.py +0 -68
- /package/src/pytorch/{python → templates/cpu-minimal}/__main__.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/backends/__init__.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/cancel.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/capabilities.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/completion.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/format_test.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/render.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/tokenize.py +0 -0
- /package/src/pytorch/{python/utils → templates/cpu-minimal/tests}/__init__.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/utils/chat_template_constraints.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/utils/prompt_builder.py +0 -0
- /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
|
|