Built a universal manager and installer for all operating systems and environments. Bumped a lot of versions of things. Still needs deep testing.
This commit is contained in:
3
.gitignore
vendored
3
.gitignore
vendored
@@ -124,6 +124,9 @@ celerybeat.pid
|
||||
.venv
|
||||
.python
|
||||
.node
|
||||
.ffmpeg
|
||||
.mingit
|
||||
.uv
|
||||
env/
|
||||
venv/
|
||||
ENV/
|
||||
|
||||
@@ -81,7 +81,7 @@ cd ai-toolkit
|
||||
python3 -m venv venv
|
||||
source venv/bin/activate
|
||||
# install torch first
|
||||
pip3 install --no-cache-dir torch==2.9.1 torchvision==0.24.1 torchaudio==2.9.1 --index-url https://download.pytorch.org/whl/cu128
|
||||
pip3 install --no-cache-dir torch==2.13.0 torchvision==0.28.0 torchaudio==2.11.0 --index-url https://download.pytorch.org/whl/cu130
|
||||
pip3 install -r requirements.txt
|
||||
```
|
||||
|
||||
@@ -97,7 +97,7 @@ git clone https://github.com/ostris/ai-toolkit.git
|
||||
cd ai-toolkit
|
||||
python -m venv venv
|
||||
.\venv\Scripts\activate
|
||||
pip install --no-cache-dir torch==2.9.1 torchvision==0.24.1 torchaudio==2.9.1 --index-url https://download.pytorch.org/whl/cu128
|
||||
pip install --no-cache-dir torch==2.13.0 torchvision==0.28.0 torchaudio==2.11.0 --index-url https://download.pytorch.org/whl/cu130
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
|
||||
@@ -37,7 +37,7 @@ conda activate ai-toolkit
|
||||
**2) Install PyTorch**
|
||||
|
||||
```
|
||||
pip3 install torch==2.9.1 torchvision==0.24.1 torchaudio==2.9.1 --index-url https://download.pytorch.org/whl/cu130
|
||||
pip3 install torch==2.13.0 torchvision==0.28.0 torchaudio==2.11.0 --index-url https://download.pytorch.org/whl/cu130
|
||||
```
|
||||
|
||||
|
||||
|
||||
70
manager/README.md
Normal file
70
manager/README.md
Normal file
@@ -0,0 +1,70 @@
|
||||
# AI Toolkit Manager
|
||||
|
||||
Self-contained install/update manager for this checkout of AI Toolkit. Runs
|
||||
with any Python >= 3.8 and **no dependencies**, so it works before the
|
||||
training environment exists.
|
||||
|
||||
```bash
|
||||
python3 -m manager install # first-time setup: venv + torch + requirements
|
||||
python3 -m manager check # is an update available / are deps out of sync?
|
||||
python3 -m manager update # git pull, then sync deps + run migrations
|
||||
python3 -m manager launch # start the web UI (http://localhost:8675)
|
||||
python3 -m manager doctor # diagnose problems
|
||||
```
|
||||
|
||||
## Design
|
||||
|
||||
- **The install logic lives in the repo it installs.** Every commit knows how
|
||||
to install itself; external frontends (the desktop launcher, `install.sh`,
|
||||
`install.ps1`) just shell out to this CLI and stay dumb. Machine-readable
|
||||
output via `check --json` / `detect --json`.
|
||||
- **Hardware → spec mapping** is in [spec.py](spec.py). One universal torch
|
||||
pin (2.12.0 / torchvision 0.27.0 / torchaudio 2.11.0) on every platform:
|
||||
cu130 wheels when the driver supports CUDA 13 (cu126 fallback for older
|
||||
drivers, refused outright on Blackwell GPUs which need cu130), same stack +
|
||||
Python 3.11 + `dgx_requirements.txt` on DGX/Grace, PyPI wheels on Mac,
|
||||
rocm7.1 (experimental) for AMD, `--cpu` to force a CPU install. **Torch
|
||||
pins there must be updated together with the README install instructions,
|
||||
run_mac.zsh, and dgx_instructions.md.**
|
||||
- **Accelerators everywhere wheels exist**, via per-spec `extra_packages`
|
||||
(installed after requirements with `--upgrade` so they override pins) and
|
||||
`optional_packages` (installed one-by-one, warn-only on failure):
|
||||
`torchcodec==0.15.0` on all platforms; flash-attn 2.8.3 prebuilt wheels
|
||||
(mjun0812) on Linux x86_64/aarch64 + Windows; NATTEN 0.21.7 wheels
|
||||
(whl.natten.org) on Linux both arches; triton bundled with torch on Linux
|
||||
and `triton-windows` 3.7.x on Windows. No flash-attn/NATTEN/triton on Mac,
|
||||
no NATTEN on Windows (no wheels exist).
|
||||
- **Nothing global is ever installed.** FFmpeg (shared builds — the libs
|
||||
torchcodec dlopens) goes to `.ffmpeg/` ([ffmpeg.py](ffmpeg.py)), Node
|
||||
(when the system lacks >= 20) to `.node/` ([nodejs.py](nodejs.py)), the uv
|
||||
binary (when absent) to `.uv/` ([uvbin.py](uvbin.py)) with uv-managed
|
||||
Pythons kept in `.uv/python/` via `UV_PYTHON_INSTALL_DIR`, and on Windows
|
||||
without git, portable MinGit to `.mingit/` ([gitwin.py](gitwin.py)) — all
|
||||
inside the repo and gitignored. The first clone on a git-less Windows
|
||||
box is handled by the bootstrap layer (install.ps1 / desktop launcher),
|
||||
which downloads MinGit itself and moves it into the checkout afterwards. `manager launch` puts them on PATH (and
|
||||
LD_LIBRARY_PATH on Linux) for the whole UI/training process tree, and a
|
||||
generated `sitecustomize.py` in the venv exposes ffmpeg to any direct use
|
||||
of the venv python (plus `os.add_dll_directory` on Windows).
|
||||
- **Hostile-environment hardening** (learned from the community Windows
|
||||
installer): every python/pip subprocess runs with PYTHONPATH/PYTHONHOME/
|
||||
CONDA/PYENV/PIP_* scrubbed from the env; git runs with
|
||||
`GIT_LFS_SKIP_SMUDGE=1`; git-pinned requirements (diffusers) are
|
||||
force-reinstalled when requirements change since pip skips unchanged
|
||||
version numbers; `launch` polls the UI port and opens the browser when
|
||||
ready (`--no-browser` to disable, auto-skipped on headless boxes).
|
||||
- **uv is used when present** (fast installs, auto-downloads the right
|
||||
Python); plain `venv` + `pip` otherwise. The venv is created at `.venv/`
|
||||
(an existing `venv/` is also respected, matching `ui/cron/pythonPath.ts`).
|
||||
- **State** (requirements hash, applied migrations) lives inside the venv
|
||||
(`aitk_manager_state.json`) — deleting the venv resets everything.
|
||||
- **Update flow**: `update` pulls fast-forward only, then **re-execs**
|
||||
`python -m manager sync` so the freshly pulled manager code — not the stale
|
||||
in-memory copy — performs its own dependency sync and migrations.
|
||||
**Local work is never overwritten**: a dirty tree aborts the update by
|
||||
default (untracked files don't count), `--auto` (used by the run_* scripts)
|
||||
warns and skips the pull instead so launching still works, and there is no
|
||||
reset/clean anywhere — even `--force` relies on git itself refusing to
|
||||
clobber modified files.
|
||||
- **Migrations** ([migrations.py](migrations.py)): one-time post-update steps,
|
||||
each applied at most once per environment.
|
||||
6
manager/__init__.py
Normal file
6
manager/__init__.py
Normal file
@@ -0,0 +1,6 @@
|
||||
"""AI Toolkit install / update manager. Stdlib-only; see __main__.py."""
|
||||
|
||||
# Version of the CLI contract consumed by external frontends (desktop
|
||||
# launcher, install scripts). Bump only on breaking changes to command
|
||||
# names/flags or --json output shapes.
|
||||
CLI_CONTRACT_VERSION = 1
|
||||
249
manager/__main__.py
Normal file
249
manager/__main__.py
Normal file
@@ -0,0 +1,249 @@
|
||||
"""AI Toolkit manager CLI.
|
||||
|
||||
Runs with any Python >= 3.8 and no dependencies, so it works before the
|
||||
training environment exists. This is the single entry point every installer
|
||||
frontend (shell scripts, the desktop launcher, the web UI) shells out to.
|
||||
|
||||
python3 -m manager install first-time environment setup
|
||||
python3 -m manager check [--json] is an update / dep sync needed?
|
||||
python3 -m manager update git pull + dependency sync + migrations
|
||||
python3 -m manager sync dependency sync only (no git pull)
|
||||
python3 -m manager launch start the web UI
|
||||
python3 -m manager detect [--json] show detected hardware
|
||||
python3 -m manager doctor full environment diagnostics
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
# allow `python manager/__main__.py` as well as `python -m manager`
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
from manager import detect as detect_mod
|
||||
from manager import env, gitops, launch, spec as spec_mod, util
|
||||
from manager.util import die, info, ok, print_json, warn
|
||||
|
||||
|
||||
def _resolve_spec(args):
|
||||
detection = detect_mod.detect()
|
||||
try:
|
||||
return detection, spec_mod.build_spec(
|
||||
detection, allow_cpu=getattr(args, "cpu", False)
|
||||
)
|
||||
except RuntimeError as e:
|
||||
die(str(e))
|
||||
|
||||
|
||||
def cmd_detect(args):
|
||||
detection = detect_mod.detect()
|
||||
try:
|
||||
s = spec_mod.build_spec(detection, allow_cpu=True)
|
||||
detection["spec"] = s.as_dict()
|
||||
except RuntimeError as e:
|
||||
detection["spec_error"] = str(e)
|
||||
if args.json:
|
||||
print_json(detection)
|
||||
else:
|
||||
backend = detection.get("spec", {}).get("backend", "unknown")
|
||||
info("os=%s arch=%s backend=%s" % (detection["os"], detection["arch"], backend))
|
||||
if detection["nvidia"]:
|
||||
for gpu in detection["nvidia"]["gpus"]:
|
||||
info("gpu: %s (%s)" % (gpu["name"], gpu["memory"]))
|
||||
|
||||
|
||||
def cmd_install(args):
|
||||
detection, s = _resolve_spec(args)
|
||||
env.sync(s, detection, dry_run=args.dry_run, force=args.force)
|
||||
if not args.dry_run:
|
||||
ok("Install complete. Start the UI with: python3 -m manager launch")
|
||||
|
||||
|
||||
def cmd_sync(args):
|
||||
detection, s = _resolve_spec(args)
|
||||
env.sync(s, detection, dry_run=args.dry_run, force=args.force)
|
||||
|
||||
|
||||
def cmd_check(args):
|
||||
_, s = _resolve_spec(args)
|
||||
fetched = gitops.fetch()
|
||||
behind = gitops.behind_count()
|
||||
data = {
|
||||
"version": _toolkit_version(),
|
||||
"branch": gitops.current_branch(),
|
||||
"commit": gitops.current_commit(),
|
||||
"remote_commit": gitops.remote_commit(),
|
||||
"dirty": gitops.is_dirty(),
|
||||
"fetch_ok": fetched,
|
||||
"behind": behind,
|
||||
"incoming": gitops.incoming_log(),
|
||||
"venv": env.venv_exists(),
|
||||
"deps_in_sync": env.venv_exists()
|
||||
and env.torch_matches(s)
|
||||
and env.requirements_in_sync(s),
|
||||
"backend": s.backend,
|
||||
}
|
||||
data["update_available"] = bool(behind) or not data["deps_in_sync"]
|
||||
if args.json:
|
||||
print_json(data)
|
||||
return
|
||||
info("AI Toolkit %s (%s @ %s)" % (data["version"], data["branch"], data["commit"]))
|
||||
if not fetched:
|
||||
warn("Could not reach the remote (offline?) — update status may be stale.")
|
||||
if behind:
|
||||
info("Update available: %d new commit(s)." % behind)
|
||||
for line in data["incoming"]:
|
||||
print(" " + line)
|
||||
elif behind == 0:
|
||||
ok("Code is up to date.")
|
||||
if not data["deps_in_sync"]:
|
||||
warn("Dependencies are out of sync. Run: python3 -m manager sync")
|
||||
elif behind == 0:
|
||||
ok("Dependencies are in sync.")
|
||||
|
||||
|
||||
def cmd_update(args):
|
||||
"""git pull (never destructive) + dependency sync.
|
||||
|
||||
Local work is sacred: a dirty tree either aborts (default), or with
|
||||
--auto is skipped with a warning so run scripts can continue to launch.
|
||||
We never reset/clean; even a forced pull is --ff-only, which git itself
|
||||
aborts rather than overwriting local changes.
|
||||
"""
|
||||
auto = getattr(args, "auto", False)
|
||||
skip_pull = False
|
||||
|
||||
if gitops.is_dirty() and not args.force:
|
||||
if auto:
|
||||
warn(
|
||||
"Local changes detected — skipping the code update to protect "
|
||||
"your work. Commit or stash your changes to receive updates."
|
||||
)
|
||||
skip_pull = True
|
||||
else:
|
||||
die(
|
||||
"You have local changes to tracked files. Commit or stash them, "
|
||||
"or re-run with --force to attempt the update anyway (git will "
|
||||
"still refuse rather than overwrite your changes)."
|
||||
)
|
||||
|
||||
if not skip_pull and not gitops.fetch():
|
||||
if auto:
|
||||
warn("Could not reach the git remote — skipping the update check.")
|
||||
skip_pull = True
|
||||
else:
|
||||
die("Could not reach the git remote. Check your network and try again.")
|
||||
|
||||
if not skip_pull:
|
||||
behind = gitops.behind_count()
|
||||
if behind is None:
|
||||
warn("Current branch has no upstream; skipping git pull.")
|
||||
elif behind == 0:
|
||||
ok("Code already up to date.")
|
||||
else:
|
||||
info("Pulling %d new commit(s)..." % behind)
|
||||
gitops.pull_ff()
|
||||
ok("Code updated to %s." % gitops.current_commit())
|
||||
# Re-exec so the freshly pulled manager code runs its own dependency
|
||||
# sync and migrations (the in-memory copy of this module is stale now).
|
||||
cmd = [sys.executable, "-m", "manager", "sync"]
|
||||
if args.dry_run:
|
||||
cmd.append("--dry-run")
|
||||
sys.exit(subprocess.call(cmd, cwd=util.REPO_ROOT))
|
||||
# nothing was pulled — safe to sync with the code already loaded
|
||||
detection, s = _resolve_spec(args)
|
||||
env.sync(s, detection, dry_run=args.dry_run)
|
||||
|
||||
|
||||
def cmd_launch(args):
|
||||
sys.exit(launch.launch_ui(open_browser=not args.no_browser))
|
||||
|
||||
|
||||
def cmd_doctor(args):
|
||||
from manager import doctor
|
||||
|
||||
doctor.run_doctor()
|
||||
|
||||
|
||||
def cmd_version(args):
|
||||
print(_toolkit_version())
|
||||
|
||||
|
||||
def _toolkit_version():
|
||||
version = {}
|
||||
try:
|
||||
with open(os.path.join(util.REPO_ROOT, "version.py")) as f:
|
||||
exec(f.read(), version)
|
||||
return version.get("VERSION", "unknown")
|
||||
except OSError:
|
||||
return "unknown"
|
||||
|
||||
|
||||
def main(argv=None):
|
||||
parser = argparse.ArgumentParser(
|
||||
prog="manager", description="AI Toolkit install / update manager"
|
||||
)
|
||||
sub = parser.add_subparsers(dest="command")
|
||||
|
||||
def add(name, fn, **kwargs):
|
||||
p = sub.add_parser(name, **kwargs)
|
||||
p.set_defaults(fn=fn)
|
||||
return p
|
||||
|
||||
p = add("detect", cmd_detect, help="show detected hardware and env spec")
|
||||
p.add_argument("--json", action="store_true")
|
||||
|
||||
for name, fn, help_text in (
|
||||
("install", cmd_install, "first-time environment setup"),
|
||||
("sync", cmd_sync, "sync dependencies for the current checkout"),
|
||||
):
|
||||
p = add(name, fn, help=help_text)
|
||||
p.add_argument("--cpu", action="store_true", help="allow CPU-only install")
|
||||
p.add_argument("--dry-run", action="store_true")
|
||||
p.add_argument(
|
||||
"--force",
|
||||
action="store_true",
|
||||
help="reinstall requirements even if in sync",
|
||||
)
|
||||
|
||||
p = add("check", cmd_check, help="check for updates (use --json for machines)")
|
||||
p.add_argument("--json", action="store_true")
|
||||
p.add_argument("--cpu", action="store_true", help=argparse.SUPPRESS)
|
||||
|
||||
p = add("update", cmd_update, help="git pull + dependency sync + migrations")
|
||||
p.add_argument("--cpu", action="store_true", help="allow CPU-only install")
|
||||
p.add_argument("--dry-run", action="store_true")
|
||||
p.add_argument(
|
||||
"--force", action="store_true", help="update even with local changes"
|
||||
)
|
||||
p.add_argument(
|
||||
"--auto",
|
||||
action="store_true",
|
||||
help="unattended mode (run scripts): on local changes or an unreachable "
|
||||
"remote, warn and skip the code update instead of failing; deps still sync",
|
||||
)
|
||||
|
||||
p = add("launch", cmd_launch, help="start the web UI")
|
||||
p.add_argument(
|
||||
"--no-browser",
|
||||
action="store_true",
|
||||
help="do not open a browser when the UI is ready",
|
||||
)
|
||||
add("doctor", cmd_doctor, help="diagnose the environment")
|
||||
add("version", cmd_version, help="print the toolkit version")
|
||||
|
||||
args = parser.parse_args(argv)
|
||||
if not getattr(args, "command", None):
|
||||
parser.print_help()
|
||||
return 1
|
||||
util.set_json_mode(bool(getattr(args, "json", False)))
|
||||
try:
|
||||
args.fn(args)
|
||||
except KeyboardInterrupt:
|
||||
return 130
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
110
manager/detect.py
Normal file
110
manager/detect.py
Normal file
@@ -0,0 +1,110 @@
|
||||
"""Hardware / platform detection. Stdlib only, safe to run anywhere."""
|
||||
|
||||
import os
|
||||
import platform
|
||||
import re
|
||||
import subprocess
|
||||
|
||||
from .util import which, IS_WINDOWS, IS_MAC, IS_LINUX
|
||||
|
||||
|
||||
def _run_quiet(cmd):
|
||||
try:
|
||||
out = subprocess.run(
|
||||
cmd, stdout=subprocess.PIPE, stderr=subprocess.DEVNULL, timeout=15
|
||||
)
|
||||
if out.returncode != 0:
|
||||
return None
|
||||
return out.stdout.decode("utf-8", errors="replace")
|
||||
except (FileNotFoundError, subprocess.TimeoutExpired, OSError):
|
||||
return None
|
||||
|
||||
|
||||
def detect_nvidia():
|
||||
"""Returns dict with gpus + driver info, or None if no working nvidia-smi."""
|
||||
smi = which("nvidia-smi")
|
||||
if not smi:
|
||||
return None
|
||||
fields = "name,memory.total,driver_version,compute_cap"
|
||||
csv = _run_quiet([smi, "--query-gpu=" + fields, "--format=csv,noheader"])
|
||||
if not csv:
|
||||
# older drivers don't know the compute_cap field
|
||||
csv = _run_quiet(
|
||||
[
|
||||
smi,
|
||||
"--query-gpu=name,memory.total,driver_version",
|
||||
"--format=csv,noheader",
|
||||
]
|
||||
)
|
||||
if not csv:
|
||||
return None
|
||||
gpus = []
|
||||
driver = None
|
||||
for line in csv.strip().splitlines():
|
||||
parts = [p.strip() for p in line.split(",")]
|
||||
if len(parts) >= 3:
|
||||
gpu = {"name": parts[0], "memory": parts[1]}
|
||||
if len(parts) >= 4:
|
||||
gpu["compute_cap"] = parts[3]
|
||||
gpus.append(gpu)
|
||||
driver = parts[2]
|
||||
if not gpus:
|
||||
return None
|
||||
# Max CUDA version the driver supports only appears in the banner output
|
||||
banner = _run_quiet([smi]) or ""
|
||||
m = re.search(r"CUDA Version:\s*([0-9]+\.[0-9]+)", banner)
|
||||
cuda_version = m.group(1) if m else None
|
||||
return {"gpus": gpus, "driver": driver, "cuda_version": cuda_version}
|
||||
|
||||
|
||||
def detect_rocm():
|
||||
"""Returns dict if an AMD ROCm stack is present, else None."""
|
||||
if not IS_LINUX:
|
||||
return None
|
||||
smi = which("rocm-smi")
|
||||
if not smi and not os.path.isdir("/opt/rocm"):
|
||||
return None
|
||||
gpus = []
|
||||
out = _run_quiet([smi, "--showproductname"]) if smi else None
|
||||
if out:
|
||||
for m in re.finditer(r"Card [Ss]eries:\s*(.+)", out):
|
||||
gpus.append({"name": m.group(1).strip()})
|
||||
return {"gpus": gpus}
|
||||
|
||||
|
||||
def detect():
|
||||
"""Full platform detection. Returns a plain dict (json-serializable)."""
|
||||
system = platform.system() # Linux / Darwin / Windows
|
||||
arch = platform.machine().lower() # x86_64 / amd64 / arm64 / aarch64
|
||||
if arch == "amd64":
|
||||
arch = "x86_64"
|
||||
if arch == "arm64" and not IS_MAC:
|
||||
arch = "aarch64"
|
||||
|
||||
result = {
|
||||
"os": {"Linux": "linux", "Darwin": "mac", "Windows": "windows"}.get(
|
||||
system, system.lower()
|
||||
),
|
||||
"arch": arch,
|
||||
"python": platform.python_version(),
|
||||
"nvidia": None,
|
||||
"rocm": None,
|
||||
"backend": "cpu",
|
||||
}
|
||||
|
||||
nvidia = detect_nvidia()
|
||||
if nvidia:
|
||||
result["nvidia"] = nvidia
|
||||
result["backend"] = "cuda"
|
||||
else:
|
||||
rocm = detect_rocm()
|
||||
if rocm:
|
||||
result["rocm"] = rocm
|
||||
result["backend"] = "rocm"
|
||||
|
||||
if IS_MAC:
|
||||
result["backend"] = "mps" if arch == "arm64" else "cpu"
|
||||
|
||||
# DGX OS / Grace (GB10, DGX Spark): NVIDIA GPU on aarch64 Linux
|
||||
result["is_dgx"] = bool(result["os"] == "linux" and arch == "aarch64" and nvidia)
|
||||
return result
|
||||
127
manager/doctor.py
Normal file
127
manager/doctor.py
Normal file
@@ -0,0 +1,127 @@
|
||||
"""Environment diagnostics: `python -m manager doctor`."""
|
||||
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
from . import detect as detect_mod
|
||||
from . import env, ffmpeg, gitops, nodejs
|
||||
from .util import REPO_ROOT, clean_env, find_uv, venv_dir, venv_python
|
||||
|
||||
|
||||
def _check(label, passed, detail=""):
|
||||
if sys.stdout.isatty():
|
||||
mark = "\033[32mOK\033[0m " if passed else "\033[31mFAIL\033[0m"
|
||||
else:
|
||||
mark = "OK " if passed else "FAIL"
|
||||
print(" [%s] %-18s %s" % (mark, label, detail))
|
||||
return passed
|
||||
|
||||
|
||||
def run_doctor():
|
||||
print("AI Toolkit doctor\n")
|
||||
d = detect_mod.detect()
|
||||
|
||||
from . import gitwin
|
||||
|
||||
_check("os / arch", True, "%s %s" % (d["os"], d["arch"]))
|
||||
git = gitwin.find_git()
|
||||
_check(
|
||||
"git",
|
||||
git is not None,
|
||||
git or "not found (manager sync installs a local copy on Windows)",
|
||||
)
|
||||
uv = find_uv()
|
||||
_check("uv", True, uv or "not found (optional, recommended)")
|
||||
|
||||
if d["nvidia"]:
|
||||
names = ", ".join(g["name"] for g in d["nvidia"]["gpus"])
|
||||
_check(
|
||||
"gpu",
|
||||
True,
|
||||
"%s (driver %s, CUDA %s)"
|
||||
% (names, d["nvidia"]["driver"], d["nvidia"]["cuda_version"]),
|
||||
)
|
||||
elif d["rocm"]:
|
||||
_check("gpu", True, "AMD ROCm (experimental)")
|
||||
elif d["backend"] == "mps":
|
||||
_check("gpu", True, "Apple Silicon (MPS)")
|
||||
else:
|
||||
_check("gpu", False, "no supported GPU detected")
|
||||
|
||||
has_venv = env.venv_exists()
|
||||
_check("venv", has_venv, venv_dir() if has_venv else "not created yet")
|
||||
if has_venv:
|
||||
torch = env.installed_torch()
|
||||
_check("torch", torch is not None, torch or "not installed")
|
||||
if torch and d["backend"] == "cuda":
|
||||
try:
|
||||
out = subprocess.run(
|
||||
[
|
||||
venv_python(),
|
||||
"-c",
|
||||
"import torch; print(torch.cuda.is_available())",
|
||||
],
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=120,
|
||||
env=clean_env(),
|
||||
)
|
||||
avail = out.stdout.decode().strip() == "True"
|
||||
_check(
|
||||
"torch sees gpu",
|
||||
avail,
|
||||
"" if avail else "torch.cuda.is_available() is False",
|
||||
)
|
||||
except (OSError, subprocess.TimeoutExpired):
|
||||
_check("torch sees gpu", False, "could not query")
|
||||
|
||||
node_exe, node_major = nodejs.have_usable_node()
|
||||
if node_exe:
|
||||
_check("node", True, "%s (v%s)" % (node_exe, node_major))
|
||||
else:
|
||||
_check(
|
||||
"node",
|
||||
False,
|
||||
"none >= %d found (manager sync installs a local copy)"
|
||||
% nodejs.MIN_NODE_MAJOR,
|
||||
)
|
||||
|
||||
if os.path.isfile(ffmpeg.ffmpeg_exe()):
|
||||
# run it with the launch env so missing shared libs are caught
|
||||
from . import launch
|
||||
|
||||
try:
|
||||
out = subprocess.run(
|
||||
[ffmpeg.ffmpeg_exe(), "-version"],
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=30,
|
||||
env=launch.build_env(),
|
||||
)
|
||||
works = out.returncode == 0
|
||||
detail = (
|
||||
out.stdout.decode().splitlines()[0]
|
||||
if works
|
||||
else "installed but fails to run"
|
||||
)
|
||||
except (OSError, subprocess.TimeoutExpired):
|
||||
works, detail = False, "installed but fails to run"
|
||||
_check("ffmpeg (local)", works, detail)
|
||||
else:
|
||||
_check("ffmpeg (local)", False, "not installed (manager sync installs it)")
|
||||
|
||||
try:
|
||||
free_gb = shutil.disk_usage(REPO_ROOT).free / (1024**3)
|
||||
_check("disk space", free_gb > 30, "%.0f GB free" % free_gb)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
branch = gitops.current_branch()
|
||||
_check(
|
||||
"git checkout",
|
||||
True,
|
||||
"%s @ %s%s"
|
||||
% (branch, gitops.current_commit(), " (dirty)" if gitops.is_dirty() else ""),
|
||||
)
|
||||
454
manager/env.py
Normal file
454
manager/env.py
Normal file
@@ -0,0 +1,454 @@
|
||||
"""Python environment provisioning and dependency sync.
|
||||
|
||||
Strategy:
|
||||
- If a venv already exists (.venv or venv), use it.
|
||||
- Otherwise create one: prefer uv (downloads the exact Python version needed),
|
||||
fall back to the running Python's venv module if it is new enough.
|
||||
- Installs go through `uv pip` when uv is available (much faster), else pip.
|
||||
|
||||
State (torch backend, requirements hash, applied migrations) is stored inside
|
||||
the venv so a deleted venv means a clean slate — which is correct.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
from .util import (
|
||||
REPO_ROOT,
|
||||
IS_WINDOWS,
|
||||
clean_env,
|
||||
die,
|
||||
file_hash,
|
||||
find_uv,
|
||||
info,
|
||||
ok,
|
||||
run,
|
||||
venv_dir,
|
||||
venv_python,
|
||||
warn,
|
||||
)
|
||||
|
||||
STATE_FILE = "aitk_manager_state.json"
|
||||
|
||||
MIN_SYSTEM_PYTHON = (3, 10)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- state
|
||||
|
||||
|
||||
def state_path(venv=None):
|
||||
return os.path.join(venv or venv_dir(), STATE_FILE)
|
||||
|
||||
|
||||
def load_state():
|
||||
try:
|
||||
with open(state_path(), "r") as f:
|
||||
return json.load(f)
|
||||
except (OSError, ValueError):
|
||||
return {}
|
||||
|
||||
|
||||
def save_state(state):
|
||||
with open(state_path(), "w") as f:
|
||||
json.dump(state, f, indent=2)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- venv
|
||||
|
||||
|
||||
def venv_exists():
|
||||
return os.path.isfile(venv_python())
|
||||
|
||||
|
||||
def ensure_venv(spec, dry_run=False):
|
||||
"""Create the venv if missing. Returns path to the venv python."""
|
||||
if venv_exists():
|
||||
return venv_python()
|
||||
|
||||
target = venv_dir()
|
||||
uv = find_uv()
|
||||
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")
|
||||
)
|
||||
return venv_python(target)
|
||||
|
||||
if uv:
|
||||
info("Creating venv with uv (python %s) at %s" % (spec.python_version, target))
|
||||
run(
|
||||
[uv, "venv", target, "--python", spec.python_version, "--seed"],
|
||||
env=clean_env(),
|
||||
)
|
||||
else:
|
||||
if sys.version_info < MIN_SYSTEM_PYTHON:
|
||||
die(
|
||||
"Python %d.%d is too old (need >= %d.%d) and uv is not installed.\n"
|
||||
"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 pyver != spec.python_version:
|
||||
warn(
|
||||
"Recommended Python is %s but using system Python %s "
|
||||
"(install uv to get the exact version automatically)."
|
||||
% (spec.python_version, pyver)
|
||||
)
|
||||
info("Creating venv at %s" % target)
|
||||
run([sys.executable, "-m", "venv", target])
|
||||
ok("Virtual environment ready.")
|
||||
return venv_python(target)
|
||||
|
||||
|
||||
def _pip_install(args, dry_run=False, upgrade=False, check=True):
|
||||
"""Install into the venv, via uv pip if available. Returns exit code."""
|
||||
uv = find_uv()
|
||||
if uv:
|
||||
cmd = [uv, "pip", "install", "--python", venv_python()]
|
||||
else:
|
||||
cmd = [venv_python(), "-m", "pip", "install"]
|
||||
if upgrade:
|
||||
cmd.append("--upgrade")
|
||||
cmd += args
|
||||
if dry_run:
|
||||
info("[dry-run] would run: %s" % " ".join(cmd))
|
||||
return 0
|
||||
code, _ = run(cmd, stream=True, env=clean_env(), check=check)
|
||||
return code
|
||||
|
||||
|
||||
def _pip_uninstall(packages, dry_run=False):
|
||||
uv = find_uv()
|
||||
if uv:
|
||||
cmd = [uv, "pip", "uninstall", "--python", venv_python()] + packages
|
||||
else:
|
||||
cmd = [venv_python(), "-m", "pip", "uninstall", "-y"] + packages
|
||||
if dry_run:
|
||||
info("[dry-run] would run: %s" % " ".join(cmd))
|
||||
return
|
||||
# non-fatal: the package may simply not be installed yet
|
||||
run(cmd, check=False, env=clean_env())
|
||||
|
||||
|
||||
def venv_python_version():
|
||||
"""'3.12' etc. from the venv interpreter, or None."""
|
||||
if not venv_exists():
|
||||
return None
|
||||
try:
|
||||
out = subprocess.run(
|
||||
[venv_python(), "-c", "import sys; print('%d.%d' % sys.version_info[:2])"],
|
||||
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
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- torch
|
||||
|
||||
|
||||
def installed_torch():
|
||||
"""Returns torch.__version__ from the venv (e.g. '2.9.1+cu128'), or None."""
|
||||
if not venv_exists():
|
||||
return None
|
||||
try:
|
||||
out = subprocess.run(
|
||||
[venv_python(), "-c", "import torch; print(torch.__version__)"],
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=120,
|
||||
env=clean_env(),
|
||||
)
|
||||
if out.returncode != 0:
|
||||
return None
|
||||
return out.stdout.decode().strip() or None
|
||||
except (OSError, subprocess.TimeoutExpired):
|
||||
return None
|
||||
|
||||
|
||||
def torch_matches(spec):
|
||||
current = installed_torch()
|
||||
if current is None:
|
||||
return False
|
||||
want = spec.torch_packages["torch"]
|
||||
# local version tag carries the backend: "2.9.1+cu128"
|
||||
if "+" in current:
|
||||
version, local = current.split("+", 1)
|
||||
return version == want and local == spec.backend
|
||||
# PyPI wheels (mac) have no local tag
|
||||
return current == want and spec.backend in ("mps", "cpu")
|
||||
|
||||
|
||||
def ensure_torch(spec, dry_run=False):
|
||||
if torch_matches(spec):
|
||||
ok(
|
||||
"PyTorch %s (%s) already installed."
|
||||
% (spec.torch_packages["torch"], spec.backend)
|
||||
)
|
||||
return False
|
||||
current = installed_torch()
|
||||
if current:
|
||||
info(
|
||||
"PyTorch %s installed, need %s (%s) — reinstalling."
|
||||
% (current, spec.torch_packages["torch"], spec.backend)
|
||||
)
|
||||
else:
|
||||
info(
|
||||
"Installing PyTorch %s (%s)..."
|
||||
% (spec.torch_packages["torch"], spec.backend)
|
||||
)
|
||||
_pip_install(spec.torch_args(), dry_run=dry_run)
|
||||
return True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- requirements
|
||||
|
||||
|
||||
def requirements_hash(spec):
|
||||
"""Hash of every requirements file plus the spec itself."""
|
||||
req_files = [
|
||||
os.path.join(REPO_ROOT, f)
|
||||
for f in os.listdir(REPO_ROOT)
|
||||
if f.startswith("requirements") and f.endswith(".txt")
|
||||
]
|
||||
req_files.append(os.path.join(REPO_ROOT, "dgx_requirements.txt"))
|
||||
base = file_hash(req_files)
|
||||
import hashlib
|
||||
|
||||
h = hashlib.sha256()
|
||||
h.update(base.encode())
|
||||
h.update(json.dumps(spec.as_dict(), sort_keys=True).encode())
|
||||
return h.hexdigest()
|
||||
|
||||
|
||||
def requirements_in_sync(spec):
|
||||
if not venv_exists():
|
||||
return False
|
||||
return load_state().get("req_hash") == requirements_hash(spec)
|
||||
|
||||
|
||||
def _git_pinned_packages(spec):
|
||||
"""{package_name: full git+ requirement line} from the requirements files.
|
||||
|
||||
pip skips reinstalling a git pin whose version number didn't change even
|
||||
when the commit hash did, so pins whose URL changed since the last sync
|
||||
get uninstalled first to force the new commit (the trick the community
|
||||
Windows installer uses for diffusers).
|
||||
"""
|
||||
pins = {}
|
||||
seen_files = set()
|
||||
|
||||
def scan(path):
|
||||
if path in seen_files or not os.path.isfile(path):
|
||||
return
|
||||
seen_files.add(path)
|
||||
with open(path) as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if line.startswith("-r "):
|
||||
scan(os.path.join(os.path.dirname(path), line[3:].strip()))
|
||||
elif "git+" in line and not line.startswith("#"):
|
||||
# e.g. git+https://github.com/huggingface/diffusers.git@<sha>
|
||||
tail = line.split("/")[-1]
|
||||
name = tail.split(".git")[0].split("@")[0]
|
||||
if name:
|
||||
pins[name] = line
|
||||
|
||||
scan(spec.requirements_path())
|
||||
return pins
|
||||
|
||||
|
||||
def _stale_git_pins(spec):
|
||||
"""Git-pinned packages whose pin (commit) changed since the last sync."""
|
||||
current = _git_pinned_packages(spec)
|
||||
state = load_state()
|
||||
stored = state.get("git_pins")
|
||||
if stored is None:
|
||||
# no record of what's installed (pre-tracking env): if deps were ever
|
||||
# installed here, play it safe and force-reinstall all git pins once
|
||||
return list(current) if state.get("req_hash") else []
|
||||
return [name for name, line in current.items() if stored.get(name) != line]
|
||||
|
||||
|
||||
def _optional_import_name(pkg):
|
||||
"""Importable module name for an optional package spec or wheel URL."""
|
||||
if "://" in pkg:
|
||||
name = os.path.basename(pkg).split("-")[0]
|
||||
else:
|
||||
name = pkg
|
||||
for sep in ("==", ">=", "<=", "<", ">", "["):
|
||||
name = name.split(sep)[0]
|
||||
return name.strip().replace("-", "_")
|
||||
|
||||
|
||||
def _venv_import_ok(module_name):
|
||||
try:
|
||||
out = subprocess.run(
|
||||
[venv_python(), "-c", "import %s" % module_name],
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=180,
|
||||
env=clean_env(),
|
||||
)
|
||||
return out.returncode == 0
|
||||
except (OSError, subprocess.TimeoutExpired):
|
||||
return False
|
||||
|
||||
|
||||
def _filter_extras(extras):
|
||||
"""Drop wheel URLs whose cpXY tag doesn't match the venv python."""
|
||||
pyver = venv_python_version()
|
||||
cp_tag = "cp" + pyver.replace(".", "") if pyver else None
|
||||
kept = []
|
||||
for pkg in extras:
|
||||
if "cp3" in pkg and cp_tag and cp_tag not in pkg:
|
||||
warn(
|
||||
"Skipping %s (built for a different python than venv %s)."
|
||||
% (os.path.basename(pkg), pyver)
|
||||
)
|
||||
continue
|
||||
kept.append(pkg)
|
||||
return kept
|
||||
|
||||
|
||||
def ensure_requirements(spec, dry_run=False, force=False):
|
||||
if not force and requirements_in_sync(spec):
|
||||
ok("Requirements already in sync.")
|
||||
return False
|
||||
# force git-pinned deps (diffusers) onto a newly pinned commit — pip won't
|
||||
# reinstall them on its own because the version number stays the same
|
||||
stale = _stale_git_pins(spec)
|
||||
if stale:
|
||||
info("Git pin changed — reinstalling: %s" % ", ".join(stale))
|
||||
_pip_uninstall(stale, dry_run=dry_run)
|
||||
info("Installing requirements from %s..." % spec.requirements_file)
|
||||
_pip_install(["-r", spec.requirements_path()], dry_run=dry_run)
|
||||
find_links = []
|
||||
for url in spec.find_links:
|
||||
find_links += ["--find-links", url]
|
||||
extras = _filter_extras(spec.extra_packages)
|
||||
if extras:
|
||||
info("Installing platform extras...")
|
||||
_pip_install(extras + find_links, dry_run=dry_run, upgrade=True)
|
||||
# accelerators (flash-attn, NATTEN, ...): install one-by-one, warn on
|
||||
# failure — training works without them, so never fail the whole install
|
||||
for pkg in _filter_extras(spec.optional_packages):
|
||||
label = os.path.basename(pkg) if "://" in pkg else pkg
|
||||
info("Installing optional accelerator: %s" % label)
|
||||
code = _pip_install(
|
||||
[pkg] + find_links, dry_run=dry_run, upgrade=True, check=False
|
||||
)
|
||||
if code != 0:
|
||||
warn("Optional package failed to install (continuing): %s" % label)
|
||||
continue
|
||||
if dry_run:
|
||||
continue
|
||||
# prebuilt accelerator wheels are sometimes built against a torch
|
||||
# nightly and fail to load against the release ABI — verify the
|
||||
# import and roll back rather than leaving a broken wheel installed
|
||||
name = _optional_import_name(pkg)
|
||||
if not _venv_import_ok(name):
|
||||
warn(
|
||||
"%s installed but fails to import against this torch build — "
|
||||
"removing it (training falls back to native attention)." % name
|
||||
)
|
||||
_pip_uninstall([name])
|
||||
if not dry_run:
|
||||
state = load_state()
|
||||
state["req_hash"] = requirements_hash(spec)
|
||||
state["backend"] = spec.backend
|
||||
state["git_pins"] = _git_pinned_packages(spec)
|
||||
save_state(state)
|
||||
return True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- sitecustomize
|
||||
|
||||
|
||||
def write_sitecustomize(dry_run=False):
|
||||
"""Drop a sitecustomize.py into the venv that exposes the local ffmpeg.
|
||||
|
||||
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.
|
||||
"""
|
||||
from . import ffmpeg
|
||||
|
||||
if not venv_exists():
|
||||
return
|
||||
try:
|
||||
out = subprocess.run(
|
||||
[
|
||||
venv_python(),
|
||||
"-c",
|
||||
"import sysconfig; print(sysconfig.get_paths()['purelib'])",
|
||||
],
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=30,
|
||||
env=clean_env(),
|
||||
)
|
||||
site_packages = out.stdout.decode().strip()
|
||||
except (OSError, subprocess.TimeoutExpired):
|
||||
site_packages = ""
|
||||
if not site_packages or not os.path.isdir(site_packages):
|
||||
warn("Could not locate venv site-packages — skipping sitecustomize.")
|
||||
return
|
||||
target = os.path.join(site_packages, "sitecustomize.py")
|
||||
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"
|
||||
"_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"
|
||||
"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"
|
||||
" _prev = os.environ.get('LD_LIBRARY_PATH', '')\n"
|
||||
" if _FFMPEG_LIB not in _prev.split(os.pathsep):\n"
|
||||
" os.environ['LD_LIBRARY_PATH'] = (\n"
|
||||
" _FFMPEG_LIB + ((os.pathsep + _prev) if _prev else '')\n"
|
||||
" )\n"
|
||||
) % (ffmpeg.bin_dir(), ffmpeg.lib_dir())
|
||||
if dry_run:
|
||||
info("[dry-run] would write %s" % target)
|
||||
return
|
||||
with open(target, "w") as f:
|
||||
f.write(content)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- sync
|
||||
|
||||
|
||||
def sync(spec, detection, dry_run=False, force=False):
|
||||
"""Bring the environment fully up to date for this checkout."""
|
||||
from . import ffmpeg, gitwin, migrations, nodejs, uvbin
|
||||
|
||||
for note in spec.notes:
|
||||
warn(note)
|
||||
uvbin.ensure_uv(dry_run=dry_run)
|
||||
gitwin.ensure_git(dry_run=dry_run)
|
||||
ensure_venv(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
|
||||
ensure_requirements(spec, dry_run=dry_run, force=force or changed_torch)
|
||||
ffmpeg.ensure_ffmpeg(detection, dry_run=dry_run)
|
||||
nodejs.ensure_node(detection, dry_run=dry_run)
|
||||
write_sitecustomize(dry_run=dry_run)
|
||||
migrations.run_pending(dry_run=dry_run)
|
||||
ok("Environment is up to date.")
|
||||
165
manager/ffmpeg.py
Normal file
165
manager/ffmpeg.py
Normal file
@@ -0,0 +1,165 @@
|
||||
"""Local (never global) FFmpeg provisioning into <repo>/.ffmpeg/.
|
||||
|
||||
Why: the toolkit shells out to ffmpeg/ffprobe for video work, and torchcodec
|
||||
dlopens the FFmpeg *shared libraries* at runtime. Installing FFmpeg system-wide
|
||||
(apt/winget/brew) is exactly what we want to avoid, so we download a portable
|
||||
build next to the repo:
|
||||
|
||||
- Linux / Windows: BtbN shared builds (bin/ + lib/ with .so/.dll) — the shared
|
||||
libs are what torchcodec needs. The FFmpeg major version must be one the
|
||||
pinned torchcodec supports.
|
||||
- macOS: Martin Riedl static ffmpeg/ffprobe executables (no shared libs
|
||||
published; torchcodec on mac keeps whatever it uses today).
|
||||
|
||||
Exposure to the rest of the system:
|
||||
- `manager launch` prepends .ffmpeg/bin to PATH (and .ffmpeg/lib to
|
||||
LD_LIBRARY_PATH on Linux) so the UI and every training job it spawns see it.
|
||||
- env.py writes a sitecustomize.py into the venv that prepends .ffmpeg/bin to
|
||||
PATH and (on Windows) calls os.add_dll_directory — so torchcodec finds the
|
||||
DLLs in ANY use of the venv python, not just via `manager launch`.
|
||||
"""
|
||||
|
||||
import os
|
||||
import shutil
|
||||
import stat
|
||||
import tempfile
|
||||
|
||||
from .util import (
|
||||
IS_MAC,
|
||||
IS_WINDOWS,
|
||||
REPO_ROOT,
|
||||
download,
|
||||
extract_archive,
|
||||
info,
|
||||
ok,
|
||||
warn,
|
||||
)
|
||||
|
||||
FFMPEG_DIR = os.path.join(REPO_ROOT, ".ffmpeg")
|
||||
|
||||
# FFmpeg 8 on all BtbN platforms — torchcodec 0.15 (pinned in spec.py)
|
||||
# supports ffmpeg up to 8. Bump these together with the torchcodec pin.
|
||||
_BTBN = "https://github.com/BtbN/FFmpeg-Builds/releases/download/latest/"
|
||||
_RIEDL = (
|
||||
"https://ffmpeg.martin-riedl.de/redirect/latest/macos/{arch}/release/{tool}.zip"
|
||||
)
|
||||
|
||||
_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",
|
||||
}
|
||||
|
||||
|
||||
def bin_dir():
|
||||
return os.path.join(FFMPEG_DIR, "bin")
|
||||
|
||||
|
||||
def lib_dir():
|
||||
return os.path.join(FFMPEG_DIR, "lib")
|
||||
|
||||
|
||||
def ffmpeg_exe():
|
||||
return os.path.join(bin_dir(), "ffmpeg.exe" if IS_WINDOWS else "ffmpeg")
|
||||
|
||||
|
||||
def is_installed(source_url):
|
||||
marker = os.path.join(FFMPEG_DIR, ".source")
|
||||
if not os.path.isfile(ffmpeg_exe()) or not os.path.isfile(marker):
|
||||
return False
|
||||
with open(marker) as f:
|
||||
return f.read().strip() == source_url
|
||||
|
||||
|
||||
def _mark_installed(source_url):
|
||||
with open(os.path.join(FFMPEG_DIR, ".source"), "w") as f:
|
||||
f.write(source_url)
|
||||
|
||||
|
||||
def _install_btbn(url):
|
||||
tmp = tempfile.mkdtemp(prefix="aitk_ffmpeg_")
|
||||
try:
|
||||
archive = os.path.join(tmp, os.path.basename(url))
|
||||
download(url, archive, label="ffmpeg")
|
||||
extract_archive(archive, tmp)
|
||||
# archives contain a single top-level dir with bin/ lib/ include/
|
||||
inner = None
|
||||
for name in os.listdir(tmp):
|
||||
path = os.path.join(tmp, name)
|
||||
if os.path.isdir(path) and os.path.isdir(os.path.join(path, "bin")):
|
||||
inner = path
|
||||
break
|
||||
if inner is None:
|
||||
warn("Unexpected ffmpeg archive layout — skipping ffmpeg install.")
|
||||
return False
|
||||
if os.path.isdir(FFMPEG_DIR):
|
||||
shutil.rmtree(FFMPEG_DIR)
|
||||
shutil.move(inner, FFMPEG_DIR)
|
||||
return True
|
||||
finally:
|
||||
shutil.rmtree(tmp, ignore_errors=True)
|
||||
|
||||
|
||||
def _install_mac(detection):
|
||||
arch = "arm64" if detection["arch"] == "arm64" else "amd64"
|
||||
tmp = tempfile.mkdtemp(prefix="aitk_ffmpeg_")
|
||||
try:
|
||||
os.makedirs(bin_dir(), exist_ok=True)
|
||||
for tool in ("ffmpeg", "ffprobe"):
|
||||
url = _RIEDL.format(arch=arch, tool=tool)
|
||||
archive = os.path.join(tmp, tool + ".zip")
|
||||
download(url, archive, label=tool)
|
||||
extract_archive(archive, tmp)
|
||||
src = os.path.join(tmp, tool)
|
||||
if not os.path.isfile(src):
|
||||
warn("Unexpected %s archive layout — skipping." % tool)
|
||||
return False
|
||||
dest = os.path.join(bin_dir(), tool)
|
||||
shutil.move(src, dest)
|
||||
os.chmod(
|
||||
dest, os.stat(dest).st_mode | stat.S_IXUSR | stat.S_IXGRP | stat.S_IXOTH
|
||||
)
|
||||
return True
|
||||
finally:
|
||||
shutil.rmtree(tmp, ignore_errors=True)
|
||||
|
||||
|
||||
def source_url(detection):
|
||||
if detection["os"] == "mac":
|
||||
arch = "arm64" if detection["arch"] == "arm64" else "amd64"
|
||||
return _RIEDL.format(arch=arch, tool="ffmpeg")
|
||||
return _SOURCES.get((detection["os"], detection["arch"]))
|
||||
|
||||
|
||||
def ensure_ffmpeg(detection, dry_run=False):
|
||||
url = source_url(detection)
|
||||
if url is None:
|
||||
warn(
|
||||
"No portable FFmpeg source for %s/%s — skipping local ffmpeg."
|
||||
% (detection["os"], detection["arch"])
|
||||
)
|
||||
return False
|
||||
if is_installed(url):
|
||||
ok("Local FFmpeg already installed (.ffmpeg/).")
|
||||
return False
|
||||
if dry_run:
|
||||
info("[dry-run] would install local FFmpeg from %s into %s" % (url, FFMPEG_DIR))
|
||||
return False
|
||||
installed = (
|
||||
_install_mac(detection) if detection["os"] == "mac" else _install_btbn(url)
|
||||
)
|
||||
if installed:
|
||||
_mark_installed(url)
|
||||
ok("Local FFmpeg installed at %s" % FFMPEG_DIR)
|
||||
return installed
|
||||
|
||||
|
||||
def env_additions():
|
||||
"""(path_dirs, ld_library_dirs) to prepend when launching anything."""
|
||||
paths = []
|
||||
lib_paths = []
|
||||
if os.path.isdir(bin_dir()):
|
||||
paths.append(bin_dir())
|
||||
if not IS_WINDOWS and not IS_MAC and os.path.isdir(lib_dir()):
|
||||
lib_paths.append(lib_dir())
|
||||
return paths, lib_paths
|
||||
92
manager/gitops.py
Normal file
92
manager/gitops.py
Normal file
@@ -0,0 +1,92 @@
|
||||
"""Git operations for update checking and pulling. Stdlib only."""
|
||||
|
||||
import os
|
||||
|
||||
from . import gitwin
|
||||
from .util import run, REPO_ROOT, die
|
||||
|
||||
|
||||
def _git(args, capture=True, check=True):
|
||||
git = gitwin.find_git()
|
||||
if not git:
|
||||
die(
|
||||
"git was not found. On Windows run `python -m manager sync` to "
|
||||
"install a local copy; elsewhere install git with your package "
|
||||
"manager (or xcode-select --install on macOS)."
|
||||
)
|
||||
# skip LFS payloads — the repo may carry LFS files not needed at runtime
|
||||
env = os.environ.copy()
|
||||
env["GIT_LFS_SKIP_SMUDGE"] = "1"
|
||||
return run([git] + args, cwd=REPO_ROOT, capture=capture, check=check, env=env)
|
||||
|
||||
|
||||
def current_branch():
|
||||
_, out = _git(["rev-parse", "--abbrev-ref", "HEAD"])
|
||||
return out
|
||||
|
||||
|
||||
def current_commit():
|
||||
_, out = _git(["rev-parse", "--short", "HEAD"])
|
||||
return out
|
||||
|
||||
|
||||
def is_dirty():
|
||||
_, out = _git(["status", "--porcelain"])
|
||||
# untracked files don't block updates; modified/staged tracked files do
|
||||
for line in (out or "").splitlines():
|
||||
if not line.startswith("??"):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def fetch():
|
||||
code, _ = _git(["fetch", "--quiet"], capture=False, check=False)
|
||||
return code == 0
|
||||
|
||||
|
||||
def upstream():
|
||||
code, out = _git(
|
||||
["rev-parse", "--abbrev-ref", "--symbolic-full-name", "@{u}"], check=False
|
||||
)
|
||||
if code != 0:
|
||||
return None
|
||||
return out
|
||||
|
||||
|
||||
def behind_count():
|
||||
"""Commits the local branch is behind its upstream. None if no upstream."""
|
||||
up = upstream()
|
||||
if not up:
|
||||
return None
|
||||
code, out = _git(["rev-list", "--count", "HEAD..@{u}"], check=False)
|
||||
if code != 0 or out is None:
|
||||
return None
|
||||
try:
|
||||
return int(out)
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
|
||||
def remote_commit():
|
||||
code, out = _git(["rev-parse", "--short", "@{u}"], check=False)
|
||||
return out if code == 0 else None
|
||||
|
||||
|
||||
def incoming_log(limit=15):
|
||||
code, out = _git(
|
||||
["log", "--oneline", "HEAD..@{u}", "--max-count=%d" % limit], check=False
|
||||
)
|
||||
if code != 0 or not out:
|
||||
return []
|
||||
return out.splitlines()
|
||||
|
||||
|
||||
def pull_ff():
|
||||
"""Fast-forward pull. Dies with guidance on failure."""
|
||||
code, _ = _git(["pull", "--ff-only"], capture=False, check=False)
|
||||
if code != 0:
|
||||
die(
|
||||
"git pull --ff-only failed. Your branch has local commits or "
|
||||
"conflicts with the remote. Resolve manually (git stash / git rebase) "
|
||||
"and re-run the update."
|
||||
)
|
||||
65
manager/gitwin.py
Normal file
65
manager/gitwin.py
Normal file
@@ -0,0 +1,65 @@
|
||||
"""Local (never global) Git provisioning for Windows into <repo>/.mingit/.
|
||||
|
||||
MinGit is the official minimal, portable Git for Windows — a plain zip meant
|
||||
for embedding in tools (no installer, no registry, no PATH changes). On
|
||||
Windows, if the system has no git, `manager sync` drops one here so updates
|
||||
keep working from any terminal. The first clone (before this repo exists) is
|
||||
handled the same way by the bootstrap layer (install.ps1 / desktop launcher),
|
||||
which then moves its MinGit into the fresh checkout as .mingit/.
|
||||
|
||||
Linux and macOS have no sane portable git (glibc / Xcode CLT entanglement),
|
||||
and git is effectively always available there — so this is Windows-only and
|
||||
the manager just errors with install instructions elsewhere.
|
||||
"""
|
||||
|
||||
import os
|
||||
import shutil
|
||||
import tempfile
|
||||
|
||||
from .util import IS_WINDOWS, REPO_ROOT, download, extract_archive, info, ok, warn
|
||||
|
||||
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"
|
||||
)
|
||||
|
||||
|
||||
def local_git_exe():
|
||||
return os.path.join(MINGIT_DIR, "cmd", "git.exe")
|
||||
|
||||
|
||||
def find_git():
|
||||
"""Path/command for git: repo-local MinGit first (Windows), then system."""
|
||||
if IS_WINDOWS and os.path.isfile(local_git_exe()):
|
||||
return local_git_exe()
|
||||
return shutil.which("git")
|
||||
|
||||
|
||||
def ensure_git(dry_run=False):
|
||||
"""Windows-only: provision .mingit/ when the system has no git."""
|
||||
if not IS_WINDOWS:
|
||||
return False
|
||||
if find_git():
|
||||
return False
|
||||
if dry_run:
|
||||
info("[dry-run] would install MinGit into %s" % MINGIT_DIR)
|
||||
return False
|
||||
tmp = tempfile.mkdtemp(prefix="aitk_mingit_")
|
||||
try:
|
||||
archive = os.path.join(tmp, "mingit.zip")
|
||||
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)
|
||||
if not os.path.isfile(os.path.join(extracted, "cmd", "git.exe")):
|
||||
warn("Unexpected MinGit archive layout — skipping local git install.")
|
||||
return False
|
||||
if os.path.isdir(MINGIT_DIR):
|
||||
shutil.rmtree(MINGIT_DIR)
|
||||
shutil.move(extracted, MINGIT_DIR)
|
||||
ok("Local Git (MinGit) installed at %s" % MINGIT_DIR)
|
||||
return True
|
||||
finally:
|
||||
shutil.rmtree(tmp, ignore_errors=True)
|
||||
116
manager/launch.py
Normal file
116
manager/launch.py
Normal file
@@ -0,0 +1,116 @@
|
||||
"""Launch the AI Toolkit web UI (ui/ -> npm run build_and_start)."""
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import threading
|
||||
|
||||
from . import ffmpeg, nodejs
|
||||
from .util import (
|
||||
IS_LINUX,
|
||||
IS_WINDOWS,
|
||||
REPO_ROOT,
|
||||
clean_env,
|
||||
die,
|
||||
info,
|
||||
venv_dir,
|
||||
venv_python,
|
||||
)
|
||||
|
||||
UI_PORT = 8675
|
||||
UI_URL = "http://localhost:%d" % UI_PORT
|
||||
BROWSER_POLL_SECONDS = 300
|
||||
|
||||
|
||||
def build_env():
|
||||
"""Scrubbed env with local node, ffmpeg, and the venv on PATH.
|
||||
|
||||
Everything the UI worker spawns (training jobs) inherits this, so the
|
||||
local ffmpeg/node are visible to the whole process tree.
|
||||
"""
|
||||
env = clean_env()
|
||||
path_dirs = []
|
||||
if os.path.isdir(nodejs.node_bin_dir()):
|
||||
path_dirs.append(nodejs.node_bin_dir())
|
||||
ff_paths, ff_libs = ffmpeg.env_additions()
|
||||
path_dirs += ff_paths
|
||||
vbin = os.path.join(venv_dir(), "Scripts" if IS_WINDOWS else "bin")
|
||||
if os.path.isdir(vbin):
|
||||
path_dirs.append(vbin)
|
||||
if path_dirs:
|
||||
env["PATH"] = os.pathsep.join(path_dirs) + os.pathsep + env.get("PATH", "")
|
||||
if ff_libs:
|
||||
env["LD_LIBRARY_PATH"] = os.pathsep.join(
|
||||
ff_libs + [env.get("LD_LIBRARY_PATH", "")]
|
||||
).rstrip(os.pathsep)
|
||||
return env
|
||||
|
||||
|
||||
def _find_npm(env):
|
||||
import shutil
|
||||
|
||||
return shutil.which("npm.cmd" if IS_WINDOWS else "npm", path=env.get("PATH"))
|
||||
|
||||
|
||||
def _headless():
|
||||
return IS_LINUX and not (
|
||||
os.environ.get("DISPLAY") or os.environ.get("WAYLAND_DISPLAY")
|
||||
)
|
||||
|
||||
|
||||
def _open_browser_when_ready(stop_event):
|
||||
"""Poll the UI port in the background; open the browser once it responds."""
|
||||
import time
|
||||
import urllib.request
|
||||
import webbrowser
|
||||
|
||||
waited = 0
|
||||
while not stop_event.is_set() and waited < BROWSER_POLL_SECONDS:
|
||||
try:
|
||||
urllib.request.urlopen(UI_URL, timeout=2).close()
|
||||
webbrowser.open(UI_URL)
|
||||
return
|
||||
except OSError:
|
||||
time.sleep(2)
|
||||
waited += 2
|
||||
|
||||
|
||||
def launch_ui(open_browser=True):
|
||||
if not os.path.isfile(venv_python()):
|
||||
die("No Python environment found. Run: python3 -m manager install")
|
||||
|
||||
env = build_env()
|
||||
npm = _find_npm(env)
|
||||
if not npm:
|
||||
die(
|
||||
"Node.js was not found. Run `python3 -m manager sync` to install a "
|
||||
"local copy, or install Node.js >= %d from https://nodejs.org."
|
||||
% nodejs.MIN_NODE_MAJOR
|
||||
)
|
||||
_, major = nodejs.have_usable_node()
|
||||
if major is not None and major < nodejs.MIN_NODE_MAJOR:
|
||||
die(
|
||||
"Node.js v%d found, but >= %d is required. Run `python3 -m manager sync`."
|
||||
% (major, nodejs.MIN_NODE_MAJOR)
|
||||
)
|
||||
|
||||
info("Starting AI Toolkit UI (%s) ..." % UI_URL)
|
||||
stop_event = threading.Event()
|
||||
if open_browser and not _headless():
|
||||
threading.Thread(
|
||||
target=_open_browser_when_ready, args=(stop_event,), daemon=True
|
||||
).start()
|
||||
|
||||
proc = subprocess.Popen(
|
||||
[npm, "run", "build_and_start"], cwd=os.path.join(REPO_ROOT, "ui"), env=env
|
||||
)
|
||||
try:
|
||||
return proc.wait()
|
||||
except KeyboardInterrupt:
|
||||
proc.terminate()
|
||||
try:
|
||||
proc.wait(timeout=10)
|
||||
except subprocess.TimeoutExpired:
|
||||
proc.kill()
|
||||
return 130
|
||||
finally:
|
||||
stop_event.set()
|
||||
34
manager/migrations.py
Normal file
34
manager/migrations.py
Normal file
@@ -0,0 +1,34 @@
|
||||
"""One-time migration steps that run after an update.
|
||||
|
||||
Add a migration when an update needs more than a dependency sync (moving
|
||||
files, converting configs, clearing caches, etc). Each runs at most once per
|
||||
environment; applied ids are recorded in the venv state file.
|
||||
|
||||
def _example(dry_run):
|
||||
...
|
||||
|
||||
MIGRATIONS = [
|
||||
{"id": "2026-07-example-cache-move", "run": _example},
|
||||
]
|
||||
"""
|
||||
|
||||
from .util import info
|
||||
from . import env
|
||||
|
||||
MIGRATIONS = []
|
||||
|
||||
|
||||
def run_pending(dry_run=False):
|
||||
if not MIGRATIONS:
|
||||
return
|
||||
state = env.load_state()
|
||||
applied = set(state.get("migrations", []))
|
||||
for migration in MIGRATIONS:
|
||||
if migration["id"] in applied:
|
||||
continue
|
||||
info("Running migration: %s" % migration["id"])
|
||||
if not dry_run:
|
||||
migration["run"](dry_run)
|
||||
applied.add(migration["id"])
|
||||
state["migrations"] = sorted(applied)
|
||||
env.save_state(state)
|
||||
121
manager/nodejs.py
Normal file
121
manager/nodejs.py
Normal file
@@ -0,0 +1,121 @@
|
||||
"""Local (never global) Node.js provisioning into <repo>/.node/.
|
||||
|
||||
A system Node >= 20 is used when present; otherwise an official portable
|
||||
build is downloaded next to the repo (same approach run_mac.zsh already uses).
|
||||
Nothing is ever installed system-wide.
|
||||
"""
|
||||
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
import tempfile
|
||||
|
||||
from .util import (
|
||||
IS_MAC,
|
||||
IS_WINDOWS,
|
||||
REPO_ROOT,
|
||||
download,
|
||||
extract_archive,
|
||||
info,
|
||||
ok,
|
||||
warn,
|
||||
which,
|
||||
)
|
||||
|
||||
NODE_DIR = os.path.join(REPO_ROOT, ".node")
|
||||
# Node 24 is the current LTS line and matches the dgx_instructions.md guidance.
|
||||
NODE_VERSION = "24.11.1"
|
||||
MIN_NODE_MAJOR = 20
|
||||
|
||||
|
||||
def node_bin_dir():
|
||||
# windows zips have node.exe/npm.cmd at the archive root; unix under bin/
|
||||
return NODE_DIR if IS_WINDOWS else os.path.join(NODE_DIR, "bin")
|
||||
|
||||
|
||||
def local_node_exe():
|
||||
return os.path.join(node_bin_dir(), "node.exe" if IS_WINDOWS else "node")
|
||||
|
||||
|
||||
def _node_major(exe):
|
||||
try:
|
||||
out = subprocess.run(
|
||||
[exe, "--version"],
|
||||
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])
|
||||
except (OSError, subprocess.TimeoutExpired, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def _dist_url(detection):
|
||||
arch = detection["arch"]
|
||||
if detection["os"] == "linux":
|
||||
plat = "linux-arm64" if arch == "aarch64" else "linux-x64"
|
||||
ext = "tar.xz"
|
||||
elif detection["os"] == "mac":
|
||||
plat = "darwin-arm64" if arch == "arm64" else "darwin-x64"
|
||||
ext = "tar.gz"
|
||||
elif detection["os"] == "windows":
|
||||
plat = "win-x64"
|
||||
ext = "zip"
|
||||
else:
|
||||
return None, None
|
||||
name = "node-v%s-%s" % (NODE_VERSION, plat)
|
||||
return "https://nodejs.org/dist/v%s/%s.%s" % (NODE_VERSION, name, ext), name
|
||||
|
||||
|
||||
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:
|
||||
return local, major
|
||||
system = which("node")
|
||||
if system:
|
||||
major = _node_major(system)
|
||||
if major is not None and major >= MIN_NODE_MAJOR:
|
||||
return system, major
|
||||
return None, None
|
||||
|
||||
|
||||
def ensure_node(detection, dry_run=False):
|
||||
exe, major = have_usable_node()
|
||||
if exe:
|
||||
ok("Node.js v%d found (%s)." % (major, exe))
|
||||
return False
|
||||
url, inner_name = _dist_url(detection)
|
||||
if url is None:
|
||||
warn(
|
||||
"No portable Node.js build for this platform — install Node >= %d manually."
|
||||
% MIN_NODE_MAJOR
|
||||
)
|
||||
return False
|
||||
if dry_run:
|
||||
info(
|
||||
"[dry-run] would install portable Node.js v%s into %s"
|
||||
% (NODE_VERSION, NODE_DIR)
|
||||
)
|
||||
return False
|
||||
tmp = tempfile.mkdtemp(prefix="aitk_node_")
|
||||
try:
|
||||
archive = os.path.join(tmp, os.path.basename(url))
|
||||
download(url, archive, label="node v%s" % NODE_VERSION)
|
||||
extract_archive(archive, tmp)
|
||||
inner = os.path.join(tmp, inner_name)
|
||||
if not os.path.isdir(inner):
|
||||
warn("Unexpected Node.js archive layout — skipping node install.")
|
||||
return False
|
||||
if os.path.isdir(NODE_DIR):
|
||||
shutil.rmtree(NODE_DIR)
|
||||
shutil.move(inner, NODE_DIR)
|
||||
ok("Portable Node.js v%s installed at %s" % (NODE_VERSION, NODE_DIR))
|
||||
return True
|
||||
finally:
|
||||
shutil.rmtree(tmp, ignore_errors=True)
|
||||
259
manager/spec.py
Normal file
259
manager/spec.py
Normal file
@@ -0,0 +1,259 @@
|
||||
"""Maps detected hardware to an environment spec.
|
||||
|
||||
This is the single source of truth for "what does this machine need to run
|
||||
this commit of AI Toolkit". One universal torch version across all platforms;
|
||||
per-platform accelerator extras (flash-attn, NATTEN, triton) wherever prebuilt
|
||||
wheels exist. **Update the pins below together with the README install
|
||||
instructions and run_mac.zsh.**
|
||||
|
||||
Wheel coverage for the pinned set (verified 2026-07, flash-attn + NATTEN GPU
|
||||
kernels smoke-tested on an RTX 5090 / sm120 with torch 2.13.0+cu130):
|
||||
- torch 2.13.0: cu126/cu130 wheels for linux x86_64 + aarch64 + windows; PyPI
|
||||
wheels for mac arm64. torchaudio is in maintenance mode — 2.11.0 is the
|
||||
current release and is torch-version-agnostic (no torch dep in metadata).
|
||||
- torchcodec 0.15.0: supports torch >= 2.11, wheels on all platforms.
|
||||
- flash-attn 2.8.3: prebuilt by mjun0812/flash-attention-prebuild-wheels for
|
||||
{cu126,cu130} x {cp310..cp314} x {linux x86_64, linux aarch64, windows}.
|
||||
NOTE: the torch2.12 linux wheels there were built against a torch nightly
|
||||
and fail to import on 2.12.0 final — the torch2.13 batches (v0.9.47+) are
|
||||
verified good. Re-verify imports whenever bumping torch.
|
||||
- NATTEN 0.21.7: prebuilt at whl.natten.org for {cu126,cu130,cu132} x
|
||||
{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.
|
||||
|
||||
extra_packages are installed AFTER requirements.txt with --upgrade so they can
|
||||
override requirement pins (e.g. torchcodec). optional_packages are installed
|
||||
one-by-one and only warn on failure (accelerators the training code can live
|
||||
without). Wheel URLs containing a cpXY tag are skipped automatically if the
|
||||
venv python doesn't match.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
from .util import REPO_ROOT
|
||||
|
||||
# ---- version pins (edit these to move the fleet forward) -------------------
|
||||
|
||||
TORCH = {"torch": "2.13.0", "torchvision": "0.28.0", "torchaudio": "2.11.0"}
|
||||
TORCH_TAG = "2.13" # as it appears in flash-attn / natten wheel names
|
||||
|
||||
TORCHCODEC = "torchcodec==0.15.0"
|
||||
TRITON_WINDOWS = "triton-windows>=3.7,<3.8"
|
||||
|
||||
NATTEN_VERSION = "0.21.7"
|
||||
NATTEN_FIND_LINKS = "https://whl.natten.org"
|
||||
|
||||
FLASH_ATTN_VERSION = "2.8.3"
|
||||
_FA_BASE = (
|
||||
"https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/"
|
||||
)
|
||||
# (os, arch) -> (release tag, wheel platform tag) — tags are per torch
|
||||
# version; these carry the torch2.13 builds
|
||||
_FA_BUILDS = {
|
||||
("linux", "x86_64"): ("v0.9.47", "manylinux_2_24_x86_64.manylinux_2_28_x86_64"),
|
||||
("linux", "aarch64"): ("v0.9.48", "manylinux_2_34_aarch64"),
|
||||
("windows", "x86_64"): ("v0.9.52", "win_amd64"),
|
||||
}
|
||||
|
||||
# helper build tools some sdists need on Windows
|
||||
_WIN_HELPERS = ["wheel", "setuptools", "poetry-core", "hf_xet"]
|
||||
|
||||
PYTORCH_INDEX = "https://download.pytorch.org/whl/"
|
||||
|
||||
|
||||
class EnvSpec(object):
|
||||
def __init__(
|
||||
self,
|
||||
backend,
|
||||
torch_packages,
|
||||
torch_index=None,
|
||||
python_version="3.12",
|
||||
requirements_file="requirements.txt",
|
||||
extra_packages=None,
|
||||
optional_packages=None,
|
||||
find_links=None,
|
||||
notes=None,
|
||||
):
|
||||
self.backend = backend # 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
|
||||
self.requirements_file = requirements_file
|
||||
self.extra_packages = extra_packages or []
|
||||
self.optional_packages = optional_packages or []
|
||||
self.find_links = find_links or []
|
||||
self.notes = notes or []
|
||||
|
||||
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]
|
||||
return args
|
||||
|
||||
def requirements_path(self):
|
||||
return os.path.join(REPO_ROOT, self.requirements_file)
|
||||
|
||||
def as_dict(self):
|
||||
return {
|
||||
"backend": self.backend,
|
||||
"torch_packages": self.torch_packages,
|
||||
"torch_index": self.torch_index,
|
||||
"python_version": self.python_version,
|
||||
"requirements_file": self.requirements_file,
|
||||
"extra_packages": self.extra_packages,
|
||||
"optional_packages": self.optional_packages,
|
||||
"find_links": self.find_links,
|
||||
"notes": self.notes,
|
||||
}
|
||||
|
||||
|
||||
def _flash_attn_url(flavor, os_name, arch, python_version):
|
||||
build = _FA_BUILDS.get((os_name, arch))
|
||||
if build is None:
|
||||
return None
|
||||
tag, plat = build
|
||||
cp = "cp" + python_version.replace(".", "")
|
||||
return "%s%s/flash_attn-%s+%storch%s-%s-%s-%s.whl" % (
|
||||
_FA_BASE,
|
||||
tag,
|
||||
FLASH_ATTN_VERSION,
|
||||
flavor,
|
||||
TORCH_TAG,
|
||||
cp,
|
||||
cp,
|
||||
plat,
|
||||
)
|
||||
|
||||
|
||||
def _natten_pin(flavor):
|
||||
# natten wheel local tags use the full torch version without dots: torch2120cu130
|
||||
return "natten==%s+torch%s%s" % (
|
||||
NATTEN_VERSION,
|
||||
TORCH["torch"].replace(".", ""),
|
||||
flavor.replace(".", ""),
|
||||
)
|
||||
|
||||
|
||||
def _cuda_flavor(detection):
|
||||
"""Pick a cuda wheel flavor the installed driver can actually run."""
|
||||
nvidia = detection.get("nvidia") or {}
|
||||
cuda = None
|
||||
if nvidia.get("cuda_version"):
|
||||
try:
|
||||
cuda = tuple(int(x) for x in nvidia["cuda_version"].split("."))
|
||||
except ValueError:
|
||||
cuda = None
|
||||
if cuda is None:
|
||||
# driver present but version unknown — assume current
|
||||
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)
|
||||
if cuda >= (12, 6):
|
||||
if has_blackwell:
|
||||
raise RuntimeError(
|
||||
"Blackwell GPU detected but the NVIDIA driver only supports "
|
||||
"CUDA %s. Blackwell needs the cu130 build — update your "
|
||||
"driver to 580+ and re-run." % nvidia["cuda_version"]
|
||||
)
|
||||
return "cu126", [
|
||||
"NVIDIA driver only supports CUDA %s — installing cu126 wheels. "
|
||||
"Updating your driver is recommended." % nvidia["cuda_version"]
|
||||
]
|
||||
raise RuntimeError(
|
||||
"NVIDIA driver only supports CUDA %s, which is too old for the pinned "
|
||||
"torch build. Update your NVIDIA driver, then re-run install."
|
||||
% nvidia["cuda_version"]
|
||||
)
|
||||
|
||||
|
||||
def _cuda_spec(detection):
|
||||
os_name = detection["os"]
|
||||
arch = detection["arch"]
|
||||
flavor, notes = _cuda_flavor(detection)
|
||||
python_version = "3.12"
|
||||
requirements = (
|
||||
"dgx_requirements.txt" if detection.get("is_dgx") else "requirements.txt"
|
||||
)
|
||||
if detection.get("is_dgx"):
|
||||
# the old "Python 3.11 on DGX OS" constraint was for conda/system
|
||||
# installs; uv provisions 3.12 and all aarch64 cp312 wheels exist now
|
||||
notes = notes + [
|
||||
"DGX OS / Grace detected: using %s wheels and dgx_requirements.txt."
|
||||
% flavor
|
||||
]
|
||||
|
||||
extras = [TORCHCODEC]
|
||||
optional = []
|
||||
find_links = []
|
||||
|
||||
fa_url = _flash_attn_url(flavor, os_name, arch, python_version)
|
||||
if fa_url:
|
||||
optional.append(fa_url)
|
||||
|
||||
if os_name == "linux":
|
||||
optional.append(_natten_pin(flavor))
|
||||
find_links.append(NATTEN_FIND_LINKS)
|
||||
elif os_name == "windows":
|
||||
extras = _WIN_HELPERS + extras + [TRITON_WINDOWS]
|
||||
notes = notes + ["NATTEN has no Windows wheels — skipping it."]
|
||||
|
||||
return EnvSpec(
|
||||
flavor,
|
||||
TORCH,
|
||||
torch_index=PYTORCH_INDEX + flavor,
|
||||
python_version=python_version,
|
||||
requirements_file=requirements,
|
||||
extra_packages=extras,
|
||||
optional_packages=optional,
|
||||
find_links=find_links,
|
||||
notes=notes,
|
||||
)
|
||||
|
||||
|
||||
def build_spec(detection, allow_cpu=False):
|
||||
"""Returns EnvSpec, or raises RuntimeError with a user-facing message."""
|
||||
os_name = detection["os"]
|
||||
|
||||
if os_name == "mac":
|
||||
notes = ["flash-attn / NATTEN / triton are unavailable on macOS."]
|
||||
if detection["backend"] != "mps":
|
||||
notes.append("Intel Mac detected — training will be extremely slow.")
|
||||
return EnvSpec(
|
||||
"mps",
|
||||
TORCH,
|
||||
python_version="3.12",
|
||||
extra_packages=[TORCHCODEC],
|
||||
notes=notes,
|
||||
)
|
||||
|
||||
if detection["backend"] == "cuda":
|
||||
return _cuda_spec(detection)
|
||||
|
||||
if detection["backend"] == "rocm":
|
||||
return EnvSpec(
|
||||
"rocm7.1",
|
||||
TORCH,
|
||||
torch_index=PYTORCH_INDEX + "rocm7.1",
|
||||
extra_packages=[TORCHCODEC],
|
||||
notes=[
|
||||
"AMD ROCm support is experimental and largely untested.",
|
||||
"flash-attn / NATTEN prebuilt wheels are unavailable for ROCm.",
|
||||
],
|
||||
)
|
||||
|
||||
# CPU fallback
|
||||
if not allow_cpu:
|
||||
raise RuntimeError(
|
||||
"No supported GPU detected (NVIDIA CUDA, AMD ROCm, or Apple Silicon). "
|
||||
"Training on CPU is not practical. Pass --cpu to install anyway."
|
||||
)
|
||||
return EnvSpec(
|
||||
"cpu",
|
||||
TORCH,
|
||||
torch_index=PYTORCH_INDEX + "cpu",
|
||||
extra_packages=[TORCHCODEC],
|
||||
notes=["CPU-only install: training will be impractically slow."],
|
||||
)
|
||||
230
manager/util.py
Normal file
230
manager/util.py
Normal file
@@ -0,0 +1,230 @@
|
||||
"""Shared helpers for the AI Toolkit manager. Stdlib only."""
|
||||
|
||||
import json
|
||||
import os
|
||||
import platform
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
IS_WINDOWS = platform.system() == "Windows"
|
||||
IS_MAC = platform.system() == "Darwin"
|
||||
IS_LINUX = platform.system() == "Linux"
|
||||
|
||||
# When --json is used, human output goes to stderr so stdout stays machine-readable
|
||||
_json_mode = False
|
||||
|
||||
|
||||
def set_json_mode(enabled):
|
||||
global _json_mode
|
||||
_json_mode = enabled
|
||||
|
||||
|
||||
def _supports_color(stream):
|
||||
if os.environ.get("NO_COLOR"):
|
||||
return False
|
||||
return hasattr(stream, "isatty") and stream.isatty()
|
||||
|
||||
|
||||
def _emit(prefix, msg, color):
|
||||
stream = sys.stderr if _json_mode else sys.stdout
|
||||
if _supports_color(stream):
|
||||
stream.write("\033[%sm%s\033[0m %s\n" % (color, prefix, msg))
|
||||
else:
|
||||
stream.write("%s %s\n" % (prefix, msg))
|
||||
stream.flush()
|
||||
|
||||
|
||||
def info(msg):
|
||||
_emit("[*]", msg, "36")
|
||||
|
||||
|
||||
def ok(msg):
|
||||
_emit("[+]", msg, "32")
|
||||
|
||||
|
||||
def warn(msg):
|
||||
_emit("[!]", msg, "33")
|
||||
|
||||
|
||||
def error(msg):
|
||||
_emit("[x]", msg, "31")
|
||||
|
||||
|
||||
def die(msg, code=1):
|
||||
error(msg)
|
||||
sys.exit(code)
|
||||
|
||||
|
||||
def print_json(data):
|
||||
sys.stdout.write(json.dumps(data, indent=2) + "\n")
|
||||
sys.stdout.flush()
|
||||
|
||||
|
||||
def run(cmd, cwd=None, capture=False, check=True, env=None, stream=False):
|
||||
"""Run a command. capture=True returns stdout text (stripped).
|
||||
|
||||
stream=True inherits stdio so the user sees live output.
|
||||
Returns (returncode, stdout_or_None).
|
||||
"""
|
||||
kwargs = {"cwd": cwd or REPO_ROOT}
|
||||
if env is not None:
|
||||
kwargs["env"] = env
|
||||
if capture:
|
||||
kwargs["stdout"] = subprocess.PIPE
|
||||
kwargs["stderr"] = subprocess.PIPE
|
||||
elif _json_mode and not stream:
|
||||
# keep stdout clean in json mode
|
||||
kwargs["stdout"] = sys.stderr
|
||||
|
||||
try:
|
||||
proc = subprocess.run(cmd, **kwargs)
|
||||
except FileNotFoundError:
|
||||
if check:
|
||||
die("Command not found: %s" % cmd[0])
|
||||
return 127, None
|
||||
|
||||
out = None
|
||||
if capture:
|
||||
out = proc.stdout.decode("utf-8", errors="replace").strip()
|
||||
if check and proc.returncode != 0:
|
||||
detail = ""
|
||||
if capture and proc.stderr:
|
||||
detail = "\n" + proc.stderr.decode("utf-8", errors="replace").strip()
|
||||
die("Command failed (%d): %s%s" % (proc.returncode, " ".join(cmd), detail))
|
||||
return proc.returncode, out
|
||||
|
||||
|
||||
def which(name):
|
||||
return shutil.which(name)
|
||||
|
||||
|
||||
def find_uv():
|
||||
"""Find uv: repo-local .uv/ first, then PATH, then common install dirs."""
|
||||
home = os.path.expanduser("~")
|
||||
candidates = [
|
||||
os.path.join(REPO_ROOT, ".uv", "uv.exe" if IS_WINDOWS else "uv"),
|
||||
]
|
||||
for c in candidates:
|
||||
if os.path.isfile(c) and os.access(c, os.X_OK):
|
||||
return c
|
||||
uv = shutil.which("uv")
|
||||
if uv:
|
||||
return uv
|
||||
candidates = [
|
||||
os.path.join(home, ".local", "bin", "uv"),
|
||||
os.path.join(home, ".cargo", "bin", "uv"),
|
||||
]
|
||||
if IS_WINDOWS:
|
||||
local = os.environ.get("LOCALAPPDATA", "")
|
||||
if local:
|
||||
candidates.append(os.path.join(local, "uv", "uv.exe"))
|
||||
candidates.append(os.path.join(home, ".local", "bin", "uv.exe"))
|
||||
for c in candidates:
|
||||
if os.path.isfile(c) and os.access(c, os.X_OK):
|
||||
return c
|
||||
return None
|
||||
|
||||
|
||||
def venv_dir():
|
||||
"""Existing venv dir (.venv preferred, matching ui/cron/pythonPath.ts), else default target."""
|
||||
for name in (".venv", "venv"):
|
||||
d = os.path.join(REPO_ROOT, name)
|
||||
if os.path.isdir(d):
|
||||
return d
|
||||
return os.path.join(REPO_ROOT, ".venv")
|
||||
|
||||
|
||||
def venv_python(venv=None):
|
||||
venv = venv or venv_dir()
|
||||
if IS_WINDOWS:
|
||||
return os.path.join(venv, "Scripts", "python.exe")
|
||||
return os.path.join(venv, "bin", "python3")
|
||||
|
||||
|
||||
# Env vars that let a system/conda/pyenv Python leak into our subprocesses.
|
||||
# Scrubbed from every python/pip/node invocation (mirrors what the community
|
||||
# Windows installer learned the hard way).
|
||||
_SCRUB_VARS = (
|
||||
"PYTHONPATH",
|
||||
"PYTHONHOME",
|
||||
"PYTHON",
|
||||
"PYTHONSTARTUP",
|
||||
"PYTHONUSERBASE",
|
||||
"PYTHONEXECUTABLE",
|
||||
"PIP_CONFIG_FILE",
|
||||
"PIP_REQUIRE_VIRTUALENV",
|
||||
"VIRTUAL_ENV",
|
||||
"CONDA_PREFIX",
|
||||
"CONDA_DEFAULT_ENV",
|
||||
"PYENV_ROOT",
|
||||
"PYENV_VERSION",
|
||||
)
|
||||
|
||||
|
||||
def clean_env(extra=None):
|
||||
"""os.environ copy with Python-hijacking vars removed.
|
||||
|
||||
Also points uv's managed-python store into the repo (.uv/python) so
|
||||
interpreter downloads never land outside the checkout.
|
||||
"""
|
||||
env = os.environ.copy()
|
||||
for var in _SCRUB_VARS:
|
||||
env.pop(var, None)
|
||||
env.setdefault("UV_PYTHON_INSTALL_DIR", os.path.join(REPO_ROOT, ".uv", "python"))
|
||||
if extra:
|
||||
env.update(extra)
|
||||
return env
|
||||
|
||||
|
||||
def download(url, dest, label=None):
|
||||
"""Download url to dest (stdlib only), logging progress every ~10%."""
|
||||
import urllib.request
|
||||
|
||||
label = label or os.path.basename(dest)
|
||||
info("Downloading %s ..." % label)
|
||||
last = [-1]
|
||||
|
||||
def hook(blocks, block_size, total):
|
||||
if total <= 0:
|
||||
return
|
||||
pct = min(100, int(blocks * block_size * 100 / total))
|
||||
if pct >= last[0] + 10:
|
||||
last[0] = pct
|
||||
info(" %s: %d%%" % (label, pct))
|
||||
|
||||
tmp = dest + ".part"
|
||||
try:
|
||||
urllib.request.urlretrieve(url, tmp, reporthook=hook)
|
||||
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))
|
||||
os.replace(tmp, dest)
|
||||
|
||||
|
||||
def extract_archive(archive, dest_dir):
|
||||
"""Extract .zip / .tar.* into dest_dir (created if needed)."""
|
||||
import tarfile
|
||||
import zipfile
|
||||
|
||||
os.makedirs(dest_dir, exist_ok=True)
|
||||
if archive.endswith(".zip"):
|
||||
with zipfile.ZipFile(archive) as z:
|
||||
z.extractall(dest_dir)
|
||||
else:
|
||||
with tarfile.open(archive) as t:
|
||||
t.extractall(dest_dir)
|
||||
|
||||
|
||||
def file_hash(paths):
|
||||
import hashlib
|
||||
|
||||
h = hashlib.sha256()
|
||||
for p in sorted(paths):
|
||||
if os.path.isfile(p):
|
||||
with open(p, "rb") as f:
|
||||
h.update(f.read())
|
||||
return h.hexdigest()
|
||||
78
manager/uvbin.py
Normal file
78
manager/uvbin.py
Normal file
@@ -0,0 +1,78 @@
|
||||
"""Local (never system) uv provisioning into <repo>/.uv/.
|
||||
|
||||
uv is a single static binary (no deps, no Python needed). Keeping it inside
|
||||
the repo means fast installs + exact-version Python provisioning with zero
|
||||
footprint outside the checkout. The run scripts also install into this same
|
||||
dir (via the astral installer with UV_INSTALL_DIR) for the no-python
|
||||
bootstrap case; util.find_uv() checks here first either way.
|
||||
"""
|
||||
|
||||
import os
|
||||
import platform
|
||||
import shutil
|
||||
import stat
|
||||
import tempfile
|
||||
|
||||
from .util import IS_MAC, IS_WINDOWS, REPO_ROOT, download, extract_archive, ok, warn
|
||||
|
||||
UV_DIR = os.path.join(REPO_ROOT, ".uv")
|
||||
|
||||
_RELEASE = "https://github.com/astral-sh/uv/releases/latest/download/"
|
||||
|
||||
|
||||
def _asset():
|
||||
arch = platform.machine().lower()
|
||||
if arch == "amd64":
|
||||
arch = "x86_64"
|
||||
if IS_WINDOWS:
|
||||
return "uv-x86_64-pc-windows-msvc.zip"
|
||||
if IS_MAC:
|
||||
return "uv-%s-apple-darwin.tar.gz" % (
|
||||
"aarch64" if arch == "arm64" else "x86_64"
|
||||
)
|
||||
return "uv-%s-unknown-linux-gnu.tar.gz" % (
|
||||
"aarch64" if arch in ("aarch64", "arm64") else "x86_64"
|
||||
)
|
||||
|
||||
|
||||
def local_uv_exe():
|
||||
return os.path.join(UV_DIR, "uv.exe" if IS_WINDOWS else "uv")
|
||||
|
||||
|
||||
def ensure_uv(dry_run=False):
|
||||
"""Download the uv binary into .uv/ if it isn't available anywhere."""
|
||||
from .util import find_uv
|
||||
|
||||
if find_uv():
|
||||
return False
|
||||
if dry_run:
|
||||
from .util import info
|
||||
|
||||
info("[dry-run] would download uv into %s" % UV_DIR)
|
||||
return False
|
||||
tmp = tempfile.mkdtemp(prefix="aitk_uv_")
|
||||
try:
|
||||
asset = _asset()
|
||||
archive = os.path.join(tmp, asset)
|
||||
download(_RELEASE + asset, archive, label="uv")
|
||||
extract_archive(archive, tmp)
|
||||
# zip: uv.exe at root; tarballs: uv-<triple>/uv — search the tree
|
||||
exe_name = "uv.exe" if IS_WINDOWS else "uv"
|
||||
found = None
|
||||
for root, _dirs, files in os.walk(tmp):
|
||||
if exe_name in files:
|
||||
found = os.path.join(root, exe_name)
|
||||
break
|
||||
if not found:
|
||||
warn("Unexpected uv archive layout — continuing without uv.")
|
||||
return False
|
||||
os.makedirs(UV_DIR, exist_ok=True)
|
||||
dest = local_uv_exe()
|
||||
shutil.move(found, dest)
|
||||
os.chmod(
|
||||
dest, os.stat(dest).st_mode | stat.S_IXUSR | stat.S_IXGRP | stat.S_IXOTH
|
||||
)
|
||||
ok("uv installed at %s" % dest)
|
||||
return True
|
||||
finally:
|
||||
shutil.rmtree(tmp, ignore_errors=True)
|
||||
@@ -1,4 +1,4 @@
|
||||
torchao==0.10.0
|
||||
torchao==0.17.0
|
||||
safetensors
|
||||
git+https://github.com/huggingface/diffusers.git@c943837899b16cbae2f619b8dd4f7bb6f07dd81a
|
||||
#pip install git+https://github.com/huggingface/diffusers.git@refs/pull/13432/head
|
||||
|
||||
58
run_linux.sh
Executable file
58
run_linux.sh
Executable file
@@ -0,0 +1,58 @@
|
||||
#!/usr/bin/env bash
|
||||
# Update-and-run script for Linux — thin bootstrap over the in-repo manager.
|
||||
#
|
||||
# Everything (venv via uv-managed Python, torch for your GPU, requirements,
|
||||
# portable Node.js and FFmpeg, dependency updates) is handled by
|
||||
# `python -m manager`; this script only makes sure uv + a Python interpreter
|
||||
# exist, then delegates. Works on desktop and headless boxes alike.
|
||||
set -euo pipefail
|
||||
|
||||
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)"
|
||||
|
||||
# ── Banner ─────────────────────────────────────────────────────────
|
||||
echo ""
|
||||
printf '\033[36m'
|
||||
cat << 'BANNER'
|
||||
_ ___ _____ _ _ _ _
|
||||
/ \ |_ _| |_ _| ___ ___ | || | __(_)| |_
|
||||
/ _ \ | | | | / _ \ / _ \| || |/ /| || __|
|
||||
/ ___ \ | | | | | (_) || (_) | || < | || |_
|
||||
/_/ \_\|___| |_| \___/ \___/|_||_|\_\|_| \__|
|
||||
BANNER
|
||||
printf '\033[0m'
|
||||
printf '\033[90m Linux Setup & Launcher\033[0m\n'
|
||||
echo ""
|
||||
|
||||
# ── 1. Ensure uv (prebuilt static binary, kept inside the repo) ─────
|
||||
export PATH="$SCRIPT_DIR/.uv:$PATH"
|
||||
export UV_PYTHON_INSTALL_DIR="$SCRIPT_DIR/.uv/python"
|
||||
if ! command -v uv >/dev/null 2>&1; then
|
||||
echo "Downloading uv (package/python manager) into .uv/ ..."
|
||||
curl -LsSf https://astral.sh/uv/install.sh | \
|
||||
UV_INSTALL_DIR="$SCRIPT_DIR/.uv" UV_NO_MODIFY_PATH=1 sh
|
||||
fi
|
||||
|
||||
# ── 2. Find a Python to run the manager (stdlib-only, needs >= 3.9) ─
|
||||
find_python() {
|
||||
for cmd in python3 python; do
|
||||
if command -v "$cmd" >/dev/null 2>&1; then
|
||||
if "$cmd" -c 'import sys; sys.exit(0 if sys.version_info >= (3, 9) else 1)' 2>/dev/null; then
|
||||
echo "$cmd"
|
||||
return 0
|
||||
fi
|
||||
fi
|
||||
done
|
||||
return 1
|
||||
}
|
||||
|
||||
PYTHON="$(find_python || true)"
|
||||
if [[ -z "$PYTHON" ]]; then
|
||||
echo "No system Python found — provisioning one with uv..."
|
||||
uv python install 3.12
|
||||
PYTHON="$(uv python find 3.12)"
|
||||
fi
|
||||
|
||||
# ── 3. Sync the environment and start the UI ────────────────────────
|
||||
cd "$SCRIPT_DIR"
|
||||
"$PYTHON" -m manager update --auto
|
||||
exec "$PYTHON" -m manager launch
|
||||
171
run_mac.zsh
171
run_mac.zsh
@@ -1,5 +1,9 @@
|
||||
#!/usr/bin/env zsh
|
||||
# Update-and-run script for macOS — portable Python 3.12 + PyTorch
|
||||
# Update-and-run script for macOS — thin bootstrap over the in-repo manager.
|
||||
#
|
||||
# Everything (venv via uv-managed Python, torch, requirements, portable
|
||||
# Node.js and FFmpeg, dependency updates) is handled by `python -m manager`;
|
||||
# this script only makes sure uv + a Python interpreter exist, then delegates.
|
||||
set -euo pipefail
|
||||
|
||||
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)"
|
||||
@@ -17,150 +21,37 @@ BANNER
|
||||
echo "\033[0m"
|
||||
echo "\033[90m macOS Setup & Launcher\033[0m"
|
||||
echo ""
|
||||
VENV_DIR="$SCRIPT_DIR/.venv"
|
||||
PIP="$VENV_DIR/bin/pip"
|
||||
PYTHON="$VENV_DIR/bin/python3"
|
||||
PYTHON_VERSION="3.12.8"
|
||||
RELEASE_TAG="20241219"
|
||||
|
||||
# --- Package versions (update these as needed) ---
|
||||
NODE_VERSION="23.11.1"
|
||||
TORCH_VERSION="2.11.0"
|
||||
TORCHVISION_VERSION="0.26.0"
|
||||
TORCHAUDIO_VERSION="2.11.0"
|
||||
|
||||
# Detect architecture
|
||||
ARCH="$(uname -m)"
|
||||
if [[ "$ARCH" == "arm64" ]]; then
|
||||
PLATFORM="aarch64-apple-darwin"
|
||||
elif [[ "$ARCH" == "x86_64" ]]; then
|
||||
PLATFORM="x86_64-apple-darwin"
|
||||
else
|
||||
echo "Error: Unsupported architecture: $ARCH"
|
||||
exit 1
|
||||
# ── 1. Ensure uv (prebuilt static binary, kept inside the repo) ─────
|
||||
export PATH="$SCRIPT_DIR/.uv:$PATH"
|
||||
export UV_PYTHON_INSTALL_DIR="$SCRIPT_DIR/.uv/python"
|
||||
if ! command -v uv >/dev/null 2>&1; then
|
||||
echo "Downloading uv (package/python manager) into .uv/ ..."
|
||||
curl -LsSf https://astral.sh/uv/install.sh | \
|
||||
UV_INSTALL_DIR="$SCRIPT_DIR/.uv" UV_NO_MODIFY_PATH=1 sh
|
||||
fi
|
||||
|
||||
# ── 1. Download standalone Python if needed ─────────────────────────
|
||||
PYTHON_DIR="$SCRIPT_DIR/.python"
|
||||
PYTHON_BIN="$PYTHON_DIR/bin/python3"
|
||||
|
||||
if [[ ! -x "$PYTHON_BIN" ]]; then
|
||||
TARBALL="cpython-${PYTHON_VERSION}+${RELEASE_TAG}-${PLATFORM}-install_only.tar.gz"
|
||||
URL="https://github.com/indygreg/python-build-standalone/releases/download/${RELEASE_TAG}/${TARBALL}"
|
||||
|
||||
TMPDIR_DL="$(mktemp -d)"
|
||||
trap 'rm -rf "$TMPDIR_DL"' EXIT
|
||||
|
||||
echo "Downloading standalone Python ${PYTHON_VERSION} (${PLATFORM})..."
|
||||
curl -fSL --progress-bar -o "$TMPDIR_DL/$TARBALL" "$URL"
|
||||
|
||||
echo "Extracting..."
|
||||
tar -xzf "$TMPDIR_DL/$TARBALL" -C "$TMPDIR_DL"
|
||||
|
||||
# Move to permanent location (the archive extracts to a "python" folder)
|
||||
rm -rf "$PYTHON_DIR"
|
||||
mv "$TMPDIR_DL/python" "$PYTHON_DIR"
|
||||
|
||||
rm -rf "$TMPDIR_DL"
|
||||
trap - EXIT
|
||||
|
||||
echo "Standalone Python installed to $PYTHON_DIR"
|
||||
fi
|
||||
|
||||
# ── 2. Create venv if it doesn't exist ──────────────────────────────
|
||||
if [[ ! -d "$VENV_DIR" ]]; then
|
||||
echo "Creating virtual environment at $VENV_DIR..."
|
||||
"$PYTHON_BIN" -m venv "$VENV_DIR"
|
||||
echo "Virtual environment created."
|
||||
fi
|
||||
|
||||
# ── 3. Download / update portable Node.js ──────────────────────────
|
||||
NODE_DIR="$SCRIPT_DIR/.node"
|
||||
NODE_BIN="$NODE_DIR/bin/node"
|
||||
|
||||
NEED_NODE=false
|
||||
if [[ ! -x "$NODE_BIN" ]]; then
|
||||
NEED_NODE=true
|
||||
elif [[ "$("$NODE_BIN" --version 2>/dev/null)" != "v${NODE_VERSION}" ]]; then
|
||||
echo "Node.js version mismatch (want v${NODE_VERSION}, have $("$NODE_BIN" --version))."
|
||||
NEED_NODE=true
|
||||
fi
|
||||
|
||||
if $NEED_NODE; then
|
||||
if [[ "$ARCH" == "arm64" ]]; then
|
||||
NODE_ARCH="arm64"
|
||||
else
|
||||
NODE_ARCH="x64"
|
||||
# ── 2. Find a Python to run the manager (stdlib-only, needs >= 3.9) ─
|
||||
find_python() {
|
||||
for cmd in python3 python; do
|
||||
if command -v "$cmd" >/dev/null 2>&1; then
|
||||
if "$cmd" -c 'import sys; sys.exit(0 if sys.version_info >= (3, 9) else 1)' 2>/dev/null; then
|
||||
echo "$cmd"
|
||||
return 0
|
||||
fi
|
||||
|
||||
NODE_TARBALL="node-v${NODE_VERSION}-darwin-${NODE_ARCH}.tar.gz"
|
||||
NODE_URL="https://nodejs.org/dist/v${NODE_VERSION}/${NODE_TARBALL}"
|
||||
|
||||
TMPDIR_DL="$(mktemp -d)"
|
||||
trap 'rm -rf "$TMPDIR_DL"' EXIT
|
||||
|
||||
echo "Downloading Node.js v${NODE_VERSION} (darwin-${NODE_ARCH})..."
|
||||
curl -fSL --progress-bar -o "$TMPDIR_DL/$NODE_TARBALL" "$NODE_URL"
|
||||
|
||||
echo "Extracting..."
|
||||
tar -xzf "$TMPDIR_DL/$NODE_TARBALL" -C "$TMPDIR_DL"
|
||||
|
||||
rm -rf "$NODE_DIR"
|
||||
mv "$TMPDIR_DL/node-v${NODE_VERSION}-darwin-${NODE_ARCH}" "$NODE_DIR"
|
||||
|
||||
rm -rf "$TMPDIR_DL"
|
||||
trap - EXIT
|
||||
|
||||
echo "Node.js v${NODE_VERSION} installed to $NODE_DIR"
|
||||
else
|
||||
echo "Node.js v${NODE_VERSION} is up to date."
|
||||
fi
|
||||
|
||||
# ── 4. Install / update PyTorch packages ────────────────────────────
|
||||
# Helper: returns 0 if the package is installed at the exact version
|
||||
pkg_ok() {
|
||||
local pkg="$1" want="$2"
|
||||
local got
|
||||
got="$("$PIP" show "$pkg" 2>/dev/null | awk '/^Version:/{print $2}')" || true
|
||||
[[ "$got" == "$want" ]]
|
||||
fi
|
||||
done
|
||||
return 1
|
||||
}
|
||||
|
||||
PKGS_TO_INSTALL=()
|
||||
|
||||
pkg_ok "torch" "$TORCH_VERSION" || PKGS_TO_INSTALL+=("torch==$TORCH_VERSION")
|
||||
pkg_ok "torchvision" "$TORCHVISION_VERSION" || PKGS_TO_INSTALL+=("torchvision==$TORCHVISION_VERSION")
|
||||
pkg_ok "torchaudio" "$TORCHAUDIO_VERSION" || PKGS_TO_INSTALL+=("torchaudio==$TORCHAUDIO_VERSION")
|
||||
|
||||
if (( ${#PKGS_TO_INSTALL[@]} )); then
|
||||
echo "Installing / updating: ${PKGS_TO_INSTALL[*]}"
|
||||
"$PIP" install "${PKGS_TO_INSTALL[@]}"
|
||||
else
|
||||
echo "PyTorch packages are up to date."
|
||||
PYTHON="$(find_python || true)"
|
||||
if [[ -z "$PYTHON" ]]; then
|
||||
echo "No system Python found — provisioning one with uv..."
|
||||
uv python install 3.12
|
||||
PYTHON="$(uv python find 3.12)"
|
||||
fi
|
||||
|
||||
# ── 5. Install / update requirements.txt ────────────────────────────
|
||||
REQUIREMENTS="$SCRIPT_DIR/requirements.txt"
|
||||
REQ_HASH_FILE="$VENV_DIR/.requirements_hash"
|
||||
|
||||
if [[ -f "$REQUIREMENTS" ]]; then
|
||||
# Hash all requirements files (follows -r includes)
|
||||
CURRENT_HASH="$(cat "$SCRIPT_DIR"/requirements*.txt 2>/dev/null | shasum -a 256 | awk '{print $1}')"
|
||||
STORED_HASH=""
|
||||
[[ -f "$REQ_HASH_FILE" ]] && STORED_HASH="$(cat "$REQ_HASH_FILE")"
|
||||
|
||||
if [[ "$CURRENT_HASH" != "$STORED_HASH" ]]; then
|
||||
echo "Installing / updating requirements.txt..."
|
||||
"$PIP" install -r "$REQUIREMENTS"
|
||||
echo "$CURRENT_HASH" > "$REQ_HASH_FILE"
|
||||
else
|
||||
echo "Requirements are up to date."
|
||||
fi
|
||||
fi
|
||||
|
||||
# ── 6. Build and start the UI ───────────────────────────────────────
|
||||
export PATH="$NODE_DIR/bin:$VENV_DIR/bin:$PATH"
|
||||
|
||||
echo ""
|
||||
echo "Starting UI..."
|
||||
cd "$SCRIPT_DIR/ui"
|
||||
npm run build_and_start
|
||||
# ── 3. Sync the environment and start the UI ────────────────────────
|
||||
cd "$SCRIPT_DIR"
|
||||
"$PYTHON" -m manager update --auto
|
||||
exec "$PYTHON" -m manager launch
|
||||
|
||||
@@ -55,7 +55,6 @@ image = (
|
||||
"toml",
|
||||
"pydantic",
|
||||
"omegaconf",
|
||||
"k-diffusion",
|
||||
"open_clip_torch",
|
||||
"timm",
|
||||
"prodigyopt",
|
||||
|
||||
77
run_windows.bat
Normal file
77
run_windows.bat
Normal file
@@ -0,0 +1,77 @@
|
||||
@echo off&&cd /d %~dp0
|
||||
REM Update-and-run script for Windows - thin bootstrap over the in-repo manager.
|
||||
REM
|
||||
REM Everything (venv via uv-managed Python, torch for your GPU, requirements,
|
||||
REM portable Node.js / FFmpeg / Git, dependency updates) is handled by
|
||||
REM `python -m manager`; this script only makes sure uv + a Python interpreter
|
||||
REM exist, then delegates.
|
||||
setlocal EnableDelayedExpansion
|
||||
Title AI Toolkit
|
||||
|
||||
echo.
|
||||
echo _ ___ _____ _ _ _ _
|
||||
echo / \ ^|_ _^| ^|_ _^| ___ ___ ^| ^|^| ^| __(_)^| ^|_
|
||||
echo / _ \ ^| ^| ^| ^| / _ \ / _ \^| ^|^| ^|/ /^| ^|^| __^|
|
||||
echo / ___ \ ^| ^| ^| ^| ^| (_) ^|^| (_) ^| ^|^| ^< ^| ^|^| ^|_
|
||||
echo /_/ \_\^|___^| ^|_^| \___/ \___/^|_^|^|_^|\_\^|_^| \__^|
|
||||
echo.
|
||||
echo Windows Setup ^& Launcher
|
||||
echo.
|
||||
|
||||
REM Clear env vars that let a stray conda/pyenv/system Python hijack things
|
||||
set PYTHONPATH=
|
||||
set PYTHONHOME=
|
||||
set PYTHONSTARTUP=
|
||||
set PYTHONUSERBASE=
|
||||
set PIP_CONFIG_FILE=
|
||||
set VIRTUAL_ENV=
|
||||
set CONDA_PREFIX=
|
||||
set CONDA_DEFAULT_ENV=
|
||||
set PYENV_ROOT=
|
||||
set PYENV_VERSION=
|
||||
|
||||
REM ---- 1. Ensure uv (prebuilt static binary, kept inside the repo) ----
|
||||
set "PATH=%~dp0.uv;%PATH%"
|
||||
set "UV_PYTHON_INSTALL_DIR=%~dp0.uv\python"
|
||||
where uv.exe >nul 2>&1
|
||||
if errorlevel 1 (
|
||||
echo Downloading uv ^(package/python manager^) into .uv\ ...
|
||||
powershell -NoProfile -ExecutionPolicy ByPass -Command ^
|
||||
"$env:UV_INSTALL_DIR = Join-Path '%~dp0' '.uv'; $env:UV_NO_MODIFY_PATH = '1'; irm https://astral.sh/uv/install.ps1 | iex"
|
||||
where uv.exe >nul 2>&1
|
||||
if errorlevel 1 (
|
||||
echo ERROR: uv download failed. See https://docs.astral.sh/uv/
|
||||
pause
|
||||
exit /b 1
|
||||
)
|
||||
)
|
||||
|
||||
REM ---- 2. Find a Python to run the manager (stdlib-only, needs 3.9+) ----
|
||||
set "PY="
|
||||
for %%C in (python.exe py.exe) do (
|
||||
if not defined PY (
|
||||
%%C -c "import sys; sys.exit(0 if sys.version_info >= (3, 9) else 1)" >nul 2>&1
|
||||
if not errorlevel 1 set "PY=%%C"
|
||||
)
|
||||
)
|
||||
if not defined PY (
|
||||
echo No system Python found - provisioning one with uv...
|
||||
uv python install 3.12
|
||||
for /f "delims=" %%P in ('uv python find 3.12') do set "PY=%%P"
|
||||
)
|
||||
if not defined PY (
|
||||
echo ERROR: could not find or install a Python interpreter.
|
||||
pause
|
||||
exit /b 1
|
||||
)
|
||||
|
||||
REM ---- 3. Sync the environment and start the UI ----
|
||||
"%PY%" -m manager update --auto
|
||||
if errorlevel 1 (
|
||||
echo.
|
||||
echo Setup failed - see output above.
|
||||
pause
|
||||
exit /b 1
|
||||
)
|
||||
"%PY%" -m manager launch
|
||||
pause
|
||||
Reference in New Issue
Block a user