@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
|
@@ -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
|
|
62
|
-
|
|
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
|
-
|
|
70
|
-
|
|
71
|
-
|
|
72
|
-
|
|
73
|
-
|
|
74
|
-
|
|
75
|
-
|
|
76
|
-
|
|
77
|
-
|
|
78
|
-
|
|
79
|
-
|
|
80
|
-
|
|
81
|
-
|
|
82
|
-
|
|
83
|
-
|
|
84
|
-
|
|
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)
|
|
@@ -0,0 +1,41 @@
|
|
|
1
|
+
"""Streaming zip storage for MLX prompt caches."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import zipfile
|
|
6
|
+
from collections.abc import Callable
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
CACHE_ENTRY_NAME = "prompt_cache.safetensors"
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def save_prompt_cache_zip(
|
|
14
|
+
file_name: str,
|
|
15
|
+
cache: Any,
|
|
16
|
+
save_impl: Callable[[Any, Any], None],
|
|
17
|
+
) -> None:
|
|
18
|
+
"""Save a prompt cache into a compressed zip member as it is produced.
|
|
19
|
+
|
|
20
|
+
``save_impl`` receives the zip member's writable stream. It must write
|
|
21
|
+
the safetensors payload to that stream instead of creating an intermediate
|
|
22
|
+
uncompressed file.
|
|
23
|
+
"""
|
|
24
|
+
with zipfile.ZipFile(
|
|
25
|
+
file_name,
|
|
26
|
+
mode="w",
|
|
27
|
+
compression=zipfile.ZIP_DEFLATED,
|
|
28
|
+
allowZip64=True,
|
|
29
|
+
) as archive:
|
|
30
|
+
with archive.open(CACHE_ENTRY_NAME, mode="w", force_zip64=True) as cache_file:
|
|
31
|
+
save_impl(cache_file, cache)
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def load_prompt_cache_zip(
|
|
35
|
+
file_name: str,
|
|
36
|
+
load_impl: Callable[[Any], Any],
|
|
37
|
+
) -> Any:
|
|
38
|
+
"""Load a prompt cache from the safetensors member of a zip archive."""
|
|
39
|
+
with zipfile.ZipFile(file_name, mode="r") as archive:
|
|
40
|
+
with archive.open(CACHE_ENTRY_NAME, mode="r") as cache_file:
|
|
41
|
+
return load_impl(cache_file)
|
|
@@ -7,13 +7,29 @@ from typing import Any, Iterator
|
|
|
7
7
|
|
|
8
8
|
from mlx_lm import load as mlx_lm_load
|
|
9
9
|
from mlx_lm import stream_generate as mlx_lm_stream_generate
|
|
10
|
-
from mlx_lm.models.cache import
|
|
10
|
+
from mlx_lm.models.cache import (
|
|
11
|
+
make_prompt_cache,
|
|
12
|
+
save_prompt_cache as mlx_save_prompt_cache,
|
|
13
|
+
load_prompt_cache as mlx_load_prompt_cache,
|
|
14
|
+
trim_prompt_cache,
|
|
15
|
+
)
|
|
11
16
|
from mlx_lm.sample_utils import make_sampler
|
|
12
17
|
|
|
13
18
|
from backends.base import ModelBackend
|
|
19
|
+
from backends.cache_archive import load_prompt_cache_zip, save_prompt_cache_zip
|
|
14
20
|
from utils.token_utils import is_eod_token
|
|
15
21
|
|
|
16
22
|
|
|
23
|
+
def save_prompt_cache(file_name: str, cache: Any) -> None:
|
|
24
|
+
"""Save a prompt cache as a compressed archive without an uncompressed copy."""
|
|
25
|
+
save_prompt_cache_zip(file_name, cache, mlx_save_prompt_cache)
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def load_prompt_cache(file_name: str) -> list:
|
|
29
|
+
"""Load a prompt cache from its compressed safetensors archive."""
|
|
30
|
+
return load_prompt_cache_zip(file_name, mlx_load_prompt_cache)
|
|
31
|
+
|
|
32
|
+
|
|
17
33
|
class MlxLmBackend(ModelBackend):
|
|
18
34
|
"""`mlx_lm` backend for text-only models."""
|
|
19
35
|
|
|
@@ -62,12 +78,29 @@ class MlxLmBackend(ModelBackend):
|
|
|
62
78
|
|
|
63
79
|
# get_cache_offset is inherited from ModelBackend base class
|
|
64
80
|
|
|
65
|
-
def _tokenize_prompt(
|
|
81
|
+
def _tokenize_prompt(
|
|
82
|
+
self,
|
|
83
|
+
prompt: str,
|
|
84
|
+
images: list | None = None,
|
|
85
|
+
max_image_size: int = 768,
|
|
86
|
+
) -> list[int]:
|
|
66
87
|
"""Tokenize a prompt string using the same logic as stream_generate."""
|
|
67
88
|
add_special = self.tokenizer.bos_token is None or not prompt.startswith(
|
|
68
89
|
self.tokenizer.bos_token
|
|
69
90
|
)
|
|
70
|
-
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)
|
|
71
104
|
|
|
72
105
|
@staticmethod
|
|
73
106
|
def _write_cache_meta(
|
|
@@ -106,6 +139,8 @@ class MlxLmBackend(ModelBackend):
|
|
|
106
139
|
trim_to_tokens: int | None = None,
|
|
107
140
|
prefix_offsets: list[int] | None = None,
|
|
108
141
|
prefix_hashes: list[str] | None = None,
|
|
142
|
+
images: list | None = None,
|
|
143
|
+
max_image_size: int = 768,
|
|
109
144
|
) -> dict:
|
|
110
145
|
if self.model is None or self.tokenizer is None:
|
|
111
146
|
raise RuntimeError("Model is not loaded")
|
|
@@ -184,7 +219,13 @@ class MlxLmBackend(ModelBackend):
|
|
|
184
219
|
sys.stderr.write(f"Cache created: {cache_path} ({token_count} tokens)\n")
|
|
185
220
|
return {"cache_path": cache_path, "token_count": token_count}
|
|
186
221
|
|
|
187
|
-
def load_cache_from_file(
|
|
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:
|
|
188
229
|
try:
|
|
189
230
|
return load_prompt_cache(cache_path)
|
|
190
231
|
except FileNotFoundError:
|