@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
@@ -4,13 +4,14 @@
4
4
  * Python runtime 管理 CLI
5
5
  *
6
6
  * setup mlx — ~/.modular-prompt/runtimes/mlx に venv を作成
7
- * setup pytorch — ~/.modular-prompt/runtimes/pytorch に cpu-minimal venv を作成
7
+ * setup pytorch — ~/.modular-prompt/runtimes/pytorch に選択した variant の venv を作成
8
+ * sync pytorch — PyTorch runtime のコードと依存を更新
8
9
  * setup --status
9
10
  * cleanup mlx [--yes]
10
11
  * cleanup --all [--yes]
11
12
  */
12
13
 
13
- import { execSync } from 'child_process';
14
+ import { execFileSync, execSync } from 'child_process';
14
15
  import { existsSync, readFileSync, rmSync, mkdirSync } from 'fs';
15
16
  import { createInterface } from 'readline';
16
17
  import { dirname, join } from 'path';
@@ -32,16 +33,29 @@ const {
32
33
  getRuntimeDir,
33
34
  getVenvPath,
34
35
  getMlxPythonDir,
35
- getPytorchPythonDir,
36
+ getPytorchRuntimePythonDir,
37
+ getPytorchTemplateDir,
38
+ PYTORCH_DEFAULT_VARIANT,
36
39
  isRuntimeReady,
37
40
  } = await import(runtimeModuleUrl('paths-core.mjs'));
38
41
 
42
+ const {
43
+ SETUP_MLX_MONOREPO,
44
+ SETUP_PYTORCH_MONOREPO,
45
+ SYNC_PYTORCH_CLI,
46
+ } = await import(runtimeModuleUrl('setup-commands-core.mjs'));
47
+
39
48
  const {
40
49
  collectInstalledPackages,
41
50
  readManifest,
42
51
  writeManifest,
43
52
  } = await import(runtimeModuleUrl('manifest-core.mjs'));
44
53
 
54
+ const {
55
+ seedPytorchTemplate,
56
+ syncPytorchTemplate,
57
+ } = await import(runtimeModuleUrl('pytorch-template-core.mjs'));
58
+
45
59
  function readPackageVersion() {
46
60
  try {
47
61
  const pkg = JSON.parse(readFileSync(join(packageRoot, 'package.json'), 'utf8'));
@@ -90,7 +104,7 @@ function setupMlx() {
90
104
  };
91
105
 
92
106
  try {
93
- execSync('uv venv --python 3.13', { cwd: pythonDir, stdio: 'inherit', env });
107
+ execSync('uv venv --clear --python 3.13', { cwd: pythonDir, stdio: 'inherit', env });
94
108
  execSync('uv pip install -e .', { cwd: pythonDir, stdio: 'inherit', env });
95
109
 
96
110
  writeManifest('mlx', {
@@ -113,22 +127,171 @@ function setupMlx() {
113
127
  }
114
128
 
115
129
  const PYTORCH_CPU_INDEX = 'https://download.pytorch.org/whl/cpu';
130
+ const PYTORCH_CUDA_INDEX_BASE = 'https://download.pytorch.org/whl';
131
+ const PYTORCH_DEFAULT_CUDA_VERSION = '12.4';
132
+ const PYTORCH_TORCH_VERSION = '2.9.1';
116
133
  const PYTORCH_PYTHON_VERSION = '3.12';
117
134
 
118
- function setupPytorch() {
119
- const pythonDir = getPytorchPythonDir(packageRoot);
120
- if (!existsSync(pythonDir)) {
121
- console.error(`❌ PyTorch Python project not found: ${pythonDir}`);
135
+ function resolveVenvPython(venvPath) {
136
+ return process.platform === 'win32'
137
+ ? join(venvPath, 'Scripts', 'python.exe')
138
+ : join(venvPath, 'bin', 'python');
139
+ }
140
+
141
+ function resolveCudaVersion(value = PYTORCH_DEFAULT_CUDA_VERSION) {
142
+ const raw = String(value).trim().toLowerCase();
143
+ const match =
144
+ raw.match(/^(?:cu)?(\d{1,2})\.(\d{1,2})$/) ??
145
+ raw.match(/^(?:cu)?(\d{2})(\d{1,2})$/);
146
+ if (!match) {
147
+ throw new Error(
148
+ `Invalid CUDA version "${value}". Use a version such as 12.4 or cu124.`,
149
+ );
150
+ }
151
+
152
+ const major = Number(match[1]);
153
+ const minor = Number(match[2]);
154
+ if (major < 1 || minor > 99) {
155
+ throw new Error(
156
+ `Invalid CUDA version "${value}". Use a version such as 12.4 or cu124.`,
157
+ );
158
+ }
159
+
160
+ return {
161
+ version: `${major}.${minor}`,
162
+ index: `${PYTORCH_CUDA_INDEX_BASE}/cu${major}${minor}`,
163
+ };
164
+ }
165
+
166
+ function resolvePytorchIndex(variant, cudaVersion) {
167
+ if (variant !== 'cuda') {
168
+ return { index: PYTORCH_CPU_INDEX };
169
+ }
170
+ return resolveCudaVersion(cudaVersion);
171
+ }
172
+
173
+ function hasNvidiaGpu() {
174
+ try {
175
+ const output = execFileSync(
176
+ 'nvidia-smi',
177
+ ['--query-gpu=name', '--format=csv,noheader,nounits'],
178
+ { encoding: 'utf8', stdio: ['ignore', 'pipe', 'ignore'] },
179
+ );
180
+ return output.trim().length > 0;
181
+ } catch {
182
+ return false;
183
+ }
184
+ }
185
+
186
+ function checkCudaAvailability(venvPath) {
187
+ const venvPython = resolveVenvPython(venvPath);
188
+ if (!existsSync(venvPython)) {
189
+ return null;
190
+ }
191
+
192
+ try {
193
+ const output = execFileSync(
194
+ venvPython,
195
+ ['-c', 'import torch; print(torch.cuda.is_available())'],
196
+ { encoding: 'utf8', stdio: ['ignore', 'pipe', 'ignore'] },
197
+ );
198
+ const value = output.trim().split(/\s+/).at(-1)?.toLowerCase();
199
+ if (value === 'true') {
200
+ return true;
201
+ }
202
+ if (value === 'false') {
203
+ return false;
204
+ }
205
+ } catch {
206
+ // A partially installed runtime or an unavailable torch import is reported as unknown.
207
+ }
208
+ return null;
209
+ }
210
+
211
+ function installPytorchProject(
212
+ pythonDir,
213
+ venvPath,
214
+ env,
215
+ { installTorch = false, torchIndex = PYTORCH_CPU_INDEX } = {},
216
+ ) {
217
+ const venvPython = resolveVenvPython(venvPath);
218
+ if (installTorch) {
219
+ execFileSync(
220
+ 'uv',
221
+ [
222
+ 'pip',
223
+ 'install',
224
+ '--python',
225
+ venvPython,
226
+ `torch==${PYTORCH_TORCH_VERSION}`,
227
+ '--index-url',
228
+ torchIndex,
229
+ ],
230
+ { cwd: pythonDir, stdio: 'inherit', env },
231
+ );
232
+ }
233
+ execFileSync(
234
+ 'uv',
235
+ ['pip', 'install', '--python', venvPython, '.'],
236
+ { cwd: pythonDir, stdio: 'inherit', env },
237
+ );
238
+ }
239
+
240
+ function writePytorchManifest(previousManifest, variant, packages, cudaVersion) {
241
+ const manifest = {
242
+ ...(previousManifest ?? {}),
243
+ profile: 'pytorch',
244
+ variant,
245
+ driverVersion,
246
+ platform: previousManifest?.platform ?? process.platform,
247
+ pythonVersion: previousManifest?.pythonVersion ?? PYTORCH_PYTHON_VERSION,
248
+ createdAt: new Date().toISOString(),
249
+ };
250
+
251
+ if (variant === 'cuda') {
252
+ manifest.cudaVersion = resolveCudaVersion(
253
+ cudaVersion ?? previousManifest?.cudaVersion ?? PYTORCH_DEFAULT_CUDA_VERSION,
254
+ ).version;
255
+ } else {
256
+ delete manifest.cudaVersion;
257
+ }
258
+
259
+ if (packages) {
260
+ if (packages.torch) {
261
+ manifest.torchVersion = packages.torch;
262
+ }
263
+ manifest.packages = packages;
264
+ }
265
+ if (!manifest.torchVersion) {
266
+ manifest.torchVersion = PYTORCH_TORCH_VERSION;
267
+ }
268
+ writeManifest('pytorch', manifest);
269
+ }
270
+
271
+ function setupPytorch(variant = PYTORCH_DEFAULT_VARIANT, cudaVersion) {
272
+ const templateDir = getPytorchTemplateDir(packageRoot, variant);
273
+ if (!existsSync(templateDir)) {
274
+ console.error(`❌ PyTorch template not found for variant ${variant}: ${templateDir}`);
122
275
  process.exit(1);
123
276
  }
124
277
 
278
+ const pytorchIndex = resolvePytorchIndex(variant, cudaVersion);
279
+ if (variant === 'cuda' && !hasNvidiaGpu()) {
280
+ console.warn(
281
+ '⚠️ NVIDIA GPU/driver was not detected. Continuing with the CUDA runtime; ' +
282
+ 'verify torch.cuda.is_available() before running inference.',
283
+ );
284
+ }
285
+
286
+ const pythonDir = getPytorchRuntimePythonDir();
125
287
  const venvPath = getVenvPath('pytorch');
126
288
  const runtimeDir = getRuntimeDir('pytorch');
127
289
 
128
- console.log('🚀 Setting up PyTorch runtime (cpu-minimal)...\n');
290
+ console.log(`🚀 Setting up PyTorch runtime (${variant})...\n`);
291
+ console.log(`📁 Template: ${templateDir}`);
129
292
  console.log(`📁 Python project: ${pythonDir}`);
130
293
  console.log(`📁 Runtime venv: ${venvPath}`);
131
- console.log(`📦 torch index: ${PYTORCH_CPU_INDEX}\n`);
294
+ console.log(`📦 torch index: ${pytorchIndex.index}\n`);
132
295
 
133
296
  ensureUv();
134
297
  mkdirSync(runtimeDir, { recursive: true });
@@ -139,34 +302,24 @@ function setupPytorch() {
139
302
  };
140
303
 
141
304
  try {
142
- execSync(`uv venv --python ${PYTORCH_PYTHON_VERSION}`, { cwd: pythonDir, stdio: 'inherit', env });
143
- const venvPython =
144
- process.platform === 'win32'
145
- ? join(venvPath, 'Scripts', 'python.exe')
146
- : join(venvPath, 'bin', 'python');
147
- execSync(`uv pip install --python "${venvPython}" "torch==2.9.1" --index-url ${PYTORCH_CPU_INDEX}`, {
148
- cwd: pythonDir,
149
- stdio: 'inherit',
150
- env,
305
+ seedPytorchTemplate(templateDir, pythonDir);
306
+ execSync(`uv venv --clear --python ${PYTORCH_PYTHON_VERSION}`, { cwd: pythonDir, stdio: 'inherit', env });
307
+ installPytorchProject(pythonDir, venvPath, env, {
308
+ installTorch: true,
309
+ torchIndex: pytorchIndex.index,
151
310
  });
152
- execSync(`uv pip install --python "${venvPython}" -e .`, { cwd: pythonDir, stdio: 'inherit', env });
153
311
 
154
312
  const packages = collectInstalledPackages(pythonDir, venvPath);
155
- writeManifest('pytorch', {
156
- profile: 'pytorch',
157
- variant: 'cpu-minimal',
158
- driverVersion,
159
- platform: process.platform,
160
- pythonVersion: PYTORCH_PYTHON_VERSION,
161
- torchVersion: packages?.torch,
162
- createdAt: new Date().toISOString(),
163
- packages,
164
- });
313
+ writePytorchManifest(null, variant, packages, pytorchIndex.version);
165
314
 
166
- console.log('\n✅ PyTorch runtime setup completed (cpu-minimal).');
315
+ console.log(`\n✅ PyTorch runtime setup completed (${variant}).`);
167
316
  console.log(` Home: ${getModularPromptHome()}`);
168
317
  console.log(' You can now use PyTorchDriver from @modular-prompt/driver');
169
- console.log(' For CUDA / custom environments, see docs/LOCAL_MODEL_SETUP.md');
318
+ const localModelSetupDoc = join(packageRoot, 'docs', 'LOCAL_MODEL_SETUP.md');
319
+ const docHint = existsSync(localModelSetupDoc)
320
+ ? localModelSetupDoc
321
+ : './docs/LOCAL_MODEL_SETUP.md in @modular-prompt/driver';
322
+ console.log(` For CUDA / custom environments, see ${docHint}`);
170
323
  } catch (error) {
171
324
  const message = error instanceof Error ? error.message : String(error);
172
325
  console.error('❌ Failed to setup PyTorch runtime:', message);
@@ -174,11 +327,71 @@ function setupPytorch() {
174
327
  }
175
328
  }
176
329
 
330
+ function syncPytorch(requestedVariant) {
331
+ const pythonDir = getPytorchRuntimePythonDir();
332
+ const venvPath = getVenvPath('pytorch');
333
+ if (!isRuntimeReady('pytorch')) {
334
+ console.error(
335
+ `❌ PyTorch runtime is not ready at ${getRuntimeDir('pytorch')}. ` +
336
+ 'Run: modular-prompt-runtime setup pytorch',
337
+ );
338
+ process.exit(1);
339
+ }
340
+
341
+ const manifest = readManifest('pytorch');
342
+ if (
343
+ requestedVariant &&
344
+ manifest?.variant &&
345
+ requestedVariant !== manifest.variant
346
+ ) {
347
+ console.error(
348
+ `❌ PyTorch runtime variant mismatch: runtime is ${manifest.variant}, ` +
349
+ `but ${requestedVariant} was requested. Sync without --variant or rerun setup pytorch.`,
350
+ );
351
+ process.exit(1);
352
+ }
353
+ const variant = manifest?.variant ?? requestedVariant ?? PYTORCH_DEFAULT_VARIANT;
354
+ const templateDir = getPytorchTemplateDir(packageRoot, variant);
355
+ if (!existsSync(templateDir)) {
356
+ console.error(`❌ PyTorch template not found for variant ${variant}: ${templateDir}`);
357
+ process.exit(1);
358
+ }
359
+
360
+ console.log(`🔄 Syncing PyTorch runtime (${variant})...\n`);
361
+ console.log(`📁 Template: ${templateDir}`);
362
+ console.log(`📁 Python project: ${pythonDir}`);
363
+ console.log(`📁 Runtime venv: ${venvPath}\n`);
364
+
365
+ ensureUv();
366
+ const env = {
367
+ ...process.env,
368
+ UV_PROJECT_ENVIRONMENT: venvPath,
369
+ };
370
+
371
+ try {
372
+ syncPytorchTemplate(templateDir, pythonDir);
373
+ installPytorchProject(pythonDir, venvPath, env);
374
+
375
+ const packages = collectInstalledPackages(pythonDir, venvPath);
376
+ writePytorchManifest(manifest, variant, packages);
377
+
378
+ console.log('\n✅ PyTorch runtime sync completed.');
379
+ console.log(` Runtime driver version: ${driverVersion}`);
380
+ } catch (error) {
381
+ const message = error instanceof Error ? error.message : String(error);
382
+ console.error('❌ Failed to sync PyTorch runtime:', message);
383
+ process.exit(1);
384
+ }
385
+ }
386
+
177
387
  function formatManifestDetail(manifest) {
178
388
  const parts = [`driver ${manifest.driverVersion}`];
179
389
  if (manifest.variant) {
180
390
  parts.push(`variant ${manifest.variant}`);
181
391
  }
392
+ if (manifest.cudaVersion) {
393
+ parts.push(`CUDA ${manifest.cudaVersion}`);
394
+ }
182
395
  const torchVersion = manifest.torchVersion ?? manifest.packages?.torch;
183
396
  if (torchVersion) {
184
397
  parts.push(`torch ${torchVersion}`);
@@ -191,19 +404,35 @@ function printStatus() {
191
404
  console.log(`modular-prompt home: ${getModularPromptHome()}\n`);
192
405
  for (const profile of RUNTIME_PROFILES) {
193
406
  const ready = isRuntimeReady(profile);
194
- const manifest = ready ? readManifest(profile) : null;
195
- const detail = manifest ? formatManifestDetail(manifest) : '';
407
+ const manifest = readManifest(profile);
408
+ const detail = ready && manifest ? formatManifestDetail(manifest) : '';
196
409
  const icon = ready ? '✅' : '❌';
197
410
  const runtimePath = getRuntimeDir(profile);
198
411
  console.log(`${icon} ${profile}: ${ready ? 'ready' : 'not installed'}${detail}`);
199
412
  console.log(` ${runtimePath}`);
413
+ if (profile === 'pytorch' && manifest?.variant === 'cuda') {
414
+ const cudaAvailable = ready ? checkCudaAvailability(getVenvPath('pytorch')) : null;
415
+ const cudaStatus =
416
+ cudaAvailable === true
417
+ ? 'available'
418
+ : cudaAvailable === false
419
+ ? 'unavailable'
420
+ : 'unknown';
421
+ console.log(` CUDA: ${cudaStatus}`);
422
+ }
423
+ if (profile === 'pytorch' && manifest && manifest.driverVersion !== driverVersion) {
424
+ console.log(
425
+ ` ⚠️ driver version differs (installed ${manifest.driverVersion}, current ${driverVersion}). ` +
426
+ `Run: ${SYNC_PYTORCH_CLI}`,
427
+ );
428
+ }
200
429
  }
201
430
  const setupHints = [];
202
431
  if (!isRuntimeReady('mlx') && process.platform === 'darwin') {
203
- setupHints.push('pnpm run setup-mlx -w @modular-prompt/driver');
432
+ setupHints.push(SETUP_MLX_MONOREPO);
204
433
  }
205
434
  if (!isRuntimeReady('pytorch')) {
206
- setupHints.push('pnpm run setup-pytorch -w @modular-prompt/driver');
435
+ setupHints.push(SETUP_PYTORCH_MONOREPO);
207
436
  }
208
437
  if (setupHints.length > 0) {
209
438
  console.log(`\nRun: ${setupHints.join(' or ')}`);
@@ -253,19 +482,49 @@ async function cleanupAll() {
253
482
 
254
483
  function printUsage() {
255
484
  console.log(`Usage:
256
- node scripts/runtime-cli.js setup mlx Set up MLX Python runtime (macOS only)
257
- node scripts/runtime-cli.js setup pytorch Set up PyTorch runtime (cpu-minimal)
258
- node scripts/runtime-cli.js setup --status Show runtime status
259
- node scripts/runtime-cli.js cleanup mlx Remove MLX runtime
260
- node scripts/runtime-cli.js cleanup pytorch Remove PyTorch runtime
261
- node scripts/runtime-cli.js cleanup --all Remove entire ~/.modular-prompt
262
- node scripts/runtime-cli.js cleanup ... --yes Skip confirmation
263
-
264
- npm scripts: setup-mlx, setup-pytorch, runtime:status, runtime:cleanup`);
485
+ modular-prompt-runtime setup mlx Set up MLX Python runtime (macOS only)
486
+ modular-prompt-runtime setup pytorch Set up PyTorch runtime (cpu-minimal)
487
+ modular-prompt-runtime setup pytorch --variant <variant> [--cuda <version>]
488
+ modular-prompt-runtime sync pytorch Sync PyTorch code and dependencies
489
+ modular-prompt-runtime sync pytorch --variant <variant>
490
+ modular-prompt-runtime setup --status Show runtime status
491
+ modular-prompt-runtime cleanup mlx Remove MLX runtime
492
+ modular-prompt-runtime cleanup pytorch Remove PyTorch runtime
493
+ modular-prompt-runtime cleanup --all Remove entire ~/.modular-prompt
494
+ modular-prompt-runtime cleanup ... --yes Skip confirmation
495
+
496
+ npm scripts: setup-mlx, setup-pytorch, runtime:status, runtime:sync-pytorch, runtime:cleanup`);
497
+ }
498
+
499
+ function parseVariant(args) {
500
+ return parseOptionValue(args, '--variant');
501
+ }
502
+
503
+ function parseCudaVersion(args) {
504
+ return parseOptionValue(args, '--cuda');
505
+ }
506
+
507
+ function parseOptionValue(args, option) {
508
+ const inlinePrefix = `${option}=`;
509
+ const inline = args.find((arg) => arg.startsWith(inlinePrefix));
510
+ const index = args.indexOf(option);
511
+ if (!inline && index === -1) {
512
+ return undefined;
513
+ }
514
+ if (inline && index !== -1) {
515
+ throw new Error(`Specify ${option} only once`);
516
+ }
517
+
518
+ const value = inline ? inline.slice(inlinePrefix.length) : args[index + 1];
519
+ if (!value || value.startsWith('-')) {
520
+ throw new Error(`Missing value for ${option}`);
521
+ }
522
+ return value;
265
523
  }
266
524
 
267
525
  async function main() {
268
- const [command, target] = process.argv.slice(2);
526
+ const args = process.argv.slice(2);
527
+ const [command, target] = args;
269
528
 
270
529
  if (!command || command === '--help' || command === '-h') {
271
530
  printUsage();
@@ -282,7 +541,13 @@ async function main() {
282
541
  return;
283
542
  }
284
543
  if (target === 'pytorch') {
285
- setupPytorch();
544
+ const setupArgs = args.slice(2);
545
+ const variant = parseVariant(setupArgs) ?? PYTORCH_DEFAULT_VARIANT;
546
+ const cudaVersion = parseCudaVersion(setupArgs);
547
+ if (cudaVersion && variant !== 'cuda') {
548
+ throw new Error('--cuda can only be used with --variant cuda');
549
+ }
550
+ setupPytorch(variant, cudaVersion);
286
551
  return;
287
552
  }
288
553
  console.error(`Unknown setup target: ${target ?? '(none)'}`);
@@ -290,6 +555,16 @@ async function main() {
290
555
  process.exit(1);
291
556
  }
292
557
 
558
+ if (command === 'sync') {
559
+ if (target === 'pytorch') {
560
+ syncPytorch(parseVariant(args.slice(2)));
561
+ return;
562
+ }
563
+ console.error(`Unknown sync target: ${target ?? '(none)'}`);
564
+ printUsage();
565
+ process.exit(1);
566
+ }
567
+
293
568
  if (command === 'cleanup') {
294
569
  if (target === '--all' || target === 'all') {
295
570
  await cleanupAll();
@@ -0,0 +1,163 @@
1
+ import { execFileSync } from 'node:child_process';
2
+ import {
3
+ chmodSync,
4
+ existsSync,
5
+ mkdirSync,
6
+ mkdtempSync,
7
+ readFileSync,
8
+ rmSync,
9
+ writeFileSync,
10
+ } from 'node:fs';
11
+ import { dirname, join, resolve } from 'node:path';
12
+ import { fileURLToPath } from 'node:url';
13
+ import { tmpdir } from 'node:os';
14
+ import { describe, expect, it } from 'vitest';
15
+
16
+ const packageRoot = resolve(dirname(fileURLToPath(import.meta.url)), '..');
17
+ const runtimeCli = join(packageRoot, 'scripts', 'runtime-cli.js');
18
+ const driverVersion = (
19
+ JSON.parse(readFileSync(join(packageRoot, 'package.json'), 'utf8')) as { version: string }
20
+ ).version;
21
+
22
+ describe.skipIf(process.platform === 'win32')('runtime CLI PyTorch setup and sync', () => {
23
+ it('seeds, preserves customization, and syncs the runtime project', () => {
24
+ const temporaryDirectory = mkdtempSync(join(tmpdir(), 'modular-prompt-pytorch-cli-'));
25
+ const fakeBinDirectory = join(temporaryDirectory, 'bin');
26
+ const fakeUvPath = join(fakeBinDirectory, 'uv');
27
+ const uvLogPath = join(temporaryDirectory, 'uv.log');
28
+ const runtimePythonDir = join(
29
+ temporaryDirectory,
30
+ 'runtimes',
31
+ 'pytorch',
32
+ 'python',
33
+ );
34
+
35
+ try {
36
+ mkdirSync(fakeBinDirectory, { recursive: true });
37
+ writeFileSync(
38
+ fakeUvPath,
39
+ `#!/usr/bin/env node
40
+ import { appendFileSync, mkdirSync, writeFileSync } from 'node:fs';
41
+ import { join } from 'node:path';
42
+
43
+ const args = process.argv.slice(2);
44
+ appendFileSync(process.env.FAKE_UV_LOG, args.join(' ') + '\\n');
45
+
46
+ if (args[0] === '--version') {
47
+ process.exit(0);
48
+ }
49
+ if (args[0] === 'venv') {
50
+ const environment = process.env.UV_PROJECT_ENVIRONMENT;
51
+ mkdirSync(join(environment, 'bin'), { recursive: true });
52
+ writeFileSync(join(environment, 'bin', 'python'), '');
53
+ process.exit(0);
54
+ }
55
+ if (args[0] === 'pip' && args[1] === 'list') {
56
+ process.stdout.write('[]');
57
+ process.exit(0);
58
+ }
59
+ if (args[0] === 'pip' && args[1] === 'install') {
60
+ process.exit(0);
61
+ }
62
+ process.exit(1);
63
+ `,
64
+ );
65
+ chmodSync(fakeUvPath, 0o755);
66
+
67
+ const env = {
68
+ ...process.env,
69
+ MODULAR_PROMPT_HOME: temporaryDirectory,
70
+ FAKE_UV_LOG: uvLogPath,
71
+ PATH: `${fakeBinDirectory}:${process.env.PATH ?? ''}`,
72
+ };
73
+ const runCli = (...args: string[]) =>
74
+ execFileSync(process.execPath, [runtimeCli, ...args], {
75
+ encoding: 'utf8',
76
+ env,
77
+ });
78
+
79
+ expect(() => runCli('setup', 'pytorch', '--cuda', '12.4')).toThrow(
80
+ /--cuda can only be used with --variant cuda/,
81
+ );
82
+ expect(
83
+ () => runCli('setup', 'pytorch', '--variant', 'cuda', '--cuda', '12'),
84
+ ).toThrow(/Invalid CUDA version/);
85
+ runCli('setup', 'pytorch');
86
+ expect(existsSync(join(runtimePythonDir, 'pyproject.toml'))).toBe(true);
87
+ expect(existsSync(join(runtimePythonDir, '__main__.py'))).toBe(true);
88
+ expect(existsSync(join(runtimePythonDir, 'backends', 'base.py'))).toBe(true);
89
+ const manifestPath = join(
90
+ temporaryDirectory,
91
+ 'runtimes',
92
+ 'pytorch',
93
+ 'manifest.json',
94
+ );
95
+ expect(JSON.parse(readFileSync(manifestPath, 'utf8'))).toMatchObject({
96
+ profile: 'pytorch',
97
+ variant: 'cpu-minimal',
98
+ driverVersion,
99
+ });
100
+ expect(() => runCli('sync', 'pytorch', '--variant', 'other')).toThrow(
101
+ /variant mismatch/,
102
+ );
103
+
104
+ writeFileSync(join(runtimePythonDir, 'pyproject.toml'), 'user dependencies\n');
105
+ runCli('setup', 'pytorch');
106
+ expect(readFileSync(join(runtimePythonDir, 'pyproject.toml'), 'utf8')).toBe(
107
+ 'user dependencies\n',
108
+ );
109
+
110
+ writeFileSync(join(runtimePythonDir, 'backends', 'base.py'), 'user code\n');
111
+ runCli('sync', 'pytorch');
112
+ expect(readFileSync(join(runtimePythonDir, 'pyproject.toml'), 'utf8')).toBe(
113
+ 'user dependencies\n',
114
+ );
115
+ expect(readFileSync(join(runtimePythonDir, 'backends', 'base.py'), 'utf8')).toBe(
116
+ readFileSync(
117
+ join(packageRoot, 'src', 'pytorch', 'templates', 'cpu-minimal', 'backends', 'base.py'),
118
+ 'utf8',
119
+ ),
120
+ );
121
+ expect(JSON.parse(readFileSync(manifestPath, 'utf8'))).toMatchObject({
122
+ profile: 'pytorch',
123
+ variant: 'cpu-minimal',
124
+ driverVersion,
125
+ });
126
+
127
+ runCli('setup', 'pytorch', '--variant=cuda', '--cuda=12.1');
128
+ expect(JSON.parse(readFileSync(manifestPath, 'utf8'))).toMatchObject({
129
+ profile: 'pytorch',
130
+ variant: 'cuda',
131
+ cudaVersion: '12.1',
132
+ torchVersion: '2.9.1',
133
+ driverVersion,
134
+ });
135
+ expect(readFileSync(join(runtimePythonDir, '__main__.py'), 'utf8')).toContain(
136
+ 'os.environ.get("PYTORCH_DEVICE", "cuda")',
137
+ );
138
+ expect(
139
+ readFileSync(
140
+ join(packageRoot, 'src', 'pytorch', 'templates', 'cuda', 'pyproject.toml'),
141
+ 'utf8',
142
+ ),
143
+ ).toContain('https://download.pytorch.org/whl/cu124');
144
+
145
+ runCli('setup', 'pytorch', '--variant', 'cuda');
146
+ expect(JSON.parse(readFileSync(manifestPath, 'utf8'))).toMatchObject({
147
+ variant: 'cuda',
148
+ cudaVersion: '12.4',
149
+ });
150
+
151
+ const uvLog = readFileSync(uvLogPath, 'utf8');
152
+ expect(uvLog).toContain('venv --clear --python 3.12');
153
+ expect(uvLog).toContain('pip install');
154
+ expect(uvLog).toContain('torch==2.9.1');
155
+ expect(uvLog).toContain('https://download.pytorch.org/whl/cu121');
156
+ expect(uvLog).toContain('https://download.pytorch.org/whl/cu124');
157
+ expect(uvLog).toContain(' .');
158
+ expect(uvLog).not.toContain(' -e ');
159
+ } finally {
160
+ rmSync(temporaryDirectory, { recursive: true, force: true });
161
+ }
162
+ });
163
+ });
@@ -90,7 +90,7 @@ if __name__ == "__main__":
90
90
 
91
91
  capabilities = get_capabilities(backend.get_tokenizer())
92
92
  capabilities["model_kind"] = model_kind
93
- if model_kind == "lm":
93
+ if model_kind in {"lm", "vlm"} and "cache_prefill" not in capabilities["methods"]:
94
94
  capabilities["methods"].append("cache_prefill")
95
95
 
96
96
  server = Server(backend, capabilities)