From aa762103b39373507e38fba56ccb792e44b1ce0b Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Tue, 28 Jul 2026 12:27:51 -0600 Subject: [PATCH] Add build support for Nvidia Spark --- .../captioner/AceStepCaptioner.py | 10 +- manager/doctor.py | 14 +- manager/env.py | 186 +++++++++++++-- manager/ffmpeg.py | 20 +- manager/gitwin.py | 21 +- manager/nodejs.py | 46 +++- manager/sparkdeps.py | 224 ++++++++++++++++++ manager/spec.py | 176 +++++++++++++- manager/util.py | 28 ++- spark_requirements.txt | 84 +++++++ 10 files changed, 765 insertions(+), 44 deletions(-) create mode 100644 manager/sparkdeps.py create mode 100644 spark_requirements.txt diff --git a/extensions_built_in/captioner/AceStepCaptioner.py b/extensions_built_in/captioner/AceStepCaptioner.py index ce89a2f..e1d77b3 100644 --- a/extensions_built_in/captioner/AceStepCaptioner.py +++ b/extensions_built_in/captioner/AceStepCaptioner.py @@ -1,6 +1,9 @@ from typing import Optional -import librosa +try: + import librosa +except ImportError: + librosa = None import numpy as np import torch 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): """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) duration = librosa.get_duration(y=y, sr=sr) diff --git a/manager/doctor.py b/manager/doctor.py index 2b3043b..554c7e2 100644 --- a/manager/doctor.py +++ b/manager/doctor.py @@ -25,7 +25,19 @@ def run_doctor(): 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() _check( "git", diff --git a/manager/env.py b/manager/env.py index cd0d380..5551b58 100644 --- a/manager/env.py +++ b/manager/env.py @@ -62,24 +62,79 @@ def venv_exists(): 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): """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(): return venv_python() target = venv_dir() 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: info( "[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) 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( - [uv, "venv", target, "--python", spec.python_version, "--seed"], + [uv, "venv", target, "--python", python_request, "--seed"], env=clean_env(), ) 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." % (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: warn( "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 +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): uv = find_uv() if uv: @@ -419,11 +494,13 @@ def ensure_requirements(spec, dry_run=False, force=False): _pip_uninstall(stale, dry_run=dry_run) # every pass below carries the torch pins so nothing can swap the GPU build 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) for url in spec.find_links: 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) if 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 ) _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) if not dry_run: state = load_state() @@ -467,13 +551,55 @@ def ensure_requirements(spec, dry_run=False, force=False): # ---------------------------------------------------------------- sitecustomize -def write_sitecustomize(dry_run=False): - """Drop a sitecustomize.py into the venv that exposes the local ffmpeg. +def _msvc_runtime_env(): + """{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 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 - 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 @@ -498,19 +624,29 @@ def write_sitecustomize(dry_run=False): warn("Could not locate venv site-packages — skipping sitecustomize.") return 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 = ( "# Generated by the AI Toolkit manager (manager/env.py). Do not edit;\n" "# regenerated on every `manager sync`.\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" - "if os.path.isdir(_FFMPEG_BIN):\n" - " os.environ['PATH'] = _FFMPEG_BIN + os.pathsep + os.environ.get('PATH', '')\n" - " if hasattr(os, 'add_dll_directory'):\n" - " try:\n" - " os.add_dll_directory(_FFMPEG_BIN)\n" - " except OSError:\n" - " pass\n" + "for _d in _DLL_DIRS:\n" + " if os.path.isdir(_d):\n" + " os.environ['PATH'] = _d + os.pathsep + os.environ.get('PATH', '')\n" + " if hasattr(os, 'add_dll_directory'):\n" + " try:\n" + " os.add_dll_directory(_d)\n" + " except OSError:\n" + " pass\n" "if os.path.isdir(_FFMPEG_LIB):\n" " # inherited by child processes (the ffmpeg/ffprobe executables\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" " _FFMPEG_LIB + ((os.pathsep + _prev) if _prev else '')\n" " )\n" - ) % (ffmpeg.bin_dir(), ffmpeg.lib_dir()) + ) % (runtime_env, dll_dirs, ffmpeg.lib_dir()) if dry_run: info("[dry-run] would write %s" % target) return @@ -538,13 +674,23 @@ def sync(spec, detection, dry_run=False, force=False): warn(note) uvbin.ensure_uv(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) + # 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) # a torch reinstall can clobber pinned deps; force req pass afterwards 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_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) ok("Environment is up to date.") diff --git a/manager/ffmpeg.py b/manager/ffmpeg.py index 2cc80fc..437d914 100644 --- a/manager/ffmpeg.py +++ b/manager/ffmpeg.py @@ -48,8 +48,17 @@ _SOURCES = { ("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", ("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(): return os.path.join(FFMPEG_DIR, "bin") @@ -124,15 +133,20 @@ def _install_mac(detection): shutil.rmtree(tmp, ignore_errors=True) -def source_url(detection): +def source_url(detection, spec=None): if detection["os"] == "mac": arch = "arm64" if detection["arch"] == "arm64" else "amd64" 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"])) -def ensure_ffmpeg(detection, dry_run=False): - url = source_url(detection) +def ensure_ffmpeg(detection, dry_run=False, spec=None): + url = source_url(detection, spec=spec) if url is None: warn( "No portable FFmpeg source for %s/%s — skipping local ffmpeg." diff --git a/manager/gitwin.py b/manager/gitwin.py index 72b9920..1ec39dd 100644 --- a/manager/gitwin.py +++ b/manager/gitwin.py @@ -13,6 +13,7 @@ the manager just errors with install instructions elsewhere. """ import os +import platform import shutil 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") # Update this pin together with nothing else — it's independent of torch etc. -MINGIT_URL = ( - "https://github.com/git-for-windows/git/releases/download/" - "v2.55.0.windows.3/MinGit-2.55.0.3-64-bit.zip" -) +_MINGIT_TAG = "v2.55.0.windows.3" +_MINGIT_VERSION = "2.55.0.3" + + +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(): @@ -49,7 +60,7 @@ def ensure_git(dry_run=False): tmp = tempfile.mkdtemp(prefix="aitk_mingit_") try: 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 extracted = os.path.join(tmp, "mingit") extract_archive(archive, extracted) diff --git a/manager/nodejs.py b/manager/nodejs.py index 011f1e9..787a038 100644 --- a/manager/nodejs.py +++ b/manager/nodejs.py @@ -43,20 +43,37 @@ def local_node_exe(): 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: out = subprocess.run( - [exe, "--version"], + [exe, "-p", "process.version + ' ' + process.arch"], stdout=subprocess.PIPE, stderr=subprocess.DEVNULL, timeout=15, ) if out.returncode != 0: - return None - text = out.stdout.decode().strip() # v24.11.1 - return int(text.lstrip("v").split(".")[0]) + return None, None + version, arch = out.stdout.decode().strip().split() # v24.11.1 x64 + return int(version.lstrip("v").split(".")[0]), arch 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 + if IS_WINDOWS and arch != "x64": + return None + return major def _dist_url(detection): @@ -68,6 +85,8 @@ def _dist_url(detection): plat = "darwin-arm64" if arch == "arm64" else "darwin-x64" ext = "tar.gz" 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" ext = "zip" else: @@ -80,13 +99,13 @@ def have_usable_node(): """(exe, major) for the best available node: local .node/ first, then system.""" local = local_node_exe() if os.path.isfile(local): - major = _node_major(local) - if major is not None and major >= MIN_NODE_MAJOR: + major = _usable_major(local) + if major is not None: return local, major system = which("node") if system: - major = _node_major(system) - if major is not None and major >= MIN_NODE_MAJOR: + major = _usable_major(system) + if major is not None: return system, major return None, None @@ -96,6 +115,15 @@ def ensure_node(detection, dry_run=False): if exe: ok("Node.js v%d found (%s)." % (major, exe)) 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) if url is None: warn( diff --git a/manager/sparkdeps.py b/manager/sparkdeps.py new file mode 100644 index 0000000..f5062ee --- /dev/null +++ b/manager/sparkdeps.py @@ -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() + # \bin\arm64 -> + 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).") diff --git a/manager/spec.py b/manager/spec.py index ce3d459..e06e413 100644 --- a/manager/spec.py +++ b/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. - 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. +- 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 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/" +# ---- 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): def __init__( @@ -81,8 +134,12 @@ class EnvSpec(object): optional_packages=None, find_links=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_index = torch_index # None = PyPI self.python_version = python_version @@ -91,11 +148,29 @@ class EnvSpec(object): self.optional_packages = optional_packages or [] self.find_links = find_links 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): args = ["%s==%s" % (k, v) for k, v in sorted(self.torch_packages.items())] if self.torch_index: args += ["--index-url", self.torch_index] + for links in self.torch_links: + args += ["--find-links", links] return args def torch_constraints(self): @@ -142,6 +217,11 @@ class EnvSpec(object): "optional_packages": self.optional_packages, "find_links": self.find_links, "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", [] if cuda >= (13, 0): return "cu130", [] - caps = [g.get("compute_cap") for g in nvidia.get("gpus", [])] - has_blackwell = any(c and float(c) >= 12.0 for c in caps if c) + # non-GPU rows (the NPU on ARM hybrids) report compute_cap as "[N/A]" + 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 has_blackwell: raise RuntimeError( @@ -226,7 +312,10 @@ def _cuda_spec(detection): optional = [] 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: 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): """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"] if os_name == "mac": diff --git a/manager/util.py b/manager/util.py index 5f1a50f..e7d2d61 100644 --- a/manager/util.py +++ b/manager/util.py @@ -201,10 +201,36 @@ def download(url, dest, label=None): except Exception as e: # noqa: BLE001 - surface any network failure clearly if os.path.exists(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) +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): """Extract .zip / .tar.* into dest_dir (created if needed).""" import tarfile diff --git a/spark_requirements.txt b/spark_requirements.txt new file mode 100644 index 0000000..55c6c14 --- /dev/null +++ b/spark_requirements.txt @@ -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