mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 17:15:19 -07:00
Merge remote-tracking branch 'origin/main' into prep/pr-1112
# Conflicts: # CHANGELOG.md
This commit is contained in:
@@ -6,6 +6,7 @@ voice prompt combination, and model loading progress tracking.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import platform
|
||||
from contextlib import contextmanager
|
||||
from pathlib import Path
|
||||
@@ -22,6 +23,24 @@ from ..utils.tasks import get_task_manager
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def has_in_progress_download(blobs_dir: Path) -> bool:
|
||||
"""
|
||||
Whether a HuggingFace repo's ``blobs`` dir holds a genuinely in-progress download.
|
||||
|
||||
An ``.incomplete`` blob means a download is still in progress -- unless a
|
||||
completed blob with the same hash already sits next to it, which happens
|
||||
when a retried/concurrent download leaves a stale ``.incomplete`` behind
|
||||
after the real transfer already finished. Only orphaned ``.incomplete``
|
||||
files (no matching completed blob) count as "in progress".
|
||||
"""
|
||||
if not blobs_dir.exists():
|
||||
return False
|
||||
return any(
|
||||
not incomplete.with_name(incomplete.name.removesuffix(".incomplete")).exists()
|
||||
for incomplete in blobs_dir.glob("*.incomplete")
|
||||
)
|
||||
|
||||
|
||||
def is_model_cached(
|
||||
hf_repo: str,
|
||||
*,
|
||||
@@ -48,10 +67,8 @@ def is_model_cached(
|
||||
if not repo_cache.exists():
|
||||
return False
|
||||
|
||||
# Incomplete blobs mean a download is still in progress
|
||||
blobs_dir = repo_cache / "blobs"
|
||||
if blobs_dir.exists() and any(blobs_dir.glob("*.incomplete")):
|
||||
logger.debug(f"Found .incomplete files for {hf_repo}")
|
||||
if has_in_progress_download(repo_cache / "blobs"):
|
||||
logger.debug(f"Found in-progress .incomplete file for {hf_repo}")
|
||||
return False
|
||||
|
||||
snapshots_dir = repo_cache / "snapshots"
|
||||
@@ -78,6 +95,13 @@ def is_model_cached(
|
||||
return False
|
||||
|
||||
|
||||
# Documented escape hatch (docs/content/docs/overview/gpu-acceleration.mdx):
|
||||
# users whose GPU has no compiled kernels in the bundled PyTorch set this to run
|
||||
# on CPU instead of crashing at generation time.
|
||||
FORCE_CPU_ENV_VAR = "VOICEBOX_FORCE_CPU"
|
||||
FORCE_CPU_ENABLED_VALUE = "1"
|
||||
|
||||
|
||||
def get_torch_device(
|
||||
*,
|
||||
allow_xpu: bool = False,
|
||||
@@ -93,7 +117,17 @@ def get_torch_device(
|
||||
allow_directml: Check for DirectML (Windows) support.
|
||||
allow_mps: Allow MPS (Apple Silicon). If False, MPS falls back to CPU.
|
||||
force_cpu_on_mac: Force CPU on macOS regardless of GPU availability.
|
||||
|
||||
The VOICEBOX_FORCE_CPU override wins over every other candidate, and is
|
||||
resolved before torch is imported so it still works when the installed
|
||||
build is the reason CPU is wanted.
|
||||
"""
|
||||
# Stripped: on Windows, where this override matters most, it is usually set
|
||||
# through the GUI environment editor.
|
||||
if os.environ.get(FORCE_CPU_ENV_VAR, "").strip() == FORCE_CPU_ENABLED_VALUE:
|
||||
logger.info("%s=%s set, forcing CPU device", FORCE_CPU_ENV_VAR, FORCE_CPU_ENABLED_VALUE)
|
||||
return "cpu"
|
||||
|
||||
if force_cpu_on_mac and platform.system() == "Darwin":
|
||||
return "cpu"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user