Files
voicebox/backend/services/rocm.py
e766c7cbfb feat(windows): Native AMD ROCm GPU Acceleration (Resolves #531) (#538)
* feat(windows): add native ROCm support for AMD GPUs

Implements native ROCm architecture for Windows.

- Adds backend build pipeline for voicebox-server-rocm.exe

- Detects AMD GPUs dynamically and routes PyTorch allocations

- Adds automatic download and update logic for ROCm dependencies

- Refactors UI in GpuPage.tsx and GpuAcceleration.tsx to add AMD flows

- Fixes 'Switch to CPU' lock on Windows via Tauri backend_override state

- Resolves PyInstaller/rocm_sdk UnboundLocalError silent crashes

- Resolves Numba/NumPy 2.x incompatibilities during Qwen3-TTS load

- Resolves HF_HUB_OFFLINE Catch-22 for CustomVoice processor caching

* fix(rocm): host libs archive under the app release tag, drop offline-load regression

Align the ROCm libs download with the CUDA pattern: both the server core and
the libs archive are published under the app-version release tag, with the libs
content version encoded in the filename only. The previous code fetched libs
from a separate rocm7.2-v1 tag, which disagreed with the download test.

Also revert the unrelated Qwen CustomVoice changes that wrapped model loading in
force_offline_if_cached (not imported — a NameError on load for every platform)
and re-added a Base-model cache gate. The inference-path offline guard was
deliberately removed previously.

* feat(rocm): gate download on AMD detection and persist the backend variant

The ROCm download section now only shows when the backend reports an AMD GPU on
Windows (new supports_rocm health field, backed by the memoized
is_amd_gpu_windows detection that was previously unused), or when ROCm is already
downloaded/active.

Make the backend override honor a pinned variant: set_backend_override persists
the choice to disk so it survives an app restart, start_server reads it back,
and a cuda/rocm pin now actually selects that variant instead of always
preferring ROCm. A stale pin to a deleted backend self-heals to the default
order rather than forcing CPU. Add the web no-op stub for the new method.

* chore(rocm): drop incomplete vitest harness for the unused GpuAcceleration component

GpuAcceleration.tsx is not routed anywhere (GpuPage is the live settings view),
and the added vitest setup referenced testing-library/vitest deps that were not
in the lockfile, breaking the web typecheck. Remove the dead component's test
and its scaffolding to keep this PR scoped to the ROCm feature.

* ci(rocm): add ROCm release-artifact pipeline

Mirror the CUDA packaging path for ROCm so the runtime download has artifacts to
fetch. scripts/package_rocm.py splits the PyInstaller --rocm onedir into
voicebox-server-rocm.tar.gz (core) + rocm-libs-rocm7.2-v1.tar.gz (AMD runtime:
HIP DLLs, rocBLAS Tensile data, MIOpen kernel DBs) + rocm-libs.json, matching
the names services/rocm.py expects, both under the app-version release tag.

The new build-rocm-windows job in release.yml builds on windows-latest/cp312 and
lets build_binary.py --rocm pull the official AMD Radeon wheels.

The file classifier can't be validated against a real AMD build on CI, so it has
unit coverage (test_package_rocm.py) against a synthetic onedir layout. The
prefixes/dir markers may need a tweak after the first real build on AMD
hardware — the packager hard-fails loudly if it classifies zero ROCm files.

---------

Co-authored-by: Jamie Pine <[email protected]>
2026-06-30 15:43:18 -07:00

468 lines
16 KiB
Python

"""
ROCm backend download, assembly, and verification.
Downloads two archives from GitHub Releases:
1. Server core (voicebox-server-rocm.tar.gz) — the exe + non-AMD deps,
versioned with the app.
2. ROCm libs (rocm-libs-{version}.tar.gz) — AMD runtime libraries,
versioned independently (only redownloaded on ROCm toolkit bump).
Both archives are extracted into {data_dir}/backends/rocm/ which forms the
complete PyInstaller --onedir directory structure that torch expects.
"""
import asyncio
import hashlib
import json
import logging
import os
import shutil
import sys
import tarfile
from pathlib import Path
from typing import Optional
from ..config import get_data_dir
from ..utils.progress import get_progress_manager
from .. import __version__
logger = logging.getLogger(__name__)
GITHUB_RELEASES_URL = "https://github.com/jamiepine/voicebox/releases/download"
PROGRESS_KEY = "rocm-backend"
# The current expected ROCm libs version. Bump this when we change the
# ROCm toolkit version or torch's ROCm dependency changes (e.g. rocm7.2 -> rocm7.4).
ROCM_LIBS_VERSION = "rocm7.2-v1"
# Prevents concurrent download_rocm_binary() calls from racing on the same
# temp file. The auto-update background task and the manual HTTP endpoint
# can both invoke download_rocm_binary(); without this lock the progress-
# manager status check is a TOCTOU race.
_download_lock = asyncio.Lock()
def get_backends_dir() -> Path:
"""Directory where downloaded backend binaries are stored."""
d = get_data_dir() / "backends"
d.mkdir(parents=True, exist_ok=True)
return d
def get_rocm_dir() -> Path:
"""Directory where the ROCm backend (onedir) is extracted."""
d = get_backends_dir() / "rocm"
d.mkdir(parents=True, exist_ok=True)
return d
def get_rocm_exe_name() -> str:
"""Platform-specific ROCm executable filename."""
if sys.platform == "win32":
return "voicebox-server-rocm.exe"
return "voicebox-server-rocm"
def get_rocm_binary_path() -> Optional[Path]:
"""Return path to the ROCm executable if it exists inside the onedir."""
p = get_rocm_dir() / get_rocm_exe_name()
if p.exists():
return p
return None
def get_rocm_libs_manifest_path() -> Path:
"""Path to the rocm-libs.json manifest inside the ROCm dir."""
return get_rocm_dir() / "rocm-libs.json"
def get_installed_rocm_libs_version() -> Optional[str]:
"""Read the installed ROCm libs version from rocm-libs.json, or None."""
manifest_path = get_rocm_libs_manifest_path()
if not manifest_path.exists():
return None
try:
data = json.loads(manifest_path.read_text())
return data.get("version")
except Exception as e:
logger.warning(f"Could not read rocm-libs.json: {e}")
return None
def is_rocm_active() -> bool:
"""Check if the current process is the ROCm binary.
The ROCm binary sets this env var on startup (see server.py).
"""
return os.environ.get("VOICEBOX_BACKEND_VARIANT") == "rocm"
def get_rocm_status() -> dict:
"""Get current ROCm backend status for the API."""
progress_manager = get_progress_manager()
rocm_path = get_rocm_binary_path()
progress = progress_manager.get_progress(PROGRESS_KEY)
rocm_libs_version = get_installed_rocm_libs_version()
return {
"available": rocm_path is not None,
"active": is_rocm_active(),
"binary_path": str(rocm_path) if rocm_path else None,
"rocm_libs_version": rocm_libs_version,
"downloading": progress is not None and progress.get("status") == "downloading",
"download_progress": progress,
}
def _needs_server_download(version: Optional[str] = None) -> bool:
"""Check if the server core archive needs to be (re)downloaded."""
rocm_path = get_rocm_binary_path()
if not rocm_path:
return True
# Check if the binary version matches the expected app version
installed = get_rocm_binary_version()
expected = version or __version__
if expected.startswith("v"):
expected = expected[1:]
return installed != expected
def _needs_rocm_libs_download() -> bool:
"""Check if the ROCm libs archive needs to be (re)downloaded."""
installed = get_installed_rocm_libs_version()
if installed is None:
return True
return installed != ROCM_LIBS_VERSION
async def _download_and_extract_archive(
client,
url: str,
sha256_url: Optional[str],
dest_dir: Path,
label: str,
progress_offset: int,
total_size: int,
):
"""Download a .tar.gz archive and extract it into dest_dir.
Args:
client: httpx.AsyncClient
url: URL of the .tar.gz archive
sha256_url: URL of the .sha256 checksum file (optional)
dest_dir: Directory to extract into
label: Human-readable label for progress updates
progress_offset: Byte offset for progress reporting (when downloading
multiple archives sequentially)
total_size: Total bytes across all downloads (for progress bar)
"""
progress = get_progress_manager()
temp_path = dest_dir / f".download-{label.replace(' ', '-')}.tmp"
# Clean up leftover partial download
if temp_path.exists():
temp_path.unlink()
# Fetch expected checksum (fail-fast: never extract an unverified archive)
expected_sha = None
if sha256_url:
try:
sha_resp = await client.get(sha256_url)
sha_resp.raise_for_status()
expected_sha = sha_resp.text.strip().split()[0]
logger.info(f"{label}: expected SHA-256: {expected_sha[:16]}...")
except Exception as e:
raise RuntimeError(f"{label}: failed to fetch checksum from {sha256_url}") from e
# Stream download, verify, and extract — always clean up temp file
downloaded = 0
try:
async with client.stream("GET", url) as response:
response.raise_for_status()
with open(temp_path, "wb") as f:
async for chunk in response.aiter_bytes(chunk_size=1024 * 1024):
f.write(chunk)
downloaded += len(chunk)
progress.update_progress(
PROGRESS_KEY,
current=progress_offset + downloaded,
total=total_size,
filename=f"Downloading {label}",
status="downloading",
)
# Verify integrity
if expected_sha:
progress.update_progress(
PROGRESS_KEY,
current=progress_offset + downloaded,
total=total_size,
filename=f"Verifying {label}...",
status="downloading",
)
sha256 = hashlib.sha256()
with open(temp_path, "rb") as f:
while True:
data = f.read(1024 * 1024)
if not data:
break
sha256.update(data)
actual = sha256.hexdigest()
if actual != expected_sha:
raise ValueError(
f"{label} integrity check failed: expected {expected_sha[:16]}..., got {actual[:16]}..."
)
logger.info(f"{label}: integrity verified")
# Extract (use data filter for path traversal protection on Python 3.12+)
progress.update_progress(
PROGRESS_KEY,
current=progress_offset + downloaded,
total=total_size,
filename=f"Extracting {label}...",
status="downloading",
)
with tarfile.open(temp_path, "r:gz") as tar:
tar.extractall(path=dest_dir, filter="data")
logger.info(f"{label}: extracted to {dest_dir}")
finally:
if temp_path.exists():
temp_path.unlink()
return downloaded
async def download_rocm_binary(version: Optional[str] = None):
"""Download the ROCm backend (server core + ROCm libs if needed).
Downloads both archives from GitHub Releases, extracts them into
{data_dir}/backends/rocm/, and writes the rocm-libs.json manifest.
Only downloads what's needed:
- Server core: always redownloaded (versioned with app)
- ROCm libs: only if missing or version mismatch
Args:
version: Version tag (e.g. "v0.3.0"). Defaults to current app version.
"""
if _download_lock.locked():
logger.info("ROCm download already in progress, skipping duplicate request")
return
async with _download_lock:
await _download_rocm_binary_locked(version)
async def _download_rocm_binary_locked(version: Optional[str] = None):
"""Inner implementation of download_rocm_binary, called under _download_lock."""
import httpx
if version is None:
version = f"v{__version__}"
progress = get_progress_manager()
rocm_dir = get_rocm_dir()
need_server = _needs_server_download(version)
need_libs = _needs_rocm_libs_download()
if not need_server and not need_libs:
logger.info("ROCm backend is up to date, nothing to download")
return
logger.info(
f"Starting ROCm backend download for {version} "
f"(server={'yes' if need_server else 'cached'}, "
f"libs={'yes' if need_libs else 'cached'})"
)
progress.update_progress(
PROGRESS_KEY,
current=0,
total=0,
filename="Preparing download...",
status="downloading",
)
# Server core and libs archive are both published under the app-version
# release tag; the libs content version is encoded in the filename only.
server_base_url = f"{GITHUB_RELEASES_URL}/{version}"
libs_base_url = server_base_url
server_archive = "voicebox-server-rocm.tar.gz"
libs_archive = f"rocm-libs-{ROCM_LIBS_VERSION}.tar.gz"
# Always stage when any download is needed, then atomically rename over
# rocm_dir on success. This prevents a failed mid-extraction from leaving
# rocm_dir in a partially-installed state that still passes the
# get_rocm_binary_path() existence check. Existing files are pre-copied
# into staging so partial updates (e.g. libs-only or server-only) preserve
# whatever isn't being re-downloaded.
use_staging = need_server or need_libs
staging_dir = get_backends_dir() / "rocm-staging"
if use_staging:
if staging_dir.exists():
shutil.rmtree(staging_dir)
staging_dir.mkdir(parents=True, exist_ok=True)
# Preserve existing files (server or libs) that don't need re-downloading.
# Extracted archives will overwrite only what we actually download.
if rocm_dir.exists():
shutil.copytree(rocm_dir, staging_dir, dirs_exist_ok=True)
extract_dir = staging_dir
else:
extract_dir = rocm_dir
try:
async with httpx.AsyncClient(follow_redirects=True, timeout=30.0) as client:
# Estimate total download size
total_size = 0
if need_server:
try:
head = await client.head(f"{server_base_url}/{server_archive}")
total_size += int(head.headers.get("content-length", 0))
except Exception:
pass
if need_libs:
try:
head = await client.head(f"{libs_base_url}/{libs_archive}")
total_size += int(head.headers.get("content-length", 0))
except Exception:
pass
logger.info(f"Total download size: {total_size / 1024 / 1024:.1f} MB")
offset = 0
# Download server core
if need_server:
server_downloaded = await _download_and_extract_archive(
client,
url=f"{server_base_url}/{server_archive}",
sha256_url=f"{server_base_url}/{server_archive}.sha256",
dest_dir=extract_dir,
label="ROCm server",
progress_offset=offset,
total_size=total_size,
)
offset += server_downloaded
# Make executable on Unix
exe_path = extract_dir / get_rocm_exe_name()
if sys.platform != "win32" and exe_path.exists():
exe_path.chmod(0o755)
# Download ROCm libs
if need_libs:
await _download_and_extract_archive(
client,
url=f"{libs_base_url}/{libs_archive}",
sha256_url=f"{libs_base_url}/{libs_archive}.sha256",
dest_dir=extract_dir,
label="ROCm libraries",
progress_offset=offset,
total_size=total_size,
)
# Write local rocm-libs.json manifest
manifest = {"version": ROCM_LIBS_VERSION}
(extract_dir / "rocm-libs.json").write_text(json.dumps(manifest, indent=2) + "\n")
# Atomic swap: replace rocm_dir with the fully-extracted staging dir
if use_staging:
backup_dir = get_backends_dir() / "rocm-backup"
if backup_dir.exists():
shutil.rmtree(backup_dir)
if rocm_dir.exists():
rocm_dir.rename(backup_dir)
try:
staging_dir.rename(rocm_dir)
except Exception:
if backup_dir.exists() and not rocm_dir.exists():
backup_dir.rename(rocm_dir)
raise
else:
if backup_dir.exists():
shutil.rmtree(backup_dir)
logger.info(f"ROCm backend ready at {rocm_dir}")
progress.mark_complete(PROGRESS_KEY)
except Exception as e:
if use_staging and staging_dir.exists():
shutil.rmtree(staging_dir)
logger.error(f"ROCm backend download failed: {e}")
progress.mark_error(PROGRESS_KEY, str(e))
raise
def get_rocm_binary_version() -> Optional[str]:
"""Get the version of the installed ROCm binary, or None if not installed."""
import subprocess
rocm_path = get_rocm_binary_path()
if not rocm_path:
return None
try:
result = subprocess.run(
[str(rocm_path), "--version"],
capture_output=True,
text=True,
timeout=30,
cwd=str(rocm_path.parent), # Run from the onedir directory
)
# Output format: "voicebox-server 0.3.0"
for line in result.stdout.strip().splitlines():
if "voicebox-server" in line:
return line.split()[-1]
except Exception as e:
logger.warning(f"Could not get ROCm binary version: {e}")
return None
async def check_and_update_rocm_binary():
"""Check if the ROCm binary is outdated and auto-download if so.
Called on server startup. Checks both server version and ROCm libs
version. Downloads only what's needed.
"""
rocm_path = get_rocm_binary_path()
if not rocm_path:
return # No ROCm binary installed, nothing to update
if is_rocm_active():
logger.info("ROCm backend is active; skipping auto-update to avoid replacing the running backend")
return
need_server = _needs_server_download()
need_libs = _needs_rocm_libs_download()
if not need_server and not need_libs:
logger.info(f"ROCm binary is up to date (server=v{__version__}, libs={get_installed_rocm_libs_version()})")
return
reasons = []
if need_server:
rocm_version = get_rocm_binary_version()
reasons.append(f"server v{rocm_version} != v{__version__}")
if need_libs:
installed_libs = get_installed_rocm_libs_version()
reasons.append(f"libs {installed_libs} != {ROCM_LIBS_VERSION}")
logger.info(f"ROCm backend needs update ({', '.join(reasons)}). Auto-downloading...")
try:
await download_rocm_binary()
except Exception as e:
logger.error(f"Auto-update of ROCm binary failed: {e}")
async def delete_rocm_binary() -> bool:
"""Delete the downloaded ROCm backend directory. Returns True if deleted."""
import shutil
rocm_dir = get_rocm_dir()
if rocm_dir.exists() and any(rocm_dir.iterdir()):
shutil.rmtree(rocm_dir)
logger.info(f"Deleted ROCm backend directory: {rocm_dir}")
return True
return False