@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.
Files changed (198) hide show
  1. package/README.md +124 -9
  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 +10 -3
  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 +10 -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 +18 -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 -32
  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 +2 -2
  60. package/dist/models-config/index.d.ts.map +1 -1
  61. package/dist/models-config/index.js +2 -2
  62. package/dist/models-config/index.js.map +1 -1
  63. package/dist/models-config/paths.d.ts +8 -0
  64. package/dist/models-config/paths.d.ts.map +1 -1
  65. package/dist/models-config/paths.js +16 -1
  66. package/dist/models-config/paths.js.map +1 -1
  67. package/dist/models-config/resolve.d.ts +10 -2
  68. package/dist/models-config/resolve.d.ts.map +1 -1
  69. package/dist/models-config/resolve.js +119 -6
  70. package/dist/models-config/resolve.js.map +1 -1
  71. package/dist/models-config/types.d.ts +5 -1
  72. package/dist/models-config/types.d.ts.map +1 -1
  73. package/dist/pytorch/process/index.d.ts +4 -2
  74. package/dist/pytorch/process/index.d.ts.map +1 -1
  75. package/dist/pytorch/process/index.js +24 -7
  76. package/dist/pytorch/process/index.js.map +1 -1
  77. package/dist/pytorch/pytorch-cache-controller.d.ts +84 -0
  78. package/dist/pytorch/pytorch-cache-controller.d.ts.map +1 -0
  79. package/dist/pytorch/pytorch-cache-controller.js +742 -0
  80. package/dist/pytorch/pytorch-cache-controller.js.map +1 -0
  81. package/dist/pytorch/pytorch-cache-support.d.ts +23 -0
  82. package/dist/pytorch/pytorch-cache-support.d.ts.map +1 -0
  83. package/dist/pytorch/pytorch-cache-support.js +47 -0
  84. package/dist/pytorch/pytorch-cache-support.js.map +1 -0
  85. package/dist/pytorch/pytorch-driver.d.ts +8 -1
  86. package/dist/pytorch/pytorch-driver.d.ts.map +1 -1
  87. package/dist/pytorch/pytorch-driver.js +40 -0
  88. package/dist/pytorch/pytorch-driver.js.map +1 -1
  89. package/dist/runtime/check.d.ts.map +1 -1
  90. package/dist/runtime/check.js +9 -6
  91. package/dist/runtime/check.js.map +1 -1
  92. package/dist/runtime/index.d.ts +2 -1
  93. package/dist/runtime/index.d.ts.map +1 -1
  94. package/dist/runtime/index.js +2 -1
  95. package/dist/runtime/index.js.map +1 -1
  96. package/dist/runtime/manifest-core.d.mts +1 -0
  97. package/dist/runtime/manifest-core.mjs +1 -0
  98. package/dist/runtime/manifest-core.mjs.map +1 -1
  99. package/dist/runtime/manifest.d.ts +2 -0
  100. package/dist/runtime/manifest.d.ts.map +1 -1
  101. package/dist/runtime/manifest.js.map +1 -1
  102. package/dist/runtime/paths-core.d.mts +15 -1
  103. package/dist/runtime/paths-core.d.mts.map +1 -1
  104. package/dist/runtime/paths-core.mjs +50 -5
  105. package/dist/runtime/paths-core.mjs.map +1 -1
  106. package/dist/runtime/paths.d.ts +2 -2
  107. package/dist/runtime/paths.d.ts.map +1 -1
  108. package/dist/runtime/paths.js +2 -2
  109. package/dist/runtime/paths.js.map +1 -1
  110. package/dist/runtime/pytorch-template-core.d.mts +11 -0
  111. package/dist/runtime/pytorch-template-core.d.mts.map +1 -0
  112. package/dist/runtime/pytorch-template-core.mjs +54 -0
  113. package/dist/runtime/pytorch-template-core.mjs.map +1 -0
  114. package/dist/runtime/setup-commands-core.d.mts +16 -0
  115. package/dist/runtime/setup-commands-core.d.mts.map +1 -0
  116. package/dist/runtime/setup-commands-core.mjs +18 -0
  117. package/dist/runtime/setup-commands-core.mjs.map +1 -0
  118. package/dist/runtime/setup-commands.d.ts +2 -0
  119. package/dist/runtime/setup-commands.d.ts.map +1 -0
  120. package/dist/runtime/setup-commands.js +2 -0
  121. package/dist/runtime/setup-commands.js.map +1 -0
  122. package/docs/DRIVER_API.md +455 -0
  123. package/docs/LOCAL_MODEL_SETUP.md +765 -0
  124. package/docs/mlx-api-selection.md +301 -0
  125. package/package.json +12 -5
  126. package/scripts/download-model.js +3 -2
  127. package/scripts/runtime-cli.bin.test.ts +142 -0
  128. package/scripts/runtime-cli.js +322 -47
  129. package/scripts/runtime-cli.test.ts +163 -0
  130. package/src/mlx-ml/python/__main__.py +1 -1
  131. package/src/mlx-ml/python/backends/base.py +88 -18
  132. package/src/mlx-ml/python/backends/cache_archive.py +41 -0
  133. package/src/mlx-ml/python/backends/mlx_lm.py +45 -4
  134. package/src/mlx-ml/python/backends/mlx_vlm.py +679 -2
  135. package/src/mlx-ml/python/handlers/cache.py +4 -0
  136. package/src/mlx-ml/python/handlers/generate.py +33 -10
  137. package/src/mlx-ml/python/handlers/tokenize.py +1 -4
  138. package/src/mlx-ml/python/pyproject.toml +9 -3
  139. package/src/mlx-ml/python/server.py +2 -0
  140. package/src/mlx-ml/python/uv.lock +193 -433
  141. package/src/pytorch/templates/cpu-minimal/backends/base.py +139 -0
  142. package/src/pytorch/templates/cpu-minimal/backends/transformers_lm.py +1167 -0
  143. package/src/pytorch/{python → templates/cpu-minimal}/handlers/__init__.py +1 -0
  144. package/src/pytorch/templates/cpu-minimal/handlers/cache.py +88 -0
  145. package/src/pytorch/templates/cpu-minimal/handlers/generate.py +157 -0
  146. package/src/pytorch/{python → templates/cpu-minimal}/pyproject.toml +2 -2
  147. package/src/pytorch/{python → templates/cpu-minimal}/server.py +20 -2
  148. package/src/pytorch/templates/cpu-minimal/tests/test_cache_handler.py +284 -0
  149. package/src/pytorch/templates/cpu-minimal/tests/test_capabilities.py +14 -0
  150. package/src/pytorch/templates/cpu-minimal/tests/test_server.py +141 -0
  151. package/src/pytorch/templates/cpu-minimal/tests/test_transformers_errors.py +89 -0
  152. package/src/pytorch/templates/cpu-minimal/tests/test_transformers_lm_cache.py +554 -0
  153. package/src/pytorch/templates/cpu-minimal/utils/__init__.py +0 -0
  154. package/src/pytorch/{python → templates/cpu-minimal}/utils/token_utils.py +2 -2
  155. package/src/pytorch/templates/cpu-minimal/utils/transformers_errors.py +54 -0
  156. package/src/pytorch/{python → templates/cpu-minimal}/uv.lock +149 -109
  157. package/src/pytorch/templates/cuda/__main__.py +19 -0
  158. package/src/pytorch/templates/cuda/backends/__init__.py +3 -0
  159. package/src/pytorch/{python → templates/cuda}/backends/base.py +54 -6
  160. package/src/pytorch/templates/cuda/backends/transformers_lm.py +379 -0
  161. package/src/pytorch/templates/cuda/handlers/__init__.py +7 -0
  162. package/src/pytorch/templates/cuda/handlers/cache.py +93 -0
  163. package/src/pytorch/templates/cuda/handlers/cancel.py +53 -0
  164. package/src/pytorch/templates/cuda/handlers/capabilities.py +6 -0
  165. package/src/pytorch/templates/cuda/handlers/completion.py +15 -0
  166. package/src/pytorch/templates/cuda/handlers/format_test.py +70 -0
  167. package/src/pytorch/templates/cuda/handlers/generate.py +152 -0
  168. package/src/pytorch/templates/cuda/handlers/render.py +40 -0
  169. package/src/pytorch/templates/cuda/handlers/tokenize.py +63 -0
  170. package/src/pytorch/templates/cuda/pyproject.toml +37 -0
  171. package/src/pytorch/templates/cuda/server.py +158 -0
  172. package/src/pytorch/templates/cuda/tests/__init__.py +0 -0
  173. package/src/pytorch/templates/cuda/tests/test_cache_handler.py +207 -0
  174. package/src/pytorch/templates/cuda/tests/test_capabilities.py +14 -0
  175. package/src/pytorch/templates/cuda/tests/test_server.py +145 -0
  176. package/src/pytorch/templates/cuda/tests/test_transformers_errors.py +89 -0
  177. package/src/pytorch/templates/cuda/tests/test_transformers_lm_cache.py +288 -0
  178. package/src/pytorch/templates/cuda/utils/__init__.py +0 -0
  179. package/src/pytorch/templates/cuda/utils/chat_template_constraints.py +164 -0
  180. package/src/pytorch/templates/cuda/utils/prompt_builder.py +54 -0
  181. package/src/pytorch/templates/cuda/utils/template_render.py +80 -0
  182. package/src/pytorch/templates/cuda/utils/token_utils.py +376 -0
  183. package/src/pytorch/templates/cuda/utils/transformers_errors.py +54 -0
  184. package/src/pytorch/templates/cuda/uv.lock +734 -0
  185. package/src/pytorch/python/backends/transformers_lm.py +0 -127
  186. package/src/pytorch/python/handlers/generate.py +0 -68
  187. /package/src/pytorch/{python → templates/cpu-minimal}/__main__.py +0 -0
  188. /package/src/pytorch/{python → templates/cpu-minimal}/backends/__init__.py +0 -0
  189. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/cancel.py +0 -0
  190. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/capabilities.py +0 -0
  191. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/completion.py +0 -0
  192. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/format_test.py +0 -0
  193. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/render.py +0 -0
  194. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/tokenize.py +0 -0
  195. /package/src/pytorch/{python/utils → templates/cpu-minimal/tests}/__init__.py +0 -0
  196. /package/src/pytorch/{python → templates/cpu-minimal}/utils/chat_template_constraints.py +0 -0
  197. /package/src/pytorch/{python → templates/cpu-minimal}/utils/prompt_builder.py +0 -0
  198. /package/src/pytorch/{python → templates/cpu-minimal}/utils/template_render.py +0 -0
@@ -0,0 +1,139 @@
1
+ from abc import ABC, abstractmethod
2
+ from typing import Any, Iterator
3
+
4
+
5
+ class ModelBackend(ABC):
6
+ """Abstract base class for model backends."""
7
+
8
+ @abstractmethod
9
+ def load(self, model_name: str) -> None:
10
+ """Load the target model."""
11
+ raise NotImplementedError
12
+
13
+ @abstractmethod
14
+ def get_tokenizer(self) -> Any:
15
+ """Return the tokenizer or processor."""
16
+ raise NotImplementedError
17
+
18
+ @abstractmethod
19
+ def stream_generate(
20
+ self, prompt: str | list[int], options: dict, images: list | None = None,
21
+ prompt_cache: Any | None = None,
22
+ ) -> Iterator[Any]:
23
+ """Stream generation results."""
24
+ raise NotImplementedError
25
+
26
+ @abstractmethod
27
+ def supports_vision(self) -> bool:
28
+ """Return whether image input is supported."""
29
+ raise NotImplementedError
30
+
31
+ @property
32
+ @abstractmethod
33
+ def model_kind(self) -> str:
34
+ """Return "lm" or "vlm"."""
35
+ raise NotImplementedError
36
+
37
+ def load_drafter(self, drafter_model: str) -> None:
38
+ """Load a drafter model for speculative decoding."""
39
+ raise NotImplementedError(
40
+ f"{type(self).__name__} does not support drafter models"
41
+ )
42
+
43
+ def has_drafter(self) -> bool:
44
+ """Return whether a drafter model is loaded."""
45
+ return False
46
+
47
+ def cache_prefill(
48
+ self,
49
+ cache_path: str,
50
+ prompt: str,
51
+ base_cache_path: str | None = None,
52
+ trim_to_tokens: int | None = None,
53
+ prefix_offsets: list[int] | None = None,
54
+ prefix_hashes: list[str] | None = None,
55
+ images: list | None = None,
56
+ max_image_size: int = 768,
57
+ ) -> dict:
58
+ """Build a KV cache from a prompt prefix."""
59
+ raise NotImplementedError(
60
+ f"{type(self).__name__} does not support prompt caching"
61
+ )
62
+
63
+ def consume_cache_write_tokens(self, cache_path: str) -> int:
64
+ """Return and clear write usage pending for a cache reference."""
65
+ return 0
66
+
67
+ def trim_cache(self, prompt_cache: Any, tokens: int) -> Any:
68
+ """Remove trailing tokens from a backend-owned prompt cache."""
69
+ raise NotImplementedError(
70
+ f"{type(self).__name__} does not support prompt cache trimming"
71
+ )
72
+
73
+ def tokenize_prompt(
74
+ self,
75
+ prompt: str,
76
+ images: list | None = None,
77
+ max_image_size: int = 768,
78
+ ) -> list[int]:
79
+ """Tokenize a rendered prompt using the backend's prompt rules."""
80
+ tokenizer = self.get_tokenizer()
81
+ bos_token = getattr(tokenizer, "bos_token", None)
82
+ add_special = bos_token is None or not prompt.startswith(bos_token or "")
83
+ token_ids = tokenizer.encode(prompt, add_special_tokens=add_special)
84
+ if hasattr(token_ids, "flatten"):
85
+ token_ids = token_ids.flatten().tolist()
86
+ return [int(token_id) for token_id in token_ids]
87
+
88
+ def load_cache_from_file(
89
+ self,
90
+ cache_path: str,
91
+ images: list | None = None,
92
+ max_image_size: int = 768,
93
+ prompt: str | list[int] | None = None,
94
+ prefix_token_count: int | None = None,
95
+ ) -> Any | None:
96
+ """Load a prompt cache, optionally validating only a prompt prefix."""
97
+ return None
98
+
99
+ def get_cache_offset(self, prompt_cache: Any) -> int:
100
+ """Get the number of tokens stored in a loaded prompt cache."""
101
+ if not prompt_cache:
102
+ return 0
103
+
104
+ get_seq_length = getattr(prompt_cache, "get_seq_length", None)
105
+ if callable(get_seq_length):
106
+ try:
107
+ return int(get_seq_length())
108
+ except Exception:
109
+ pass
110
+
111
+ keys = getattr(prompt_cache, "keys", None)
112
+ if keys is not None:
113
+ try:
114
+ return int(keys.shape[-2])
115
+ except Exception:
116
+ pass
117
+
118
+ layers = getattr(prompt_cache, "layers", None)
119
+ if layers:
120
+ return self.get_cache_offset(layers[0])
121
+
122
+ layer0 = prompt_cache[0]
123
+ if hasattr(layer0, 'offset'):
124
+ off = layer0.offset
125
+ return int(off.item() if hasattr(off, 'item') else off)
126
+ if hasattr(layer0, 'caches'):
127
+ for c in layer0.caches:
128
+ if hasattr(c, 'offset'):
129
+ off = c.offset
130
+ return int(off.item() if hasattr(off, 'item') else off)
131
+ try:
132
+ key = layer0[0] if isinstance(layer0, (list, tuple)) else layer0
133
+ shape = key.shape
134
+ return int(shape[-2] if len(shape) >= 2 else shape[0])
135
+ except Exception:
136
+ pass
137
+ if hasattr(layer0, 'keys') and layer0.keys is not None:
138
+ return int(layer0.keys.shape[-2])
139
+ return 0