Files
ai-toolkit/manager/env.py

455 lines
15 KiB
Python

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