Add build support for Nvidia Spark
This commit is contained in:
@@ -1,6 +1,9 @@
|
|||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
import librosa
|
try:
|
||||||
|
import librosa
|
||||||
|
except ImportError:
|
||||||
|
librosa = None
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
import torchaudio
|
import torchaudio
|
||||||
@@ -41,6 +44,11 @@ KEY_NAMES = ["C", "C#", "D", "D#", "E", "F", "F#", "G", "G#", "A", "A#", "B"]
|
|||||||
|
|
||||||
def analyze_audio(audio_path):
|
def analyze_audio(audio_path):
|
||||||
"""Extract BPM, key, and time signature from audio using librosa."""
|
"""Extract BPM, key, and time signature from audio using librosa."""
|
||||||
|
if librosa is None:
|
||||||
|
raise ImportError(
|
||||||
|
"librosa is required for the AceStep captioner but is not "
|
||||||
|
"installed (no numba/llvmlite wheels for this platform yet)."
|
||||||
|
)
|
||||||
y, sr = librosa.load(audio_path, sr=22050, mono=True)
|
y, sr = librosa.load(audio_path, sr=22050, mono=True)
|
||||||
duration = librosa.get_duration(y=y, sr=sr)
|
duration = librosa.get_duration(y=y, sr=sr)
|
||||||
|
|
||||||
|
|||||||
@@ -25,7 +25,19 @@ def run_doctor():
|
|||||||
|
|
||||||
from . import gitwin
|
from . import gitwin
|
||||||
|
|
||||||
_check("os / arch", True, "%s %s" % (d["os"], d["arch"]))
|
arch_detail = "%s %s" % (d["os"], d["arch"])
|
||||||
|
if d["os"] == "windows" and d["arch"] == "aarch64":
|
||||||
|
from . import spec as spec_mod_arch
|
||||||
|
|
||||||
|
try:
|
||||||
|
_s = spec_mod_arch.build_spec(d, allow_cpu=True)
|
||||||
|
if _s.backend == "cu134":
|
||||||
|
arch_detail += " (RTX Spark: native win_arm64 CUDA stack)"
|
||||||
|
else:
|
||||||
|
arch_detail += " (Windows-on-ARM: x64 stack via emulation)"
|
||||||
|
except RuntimeError:
|
||||||
|
pass
|
||||||
|
_check("os / arch", True, arch_detail)
|
||||||
git = gitwin.find_git()
|
git = gitwin.find_git()
|
||||||
_check(
|
_check(
|
||||||
"git",
|
"git",
|
||||||
|
|||||||
186
manager/env.py
186
manager/env.py
@@ -62,24 +62,79 @@ def venv_exists():
|
|||||||
return os.path.isfile(venv_python())
|
return os.path.isfile(venv_python())
|
||||||
|
|
||||||
|
|
||||||
|
def _venv_platform():
|
||||||
|
"""sysconfig platform of the existing venv ('win-amd64', 'win-arm64', ...)."""
|
||||||
|
if not venv_exists():
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
out = subprocess.run(
|
||||||
|
[venv_python(), "-c", "import sysconfig; print(sysconfig.get_platform())"],
|
||||||
|
stdout=subprocess.PIPE,
|
||||||
|
stderr=subprocess.DEVNULL,
|
||||||
|
timeout=30,
|
||||||
|
env=clean_env(),
|
||||||
|
)
|
||||||
|
if out.returncode != 0:
|
||||||
|
return None
|
||||||
|
return out.stdout.decode().strip() or None
|
||||||
|
except (OSError, subprocess.TimeoutExpired):
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _uv_python_platform(uv_python):
|
||||||
|
"""'win-arm64' / 'win-amd64' expected for a pinned uv interpreter request."""
|
||||||
|
if not uv_python:
|
||||||
|
return None
|
||||||
|
if "windows-aarch64" in uv_python:
|
||||||
|
return "win-arm64"
|
||||||
|
if "windows-x86_64" in uv_python:
|
||||||
|
return "win-amd64"
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
def ensure_venv(spec, dry_run=False):
|
def ensure_venv(spec, dry_run=False):
|
||||||
"""Create the venv if missing. Returns path to the venv python."""
|
"""Create the venv if missing. Returns path to the venv python."""
|
||||||
|
if venv_exists():
|
||||||
|
# Switching stacks (e.g. Spark emulated x64 <-> native arm64) needs a
|
||||||
|
# different interpreter arch; the venv is disposable by design, so
|
||||||
|
# recreate it rather than install unresolvable wheels into it.
|
||||||
|
want = _uv_python_platform(spec.uv_python)
|
||||||
|
have = _venv_platform()
|
||||||
|
if want and have and want != have:
|
||||||
|
if dry_run:
|
||||||
|
info(
|
||||||
|
"[dry-run] venv is %s but this spec needs %s — would "
|
||||||
|
"recreate the venv." % (have, want)
|
||||||
|
)
|
||||||
|
return venv_python()
|
||||||
|
warn(
|
||||||
|
"Existing venv is %s but this spec needs %s — recreating the "
|
||||||
|
"venv (all packages will be reinstalled)." % (have, want)
|
||||||
|
)
|
||||||
|
import shutil
|
||||||
|
|
||||||
|
shutil.rmtree(venv_dir(), ignore_errors=True)
|
||||||
|
else:
|
||||||
|
return venv_python()
|
||||||
if venv_exists():
|
if venv_exists():
|
||||||
return venv_python()
|
return venv_python()
|
||||||
|
|
||||||
target = venv_dir()
|
target = venv_dir()
|
||||||
uv = find_uv()
|
uv = find_uv()
|
||||||
|
# spec.uv_python pins the full interpreter build (arch included) where the
|
||||||
|
# default choice would be wrong — e.g. Windows-on-ARM must stay x86_64
|
||||||
|
python_request = spec.uv_python or spec.python_version
|
||||||
if dry_run:
|
if dry_run:
|
||||||
info(
|
info(
|
||||||
"[dry-run] would create venv at %s (python %s, via %s)"
|
"[dry-run] would create venv at %s (python %s, via %s)"
|
||||||
% (target, spec.python_version, "uv" if uv else "venv module")
|
% (target, python_request, "uv" if uv else "venv module")
|
||||||
)
|
)
|
||||||
return venv_python(target)
|
return venv_python(target)
|
||||||
|
|
||||||
if uv:
|
if uv:
|
||||||
info("Creating venv with uv (python %s) at %s" % (spec.python_version, target))
|
info("Creating venv with uv (python %s) at %s" % (python_request, target))
|
||||||
run(
|
run(
|
||||||
[uv, "venv", target, "--python", spec.python_version, "--seed"],
|
[uv, "venv", target, "--python", python_request, "--seed"],
|
||||||
env=clean_env(),
|
env=clean_env(),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -89,7 +144,13 @@ def ensure_venv(spec, dry_run=False):
|
|||||||
"Install uv (https://docs.astral.sh/uv/) or a newer Python, then re-run."
|
"Install uv (https://docs.astral.sh/uv/) or a newer Python, then re-run."
|
||||||
% (sys.version_info[:2] + MIN_SYSTEM_PYTHON)
|
% (sys.version_info[:2] + MIN_SYSTEM_PYTHON)
|
||||||
)
|
)
|
||||||
pyver = "%d.%d" % sys.version_info[:2]
|
if spec.uv_python:
|
||||||
|
warn(
|
||||||
|
"uv not found — the venv needs the %s interpreter and the "
|
||||||
|
"system Python may be a different build. Install uv if the "
|
||||||
|
"torch install below fails to resolve." % spec.uv_python
|
||||||
|
)
|
||||||
|
pyver = "%d.%d" % (sys.version_info[:2])
|
||||||
if pyver != spec.python_version:
|
if pyver != spec.python_version:
|
||||||
warn(
|
warn(
|
||||||
"Recommended Python is %s but using system Python %s "
|
"Recommended Python is %s but using system Python %s "
|
||||||
@@ -119,6 +180,20 @@ def _pip_install(args, dry_run=False, upgrade=False, check=True):
|
|||||||
return code
|
return code
|
||||||
|
|
||||||
|
|
||||||
|
def _pip_install_no_deps(pkg, dry_run=False):
|
||||||
|
"""Install a single package with --no-deps. Returns exit code."""
|
||||||
|
uv = find_uv()
|
||||||
|
if uv:
|
||||||
|
cmd = [uv, "pip", "install", "--python", venv_python(), "--no-deps", pkg]
|
||||||
|
else:
|
||||||
|
cmd = [venv_python(), "-m", "pip", "install", "--no-deps", pkg]
|
||||||
|
if dry_run:
|
||||||
|
info("[dry-run] would run: %s" % " ".join(cmd))
|
||||||
|
return 0
|
||||||
|
code, _ = run(cmd, stream=True, env=clean_env(), check=False)
|
||||||
|
return code
|
||||||
|
|
||||||
|
|
||||||
def _pip_uninstall(packages, dry_run=False):
|
def _pip_uninstall(packages, dry_run=False):
|
||||||
uv = find_uv()
|
uv = find_uv()
|
||||||
if uv:
|
if uv:
|
||||||
@@ -419,11 +494,13 @@ def ensure_requirements(spec, dry_run=False, force=False):
|
|||||||
_pip_uninstall(stale, dry_run=dry_run)
|
_pip_uninstall(stale, dry_run=dry_run)
|
||||||
# every pass below carries the torch pins so nothing can swap the GPU build
|
# every pass below carries the torch pins so nothing can swap the GPU build
|
||||||
pins = _torch_pin_args(spec, dry_run=dry_run)
|
pins = _torch_pin_args(spec, dry_run=dry_run)
|
||||||
info("Installing requirements from %s..." % spec.requirements_file)
|
|
||||||
_pip_install(["-r", spec.requirements_path()] + pins, dry_run=dry_run)
|
|
||||||
find_links = list(pins)
|
find_links = list(pins)
|
||||||
for url in spec.find_links:
|
for url in spec.find_links:
|
||||||
find_links += ["--find-links", url]
|
find_links += ["--find-links", url]
|
||||||
|
info("Installing requirements from %s..." % spec.requirements_file)
|
||||||
|
# find_links included: on Spark the requirements themselves resolve
|
||||||
|
# self-built wheels (opencv, soxr, ...) from the spark wheel set
|
||||||
|
_pip_install(["-r", spec.requirements_path()] + find_links, dry_run=dry_run)
|
||||||
extras = _filter_extras(spec.extra_packages)
|
extras = _filter_extras(spec.extra_packages)
|
||||||
if extras:
|
if extras:
|
||||||
info("Installing platform extras...")
|
info("Installing platform extras...")
|
||||||
@@ -454,6 +531,13 @@ def ensure_requirements(spec, dry_run=False, force=False):
|
|||||||
"removing it (training falls back to native attention)." % name
|
"removing it (training falls back to native attention)." % name
|
||||||
)
|
)
|
||||||
_pip_uninstall([name])
|
_pip_uninstall([name])
|
||||||
|
# packages whose dependency metadata is unsatisfiable on this platform but
|
||||||
|
# which work fine without it (e.g. tensorboard's grpcio on win_arm64)
|
||||||
|
for pkg in spec.no_deps_packages:
|
||||||
|
info("Installing (no-deps): %s" % pkg)
|
||||||
|
code = _pip_install_no_deps(pkg, dry_run=dry_run)
|
||||||
|
if code != 0:
|
||||||
|
warn("No-deps package failed to install (continuing): %s" % pkg)
|
||||||
_verify_torch(spec, dry_run=dry_run)
|
_verify_torch(spec, dry_run=dry_run)
|
||||||
if not dry_run:
|
if not dry_run:
|
||||||
state = load_state()
|
state = load_state()
|
||||||
@@ -467,13 +551,55 @@ def ensure_requirements(spec, dry_run=False, force=False):
|
|||||||
# ---------------------------------------------------------------- sitecustomize
|
# ---------------------------------------------------------------- sitecustomize
|
||||||
|
|
||||||
|
|
||||||
def write_sitecustomize(dry_run=False):
|
def _msvc_runtime_env():
|
||||||
"""Drop a sitecustomize.py into the venv that exposes the local ffmpeg.
|
"""{env: value} + [bin dirs] from vcvarsarm64, for triton's runtime JIT.
|
||||||
|
|
||||||
|
Triton compiles its kernel launcher stubs with cl.exe at runtime (cached
|
||||||
|
afterwards in ~/.triton), which needs INCLUDE/LIB and cl on PATH. Capture
|
||||||
|
the values once at sync time and bake them into sitecustomize.
|
||||||
|
"""
|
||||||
|
vcvars = (
|
||||||
|
r"C:\Program Files (x86)\Microsoft Visual Studio\2022\BuildTools"
|
||||||
|
r"\VC\Auxiliary\Build\vcvarsarm64.bat"
|
||||||
|
)
|
||||||
|
if not os.path.isfile(vcvars):
|
||||||
|
return {}, []
|
||||||
|
try:
|
||||||
|
# string form: list2cmdline would mangle the nested quoting
|
||||||
|
out = subprocess.run(
|
||||||
|
'cmd /s /c "call "%s" >nul && set"' % vcvars,
|
||||||
|
stdout=subprocess.PIPE,
|
||||||
|
stderr=subprocess.DEVNULL,
|
||||||
|
timeout=120,
|
||||||
|
)
|
||||||
|
if out.returncode != 0:
|
||||||
|
return {}, []
|
||||||
|
env = {}
|
||||||
|
for line in out.stdout.decode(errors="replace").splitlines():
|
||||||
|
if "=" in line:
|
||||||
|
k, _, v = line.partition("=")
|
||||||
|
env[k.upper()] = v
|
||||||
|
bin_dirs = [
|
||||||
|
d for d in env.get("PATH", "").split(os.pathsep)
|
||||||
|
if os.path.isfile(os.path.join(d, "cl.exe"))
|
||||||
|
][:1]
|
||||||
|
keep = {k: env[k] for k in ("INCLUDE", "LIB") if env.get(k)}
|
||||||
|
if keep and bin_dirs:
|
||||||
|
keep["CC"] = "cl"
|
||||||
|
return keep, bin_dirs
|
||||||
|
except (OSError, subprocess.TimeoutExpired):
|
||||||
|
pass
|
||||||
|
return {}, []
|
||||||
|
|
||||||
|
|
||||||
|
def write_sitecustomize(dry_run=False, spec=None):
|
||||||
|
"""Drop a sitecustomize.py into the venv exposing runtime DLL dirs.
|
||||||
|
|
||||||
sitecustomize is imported automatically at interpreter startup, so ANY use
|
sitecustomize is imported automatically at interpreter startup, so ANY use
|
||||||
of the venv python (UI-spawned training jobs, run.py from a terminal) gets
|
of the venv python (UI-spawned training jobs, run.py from a terminal) gets
|
||||||
.ffmpeg/bin on PATH — and on Windows, os.add_dll_directory so torchcodec
|
.ffmpeg/bin on PATH — and on Windows, os.add_dll_directory so torchcodec
|
||||||
finds the FFmpeg DLLs.
|
finds the FFmpeg DLLs. On Spark the spec also carries the CUDA/cuDNN/APL
|
||||||
|
bin dirs, because the native torch wheel does not bundle its DLLs.
|
||||||
"""
|
"""
|
||||||
from . import ffmpeg
|
from . import ffmpeg
|
||||||
|
|
||||||
@@ -498,19 +624,29 @@ def write_sitecustomize(dry_run=False):
|
|||||||
warn("Could not locate venv site-packages — skipping sitecustomize.")
|
warn("Could not locate venv site-packages — skipping sitecustomize.")
|
||||||
return
|
return
|
||||||
target = os.path.join(site_packages, "sitecustomize.py")
|
target = os.path.join(site_packages, "sitecustomize.py")
|
||||||
|
dll_dirs = [ffmpeg.bin_dir()] + list(getattr(spec, "runtime_dll_dirs", []) or [])
|
||||||
|
runtime_env = dict(getattr(spec, "runtime_env", {}) or {})
|
||||||
|
if getattr(spec, "backend", None) == "cu134":
|
||||||
|
# triton's runtime launcher JIT needs the MSVC environment
|
||||||
|
msvc_env, msvc_bins = _msvc_runtime_env()
|
||||||
|
runtime_env.update(msvc_env)
|
||||||
|
dll_dirs += msvc_bins
|
||||||
content = (
|
content = (
|
||||||
"# Generated by the AI Toolkit manager (manager/env.py). Do not edit;\n"
|
"# Generated by the AI Toolkit manager (manager/env.py). Do not edit;\n"
|
||||||
"# regenerated on every `manager sync`.\n"
|
"# regenerated on every `manager sync`.\n"
|
||||||
"import os\n"
|
"import os\n"
|
||||||
"_FFMPEG_BIN = %r\n"
|
"for _k, _v in %r.items():\n"
|
||||||
|
" os.environ.setdefault(_k, _v)\n"
|
||||||
|
"_DLL_DIRS = %r\n"
|
||||||
"_FFMPEG_LIB = %r\n"
|
"_FFMPEG_LIB = %r\n"
|
||||||
"if os.path.isdir(_FFMPEG_BIN):\n"
|
"for _d in _DLL_DIRS:\n"
|
||||||
" os.environ['PATH'] = _FFMPEG_BIN + os.pathsep + os.environ.get('PATH', '')\n"
|
" if os.path.isdir(_d):\n"
|
||||||
" if hasattr(os, 'add_dll_directory'):\n"
|
" os.environ['PATH'] = _d + os.pathsep + os.environ.get('PATH', '')\n"
|
||||||
" try:\n"
|
" if hasattr(os, 'add_dll_directory'):\n"
|
||||||
" os.add_dll_directory(_FFMPEG_BIN)\n"
|
" try:\n"
|
||||||
" except OSError:\n"
|
" os.add_dll_directory(_d)\n"
|
||||||
" pass\n"
|
" except OSError:\n"
|
||||||
|
" pass\n"
|
||||||
"if os.path.isdir(_FFMPEG_LIB):\n"
|
"if os.path.isdir(_FFMPEG_LIB):\n"
|
||||||
" # inherited by child processes (the ffmpeg/ffprobe executables\n"
|
" # inherited by child processes (the ffmpeg/ffprobe executables\n"
|
||||||
" # need it to find their own shared libs)\n"
|
" # need it to find their own shared libs)\n"
|
||||||
@@ -519,7 +655,7 @@ def write_sitecustomize(dry_run=False):
|
|||||||
" os.environ['LD_LIBRARY_PATH'] = (\n"
|
" os.environ['LD_LIBRARY_PATH'] = (\n"
|
||||||
" _FFMPEG_LIB + ((os.pathsep + _prev) if _prev else '')\n"
|
" _FFMPEG_LIB + ((os.pathsep + _prev) if _prev else '')\n"
|
||||||
" )\n"
|
" )\n"
|
||||||
) % (ffmpeg.bin_dir(), ffmpeg.lib_dir())
|
) % (runtime_env, dll_dirs, ffmpeg.lib_dir())
|
||||||
if dry_run:
|
if dry_run:
|
||||||
info("[dry-run] would write %s" % target)
|
info("[dry-run] would write %s" % target)
|
||||||
return
|
return
|
||||||
@@ -538,13 +674,23 @@ def sync(spec, detection, dry_run=False, force=False):
|
|||||||
warn(note)
|
warn(note)
|
||||||
uvbin.ensure_uv(dry_run=dry_run)
|
uvbin.ensure_uv(dry_run=dry_run)
|
||||||
gitwin.ensure_git(dry_run=dry_run)
|
gitwin.ensure_git(dry_run=dry_run)
|
||||||
|
if spec.backend == "cu134":
|
||||||
|
# native Spark stack: provision CUDA/cuDNN/APL runtime DLLs, VC
|
||||||
|
# redist, and (best-effort) MSVC for triton's kernel launcher JIT
|
||||||
|
from . import sparkdeps
|
||||||
|
|
||||||
|
sparkdeps.ensure_spark_runtime(dry_run=dry_run)
|
||||||
ensure_venv(spec, dry_run=dry_run)
|
ensure_venv(spec, dry_run=dry_run)
|
||||||
|
# sitecustomize must exist BEFORE any torch import check below: on Spark
|
||||||
|
# the torch wheel is unbundled and only imports once the CUDA/cuDNN/BLAS
|
||||||
|
# DLL dirs from the spec are exposed to the interpreter
|
||||||
|
write_sitecustomize(dry_run=dry_run, spec=spec)
|
||||||
changed_torch = ensure_torch(spec, dry_run=dry_run)
|
changed_torch = ensure_torch(spec, dry_run=dry_run)
|
||||||
# a torch reinstall can clobber pinned deps; force req pass afterwards
|
# a torch reinstall can clobber pinned deps; force req pass afterwards
|
||||||
ensure_requirements(spec, dry_run=dry_run, force=force or changed_torch)
|
ensure_requirements(spec, dry_run=dry_run, force=force or changed_torch)
|
||||||
ffmpeg.ensure_ffmpeg(detection, dry_run=dry_run)
|
ffmpeg.ensure_ffmpeg(detection, dry_run=dry_run, spec=spec)
|
||||||
nodejs.ensure_node(detection, dry_run=dry_run)
|
nodejs.ensure_node(detection, dry_run=dry_run)
|
||||||
nodejs.ensure_ui_deps(dry_run=dry_run)
|
nodejs.ensure_ui_deps(dry_run=dry_run)
|
||||||
write_sitecustomize(dry_run=dry_run)
|
write_sitecustomize(dry_run=dry_run, spec=spec)
|
||||||
migrations.run_pending(dry_run=dry_run)
|
migrations.run_pending(dry_run=dry_run)
|
||||||
ok("Environment is up to date.")
|
ok("Environment is up to date.")
|
||||||
|
|||||||
@@ -48,8 +48,17 @@ _SOURCES = {
|
|||||||
("linux", "x86_64"): _BTBN + "ffmpeg-n8.1-latest-linux64-gpl-shared-8.1.tar.xz",
|
("linux", "x86_64"): _BTBN + "ffmpeg-n8.1-latest-linux64-gpl-shared-8.1.tar.xz",
|
||||||
("linux", "aarch64"): _BTBN + "ffmpeg-n8.1-latest-linuxarm64-gpl-shared-8.1.tar.xz",
|
("linux", "aarch64"): _BTBN + "ffmpeg-n8.1-latest-linuxarm64-gpl-shared-8.1.tar.xz",
|
||||||
("windows", "x86_64"): _BTBN + "ffmpeg-n8.1-latest-win64-gpl-shared-8.1.zip",
|
("windows", "x86_64"): _BTBN + "ffmpeg-n8.1-latest-win64-gpl-shared-8.1.zip",
|
||||||
|
# Windows-on-ARM in the emulated-x64 stack gets the x64 build, NOT BtbN's
|
||||||
|
# winarm64 one: an x64 torchcodec can only dlopen x64 FFmpeg DLLs, and the
|
||||||
|
# exes run fine under emulation.
|
||||||
|
("windows", "aarch64"): _BTBN + "ffmpeg-n8.1-latest-win64-gpl-shared-8.1.zip",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# Native Spark stack: the self-built win_arm64 torchcodec is linked against
|
||||||
|
# (and dlopens) arm64 FFmpeg 8 — LGPL to keep the distributed torchcodec wheel
|
||||||
|
# clean, matching C:\Dev spark build scripts / the wheel-set build recipe.
|
||||||
|
_SPARK_NATIVE_SOURCE = _BTBN + "ffmpeg-n8.1-latest-winarm64-lgpl-shared-8.1.zip"
|
||||||
|
|
||||||
|
|
||||||
def bin_dir():
|
def bin_dir():
|
||||||
return os.path.join(FFMPEG_DIR, "bin")
|
return os.path.join(FFMPEG_DIR, "bin")
|
||||||
@@ -124,15 +133,20 @@ def _install_mac(detection):
|
|||||||
shutil.rmtree(tmp, ignore_errors=True)
|
shutil.rmtree(tmp, ignore_errors=True)
|
||||||
|
|
||||||
|
|
||||||
def source_url(detection):
|
def source_url(detection, spec=None):
|
||||||
if detection["os"] == "mac":
|
if detection["os"] == "mac":
|
||||||
arch = "arm64" if detection["arch"] == "arm64" else "amd64"
|
arch = "arm64" if detection["arch"] == "arm64" else "amd64"
|
||||||
return _RIEDL.format(arch=arch, tool="ffmpeg")
|
return _RIEDL.format(arch=arch, tool="ffmpeg")
|
||||||
|
if (
|
||||||
|
getattr(spec, "backend", None) == "cu134"
|
||||||
|
and (detection["os"], detection["arch"]) == ("windows", "aarch64")
|
||||||
|
):
|
||||||
|
return _SPARK_NATIVE_SOURCE
|
||||||
return _SOURCES.get((detection["os"], detection["arch"]))
|
return _SOURCES.get((detection["os"], detection["arch"]))
|
||||||
|
|
||||||
|
|
||||||
def ensure_ffmpeg(detection, dry_run=False):
|
def ensure_ffmpeg(detection, dry_run=False, spec=None):
|
||||||
url = source_url(detection)
|
url = source_url(detection, spec=spec)
|
||||||
if url is None:
|
if url is None:
|
||||||
warn(
|
warn(
|
||||||
"No portable FFmpeg source for %s/%s — skipping local ffmpeg."
|
"No portable FFmpeg source for %s/%s — skipping local ffmpeg."
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ the manager just errors with install instructions elsewhere.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import os
|
import os
|
||||||
|
import platform
|
||||||
import shutil
|
import shutil
|
||||||
import tempfile
|
import tempfile
|
||||||
|
|
||||||
@@ -20,10 +21,20 @@ from .util import IS_WINDOWS, REPO_ROOT, download, extract_archive, info, ok, wa
|
|||||||
|
|
||||||
MINGIT_DIR = os.path.join(REPO_ROOT, ".mingit")
|
MINGIT_DIR = os.path.join(REPO_ROOT, ".mingit")
|
||||||
# Update this pin together with nothing else — it's independent of torch etc.
|
# Update this pin together with nothing else — it's independent of torch etc.
|
||||||
MINGIT_URL = (
|
_MINGIT_TAG = "v2.55.0.windows.3"
|
||||||
"https://github.com/git-for-windows/git/releases/download/"
|
_MINGIT_VERSION = "2.55.0.3"
|
||||||
"v2.55.0.windows.3/MinGit-2.55.0.3-64-bit.zip"
|
|
||||||
)
|
|
||||||
|
def _mingit_url():
|
||||||
|
# git is a standalone subprocess (nothing dlopens it), so unlike
|
||||||
|
# node/ffmpeg it can be native arm64 on Windows-on-ARM
|
||||||
|
arm = platform.machine().lower() in ("arm64", "aarch64")
|
||||||
|
flavor = "arm64" if arm else "64-bit"
|
||||||
|
return "https://github.com/git-for-windows/git/releases/download/%s/MinGit-%s-%s.zip" % (
|
||||||
|
_MINGIT_TAG,
|
||||||
|
_MINGIT_VERSION,
|
||||||
|
flavor,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def local_git_exe():
|
def local_git_exe():
|
||||||
@@ -49,7 +60,7 @@ def ensure_git(dry_run=False):
|
|||||||
tmp = tempfile.mkdtemp(prefix="aitk_mingit_")
|
tmp = tempfile.mkdtemp(prefix="aitk_mingit_")
|
||||||
try:
|
try:
|
||||||
archive = os.path.join(tmp, "mingit.zip")
|
archive = os.path.join(tmp, "mingit.zip")
|
||||||
download(MINGIT_URL, archive, label="MinGit")
|
download(_mingit_url(), archive, label="MinGit")
|
||||||
# MinGit zips have no top-level folder: cmd/, mingw64/, etc at the root
|
# MinGit zips have no top-level folder: cmd/, mingw64/, etc at the root
|
||||||
extracted = os.path.join(tmp, "mingit")
|
extracted = os.path.join(tmp, "mingit")
|
||||||
extract_archive(archive, extracted)
|
extract_archive(archive, extracted)
|
||||||
|
|||||||
@@ -43,20 +43,37 @@ def local_node_exe():
|
|||||||
return os.path.join(node_bin_dir(), "node.exe" if IS_WINDOWS else "node")
|
return os.path.join(node_bin_dir(), "node.exe" if IS_WINDOWS else "node")
|
||||||
|
|
||||||
|
|
||||||
def _node_major(exe):
|
def _node_info(exe):
|
||||||
|
"""(major, arch) for a node executable, e.g. (24, 'x64'), or (None, None)."""
|
||||||
try:
|
try:
|
||||||
out = subprocess.run(
|
out = subprocess.run(
|
||||||
[exe, "--version"],
|
[exe, "-p", "process.version + ' ' + process.arch"],
|
||||||
stdout=subprocess.PIPE,
|
stdout=subprocess.PIPE,
|
||||||
stderr=subprocess.DEVNULL,
|
stderr=subprocess.DEVNULL,
|
||||||
timeout=15,
|
timeout=15,
|
||||||
)
|
)
|
||||||
if out.returncode != 0:
|
if out.returncode != 0:
|
||||||
return None
|
return None, None
|
||||||
text = out.stdout.decode().strip() # v24.11.1
|
version, arch = out.stdout.decode().strip().split() # v24.11.1 x64
|
||||||
return int(text.lstrip("v").split(".")[0])
|
return int(version.lstrip("v").split(".")[0]), arch
|
||||||
except (OSError, subprocess.TimeoutExpired, ValueError):
|
except (OSError, subprocess.TimeoutExpired, ValueError):
|
||||||
|
return None, None
|
||||||
|
|
||||||
|
|
||||||
|
def _usable_major(exe):
|
||||||
|
"""Version if this node can run the UI, else None.
|
||||||
|
|
||||||
|
On Windows only an x64 node qualifies — the UI's native modules resolve
|
||||||
|
for the node process arch and Prisma ships no windows-arm64 query engine,
|
||||||
|
so a native ARM64 node loads the UI but dies on the first DB call. The
|
||||||
|
x64 node runs under emulation on ARM hosts.
|
||||||
|
"""
|
||||||
|
major, arch = _node_info(exe)
|
||||||
|
if major is None or major < MIN_NODE_MAJOR:
|
||||||
return None
|
return None
|
||||||
|
if IS_WINDOWS and arch != "x64":
|
||||||
|
return None
|
||||||
|
return major
|
||||||
|
|
||||||
|
|
||||||
def _dist_url(detection):
|
def _dist_url(detection):
|
||||||
@@ -68,6 +85,8 @@ def _dist_url(detection):
|
|||||||
plat = "darwin-arm64" if arch == "arm64" else "darwin-x64"
|
plat = "darwin-arm64" if arch == "arm64" else "darwin-x64"
|
||||||
ext = "tar.gz"
|
ext = "tar.gz"
|
||||||
elif detection["os"] == "windows":
|
elif detection["os"] == "windows":
|
||||||
|
# win-x64 even on ARM hosts: Prisma has no windows-arm64 engine, so an
|
||||||
|
# arm64 node can't run the UI (see _usable_major); x64 node emulates fine
|
||||||
plat = "win-x64"
|
plat = "win-x64"
|
||||||
ext = "zip"
|
ext = "zip"
|
||||||
else:
|
else:
|
||||||
@@ -80,13 +99,13 @@ def have_usable_node():
|
|||||||
"""(exe, major) for the best available node: local .node/ first, then system."""
|
"""(exe, major) for the best available node: local .node/ first, then system."""
|
||||||
local = local_node_exe()
|
local = local_node_exe()
|
||||||
if os.path.isfile(local):
|
if os.path.isfile(local):
|
||||||
major = _node_major(local)
|
major = _usable_major(local)
|
||||||
if major is not None and major >= MIN_NODE_MAJOR:
|
if major is not None:
|
||||||
return local, major
|
return local, major
|
||||||
system = which("node")
|
system = which("node")
|
||||||
if system:
|
if system:
|
||||||
major = _node_major(system)
|
major = _usable_major(system)
|
||||||
if major is not None and major >= MIN_NODE_MAJOR:
|
if major is not None:
|
||||||
return system, major
|
return system, major
|
||||||
return None, None
|
return None, None
|
||||||
|
|
||||||
@@ -96,6 +115,15 @@ def ensure_node(detection, dry_run=False):
|
|||||||
if exe:
|
if exe:
|
||||||
ok("Node.js v%d found (%s)." % (major, exe))
|
ok("Node.js v%d found (%s)." % (major, exe))
|
||||||
return False
|
return False
|
||||||
|
system = which("node")
|
||||||
|
if system:
|
||||||
|
_, arch = _node_info(system)
|
||||||
|
if IS_WINDOWS and arch and arch != "x64":
|
||||||
|
info(
|
||||||
|
"System Node.js is %s but the UI needs an x64 node on Windows "
|
||||||
|
"(Prisma has no windows-arm64 engine) — installing a local x64 "
|
||||||
|
"copy." % arch
|
||||||
|
)
|
||||||
url, inner_name = _dist_url(detection)
|
url, inner_name = _dist_url(detection)
|
||||||
if url is None:
|
if url is None:
|
||||||
warn(
|
warn(
|
||||||
|
|||||||
224
manager/sparkdeps.py
Normal file
224
manager/sparkdeps.py
Normal file
@@ -0,0 +1,224 @@
|
|||||||
|
"""RTX Spark native-stack runtime provisioning (Windows on ARM, cu134).
|
||||||
|
|
||||||
|
Goal: a fresh Spark machine runs run_windows.bat and gets as close to
|
||||||
|
zero-manual-setup as licensing allows. The native wheels (torch etc.) do not
|
||||||
|
bundle CUDA / cuDNN / BLAS DLLs, and triton's launcher JIT wants MSVC. Policy:
|
||||||
|
NVIDIA components are never redistributed by us.
|
||||||
|
|
||||||
|
- CUDA 13.4 toolkit (developer preview): MANUAL install — the preview EULA
|
||||||
|
requires NVIDIA's own click-through, so the manager only detects it and
|
||||||
|
prints instructions when missing. This is the single manual step.
|
||||||
|
- cuDNN (arm64): auto-downloaded from NVIDIA's own official installer URL and
|
||||||
|
installed silently — fetched directly from NVIDIA, not redistributed.
|
||||||
|
- Arm Performance Libraries: auto-install via winget (official Arm package).
|
||||||
|
- MSVC Build Tools (triton torch.compile JIT only): auto-install via winget;
|
||||||
|
failure downgrades gracefully (training works, no torch.compile).
|
||||||
|
- VC redistributable (arm64): auto-install via winget when msvcp140 missing.
|
||||||
|
|
||||||
|
Everything is best-effort with warnings; the training stack itself only hard-
|
||||||
|
requires the CUDA + cuDNN + APL DLL dirs.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import glob
|
||||||
|
import os
|
||||||
|
import subprocess
|
||||||
|
|
||||||
|
from .util import download, info, ok, warn, which
|
||||||
|
|
||||||
|
CUDA_DOWNLOAD_PAGE = (
|
||||||
|
"https://developer.nvidia.com/cuda-13-4-0-download-archive"
|
||||||
|
"?target_os=Windows&target_arch=arm64"
|
||||||
|
)
|
||||||
|
# NVIDIA's official public installer for cuDNN on Windows arm64. Downloaded
|
||||||
|
# straight from NVIDIA at install time (we do not redistribute it). Update
|
||||||
|
# together with the wheel set when moving to a newer cuDNN.
|
||||||
|
CUDNN_INSTALLER_URL = (
|
||||||
|
"https://developer.download.nvidia.com/compute/cudnn/9.25.0/"
|
||||||
|
"local_installers/cudnn_9.25.0_windows_arm64.exe"
|
||||||
|
)
|
||||||
|
|
||||||
|
# System install roots, newest version preferred (globs, not pinned versions)
|
||||||
|
_CUDA_BIN_GLOB = r"C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\v*\bin\arm64"
|
||||||
|
_CUDNN_BIN_GLOB = r"C:\Program Files\NVIDIA\CUDNN\v*\bin\*\arm64"
|
||||||
|
_ARMPL_BIN_GLOB = r"C:\Program Files\Arm Performance Libraries\armpl_*\bin"
|
||||||
|
|
||||||
|
_VS_BUILDTOOLS = (
|
||||||
|
r"C:\Program Files (x86)\Microsoft Visual Studio\2022\BuildTools"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _newest(pattern):
|
||||||
|
matches = sorted(glob.glob(pattern))
|
||||||
|
return matches[-1] if matches else None
|
||||||
|
|
||||||
|
|
||||||
|
def cuda_bin_dir():
|
||||||
|
return _newest(_CUDA_BIN_GLOB)
|
||||||
|
|
||||||
|
|
||||||
|
def cuda_root():
|
||||||
|
d = cuda_bin_dir()
|
||||||
|
# <root>\bin\arm64 -> <root>
|
||||||
|
return os.path.dirname(os.path.dirname(d)) if d else None
|
||||||
|
|
||||||
|
|
||||||
|
def cudnn_bin_dir():
|
||||||
|
return _newest(_CUDNN_BIN_GLOB)
|
||||||
|
|
||||||
|
|
||||||
|
def armpl_bin_dir():
|
||||||
|
return _newest(_ARMPL_BIN_GLOB)
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_dll_dirs():
|
||||||
|
"""All runtime DLL dirs for the native stack (existing ones only)."""
|
||||||
|
return [d for d in (cuda_bin_dir(), cudnn_bin_dir(), armpl_bin_dir()) if d]
|
||||||
|
|
||||||
|
|
||||||
|
def runtime_complete():
|
||||||
|
return bool(cuda_bin_dir() and cudnn_bin_dir() and armpl_bin_dir())
|
||||||
|
|
||||||
|
|
||||||
|
def triton_tool_env():
|
||||||
|
"""TRITON_*_PATH env for ptxas etc. from the system CUDA install.
|
||||||
|
|
||||||
|
Our triton wheel deliberately does NOT bundle NVIDIA's compiler tools
|
||||||
|
(developer-preview licensing); resolve them from the user's toolkit.
|
||||||
|
"""
|
||||||
|
root = cuda_root()
|
||||||
|
if not root:
|
||||||
|
return {}
|
||||||
|
env = {}
|
||||||
|
for var, exe in (
|
||||||
|
("TRITON_PTXAS_PATH", "ptxas.exe"),
|
||||||
|
("TRITON_PTXAS_BLACKWELL_PATH", "ptxas.exe"),
|
||||||
|
("TRITON_CUOBJDUMP_PATH", "cuobjdump.exe"),
|
||||||
|
("TRITON_NVDISASM_PATH", "nvdisasm.exe"),
|
||||||
|
):
|
||||||
|
path = os.path.join(root, "bin", exe)
|
||||||
|
if os.path.isfile(path):
|
||||||
|
env[var] = path
|
||||||
|
return env
|
||||||
|
|
||||||
|
|
||||||
|
def check_cuda():
|
||||||
|
"""CUDA toolkit is the one manual install (preview EULA). Detect + guide."""
|
||||||
|
if cuda_bin_dir():
|
||||||
|
return True
|
||||||
|
warn(
|
||||||
|
"The CUDA 13.4 toolkit (arm64) is not installed. NVIDIA's developer "
|
||||||
|
"preview license requires installing it manually:\n"
|
||||||
|
" 1. Download from %s\n"
|
||||||
|
" 2. Install with default settings, then re-run this setup.\n"
|
||||||
|
"The RTX Spark developer driver (R616+) is required as well."
|
||||||
|
% CUDA_DOWNLOAD_PAGE
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def ensure_cudnn(dry_run=False):
|
||||||
|
"""Fetch + silently run NVIDIA's official cuDNN installer if missing."""
|
||||||
|
if cudnn_bin_dir():
|
||||||
|
return True
|
||||||
|
if dry_run:
|
||||||
|
info("[dry-run] would download and install cuDNN from NVIDIA")
|
||||||
|
return False
|
||||||
|
import tempfile
|
||||||
|
|
||||||
|
tmp = tempfile.mkdtemp(prefix="aitk_cudnn_")
|
||||||
|
try:
|
||||||
|
exe = os.path.join(tmp, os.path.basename(CUDNN_INSTALLER_URL))
|
||||||
|
download(CUDNN_INSTALLER_URL, exe, label="cuDNN (from NVIDIA)")
|
||||||
|
info("Installing cuDNN (silent)...")
|
||||||
|
code = subprocess.call([exe, "-s"])
|
||||||
|
if code != 0:
|
||||||
|
warn("cuDNN installer exited with %d." % code)
|
||||||
|
return cudnn_bin_dir() is not None
|
||||||
|
finally:
|
||||||
|
import shutil
|
||||||
|
|
||||||
|
shutil.rmtree(tmp, ignore_errors=True)
|
||||||
|
|
||||||
|
|
||||||
|
def have_msvc():
|
||||||
|
return bool(
|
||||||
|
glob.glob(os.path.join(_VS_BUILDTOOLS, "VC", "Tools", "MSVC", "*",
|
||||||
|
"bin", "Hostarm64", "arm64", "cl.exe"))
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _winget_install(args, label, dry_run=False):
|
||||||
|
winget = which("winget")
|
||||||
|
if not winget:
|
||||||
|
warn("winget not available — cannot auto-install %s." % label)
|
||||||
|
return False
|
||||||
|
if dry_run:
|
||||||
|
info("[dry-run] would winget install %s" % label)
|
||||||
|
return False
|
||||||
|
info("Installing %s (one-time, may take several minutes)..." % label)
|
||||||
|
code = subprocess.call(
|
||||||
|
[winget, "install", "--exact", "--source", "winget",
|
||||||
|
"--accept-source-agreements", "--accept-package-agreements"] + args,
|
||||||
|
stdout=subprocess.DEVNULL,
|
||||||
|
)
|
||||||
|
if code != 0:
|
||||||
|
warn("%s install failed (winget exit %d)." % (label, code))
|
||||||
|
return code == 0
|
||||||
|
|
||||||
|
|
||||||
|
def ensure_armpl(dry_run=False):
|
||||||
|
if armpl_bin_dir():
|
||||||
|
return True
|
||||||
|
return _winget_install(
|
||||||
|
["--id", "Arm.ArmPerformanceLibraries"],
|
||||||
|
"Arm Performance Libraries",
|
||||||
|
dry_run=dry_run,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def ensure_vcredist(dry_run=False):
|
||||||
|
"""VC runtime (msvcp140 etc.) — required by the native wheels."""
|
||||||
|
sysdir = os.path.join(os.environ.get("SystemRoot", r"C:\Windows"), "System32")
|
||||||
|
if os.path.isfile(os.path.join(sysdir, "msvcp140.dll")):
|
||||||
|
return True
|
||||||
|
return _winget_install(
|
||||||
|
["--id", "Microsoft.VCRedist.2015+.arm64"],
|
||||||
|
"Visual C++ Redistributable (arm64)",
|
||||||
|
dry_run=dry_run,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def ensure_msvc(dry_run=False):
|
||||||
|
"""MSVC Build Tools — only needed for triton's runtime kernel launchers.
|
||||||
|
|
||||||
|
Best effort: without it, training still works; torch.compile / triton
|
||||||
|
JIT is unavailable until the user installs Build Tools.
|
||||||
|
"""
|
||||||
|
if have_msvc():
|
||||||
|
return True
|
||||||
|
done = _winget_install(
|
||||||
|
["--id", "Microsoft.VisualStudio.2022.BuildTools", "--override",
|
||||||
|
"--quiet --wait --norestart "
|
||||||
|
"--add Microsoft.VisualStudio.Workload.VCTools "
|
||||||
|
"--add Microsoft.VisualStudio.Component.VC.Tools.ARM64 "
|
||||||
|
"--add Microsoft.VisualStudio.Component.Windows11SDK.26100"],
|
||||||
|
"MSVC Build Tools (for torch.compile/triton)",
|
||||||
|
dry_run=dry_run,
|
||||||
|
)
|
||||||
|
if not done and not dry_run:
|
||||||
|
warn(
|
||||||
|
"torch.compile/triton kernel JIT will be unavailable until MSVC "
|
||||||
|
"Build Tools are installed; training itself is unaffected."
|
||||||
|
)
|
||||||
|
return done
|
||||||
|
|
||||||
|
|
||||||
|
def ensure_spark_runtime(dry_run=False):
|
||||||
|
"""Full best-effort provisioning for the native Spark stack."""
|
||||||
|
check_cuda()
|
||||||
|
ensure_vcredist(dry_run=dry_run)
|
||||||
|
ensure_cudnn(dry_run=dry_run)
|
||||||
|
ensure_armpl(dry_run=dry_run)
|
||||||
|
ensure_msvc(dry_run=dry_run)
|
||||||
|
if runtime_complete():
|
||||||
|
ok("Spark native runtime present (CUDA + cuDNN + Arm PL).")
|
||||||
176
manager/spec.py
176
manager/spec.py
@@ -21,6 +21,14 @@ kernels smoke-tested on an RTX 5090 / sm120 with torch 2.13.0+cu130):
|
|||||||
{cp310..cp314} x {linux x86_64, linux aarch64}. No Windows/mac wheels.
|
{cp310..cp314} x {linux x86_64, linux aarch64}. No Windows/mac wheels.
|
||||||
- triton: bundled with torch on Linux (incl. aarch64; torch 2.13 bundles
|
- triton: bundled with torch on Linux (incl. aarch64; torch 2.13 bundles
|
||||||
triton 3.7.1); triton-windows 3.7.x matches on Windows; nothing for MPS.
|
triton 3.7.1); triton-windows 3.7.x matches on Windows; nothing for MPS.
|
||||||
|
- Windows-on-ARM (verified 2026-07): there are NO win_arm64 wheels for any of
|
||||||
|
the CUDA stack — torch cu130, triton-windows, flash-attn and Prisma's node
|
||||||
|
engine are all x64-only. The supported configuration is therefore the x64
|
||||||
|
stack end to end (Python, torch, Node) running under Windows' x64 emulation,
|
||||||
|
with the GPU driven natively by the NVIDIA driver. build_spec() pins the
|
||||||
|
interpreter arch explicitly so uv can never flip the venv to native aarch64
|
||||||
|
(where torch would not resolve). Revisit if pytorch.org ever publishes
|
||||||
|
win_arm64 CUDA wheels.
|
||||||
|
|
||||||
Requirements files must never pin anything torch itself depends on below what
|
Requirements files must never pin anything torch itself depends on below what
|
||||||
the pinned torch needs (torch 2.13 wants setuptools>=77.0.3). The resolver does
|
the pinned torch needs (torch 2.13 wants setuptools>=77.0.3). The resolver does
|
||||||
@@ -68,6 +76,51 @@ _WIN_HELPERS = ["wheel", "setuptools", "poetry-core", "hf_xet"]
|
|||||||
|
|
||||||
PYTORCH_INDEX = "https://download.pytorch.org/whl/"
|
PYTORCH_INDEX = "https://download.pytorch.org/whl/"
|
||||||
|
|
||||||
|
# ---- NVIDIA RTX Spark: native Windows-on-ARM CUDA ---------------------------
|
||||||
|
# The CUDA 13.4 developer preview added native win_arm64 CUDA. No public wheels
|
||||||
|
# exist for the GPU stack, so we build them ourselves (torch from main +
|
||||||
|
# pytorch/pytorch#190448, plus torchvision/torchaudio/torchcodec and the deps
|
||||||
|
# with no win_arm64 wheels). The manager installs them from a find-links
|
||||||
|
# source: the local wheels/spark/ dir during development, or the hosted URL
|
||||||
|
# once published. Without that source (or with a pre-13.4 driver) Spark
|
||||||
|
# machines fall back to the emulated x64 stack below, which also works.
|
||||||
|
SPARK_BACKEND = "cu134"
|
||||||
|
SPARK_TORCH = {
|
||||||
|
"torch": "2.14.0.dev20260727",
|
||||||
|
"torchvision": "0.29.0.dev20260727",
|
||||||
|
"torchaudio": "2.11.0.dev20260727",
|
||||||
|
}
|
||||||
|
SPARK_WHEELS_DIR = os.path.join(REPO_ROOT, "wheels", "spark")
|
||||||
|
# GitHub's expanded_assets endpoint serves plain HTML anchors — a valid
|
||||||
|
# pip/uv find-links page pointing at the release assets.
|
||||||
|
SPARK_WHEELS_URL = (
|
||||||
|
"https://github.com/ostris/ai-toolkit-spark-wheels/releases/"
|
||||||
|
"expanded_assets/cu134-20260727"
|
||||||
|
)
|
||||||
|
SPARK_UV_PYTHON = "cpython-3.12-windows-aarch64-none"
|
||||||
|
# Runtime DLL homes for the unbundled native torch (TH_BINARY_BUILD=0) are
|
||||||
|
# resolved dynamically (system installs of any version, else the downloadable
|
||||||
|
# runtime bundle) — see sparkdeps.resolve_dll_dirs().
|
||||||
|
|
||||||
|
|
||||||
|
def _spark_wheels_source():
|
||||||
|
"""find-links source holding the self-built win_arm64 wheels, or None."""
|
||||||
|
if os.path.isdir(SPARK_WHEELS_DIR):
|
||||||
|
for name in os.listdir(SPARK_WHEELS_DIR):
|
||||||
|
if name.startswith("torch-") and "win_arm64" in name:
|
||||||
|
return SPARK_WHEELS_DIR
|
||||||
|
return SPARK_WHEELS_URL
|
||||||
|
|
||||||
|
|
||||||
|
def _spark_capable(detection):
|
||||||
|
"""Driver new enough for native win_arm64 CUDA (R616+ reports CUDA 13.4)."""
|
||||||
|
nvidia = detection.get("nvidia") or {}
|
||||||
|
try:
|
||||||
|
cuda = tuple(int(x) for x in (nvidia.get("cuda_version") or "").split("."))
|
||||||
|
except ValueError:
|
||||||
|
return False
|
||||||
|
return cuda >= (13, 4)
|
||||||
|
|
||||||
|
|
||||||
class EnvSpec(object):
|
class EnvSpec(object):
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -81,8 +134,12 @@ class EnvSpec(object):
|
|||||||
optional_packages=None,
|
optional_packages=None,
|
||||||
find_links=None,
|
find_links=None,
|
||||||
notes=None,
|
notes=None,
|
||||||
|
uv_python=None,
|
||||||
|
torch_links=None,
|
||||||
|
no_deps_packages=None,
|
||||||
|
runtime_dll_dirs=None,
|
||||||
):
|
):
|
||||||
self.backend = backend # cu130 / cu126 / rocm7.1 / mps / cpu
|
self.backend = backend # cu134 / cu130 / cu126 / rocm7.1 / mps / cpu
|
||||||
self.torch_packages = torch_packages # {name: version}
|
self.torch_packages = torch_packages # {name: version}
|
||||||
self.torch_index = torch_index # None = PyPI
|
self.torch_index = torch_index # None = PyPI
|
||||||
self.python_version = python_version
|
self.python_version = python_version
|
||||||
@@ -91,11 +148,29 @@ class EnvSpec(object):
|
|||||||
self.optional_packages = optional_packages or []
|
self.optional_packages = optional_packages or []
|
||||||
self.find_links = find_links or []
|
self.find_links = find_links or []
|
||||||
self.notes = notes or []
|
self.notes = notes or []
|
||||||
|
# full uv interpreter request (e.g. "cpython-3.12-windows-x86_64-none")
|
||||||
|
# when the venv arch must not be left to uv's default; None = just
|
||||||
|
# python_version
|
||||||
|
self.uv_python = uv_python
|
||||||
|
# --find-links sources for the torch trio itself (self-built wheels);
|
||||||
|
# used when torch_index is None
|
||||||
|
self.torch_links = torch_links or []
|
||||||
|
# installed with --no-deps after everything else (e.g. tensorboard on
|
||||||
|
# Spark, whose grpcio dep has no win_arm64 wheels but is only needed
|
||||||
|
# for the server, not the log writer)
|
||||||
|
self.no_deps_packages = no_deps_packages or []
|
||||||
|
# extra DLL dirs the venv needs at runtime (unbundled CUDA/cuDNN/BLAS
|
||||||
|
# on Spark); baked into sitecustomize.py, missing dirs skipped
|
||||||
|
self.runtime_dll_dirs = runtime_dll_dirs or []
|
||||||
|
# env vars every venv python needs (sitecustomize setdefault)
|
||||||
|
self.runtime_env = {}
|
||||||
|
|
||||||
def torch_args(self):
|
def torch_args(self):
|
||||||
args = ["%s==%s" % (k, v) for k, v in sorted(self.torch_packages.items())]
|
args = ["%s==%s" % (k, v) for k, v in sorted(self.torch_packages.items())]
|
||||||
if self.torch_index:
|
if self.torch_index:
|
||||||
args += ["--index-url", self.torch_index]
|
args += ["--index-url", self.torch_index]
|
||||||
|
for links in self.torch_links:
|
||||||
|
args += ["--find-links", links]
|
||||||
return args
|
return args
|
||||||
|
|
||||||
def torch_constraints(self):
|
def torch_constraints(self):
|
||||||
@@ -142,6 +217,11 @@ class EnvSpec(object):
|
|||||||
"optional_packages": self.optional_packages,
|
"optional_packages": self.optional_packages,
|
||||||
"find_links": self.find_links,
|
"find_links": self.find_links,
|
||||||
"notes": self.notes,
|
"notes": self.notes,
|
||||||
|
"uv_python": self.uv_python,
|
||||||
|
"torch_links": self.torch_links,
|
||||||
|
"no_deps_packages": self.no_deps_packages,
|
||||||
|
"runtime_dll_dirs": self.runtime_dll_dirs,
|
||||||
|
"runtime_env": self.runtime_env,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -186,8 +266,14 @@ def _cuda_flavor(detection):
|
|||||||
return "cu130", []
|
return "cu130", []
|
||||||
if cuda >= (13, 0):
|
if cuda >= (13, 0):
|
||||||
return "cu130", []
|
return "cu130", []
|
||||||
caps = [g.get("compute_cap") for g in nvidia.get("gpus", [])]
|
# non-GPU rows (the NPU on ARM hybrids) report compute_cap as "[N/A]"
|
||||||
has_blackwell = any(c and float(c) >= 12.0 for c in caps if c)
|
caps = []
|
||||||
|
for g in nvidia.get("gpus", []):
|
||||||
|
try:
|
||||||
|
caps.append(float(g.get("compute_cap")))
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
pass
|
||||||
|
has_blackwell = any(c >= 12.0 for c in caps)
|
||||||
if cuda >= (12, 6):
|
if cuda >= (12, 6):
|
||||||
if has_blackwell:
|
if has_blackwell:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
@@ -226,7 +312,10 @@ def _cuda_spec(detection):
|
|||||||
optional = []
|
optional = []
|
||||||
find_links = []
|
find_links = []
|
||||||
|
|
||||||
fa_url = _flash_attn_url(flavor, os_name, arch, python_version)
|
# Windows-on-ARM runs the x64 wheel stack (see module docstring), so wheel
|
||||||
|
# selection uses x86_64 there regardless of the host arch.
|
||||||
|
wheel_arch = "x86_64" if os_name == "windows" else arch
|
||||||
|
fa_url = _flash_attn_url(flavor, os_name, wheel_arch, python_version)
|
||||||
if fa_url:
|
if fa_url:
|
||||||
optional.append(fa_url)
|
optional.append(fa_url)
|
||||||
|
|
||||||
@@ -250,8 +339,87 @@ def _cuda_spec(detection):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _spark_spec(detection, wheels_source):
|
||||||
|
"""Native win_arm64 CUDA 13.4 stack from self-built wheels (RTX Spark)."""
|
||||||
|
from . import sparkdeps
|
||||||
|
|
||||||
|
dll_dirs = sparkdeps.resolve_dll_dirs()
|
||||||
|
spec = _make_spark_spec(detection, wheels_source, dll_dirs)
|
||||||
|
# OpenCV's runtime CPU detection is blind to FP16/DOTPROD on Windows
|
||||||
|
# ARM64 and aborts the process at import even though the N1X supports
|
||||||
|
# both; the self-built cv2 wheels rely on this skip.
|
||||||
|
spec.runtime_env["OPENCV_SKIP_CPU_BASELINE_CHECK"] = "1"
|
||||||
|
# our triton wheel does not bundle NVIDIA's compiler tools (preview
|
||||||
|
# licensing) — point its knobs at the user's CUDA toolkit
|
||||||
|
spec.runtime_env.update(sparkdeps.triton_tool_env())
|
||||||
|
return spec
|
||||||
|
|
||||||
|
|
||||||
|
def _make_spark_spec(detection, wheels_source, dll_dirs):
|
||||||
|
return EnvSpec(
|
||||||
|
SPARK_BACKEND,
|
||||||
|
SPARK_TORCH,
|
||||||
|
torch_index=None,
|
||||||
|
python_version="3.12",
|
||||||
|
requirements_file="spark_requirements.txt",
|
||||||
|
# all pins resolve from the spark wheel set (self-built win_arm64):
|
||||||
|
# torchcodec, plus our triton port (torch.compile / compiled flex
|
||||||
|
# attention) — see the wheel-set build notes
|
||||||
|
extra_packages=_WIN_HELPERS + [
|
||||||
|
"torchcodec==0.15.0",
|
||||||
|
"triton==3.8.0+git8743423b",
|
||||||
|
],
|
||||||
|
# import-verified with rollback, like accelerators on other platforms
|
||||||
|
optional_packages=["flash-attn==2.8.3+cu134torch2.14", "natten==0.21.7"],
|
||||||
|
find_links=[wheels_source],
|
||||||
|
notes=[
|
||||||
|
"RTX Spark native mode: win_arm64 CUDA %s stack from the "
|
||||||
|
"ai-toolkit wheel set (CUDA 13.4 developer preview), including "
|
||||||
|
"self-built flash-attn, NATTEN and triton (torch.compile)."
|
||||||
|
% SPARK_BACKEND,
|
||||||
|
],
|
||||||
|
uv_python=SPARK_UV_PYTHON,
|
||||||
|
torch_links=[wheels_source],
|
||||||
|
no_deps_packages=["tensorboard"],
|
||||||
|
runtime_dll_dirs=dll_dirs,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def build_spec(detection, allow_cpu=False):
|
def build_spec(detection, allow_cpu=False):
|
||||||
"""Returns EnvSpec, or raises RuntimeError with a user-facing message."""
|
"""Returns EnvSpec, or raises RuntimeError with a user-facing message."""
|
||||||
|
if (
|
||||||
|
detection["os"] == "windows"
|
||||||
|
and detection["arch"] == "aarch64"
|
||||||
|
and detection.get("backend") == "cuda"
|
||||||
|
and os.environ.get("AITK_SPARK_NATIVE", "1") != "0"
|
||||||
|
and _spark_capable(detection)
|
||||||
|
):
|
||||||
|
from . import sparkdeps
|
||||||
|
|
||||||
|
wheels_source = _spark_wheels_source()
|
||||||
|
# the CUDA toolkit is the one manual install (preview EULA) — without
|
||||||
|
# it the native wheels cannot run, so fall back to the x64 stack
|
||||||
|
if wheels_source and sparkdeps.cuda_bin_dir():
|
||||||
|
return _spark_spec(detection, wheels_source)
|
||||||
|
|
||||||
|
spec = _build_spec(detection, allow_cpu=allow_cpu)
|
||||||
|
if detection["os"] == "windows" and detection["arch"] == "aarch64":
|
||||||
|
# Emulated-x64 fallback (no native wheel source, old driver, or
|
||||||
|
# AITK_SPARK_NATIVE=0). Pin the venv interpreter to x64 explicitly:
|
||||||
|
# uv currently defaults to an emulated x86_64 Python on arm64 hosts,
|
||||||
|
# but says it will flip to native aarch64 once it considers support
|
||||||
|
# mature — which would silently leave a venv where the cu130 torch
|
||||||
|
# wheels don't resolve.
|
||||||
|
spec.uv_python = "cpython-%s-windows-x86_64-none" % spec.python_version
|
||||||
|
spec.notes.append(
|
||||||
|
"Windows-on-ARM detected: using the x64 stack under Windows' "
|
||||||
|
"emulation (GPU work still runs natively via the NVIDIA driver). "
|
||||||
|
"Native mode needs a CUDA 13.4+ driver and the Spark wheel set."
|
||||||
|
)
|
||||||
|
return spec
|
||||||
|
|
||||||
|
|
||||||
|
def _build_spec(detection, allow_cpu=False):
|
||||||
os_name = detection["os"]
|
os_name = detection["os"]
|
||||||
|
|
||||||
if os_name == "mac":
|
if os_name == "mac":
|
||||||
|
|||||||
@@ -201,10 +201,36 @@ def download(url, dest, label=None):
|
|||||||
except Exception as e: # noqa: BLE001 - surface any network failure clearly
|
except Exception as e: # noqa: BLE001 - surface any network failure clearly
|
||||||
if os.path.exists(tmp):
|
if os.path.exists(tmp):
|
||||||
os.remove(tmp)
|
os.remove(tmp)
|
||||||
die("Download failed for %s: %s" % (url, e))
|
# Windows: python's ssl verifies against the OS cert store but never
|
||||||
|
# triggers Windows' on-demand intermediate-CA fetching, so on a fresh
|
||||||
|
# machine github downloads can fail CERTIFICATE_VERIFY_FAILED even
|
||||||
|
# though the chain is fine. curl (bundled since Win10, schannel-based)
|
||||||
|
# does fetch intermediates — fall back to it before giving up.
|
||||||
|
if not _download_with_curl(url, tmp, label):
|
||||||
|
die("Download failed for %s: %s" % (url, e))
|
||||||
os.replace(tmp, dest)
|
os.replace(tmp, dest)
|
||||||
|
|
||||||
|
|
||||||
|
def _download_with_curl(url, tmp, label):
|
||||||
|
curl = shutil.which("curl")
|
||||||
|
if not curl:
|
||||||
|
return False
|
||||||
|
info(" %s: retrying with curl..." % label)
|
||||||
|
try:
|
||||||
|
code = subprocess.call(
|
||||||
|
[curl, "-fSL", "--retry", "3", "-o", tmp, url],
|
||||||
|
stdout=subprocess.DEVNULL,
|
||||||
|
stderr=subprocess.DEVNULL,
|
||||||
|
)
|
||||||
|
except OSError:
|
||||||
|
return False
|
||||||
|
if code != 0 or not os.path.exists(tmp):
|
||||||
|
if os.path.exists(tmp):
|
||||||
|
os.remove(tmp)
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
def extract_archive(archive, dest_dir):
|
def extract_archive(archive, dest_dir):
|
||||||
"""Extract .zip / .tar.* into dest_dir (created if needed)."""
|
"""Extract .zip / .tar.* into dest_dir (created if needed)."""
|
||||||
import tarfile
|
import tarfile
|
||||||
|
|||||||
84
spark_requirements.txt
Normal file
84
spark_requirements.txt
Normal file
@@ -0,0 +1,84 @@
|
|||||||
|
# NVIDIA RTX Spark — native Windows-on-ARM (win_arm64) requirement set.
|
||||||
|
#
|
||||||
|
# Mirrors requirements_base.txt with pins adjusted to versions that actually
|
||||||
|
# publish win_arm64 wheels (or that we build ourselves — see manager/spec.py
|
||||||
|
# spark handling). Keep in sync with requirements_base.txt when bumping pins.
|
||||||
|
#
|
||||||
|
# Deliberate differences from requirements_base.txt:
|
||||||
|
# - scipy>=1.17 first release with win_arm64 wheels (also unlocks numpy 2)
|
||||||
|
# - matplotlib>=3.11 first release with win_arm64 wheels (base: 3.10.1)
|
||||||
|
# - av>=17 first release with win_arm64 wheels (base: 16.0.1)
|
||||||
|
# - numba/llvmlite rc only prerelease win_arm64 wheels exist (librosa dep)
|
||||||
|
# - tensorboard REMOVED: its grpcio dep has no win_arm64 wheels; the
|
||||||
|
# manager installs tensorboard --no-deps (writer path
|
||||||
|
# only needs protobuf) as a spark extra
|
||||||
|
# - hf_transfer REMOVED: no win_arm64 wheels; hf-xet (arm64 OK) is the
|
||||||
|
# modern replacement and installs via manager extras
|
||||||
|
# - torchcodec pin REMOVED: the manager installs the locally built
|
||||||
|
# win_arm64 torchcodec wheel as an extra
|
||||||
|
# - opencv-python, soxr, pywavelets, brotli(gradio), kornia-rs have no
|
||||||
|
# published win_arm64 wheels at any version — supplied from the ai-toolkit
|
||||||
|
# spark wheel set (self-built), resolved via --find-links
|
||||||
|
|
||||||
|
numpy>=2,<3
|
||||||
|
scipy>=1.17
|
||||||
|
|
||||||
|
torchao==0.17.0
|
||||||
|
safetensors
|
||||||
|
git+https://github.com/huggingface/diffusers.git@c943837899b16cbae2f619b8dd4f7bb6f07dd81a
|
||||||
|
transformers==5.5.3
|
||||||
|
lycoris-lora==1.8.3
|
||||||
|
flatten_json
|
||||||
|
pyyaml
|
||||||
|
oyaml
|
||||||
|
kornia
|
||||||
|
invisible-watermark
|
||||||
|
einops
|
||||||
|
accelerate
|
||||||
|
toml
|
||||||
|
albumentations==1.4.15
|
||||||
|
albucore==0.0.16
|
||||||
|
pydantic
|
||||||
|
omegaconf
|
||||||
|
open_clip_torch
|
||||||
|
timm==1.0.22
|
||||||
|
prodigyopt
|
||||||
|
controlnet_aux==0.0.10
|
||||||
|
python-dotenv
|
||||||
|
bitsandbytes
|
||||||
|
lpips
|
||||||
|
pytorch_fid
|
||||||
|
optimum-quanto==0.2.4
|
||||||
|
sentencepiece
|
||||||
|
huggingface_hub==1.23.0
|
||||||
|
peft==0.18.1
|
||||||
|
gradio
|
||||||
|
python-slugify
|
||||||
|
# exact pins: newer opencv versions exist on PyPI only as sdists (which fail
|
||||||
|
# to build on MSVC arm64 — dnn __fp16); force the self-built 4.12 wheels
|
||||||
|
opencv-python==4.12.0.88
|
||||||
|
opencv-python-headless==4.12.0.88
|
||||||
|
pytorch-wavelets==1.3.0
|
||||||
|
matplotlib>=3.11
|
||||||
|
setuptools>=77.0.3
|
||||||
|
av>=17
|
||||||
|
# librosa chain: numba/llvmlite have no cp312 win_arm64 wheels upstream — the
|
||||||
|
# rc pins below resolve from the self-built spark wheel set (llvmlite built
|
||||||
|
# against our own LLVM 22 static build)
|
||||||
|
librosa==0.11.0
|
||||||
|
numba==0.67.0rc1
|
||||||
|
llvmlite==0.49.0rc1
|
||||||
|
soxr==1.1.0
|
||||||
|
mutagen==1.47.0
|
||||||
|
soundfile
|
||||||
|
|
||||||
|
# tensorboard is installed --no-deps by the manager (its grpcio dep has no
|
||||||
|
# win_arm64 wheels; grpc is only needed by the tensorboard server, not the
|
||||||
|
# SummaryWriter path the toolkit uses). Its remaining runtime deps:
|
||||||
|
absl-py
|
||||||
|
markdown
|
||||||
|
werkzeug
|
||||||
|
tensorboard-data-server
|
||||||
|
packaging
|
||||||
|
protobuf
|
||||||
|
six
|
||||||
Reference in New Issue
Block a user