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

View File

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

View File

@@ -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,17 +624,27 @@ 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"
"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(_FFMPEG_BIN)\n"
" os.add_dll_directory(_d)\n"
" except OSError:\n"
" pass\n"
"if os.path.isdir(_FFMPEG_LIB):\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.")

View File

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

View File

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

View File

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

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

View File

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

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