mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 17:15:19 -07:00
fix(backend): honour VOICEBOX_FORCE_CPU in device selection
The override is documented in gpu-acceleration.mdx and listed as step 1 of the get_torch_device() precedence in tts-generation.mdx, but grepping the tree for VOICEBOX_FORCE_CPU matched only those two doc files - nothing read it. Users whose GPU has no compiled kernels in the bundled PyTorch had no way to fall back to CPU short of renaming the installed CUDA backend directory. Resolve it before torch is imported, so it still works when the installed build is itself the reason CPU is wanted.
This commit is contained in:
committed by
capy-ai-staging[bot]
parent
029d4d3378
commit
8b8c1429db
@@ -6,6 +6,7 @@ voice prompt combination, and model loading progress tracking.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import os
|
||||||
import platform
|
import platform
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@@ -77,6 +78,13 @@ def is_model_cached(
|
|||||||
return False
|
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(
|
def get_torch_device(
|
||||||
*,
|
*,
|
||||||
allow_xpu: bool = False,
|
allow_xpu: bool = False,
|
||||||
@@ -92,7 +100,16 @@ def get_torch_device(
|
|||||||
allow_directml: Check for DirectML (Windows) support.
|
allow_directml: Check for DirectML (Windows) support.
|
||||||
allow_mps: Allow MPS (Apple Silicon). If False, MPS falls back to CPU.
|
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.
|
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:
|
||||||
|
return "cpu"
|
||||||
|
|
||||||
if force_cpu_on_mac and platform.system() == "Darwin":
|
if force_cpu_on_mac and platform.system() == "Darwin":
|
||||||
return "cpu"
|
return "cpu"
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,80 @@
|
|||||||
|
"""
|
||||||
|
Regression tests for the VOICEBOX_FORCE_CPU environment override.
|
||||||
|
|
||||||
|
The docs promise (docs/content/docs/developer/tts-generation.mdx) that
|
||||||
|
get_torch_device() layers "VOICEBOX_FORCE_CPU environment override" ahead of
|
||||||
|
CUDA/XPU/MPS detection, and gpu-acceleration.mdx tells users to set it to fall
|
||||||
|
back to CPU when the bundled PyTorch has no kernels for their GPU.
|
||||||
|
|
||||||
|
torch is stubbed through sys.modules so these run without a torch install and
|
||||||
|
without any GPU.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python -m pytest backend/tests/test_force_cpu_env.py -v
|
||||||
|
"""
|
||||||
|
|
||||||
|
import sys
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from backend.backends.base import get_torch_device
|
||||||
|
|
||||||
|
# The documented public name and value of the override. Pinned here independently
|
||||||
|
# of the production constants so a rename of either fails these tests.
|
||||||
|
FORCE_CPU_ENV_VAR = "VOICEBOX_FORCE_CPU"
|
||||||
|
|
||||||
|
# Sentinel for "the variable is not set at all".
|
||||||
|
UNSET = None
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeCuda:
|
||||||
|
@staticmethod
|
||||||
|
def is_available() -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeTorch:
|
||||||
|
"""Minimal stand-in for a CUDA-enabled torch install."""
|
||||||
|
|
||||||
|
cuda = _FakeCuda
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def cuda_available(monkeypatch):
|
||||||
|
"""Make torch report a usable CUDA device without installing torch."""
|
||||||
|
monkeypatch.setitem(sys.modules, "torch", _FakeTorch)
|
||||||
|
|
||||||
|
|
||||||
|
def _set_override(monkeypatch, value):
|
||||||
|
if value is UNSET:
|
||||||
|
monkeypatch.delenv(FORCE_CPU_ENV_VAR, raising=False)
|
||||||
|
else:
|
||||||
|
monkeypatch.setenv(FORCE_CPU_ENV_VAR, value)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("value", ["1", " 1 "])
|
||||||
|
def test_force_cpu_wins_over_available_cuda(monkeypatch, cuda_available, value):
|
||||||
|
"""The documented value must beat an otherwise usable CUDA device.
|
||||||
|
|
||||||
|
Surrounding whitespace is tolerated: on Windows, where this override
|
||||||
|
matters most, it is typically set through the GUI environment editor."""
|
||||||
|
_set_override(monkeypatch, value)
|
||||||
|
|
||||||
|
assert get_torch_device(allow_xpu=True, allow_directml=True, allow_mps=True) == "cpu"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("value", [UNSET, "", "0"])
|
||||||
|
def test_without_override_cuda_is_still_selected(monkeypatch, cuda_available, value):
|
||||||
|
"""Unset or disabled must not disturb normal device detection."""
|
||||||
|
_set_override(monkeypatch, value)
|
||||||
|
|
||||||
|
assert get_torch_device() == "cuda"
|
||||||
|
|
||||||
|
|
||||||
|
def test_force_cpu_does_not_need_torch(monkeypatch):
|
||||||
|
"""The override is honoured before torch is imported, so it works on a
|
||||||
|
broken/incompatible torch install — which is the case it exists for."""
|
||||||
|
_set_override(monkeypatch, "1")
|
||||||
|
monkeypatch.setitem(sys.modules, "torch", None) # makes `import torch` raise
|
||||||
|
|
||||||
|
assert get_torch_device() == "cpu"
|
||||||
Reference in New Issue
Block a user