@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
package/scripts/runtime-cli.js
CHANGED
|
@@ -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
|
|
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
|
-
|
|
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
|
|
119
|
-
|
|
120
|
-
|
|
121
|
-
|
|
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(
|
|
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: ${
|
|
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
|
-
|
|
143
|
-
|
|
144
|
-
|
|
145
|
-
|
|
146
|
-
|
|
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
|
-
|
|
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(
|
|
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
|
-
|
|
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 =
|
|
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(
|
|
432
|
+
setupHints.push(SETUP_MLX_MONOREPO);
|
|
204
433
|
}
|
|
205
434
|
if (!isRuntimeReady('pytorch')) {
|
|
206
|
-
setupHints.push(
|
|
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
|
-
|
|
257
|
-
|
|
258
|
-
|
|
259
|
-
|
|
260
|
-
|
|
261
|
-
|
|
262
|
-
|
|
263
|
-
|
|
264
|
-
|
|
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
|
|
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
|
-
|
|
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
|
|
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)
|