Add build support for Nvidia Spark

This commit is contained in:
Jaret Burkett
2026-07-28 12:27:51 -06:00
parent 6d6c5a3d91
commit aa762103b3
10 changed files with 765 additions and 44 deletions

View File

@@ -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)

View File

@@ -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",

View File

@@ -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.")

View File

@@ -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."

View File

@@ -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)

View File

@@ -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
View 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).")

View File

@@ -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":

View File

@@ -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
View 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